srtree 3.0.0.3 → 3.0.0.4
raw patch · 11 files changed
+197/−98 lines, 11 filesPVP: major bump suggested
API removals or changes: PVP suggests a major version bump
API changes (from Hackage documentation)
+ Algorithm.EqSat: preLoadFrontier :: forall (m :: Type -> Type). ClassStore m => EGraphST m ()
+ Algorithm.EqSat.Build: evictOldestPct :: forall (m :: Type -> Type). ClassStore m => Int -> EGraphST m ()
+ Algorithm.EqSat.Egraph: [_canonicalCap] :: EGraphDB -> Int
+ Algorithm.EqSat.Egraph: [_nodeCap] :: EGraphDB -> Int
+ Algorithm.EqSat.Egraph: [_residentCap] :: EGraphDB -> Int
+ Algorithm.EqSat.Egraph: [cpsBulkLookup] :: EClassPageStore -> [EClassId] -> IO (IntMap EClass)
+ Algorithm.EqSat.Egraph: bulkLoad :: ClassStore m => [EClassId] -> EGraphST m ()
+ Algorithm.EqSat.Egraph: canonicalCap :: Lens' EGraphDB Int
+ Algorithm.EqSat.Egraph: nodeCap :: Lens' EGraphDB Int
+ Algorithm.EqSat.Egraph: residentCap :: Lens' EGraphDB Int
- Algorithm.EqSat.Egraph: EClassPageStore :: (EClassId -> IO (Maybe EClass)) -> (EClass -> IO ()) -> (EClassId -> IO ()) -> IO () -> IO [EClass] -> IO [EClassId] -> (SRTree () -> Int -> [EClassId] -> IO [EClassId]) -> (ENode -> EClassId -> IO ()) -> (ENode -> IO (Maybe EClassId)) -> (EClassId -> IO (Maybe EClassId)) -> (EClassId -> EClassId -> IO ()) -> IO () -> IO () -> EClassPageStore
+ Algorithm.EqSat.Egraph: EClassPageStore :: (EClassId -> IO (Maybe EClass)) -> ([EClassId] -> IO (IntMap EClass)) -> (EClass -> IO ()) -> (EClassId -> IO ()) -> IO () -> IO [EClass] -> IO [EClassId] -> (SRTree () -> Int -> [EClassId] -> IO [EClassId]) -> (ENode -> EClassId -> IO ()) -> (ENode -> IO (Maybe EClassId)) -> (EClassId -> IO (Maybe EClassId)) -> (EClassId -> EClassId -> IO ()) -> IO () -> IO () -> EClassPageStore
- Algorithm.EqSat.Egraph: EDB :: HashSet (EClassId, ENode) -> HashSet (EClassId, ENode) -> IntSet -> DB -> RangeTree Double -> RangeTree Double -> IntMap IntSet -> IntMap (RangeTree Double) -> IntMap (RangeTree Double) -> IntSet -> Int -> Bool -> Bool -> Map String (Set String) -> EGraphDB
+ Algorithm.EqSat.Egraph: EDB :: HashSet (EClassId, ENode) -> HashSet (EClassId, ENode) -> IntSet -> DB -> RangeTree Double -> RangeTree Double -> IntMap IntSet -> IntMap (RangeTree Double) -> IntMap (RangeTree Double) -> IntSet -> Int -> Bool -> Bool -> Map String (Set String) -> Int -> Int -> Int -> EGraphDB
- Algorithm.SRTree.ConfidenceIntervals: getAllProfiles :: PType -> EvalTree -> Target -> Target -> [CI] -> Double -> [ProfileT]
+ Algorithm.SRTree.ConfidenceIntervals: getAllProfiles :: PType -> EvalTree -> Target -> Target -> [CI] -> Double -> IO [ProfileT]
- Algorithm.SRTree.ConfidenceIntervals: getProfile :: EvalTree -> Target -> Double -> Double -> Int -> Either Target ProfileT
+ Algorithm.SRTree.ConfidenceIntervals: getProfile :: EvalTree -> Target -> Double -> Double -> Int -> IO ProfileT
Files
- ChangeLog.md +5/−0
- apps/Report/Main.hs +1/−1
- apps/TestCI/Main.hs +6/−6
- src/Algorithm/EqSat.hs +22/−3
- src/Algorithm/EqSat/Build.hs +23/−0
- src/Algorithm/EqSat/DB.hs +3/−3
- src/Algorithm/EqSat/Egraph.hs +43/−13
- src/Algorithm/SRTree/ConfidenceIntervals.hs +71/−58
- src/Algorithm/SRTree/Utils.hs +21/−12
- src/Text/ParseSR.hs +1/−1
- srtree.cabal +1/−1
ChangeLog.md view
@@ -1,5 +1,10 @@ # Changelog for srtree +## 3.0.0.4++- **Export `createLoss`**: expose compiled loss function for external use+- **Fix `paretoFront` type signature**: corrected return type in `SearchSR.hs`+ ## 3.0.0.3 - **Profile-likelihood CI overhaul** (`ConfidenceIntervals`):
apps/Report/Main.hs view
@@ -169,7 +169,7 @@ cis <- case raCI args of LaplaceCI -> return laplaceCIs ProfileCI -> do- let profiles = getAllProfiles ptype et thetaOpt (_stdErr stats) laplaceCIs (raAlpha args)+ profiles <- getAllProfiles ptype et thetaOpt (_stdErr stats) laplaceCIs (raAlpha args) when (raDbg args) $ forM_ (zip [0..] profiles) $ \(i, ProfileT taus thetas _ tau2theta _) -> do putStrLn $ "DEBUG Profile " ++ show i ++ " (opt=" ++ show (thetaOpt U.! i) ++ "):" putStrLn $ " tau range: [" ++ show (if U.null taus then 0 else U.head taus)
apps/TestCI/Main.hs view
@@ -208,8 +208,8 @@ putStrLn " === BATES (profile walk) ===" catch (do let estCIs = laplaceCI- profiles_bates = getAllProfiles Bates et theta_opt stdErrs estCIs 0.05- batesCI = paramCI (Profile stats profiles_bates) nSamples theta_opt 0.05+ profiles_bates <- getAllProfiles Bates et theta_opt stdErrs estCIs 0.05+ let batesCI = paramCI (Profile stats profiles_bates) nSamples theta_opt 0.05 putStrLn $ " 95% CIs:" putStrLn $ " " ++ showCIList (zip paramNames batesCI) putStrLn $ " Widths: " ++ show (map (\(CI _ l h) -> h - l) batesCI)@@ -234,8 +234,8 @@ putStrLn " === ODE (Chen & Jennrich) ===" catch (do let estCIs = laplaceCI- profiles_ode = getAllProfiles ODE et theta_opt stdErrs estCIs 0.05- odeCI = paramCI (Profile stats profiles_ode) nSamples theta_opt 0.05+ profiles_ode <- getAllProfiles ODE et theta_opt stdErrs estCIs 0.05+ let odeCI = paramCI (Profile stats profiles_ode) nSamples theta_opt 0.05 putStrLn $ " 95% CIs:" putStrLn $ " " ++ showCIList (zip paramNames odeCI) putStrLn $ " Widths: " ++ show (map (\(CI _ l h) -> h - l) odeCI)@@ -245,8 +245,8 @@ -- ---- CONSTRAINED ---- putStrLn " === CONSTRAINED (bisection) ===" catch (do- let profiles_cnstr = getAllProfiles Constrained et theta_opt stdErrs [] 0.05- cnstrCI = paramCI (Profile stats profiles_cnstr) nSamples theta_opt 0.05+ profiles_cnstr <- getAllProfiles Constrained et theta_opt stdErrs [] 0.05+ let cnstrCI = paramCI (Profile stats profiles_cnstr) nSamples theta_opt 0.05 putStrLn $ " 95% CIs:" putStrLn $ " " ++ showCIList (zip paramNames cnstrCI) putStrLn $ " Widths: " ++ show (map (\(CI _ l h) -> h - l) cnstrCI)
src/Algorithm/EqSat.hs view
@@ -28,7 +28,7 @@ import Data.List (intercalate) import Data.Map (Map) import qualified Data.Map as Map-import Data.Maybe (mapMaybe)+import Data.Maybe (mapMaybe, isJust) import Data.SRTree import Data.HashSet (HashSet) import qualified Data.HashSet as Set@@ -293,7 +293,7 @@ -- matches), this bounds a single iteration's apply/rebuild work regardless of -- graph size. iterMatchBudget :: Int-iterMatchBudget = 2000+iterMatchBudget = 500 -- | run equality saturation for a number of iterations runEqSat :: ClassStore m => CostFun -> [Rule] -> Int -> EGraphST m (Bool, Int)@@ -306,6 +306,9 @@ do -- reset dirty flag before processing this iteration modify' $ over (eDB . changed) (const False) + -- NEW: pre-load frontier transitive closure to warm the cache+ preLoadFrontier+ -- step 1: match the rules using cached compiled queries let matchSch = matchWithScheduler it adapted i (r, cq) = map (,cq) <$> matchSch i r@@ -336,7 +339,8 @@ else go (it-1) sch' compiled throttle it sch compiled = do- cleanMaps+ -- Instead of wiping all caches, evict oldest 50% to preserve warm state+ evictOldestPct 50 eClasses <- gets _eClass if IntMap.size eClasses <= 1500 then go (it-1) sch compiled@@ -345,6 +349,21 @@ if it <= 1 || not changed then pure (False, it) -- give up and return early stop else throttle (it-1) sch compiled++-- | Pre-load pages for recently-changed classes into the resident cache.+-- This ensures the matcher's hot path is cache-warm, reducing I/O during+-- the matching phase. Only does work on paged graphs.+preLoadFrontier :: ClassStore m => EGraphST m ()+preLoadFrontier = do+ hasStore <- gets (isJust . _classStore)+ if not hasStore then pure ()+ else do+ -- Load pages for all classes in the worklist and analysis set+ wl <- gets (Set.map fst . _worklist . _eDB)+ al <- gets (Set.map fst . _analysis . _eDB)+ let toLoad = IntSet.toList (Set.foldl' (flip IntSet.insert) IntSet.empty (Set.union wl al))+ -- Touch each class to trigger page load into resident cache+ mapM_ (\eid -> lookupClass eid >> pure ()) toLoad -- | apply a single step of merge-only equality saturation applySingleMergeOnlyEqSat :: ClassStore m => CostFun -> [Rule] -> EGraphST m ()
src/Algorithm/EqSat/Build.hs view
@@ -180,6 +180,12 @@ al <- gets (_analysis . _eDB) modify' $ over (eDB . worklist) (const Set.empty) . over (eDB . analysis) (const Set.empty)+ -- Batch-load all dirty class pages before processing+ -- This eliminates I/O cascades during repair/repairAnalysis+ let allIds = Set.foldl' (\s (eid, _) -> IntSet.insert eid s) IntSet.empty wl+ `IntSet.union`+ Set.foldl' (\s (eid, _) -> IntSet.insert eid s) IntSet.empty al+ bulkLoad (IntSet.toList allIds) forM_ wl (uncurry (repair costFun)) forM_ al (uncurry (repairAnalysis costFun)) {-# INLINE rebuild #-}@@ -741,3 +747,20 @@ modify' $ \eg -> eg { _eNodeToEClass = enode2eclass' , _eClass = eclassMap' } {-# INLINE cleanMaps #-}++-- | Evict the oldest @pct@ percent of entries from the resident caches.+-- For paged graphs, this selectively drops entries instead of wiping all caches+-- (which would destroy warm state). For resident graphs, this is a no-op.+evictOldestPct :: ClassStore m => Int -> EGraphST m ()+evictOldestPct pct+ | pct <= 0 || pct >= 100 = pure ()+ | otherwise = do+ hasStore <- gets (isJust . _classStore)+ when hasStore $ modify' $ \eg ->+ let m = _eClass eg+ n = IntMap.size m+ keep = n * (100 - pct) `div` 100+ in if keep < n && keep > 0+ then over eClass (const (IntMap.fromList (Prelude.drop (n - keep) (IntMap.toAscList m)))) eg+ else eg+{-# INLINE evictOldestPct #-}
src/Algorithm/EqSat/DB.hs view
@@ -292,19 +292,19 @@ -- Capping root visits bounds the *search work* independently of the result -- count. Sound: we only stop enumerating (fewer) genuine matches early. ruleRootVisit :: Int-ruleRootVisit = 512+ruleRootVisit = 256 -- | Cap on how many matches a non-n-ary rule (the cached @genericJoin@ path) -- may return per match. The n-ary matcher has 'ruleBudget'; give the cached -- path a separate (larger) budget so a single rule cannot flood the iteration. ruleMatchBudget :: Int-ruleMatchBudget = 1024+ruleMatchBudget = 256 -- | Cap on how many operator-root e-classes the streaming cached matcher visits -- per match, bounding the search work (and the page reads) independently of the -- result count, exactly as 'ruleRootVisit' does for the n-ary matcher. ruleMatchRootVisit :: Int-ruleMatchRootVisit = 2048+ruleMatchRootVisit = 512 -- | Match an n-ary pattern against every root e-class of its operator trie. --
src/Algorithm/EqSat/Egraph.hs view
@@ -24,7 +24,7 @@ import Control.Lens (element, makeLenses, view, over, (&), (+~), (-~), (.~), (^.)) --import Control.Monad (forM_, when, foldM, void) import Data.List ( intercalate, foldl' )-import Control.Monad (forM)+import Control.Monad (forM, unless) import Control.Monad.State.Strict hiding ( get, put ) import Control.Monad.IO.Class (MonadIO(..)) import Data.Functor.Identity (Identity)@@ -143,6 +143,7 @@ -- behaviour. data EClassPageStore = EClassPageStore { cpsLookup :: EClassId -> IO (Maybe EClass)+ , cpsBulkLookup :: [EClassId] -> IO (IntMap.IntMap EClass) -- ^ bulk-load pages for multiple eclasses , cpsInsert :: EClass -> IO () , cpsDelete :: EClassId -> IO () , cpsFlush :: IO () -- ^ write back all pending dirty pages@@ -178,6 +179,9 @@ , _changed :: !Bool -- dirty flag: true if modified since last check , _trackDBs :: !Bool -- maintain range DBs (False during pure simplify) , _seenMatches :: Map String (RangeSet.Set String) -- persistent (rule source -> attempted match keys)+ , _residentCap :: !Int -- resident class cache capacity (default 50000)+ , _nodeCap :: !Int -- node-to-class cache capacity (default 100000)+ , _canonicalCap :: !Int -- canonical map cache capacity (default 100000) } deriving (Show, Generic) data EClass = EClass { _eClassId :: {-# UNPACK #-} !Int -- e-class id (maybe we don't need that here)@@ -264,10 +268,11 @@ instance Binary EClassData -- Custom: keep `_trackDBs` out of the wire format so on-disk EGraphDB data -- (written before the flag existed) decodes unchanged; it defaults to True.+-- Cache cap fields are runtime-only configuration, not serialized. instance Binary EGraphDB where- put (EDB w a r p f d s sf sdl u n c _ _) =+ put (EDB w a r p f d s sf sdl u n c _ _ _ _ _) = put w >> put a >> put r >> put p >> put f >> put d >> put s >> put sf >> put sdl >> put u >> put n >> put c- get = EDB <$> get <*> get <*> get <*> get <*> get <*> get <*> get <*> get <*> get <*> get <*> get <*> get <*> pure True <*> pure Map.empty+ get = EDB <$> get <*> get <*> get <*> get <*> get <*> get <*> get <*> get <*> get <*> get <*> get <*> get <*> pure True <*> pure Map.empty <*> pure 50000 <*> pure 100000 <*> pure 100000 -- Custom: the wire format omits `_classStore` (a runtime handle to the paged -- store, never serialized); it decodes to Nothing. instance Binary EGraph where@@ -357,6 +362,13 @@ -- falls back to the store. canonicalOf :: EClassId -> EGraphST m (Maybe EClassId) canonicalOf eid = gets (IntMap.lookup eid . _canonicalMap)+ -- | Bulk-load pages for the given e-class ids into the resident cache.+ -- For paged graphs, this triggers a single SQL query instead of per-class+ -- lookups. For resident graphs, this is a no-op (all classes are already+ -- in memory). Used by 'rebuild' to warm the cache before processing the+ -- worklist, eliminating I/O cascades during repair/repairAnalysis.+ bulkLoad :: [EClassId] -> EGraphST m ()+ bulkLoad _ = pure () -- | Default candidate-root enumeration from the resident @_patDB@ trie, capped -- at @budget@ after skipping @exclude@ (used by the pure instances and as the@@ -411,7 +423,7 @@ residentClassCap :: Int residentClassCap = 50000 --- | Trim the resident @_eClass@ cache to at most 'residentClassCap' entries+-- | Trim the resident @_eClass@ cache to at most '_residentCap' entries -- by keeping the largest ids. No-op for graphs without a paged store (their -- resident map must stay complete for the pure instances). Halving on 2x keeps -- steady churn from triggering an O(n) rebuild on every insert.@@ -420,11 +432,12 @@ case _classStore eg of Nothing -> eg Just _ ->- let m = _eClass eg+ let cap = _residentCap (_eDB eg)+ m = _eClass eg n = IntMap.size m- in if n <= 2 * residentClassCap+ in if n <= 2 * cap then eg- else over eClass (const (IntMap.fromList (Prelude.drop (n - residentClassCap) (IntMap.toAscList m)))) eg+ else over eClass (const (IntMap.fromList (Prelude.drop (n - cap) (IntMap.toAscList m)))) eg -- | Bound on the resident @_eNodeToEClass@ cache on a paged graph. Beyond the -- cap (checked at 2x, halved back to cap) the map is pruned; the backing store@@ -445,11 +458,12 @@ case _classStore eg of Nothing -> eg Just _ ->- let m = _eNodeToEClass eg+ let cap = _nodeCap (_eDB eg)+ m = _eNodeToEClass eg n = HashMap.size m- in if n <= 2 * nodeCacheCap+ in if n <= 2 * cap then eg- else over eNodeToEClass (const (HashMap.fromList (Prelude.take nodeCacheCap (HashMap.toList m)))) eg+ else over eNodeToEClass (const (HashMap.fromList (Prelude.take cap (HashMap.toList m)))) eg {-# INLINE trimNodeCache #-} trimCanonicalCache :: Monad m => EGraphST m ()@@ -457,11 +471,12 @@ case _classStore eg of Nothing -> eg Just _ ->- let m = _canonicalMap eg+ let cap = _canonicalCap (_eDB eg)+ m = _canonicalMap eg n = IntMap.size m- in if n <= 2 * canonicalCacheCap+ in if n <= 2 * cap then eg- else over canonicalMap (const (IntMap.fromList (Prelude.take canonicalCacheCap (IntMap.toAscList m)))) eg+ else over canonicalMap (const (IntMap.fromList (Prelude.take cap (IntMap.toAscList m)))) eg {-# INLINE trimCanonicalCache #-} instance ClassStore Identity where@@ -603,6 +618,18 @@ trimCanonicalCache pure (Just c) Nothing -> pure Nothing+ bulkLoad eids = do+ eg <- gets id+ case _classStore eg of+ Nothing -> pure () -- resident graph: nothing to do+ Just h -> do+ -- Filter out already-cached eclasses to avoid unnecessary I/O+ let cached = IntMap.keysSet (_eClass eg)+ toLoad = filter (\eid -> not (IntSet.member eid cached)) eids+ unless (null toLoad) $ do+ pages <- liftIO (cpsBulkLookup h toLoad)+ modify' $ \eg' -> eg' { _eClass = IntMap.union (_eClass eg') pages }+ trimResidentCache -- * E-Graph basic supporting functions @@ -628,6 +655,9 @@ False True Map.empty+ 50000 -- _residentCap+ 100000 -- _nodeCap+ 100000 -- _canonicalCap {-# INLINE emptyDB #-} -- | like 'emptyDB' but skips range-DB maintenance (pure simplify mode)
src/Algorithm/SRTree/ConfidenceIntervals.hs view
@@ -27,8 +27,10 @@ import Data.Maybe ( listToMaybe ) import Algorithm.SRTree.Utils import Numeric.Optimization.NLOPT+import Data.IORef import System.IO.Unsafe ( unsafePerformIO ) import Control.Monad.Catch ( catch, SomeException )+import Control.Exception ( evaluate ) import Debug.Trace ( trace ) -- | profile likelihood algorithms: Bates (classical), ODE (faster), Constrained (fastest)@@ -190,17 +192,18 @@ k = U.length t ident = fromRowMajor k k (U.generate (k * k) (\ix -> let (i, j) = ix `divMod` k in if i == j then 1.0 else 0.0)) hess = ctHessianNLL et t- cov = unsafePerformIO $ catch (invChol hess) (\(_ :: SomeException) -> pure ident)+ cov = unsafePerformIO $ catch (invChol hess >>= evaluate) (\(_ :: SomeException) -> pure ident) covMat = toRowMajor cov stdErr = U.generate k (\ix -> sqrt $ abs (covMat U.! (ix * k + ix))) --- calculate the profile likelihood of every parameter--- restartLimit bounds recursive restarts when the optimizer finds a better point mid-profile-getAllProfiles :: PType -> EvalTree -> Target -> Target -> [CI] -> Double -> [ProfileT]+-- | Calculate the profile likelihood of every parameter.+-- restartLimit bounds recursive restarts when getProfileODE finds a better point mid-profile.+-- For Bates, getProfile handles restarts internally and never returns Left.+getAllProfiles :: PType -> EvalTree -> Target -> Target -> [CI] -> Double -> IO [ProfileT] getAllProfiles ptype et theta stdErr estCIs alpha -- Defensive: if theta is too short for the EvalTree's distribution, -- return empty profiles instead of crashing (e.g. MSE loss with Gaussian dist)- | U.length theta < 2 = []+ | U.length theta < 2 = pure [] | otherwise = go 0 et theta stdErr estCIs where restartLimit = 5 :: Int@@ -229,7 +232,7 @@ estCIs'' = if null estCIs' then let ident = U.generate (k * k) (\ix -> let (i, j) = ix `divMod` k in if i == j then 1.0 else 0.0) hess = ctHessianNLL et' theta'- cov = unsafePerformIO $ catch (invChol hess) (\(_ :: SomeException) -> pure (fromRowMajor k k ident))+ cov = unsafePerformIO $ catch (invChol hess >>= evaluate) (\(_ :: SomeException) -> pure (fromRowMajor k k ident)) covMat = toRowMajor cov se = U.generate k (\ix -> sqrt $ abs (covMat U.! (ix * k + ix))) tVal = quantile (studentT . fromIntegral $ n - k) (1 - alpha / 2.0)@@ -237,74 +240,84 @@ else estCIs' profFun ix = case ptype of- Bates -> getProfile et' theta' (stdErr' U.! ix) tau_max ix- ODE -> getProfileODE et' theta' (stdErr' U.! ix) (estCIs'' !! ix) tau_max ix- Constrained -> getProfileCnstr et' theta' (stdErr' U.! ix) tau_max' ix+ Bates -> Right <$> getProfile et' theta' (stdErr' U.! ix) tau_max ix+ ODE -> pure $ getProfileODE et' theta' (stdErr' U.! ix) (estCIs'' !! ix) tau_max ix+ Constrained -> pure $ getProfileCnstr et' theta' (stdErr' U.! ix) tau_max' ix - go' ix acc | ix == k = acc+ go' ix acc | ix == k = pure acc go' ix acc | ix == k-1 && ptype == Constrained && ctDist et' == Gaussian = case getProfileODE et' theta' (stdErr' U.! ix) (estCIs'' !! ix) tau_max ix of Left t -> let tOpt = ctOptimizer et' t; se'' = recomputeStdErr et' tOpt in go (restarts + 1) et' tOpt se'' estCIs' Right p -> go' (ix + 1) (acc <> [p])- | otherwise =- case profFun ix of+ | otherwise = do+ result <- profFun ix+ case result of Left t -> let tOpt = ctOptimizer et' t; se'' = recomputeStdErr et' tOpt in go (restarts + 1) et' tOpt se'' estCIs' Right p -> go' (ix + 1) (acc <> [p]) --- calculates the profile likelihood of a single parameter-getProfile :: EvalTree -> Target -> Double -> Double -> Int -> Either Target ProfileT+-- | Calculate the profile likelihood of a single parameter.+-- When a better optimum is found mid-walk, the walk restarts from the new MLE+-- internally (discarding previously collected points for this parameter only),+-- rather than propagating a restart to getAllProfiles.+getProfile :: EvalTree -> Target -> Double -> Double -> Int -> IO ProfileT getProfile et theta stdErr_i tau_max ix | stdErr_i == 0.0 = pure $ ProfileT (U.fromList [-tau_max, tau_max]) [theta, theta] (theta U.! ix) (const (theta U.! ix)) (const tau_max)- | otherwise =- do negDelta <- go kmax (-stdErr_i / 8) 0 1 mempty- let !negLen = length (fst negDelta)- !negTauRange = if null (fst negDelta) then (0,0) else (minimum (fst negDelta), maximum (fst negDelta))- posDelta <- go kmax (stdErr_i / 8) 0 1 p0- let !posLen = length (fst posDelta)- !posTauRange = if null (fst posDelta) then (0,0) else (minimum (fst posDelta), maximum (fst posDelta))- let (taus', thetas') = negDelta <> posDelta- taus = U.fromList taus'- thetas = thetas'- (tau2theta, theta2tau) = createSplines taus thetas stdErr_i tau_max ix optTh- pure $ ProfileT taus thetas optTh tau2theta theta2tau+ | otherwise = do+ nllOptRef <- newIORef nll_opt0+ thetaOptRef <- newIORef theta_opt0++ negDelta <- go kmax (-stdErr_i / 8) 0 1 mempty nllOptRef thetaOptRef+ thetaOpt1 <- readIORef thetaOptRef+ posDelta <- go kmax (stdErr_i / 8) 0 1 ([0], [thetaOpt1]) nllOptRef thetaOptRef++ thetaOpt2 <- readIORef thetaOptRef+ let optTh' = thetaOpt2 U.! ix+ (taus', thetas') = negDelta <> posDelta+ taus = U.fromList taus'+ thetas = thetas'+ (tau2theta, theta2tau) = createSplines taus thetas stdErr_i tau_max ix optTh'+ pure $ ProfileT taus thetas optTh' tau2theta theta2tau where- p0 = ([0], [theta_opt]) kmax = 500- nll_opt = ctNLL et theta_opt- theta_opt = ctOptimizer et theta- optTh = theta_opt U.! ix+ nll_opt0 = ctNLL et theta_opt0+ theta_opt0 = ctOptimizer et theta minimizer = ctOptimizerFixed et ix - go 0 delta _ _ acc = Right acc- go k delta t inv_slope acc@(taus, thetas)- | isNaN inv_slope = Right acc- | nll_cond < nll_opt - 1e-6 * abs nll_opt = Left theta_t- | abs tau > tau_max = Right acc'-- | otherwise = go (k-1) delta (t + inv_slope) inv_slope' acc'- where- t_delta = (theta_opt U.! ix) + delta * (t + inv_slope)- theta_delta = updateS theta_opt [(ix, t_delta)]- theta_t = minimizer theta_delta- (nll_cond, grad) = ctGradNLL et theta_t- zv = grad U.! ix- -- For LeastSquares, the correct profile likelihood statistic is- -- n * log(MSE(t)/MSE(opt)) ~ chi2_1, not 2*(MSE(t) - MSE(opt)).- tau = case ctDist et of- LeastSquares ->- let nD = fromIntegral (ctRows et) :: Double- r = max nll_cond 1e-30 / max nll_opt 1e-30- in signum delta * sqrt (max 0 (nD * log r))- _ -> signum delta * sqrt (max 0 (2*nll_cond - 2*nll_opt))- inv_slope' = if abs zv < 1e-12 * abs stdErr_i- then min 4.0 . max 0.0625 $ abs (delta * 8)- else min 4.0 . max 0.0625 . abs $ (tau / (stdErr_i * zv))- acc' = if nll_cond == nll_opt || maybe False (tau ==) (listToMaybe taus) || isNaN tau- then acc- else (tau:taus, theta_t:thetas)+ go 0 _delta _t _inv_slope acc _nllRef _thetaRef = pure acc+ go k delta t inv_slope acc@(taus, thetas) nllOptRef thetaOptRef = do+ nllOpt <- readIORef nllOptRef+ thetaOpt <- readIORef thetaOptRef+ let t_delta = (thetaOpt U.! ix) + delta * (t + inv_slope)+ theta_delta = updateS thetaOpt [(ix, t_delta)]+ validDelta = not (isNaN t_delta) && not (isInfinite t_delta)+ && not (U.any isNaN theta_delta) && not (U.any isInfinite theta_delta)+ theta_t = if validDelta then minimizer theta_delta else thetaOpt+ (nll_cond, grad) = ctGradNLL et theta_t+ zv = grad U.! ix+ tau = case ctDist et of+ LeastSquares ->+ let nD = fromIntegral (ctRows et) :: Double+ r = max nll_cond 1e-30 / max nllOpt 1e-30+ in signum delta * sqrt (max 0 (nD * log r))+ _ -> signum delta * sqrt (max 0 (2*nll_cond - 2*nllOpt))+ inv_slope' = if abs zv < 1e-12 * abs stdErr_i+ then min 4.0 . max 0.0625 $ abs (delta * 8)+ else min 4.0 . max 0.0625 . abs $ (tau / (stdErr_i * zv))+ acc' = if nll_cond == nllOpt || maybe False (tau ==) (listToMaybe taus) || isNaN tau+ then acc+ else (tau:taus, theta_t:thetas)+ if | not validDelta || isNaN inv_slope -> pure acc+ | nll_cond < nllOpt - 1e-6 * abs nllOpt -> do+ -- Better optimum found: update references and restart walk+ -- from the new MLE, discarding previously collected points.+ writeIORef nllOptRef nll_cond+ writeIORef thetaOptRef theta_t+ go kmax delta 0 1 mempty nllOptRef thetaOptRef+ | abs tau > tau_max -> pure acc'+ | otherwise -> go (k-1) delta (t + inv_slope) inv_slope' acc' nllOptRef thetaOptRef -- Based on https://insysbio.github.io/LikelihoodProfiler.jl/latest/ -- Borisov, Ivan, and Evgeny Metelkin. "Confidence intervals by constrained optimization—An algorithm and software package for practical identifiability analysis in systems biology." PLOS Computational Biology 16.12 (2020): e1008495.@@ -423,7 +436,7 @@ fexcept :: SomeException -> IO Columns fexcept _ = pure ident - covRaw = unsafePerformIO $ catch (invChol hess) fexcept+ covRaw = unsafePerformIO $ catch (invChol hess >>= evaluate) fexcept -- For LeastSquares, the Hessian code computes sum(fx*fy - res*fxy) = X^T X, -- but the actual Hessian of the Gaussian NLL profile is -1/MSE * X^T X.
src/Algorithm/SRTree/Utils.hs view
@@ -56,7 +56,11 @@ -- | Flatten list of column vectors to a row-major U.Vector Double toRowMajor :: Columns -> U.Vector Double-toRowMajor cols = U.generate (m * n) (\ix -> let (i, j) = ix `divMod` n in (cols !! j) U.! i)+toRowMajor cols = U.generate (m * n) (\ix -> let (i, j) = ix `divMod` n+ col = cols !! j+ in if i < U.length col+ then col U.! i+ else 0) -- pad with 0 for inconsistent columns where (m, n) = matSize cols -- | Restore a row-major continuous U.Vector Double back to Columns@@ -109,21 +113,26 @@ | otherwise = do l <- UM.new (m * m) let orig = toRowMajor arr+ origLen = U.length orig forM_ [0 .. m - 1] $ \i -> forM_ [0 .. m - 1] $ \j -> if i < j then unsafeWrite m l (i, j) 0 else do- let cur = orig U.! (i * m + j)- rowI = i * m- rowJ = j * m- xjj <- UM.unsafeRead l (rowJ + j)- tot <- rangedLinearDotProd rowI rowJ j l- let delta = cur - tot- if i == j- then if delta <= 0- then throwM NegDef- else UM.unsafeWrite l (rowI + j) (sqrt delta)- else UM.unsafeWrite l (rowI + j) (delta / xjj)+ let idx = i * m + j+ if idx >= origLen+ then throwM NegDef -- degenerate matrix+ else do+ let cur = orig U.! idx+ rowI = i * m+ rowJ = j * m+ xjj <- UM.unsafeRead l (rowJ + j)+ tot <- rangedLinearDotProd rowI rowJ j l+ let delta = cur - tot+ if i == j+ then if delta <= 0+ then throwM NegDef+ else UM.unsafeWrite l (rowI + j) (sqrt delta)+ else UM.unsafeWrite l (rowI + j) (delta / xjj) frozen <- U.unsafeFreeze l pure $ fromRowMajor m m frozen where (m, n) = matSize arr
src/Text/ParseSR.hs view
@@ -208,7 +208,7 @@ var = do char 'x' ix <- decimal- pure $ Fix $ Var ix+ pure $ Fix $ Var ix -- TIR is 0-based (x0, x1, x2...) <|> do char 't' ix <- decimal pure $ Fix $ Param ix
srtree.cabal view
@@ -5,7 +5,7 @@ -- see: https://github.com/sol/hpack name: srtree-version: 3.0.0.3+version: 3.0.0.4 synopsis: A general library to work with Symbolic Regression expression trees. description: A Symbolic Regression Tree data structure to work with mathematical expressions with support to first order derivative and simplification; license: BSD3