packages feed

srtree-db 0.1.3.0 → 0.1.3.2

raw patch · 4 files changed

+88/−13 lines, 4 filesPVP: minor bump suggested

API additions: PVP suggests at least a minor version bump

API changes (from Hackage documentation)

+ Algorithm.EqSat.Storage.Query: topNFiltered :: SqlBackend db => db -> Int -> Int -> [(String, String, Int)] -> IO [(EClassId, Double)]

Files

ChangeLog.md view
@@ -1,5 +1,22 @@ # Changelog for srtree-db +## 0.1.3.2++- **fitdata: NLopt restart-scale escalation** (`fitOneNLopt`): parameterized+  expressions whose valid region is far from 0 (e.g. `Log(x0+t0)` needs+  `t0 > -min(x0)`) were wrongly flagged invalid because every `[-1,1]` random+  restart landed in the NaN region. Restarts now escalate through scales+  `[1, 10, 100, 1000]` until a finite fit is found (only escalating if all+  restarts at the current scale are NaN/Inf, so normal fits are unchanged).+  Genuinely unfittable expressions (e.g. `Log(x1*t0)` when `x1` spans both+  signs) still fail at every scale and remain flagged.+- **Query helper**: added `topNFiltered` for `topN` with SQL-level filters on+  `size` and computed `n_params` (cost filters handled by the caller).++## 0.1.3.1++- **Fix**: spawning multiple threads even when using `-N` flag to limit it. + ## 0.1.3.0  - **New CLI subcommands**:
app/FitData.hs view
@@ -10,7 +10,7 @@   , runRefit   ) where -import Control.Concurrent (getNumCapabilities, threadDelay)+import Control.Concurrent (getNumCapabilities, setNumCapabilities, threadDelay) import Control.Concurrent.Async (mapConcurrently_, mapConcurrently) import Control.Monad (replicateM, when, unless, void, forM_) import Control.Exception (bracket, SomeException, catch, SomeAsyncException(..))@@ -56,6 +56,7 @@   , fitdataNIter       :: Int   , fitdataBatchSize   :: Int   , fitdataQuiet       :: Bool+  , fitdataJobs        :: Int   } deriving (Show)  fitdataParser :: Parser FitDataOpts@@ -103,11 +104,18 @@       ( long "quiet"       <> short 'q'       <> help "Suppress per-expression output; print progress every 10k expressions" )+  <*> option auto+      ( long "jobs"+      <> short 'j'+      <> value 0+      <> metavar "N"+      <> help "Number of parallel workers (0 = single-threaded, default)" )  -- | Run the fitdata sub-command. runFitData :: FitDataOpts -> IO () runFitData opts = do   let FitDataOpts{..} = opts+  when (fitdataJobs > 0) $ setNumCapabilities fitdataJobs   putStrLn $ "Loading dataset: " ++ fitdataData   hFlush stdout   ((xTrain, yTrain, _xVal, _yVal), (mYErr, _), _varnames, _target) <-@@ -188,9 +196,9 @@           -- Phase 4-5: parallel NLopt + batch write (uses fitDb)           let fitPhase pendingRef survivors = do                 let chunks = chunk nCaps survivors-                setMTPopParallel False-                results <- fmap concat $ mapConcurrently (mapM (fitOneNLopt fitdataQuiet xTrain yTrain mYErr fitdataLoss fitdataNIter fitdataNRep counter)) chunks                 setMTPopParallel True+                results <- fmap concat $ mapConcurrently (mapM (fitOneNLopt fitdataQuiet xTrain yTrain mYErr fitdataLoss fitdataNIter fitdataNRep counter)) chunks+                setMTPopParallel False                 -- Batch write all pending fits (invalid + analytical + NLopt)                 pending <- atomicModifyIORef' pendingRef (\ps -> ([], ps))                 execDb fitDb "BEGIN"@@ -343,21 +351,34 @@   writeDatasetFit db dsid (frEid fr) (Just (frFitness fr)) Nothing (frTheta fr) (frSize fr)  -- | Pure NLopt fit (no DB side effects). Returns a FitResult.+-- Restarts start at several increasing scales so a valid region far from 0+-- (e.g. ``Log(x0+t0)`` needs ``t0 > -min(x0)``) is reachable, instead of every+-- ``[-1,1]`` restart landing in the NaN region and the expression being wrongly+-- flagged invalid. Only escalates if all restarts at the current scale are+-- NaN/Inf, so normal fits are unchanged. fitOneNLopt :: Bool             -> [VU.Vector Double] -> VU.Vector Double -> Maybe (VU.Vector Double)             -> Loss -> Int -> Int -> IORef Int -> FitJob -> IO FitResult fitOneNLopt quiet xTrain yTrain mYErr loss nIter nRep counter (FitJob eid tree' _ np sz) = do   let funAndGrad = compileLossAndGrad MultiThread loss mYErr xTrain yTrain tree'-      runRestart = do-        theta0 <- VU.replicateM np (randomRIO (-1, 1))+      runRestart scale = do+        theta0 <- VU.replicateM np (randomRIO (-scale, scale))         let (theta, lossVal, _) = minimizeNLLWith funAndGrad VAR1 nIter theta0         pure (negate lossVal, theta)-  results <- replicateM nRep runRestart-  let (bestFitness, bestTheta) = maximumBy (comparing fst) results-      !thetaText = T.pack (serializeTheta [bestTheta])-  atomicModifyIORef' counter (\n -> let !n' = n + 1 in (n', ()))-  unless quiet $ putStrLn $ "  eclass " ++ show eid ++ " (" ++ takeExpr tree' ++ "): fitness=" ++ showFit bestFitness-  pure (FitResult eid bestFitness thetaText sz)+      go [] = do+        atomicModifyIORef' counter (\n -> let !n' = n + 1 in (n', ()))+        unless quiet $ putStrLn $ "  eclass " ++ show eid ++ " (" ++ takeExpr tree' ++ "): no finite fit (NaN)"+        pure (FitResult eid (negate (1/0)) (T.pack (serializeTheta [])) sz)+      go (s : rest) = do+        rs <- replicateM nRep (runRestart s)+        let (bf, bt) = maximumBy (comparing fst) rs+        if isInvalid bf+          then go rest+          else do+            atomicModifyIORef' counter (\n -> let !n' = n + 1 in (n', ()))+            unless quiet $ putStrLn $ "  eclass " ++ show eid ++ " (" ++ takeExpr tree' ++ "): fitness=" ++ showFit bf+            pure (FitResult eid bf (T.pack (serializeTheta [bt])) sz)+  go [1.0, 10.0, 100.0, 1000.0]  -- | Whether a fitness value is unusable (NaN or +/-Infinity), so it can be -- pruned and propagated to ancestors.@@ -482,6 +503,7 @@ runRefit :: FitDataOpts -> IO () runRefit opts = do   let FitDataOpts{..} = opts+  when (fitdataJobs > 0) $ setNumCapabilities fitdataJobs   putStrLn $ "Refitting dataset: " ++ fitdataDataset   putStrLn $ "Clearing previous fit data..."   hFlush stdout
src/Algorithm/EqSat/Storage/Query.hs view
@@ -15,6 +15,7 @@   , readDatasetFit   , topN   , topNIn+  , topNFiltered   , pareto   , paretoBySize   , distributionCounts@@ -117,6 +118,41 @@   pure [ (sqlToInt eid, f)        | [eid, f] <- rows        , Just f   <- [sqlToMaybeDouble f] ]++-- | Like 'topN' but with SQL-level filters on @size@ and computed @n_params@.+-- Each filter is a triple @(field, op, value)@ where field is @\"size\"@ or+-- @\"parameters\"@, op is one of @\"<\"@, @\"<=\"@, @\"=\"@, @\">=\"@, @\">\"@,+-- and value is the integer threshold.+-- Filters on @\"cost\"@ are ignored here (handled post-query by the caller).+topNFiltered :: SqlBackend db => db -> Int -> Int -> [(String, String, Int)] -> IO [(EClassId, Double)]+topNFiltered db ds n filters = do+  let sqlFilters = [ (field, op, v) | (field, op, v) <- filters, field /= "cost" ]+      (extraClauses, extraParams) = unzip $ map mkFilterClause sqlFilters+      baseQ = "SELECT eid, fitness FROM dataset_fit \+              \WHERE dataset_id = ? AND fitness IS NOT NULL"+              <> T.concat extraClauses+              <> " ORDER BY fitness DESC LIMIT ?"+      params = [SqlInteger (fromIntegral ds)] ++ extraParams ++ [SqlInteger (fromIntegral n)]+  rows <- queryDb db baseQ params+  pure [ (sqlToInt eid, f)+       | [eid, f] <- rows+       , Just f   <- [sqlToMaybeDouble f] ]++-- | Build a SQL clause fragment for a single filter.+mkFilterClause :: (String, String, Int) -> (Text, SqlValue)+mkFilterClause (field, op, v) =+  let col = case field of+              "size"       -> "size"+              "parameters" -> "(CASE WHEN theta = '' THEN 0 ELSE LENGTH(theta) - LENGTH(REPLACE(theta, ',', '')) + 1 END)"+              _            -> "size"  -- fallback+      sqlOp = case op of+                "<"  -> " < "+                "<=" -> " <= "+                "="  -> " = "+                ">=" -> " >= "+                ">"  -> " > "+                _    -> " = "+  in (" AND " <> T.pack col <> sqlOp <> "?", SqlInteger (fromIntegral v))  -- | Non-dominated classes over (max fitness, min dl) on the dataset. Returns -- the (eid, fitness, dl) triples that are not dominated by any other class.
srtree-db.cabal view
@@ -1,6 +1,6 @@ cabal-version: 2.4 name:          srtree-db-version:       0.1.3.0+version:       0.1.3.2 synopsis:      SQL persistence and querying for srtree e-graphs description:   Reusable storage layer for srtree multiset e-graphs:                a driver-neutral serialization of the Algorithm.EqSat.Store@@ -70,7 +70,7 @@   hs-source-dirs:   app   main-is:          Main.hs   other-modules:    Ingest, EqSat, FitData, Status, Export, Backfill, RandomSampler-  ghc-options:      -O2 -threaded -rtsopts -with-rtsopts=-N+  ghc-options:      -O2 -threaded -rtsopts -with-rtsopts=-N1   build-depends:       base >=4.14 && <5     , optparse-applicative >=0.17 && <0.19