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 +17/−0
- app/FitData.hs +33/−11
- src/Algorithm/EqSat/Storage/Query.hs +36/−0
- srtree-db.cabal +2/−2
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