srtree 3.0.0.2 → 3.0.0.3
raw patch · 14 files changed
+741/−69 lines, 14 filesnew-component:exe:test-ciPVP: major bump suggested
API removals or changes: PVP suggests a major version bump
API changes (from Hackage documentation)
+ Algorithm.SRTree.ConfidenceIntervals: enforceMonotonicTau :: [(Double, Double)] -> [(Double, Double)]
+ Algorithm.SRTree.ConfidenceIntervals: enforceMonotonicTheta :: Double -> [(Double, Double)] -> [(Double, Double)]
+ Algorithm.SRTree.ConfidenceIntervals: monotoneDec :: [(Double, Double)] -> [(Double, Double)]
+ Algorithm.SRTree.ConfidenceIntervals: monotoneInc :: [(Double, Double)] -> [(Double, Double)]
+ Algorithm.SRTree.ConfidenceIntervals: recomputeStdErr :: EvalTree -> Target -> Target
+ Algorithm.SRTree.NonlinearOpt: compileLossAndGrad :: ADBackEnd -> Loss -> Maybe Target -> Columns -> Target -> Fix SRTree -> Vector Double -> (Double, Vector Double)
- Algorithm.EqSat.SearchSR: paretoFront :: (Fix SRTree -> StateT EGraph (StateT StdGen IO) (Double, [Target])) -> Int -> (Int -> Int -> StateT EGraph (StateT StdGen IO) [String]) -> RndEGraph [[String]]
+ Algorithm.EqSat.SearchSR: paretoFront :: (Fix SRTree -> RndEGraph (Double, [Target])) -> Int -> (Int -> EClassId -> RndEGraph b) -> RndEGraph [b]
- Algorithm.SRTree.ConfidenceIntervals: createSplines :: Target -> Columns -> Double -> Double -> Int -> (Double -> Double, Double -> Double)
+ Algorithm.SRTree.ConfidenceIntervals: createSplines :: Target -> Columns -> Double -> Double -> Int -> Double -> (Double -> Double, Double -> Double)
- Algorithm.SRTree.ConfidenceIntervals: getEndPoint :: EvalTree -> Target -> Double -> Int -> Bool -> Double
+ Algorithm.SRTree.ConfidenceIntervals: getEndPoint :: EvalTree -> Target -> Double -> Double -> Int -> Bool -> Double
Files
- ChangeLog.md +22/−0
- apps/TestCI/Main.hs +345/−0
- src/Algorithm/EqSat/SearchSR.hs +1/−2
- src/Algorithm/EqSat/Simplify.hs +1/−0
- src/Algorithm/SRTree/AD/Unboxed.hs +0/−1
- src/Algorithm/SRTree/ConfidenceIntervals.hs +183/−50
- src/Algorithm/SRTree/Likelihoods.hs +0/−1
- src/Algorithm/SRTree/ModelSelection.hs +0/−2
- src/Algorithm/SRTree/NonlinearOpt.hs +17/−6
- src/Algorithm/SRTree/Utils.hs +0/−1
- src/Text/ParseSR.hs +0/−5
- srtree.cabal +21/−1
- test/CITests.hs +149/−0
- test/Spec.hs +2/−0
ChangeLog.md view
@@ -1,5 +1,27 @@ # Changelog for srtree +## 3.0.0.3++- **Profile-likelihood CI overhaul** (`ConfidenceIntervals`):+ - Fixed `paramCI` to use F-distribution with 1 numerator df (was k df),+ matching per-parameter profiling semantics+ - `getAllProfiles`: added restart limit (5), defensive guard for short theta,+ auto-compute Laplace CIs when estCIs is empty, and `recomputeStdErr` helper+ - `getProfile`: fixed `tau_max` to use 1 df; added `LeastSquares` profile+ statistic `n * log(MSE(t)/MSE(opt))` (was raw `2*(MSE(t)-MSE(opt))` which+ differs by factor n/(2*MSE)); increased step limit 300→500; added fallback+ for small gradient in `inv_slope'`; guarded `nll_cond < nll_opt` with epsilon+ - `getProfileCnstr`/`getEndPoint`: replaced NELDERMEAD+AugLag with robust+ bisection on the profiled NLL; wider search bounds (50× se); NaN guards+ - `getStatsFromModel`: for `LeastSquares`, scale covariance by MSE (was+ unscaled); use `max 0` for sqrt of diagonal to avoid numerical NaN+ - `createSplines`: enforce monotonicity on (tau,θ) and (θ,tau) pairs to+ prevent spline extrapolation garbage; accepts `optTh` parameter+- **Test CI executable**: new `test-ci` app in `apps/TestCI` for investigating+ profile-likelihood backends on real datasets+- **CI tests**: new `CITests` module covering monotonicity, spline, and+ negative-tau regression tests+ ## 3.0.0.2 - Added parser for NeoGP.jl
+ apps/TestCI/Main.hs view
@@ -0,0 +1,345 @@+{-# LANGUAGE OverloadedStrings #-}+{-# LANGUAGE BangPatterns #-}+{-# LANGUAGE ScopedTypeVariables #-}++module Main where++import qualified Data.ByteString.Char8 as B+import Data.SRTree+import Data.SRTree.Print (showExpr)+import Data.SRTree.Eval (Target, Columns, compile)+import Data.SRTree.Datasets (loadDataset)+import Data.SRTree.Recursion (Fix(..))+import Algorithm.SRTree.ConfidenceIntervals+import Algorithm.SRTree.Compile (compileTree, EvalTree(..))+import Algorithm.SRTree.NonlinearOpt (minimizeNLL, minimizeNLLWith, compileLossAndGrad)+import Algorithm.SRTree.Likelihoods (Distribution(..), Loss(..))+import Algorithm.SRTree.AD (ADBackEnd(..))+import Algorithm.SRTree.AD.Unboxed (setMTPopParallel)+import Numeric.Optimization.NLOPT (LocalAlgorithm(..))+import Text.ParseSR (parseSR, SRAlgs(..))+import qualified Data.Vector.Unboxed as VU+import qualified Data.Vector.Storable as VS+import qualified Data.Vector.Generic as G+import Data.List (intercalate, foldl', maximumBy, isSuffixOf)+import Data.Maybe (fromMaybe)+import System.IO (hFlush, stdout, hPutStrLn, stderr)+import System.Environment (getArgs)+import System.Random (randomRIO)+import Control.Exception (catch, SomeException, evaluate)+import Control.Monad (forM_, when)+import Data.Either (isRight)++-- | A test case: expression string, description+data TestCase = TestCase+ { tcName :: String+ , tcExpr :: String+ , tcDesc :: String+ } deriving (Show)++-- | All test cases+testCases :: [TestCase]+testCases =+ [ TestCase+ { tcName = "exp-param"+ , tcExpr = "Exp(t0 + (x1 * (t1 * ((x0 * x0) + t2))))"+ , tcDesc = "Parameterized exponential (similar shape to AGENTS.md example)"+ }+ , TestCase+ { tcName = "linear"+ , tcExpr = "t0 * x0 + t1"+ , tcDesc = "Simple linear model (well-conditioned)"+ }+ , TestCase+ { tcName = "quadratic"+ , tcExpr = "t0 * x0 * x0 + t1 * x0 + t2"+ , tcDesc = "Quadratic polynomial"+ }+ , TestCase+ { tcName = "rational"+ , tcExpr = "t0 / (t1 + x0)"+ , tcDesc = "Rational function (steep near pole)"+ }+ , TestCase+ { tcName = "sine"+ , tcExpr = "t0 * sin(t1 * x0 + t2)"+ , tcDesc = "Sinusoidal model"+ }+ , TestCase+ { tcName = "product"+ , tcExpr = "t0 * x0 * x1 + t1"+ , tcDesc = "Two-variable product"+ }+ , TestCase+ { tcName = "exp-linear"+ , tcExpr = "Exp(t0 * x0 + t1)"+ , tcDesc = "Exponential of linear (simpler than deep exp)"+ }+ , TestCase+ { tcName = "power"+ , tcExpr = "t0 * x0 ** t1"+ , tcDesc = "Power law"+ }+ , TestCase+ { tcName = "linear-mse"+ , tcExpr = "t0 * x0 + t1"+ , tcDesc = "Linear model fitted with MSE (Bates 1985 original use case)"+ }+ ]++main :: IO ()+main = do+ args <- getArgs+ let filterName = if null args then Nothing else Just (head args)++ putStrLn "========================================================================"+ putStrLn " Profile Likelihood CI Backend Investigation"+ putStrLn "========================================================================"+ putStrLn ""++ -- Load dataset+ let dataSpec = "../eggp/gaussian_train.csv:::y_noise_02:x1,x2"+ putStrLn $ "Loading dataset: " ++ dataSpec+ hFlush stdout+ ((xTr, yTr, _xVal, _yVal), (mYErr, _), _varnames, _target) <-+ loadDataset dataSpec True `catch` (\(e :: SomeException) -> do+ putStrLn $ "ERROR loading dataset: " ++ show e+ error "Failed to load dataset")++ let nSamples = VU.length yTr+ putStrLn $ " Samples: " ++ show nSamples+ putStrLn $ " Features: " ++ show (length xTr)+ putStrLn ""++ let cases = case filterName of+ Nothing -> testCases+ Just name -> filter (\tc -> tcName tc == name) testCases++ mapM_ (runTestCase xTr yTr mYErr nSamples) cases++ putStrLn ""+ putStrLn "========================================================================"+ putStrLn " Summary"+ putStrLn "========================================================================"+ putStrLn ""+ putStrLn "Key observations:"+ putStrLn " - Laplace: uses Hessian inverse; fast but may be inaccurate for nonlinear models"+ putStrLn " - Bates: classical profile walk; accurate but slow"+ putStrLn " - ODE: Chen & Jennrich ODE-based profile; fast and accurate"+ putStrLn " - Constrained: bisection on re-optimized endpoints; fast but may fail"+ putStrLn ""+ putStrLn "Issues to investigate:"+ putStrLn " 1. Constrained backend NaN on steep exponential expressions"+ putStrLn " 2. tau_max' threshold correctness for NLL Gaussian"+ putStrLn " 3. Nelder-Mead convergence in augmented Lagrangian"+ putStrLn " 4. One-sided CI failures"++-- | Run a single test case through all backends+runTestCase :: [VU.Vector Double] -> VU.Vector Double -> Maybe (VU.Vector Double) -> Int -> TestCase -> IO ()+runTestCase xTr yTr mYErr nSamples tc = do+ putStrLn "------------------------------------------------------------------------"+ putStrLn $ "Test: " ++ tcName tc+ putStrLn $ " Description: " ++ tcDesc tc+ putStrLn $ " Expression: " ++ tcExpr tc+ putStrLn ""++ -- Parse expression: convert String to ByteString for parseSR+ let parsed = parseSR TIR (B.pack "x0,x1") False (B.pack (tcExpr tc))+ case parsed of+ Left err -> putStrLn $ " PARSE ERROR: " ++ err+ Right rawTree -> do+ let tree = relabelParams rawTree+ let nParams = countParamsUniq tree+ putStrLn $ " Parsed tree: " ++ showExpr tree+ putStrLn $ " Unique params: " ++ show nParams++ if nParams == 0+ then putStrLn " SKIP: no parameters to profile"+ else do+ -- Detect MSE cases (name ends with "-mse")+ let useMSE = "-mse" `isSuffixOf` tcName tc+ dist = if useMSE then LeastSquares else Gaussian+ totalParams = if useMSE then nParams else nParams + 1 -- +1 for sigma when Gaussian++ putStrLn $ " Loss: " ++ (if useMSE then "MSE (LeastSquares)" else "NLL Gaussian")+ putStrLn $ " Total params: " ++ show totalParams+ putStrLn ""++ -- Fit with multiple restarts+ putStrLn " Fitting..."+ hFlush stdout+ setMTPopParallel True+ results <- fitMultipleRestarts dist mYErr xTr yTr tree totalParams 5+ setMTPopParallel False++ let (bestNLL, bestTheta) = maximumBy (\(a,_) (b,_) -> compare a b) results+ theta_opt = bestTheta+ negNLL = negate bestNLL++ putStrLn $ " Best loss: " ++ show negNLL+ putStrLn $ " Theta: " ++ show (VU.toList theta_opt)+ putStrLn ""++ -- Compile the EvalTree for CI computation+ let et = compileTree dist xTr yTr mYErr tree++ -- Verify the optimizer agrees+ let theta_verify = ctOptimizer et theta_opt+ nll_verify = ctNLL et theta_verify+ putStrLn $ " Verified loss (via EvalTree): " ++ show nll_verify++ -- Compute standard errors from Hessian+ let stats = getStatsFromModel dist mYErr xTr yTr tree theta_opt+ stdErrs = _stdErr stats+ putStrLn $ " Std errors (Hessian): " ++ show (VU.toList stdErrs)+ putStrLn ""++ let paramNames = [ "t" ++ show i | i <- [0 .. nParams - 1] ]+ ++ if useMSE then [] else ["sigma"]++ -- ---- LAPLACE ----+ putStrLn " === LAPLACE ==="+ let laplaceCI = paramCI (Laplace stats) nSamples theta_opt 0.05+ putStrLn $ " 95% CIs:"+ putStrLn $ " " ++ showCIList (zip paramNames laplaceCI)+ putStrLn ""++ -- ---- BATES (profile walk) ----+ 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+ putStrLn $ " 95% CIs:"+ putStrLn $ " " ++ showCIList (zip paramNames batesCI)+ putStrLn $ " Widths: " ++ show (map (\(CI _ l h) -> h - l) batesCI)+ -- Debug: show first/last profile points+ forM_ (zip [0::Int ..] profiles_bates) $ \(ix, prof) -> do+ let taus = _taus prof+ cols = _thetas prof+ nT = VU.length taus+ putStrLn $ " Profile t" ++ show ix ++ ": " ++ show nT ++ " points"+ when (nT > 0) $ do+ let firstTau = taus VU.! 0+ lastTau = taus VU.! (nT - 1)+ firstTh = (cols !! ix) VU.! 0+ lastTh = (cols !! ix) VU.! (nT - 1)+ optTh = theta_opt VU.! ix+ putStrLn $ " tau=[" ++ show firstTau ++ ", " ++ show lastTau ++ "]"+ putStrLn $ " theta=[" ++ show firstTh ++ ", " ++ show lastTh ++ "] opt=" ++ show optTh+ ) (\(e :: SomeException) -> putStrLn $ " ERROR: " ++ show e)+ putStrLn ""++ -- ---- ODE (Chen & Jennrich) ----+ 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+ putStrLn $ " 95% CIs:"+ putStrLn $ " " ++ showCIList (zip paramNames odeCI)+ putStrLn $ " Widths: " ++ show (map (\(CI _ l h) -> h - l) odeCI)+ ) (\(e :: SomeException) -> putStrLn $ " ERROR: " ++ show e)+ putStrLn ""++ -- ---- 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+ putStrLn $ " 95% CIs:"+ putStrLn $ " " ++ showCIList (zip paramNames cnstrCI)+ putStrLn $ " Widths: " ++ show (map (\(CI _ l h) -> h - l) cnstrCI)+ -- Check for NaN+ let hasNaN = any (\(CI _ l h) -> isNaN l || isNaN h) cnstrCI+ when hasNaN $ putStrLn $ " *** WARNING: NaN detected in Constrained CI ***"+ ) (\(e :: SomeException) -> putStrLn $ " ERROR: " ++ show e)+ putStrLn ""++ -- ---- Detailed profiling of Constrained for the problematic case ----+ when (nParams >= 2) $ do+ putStrLn " === DETAILED CONSTRAINED INVESTIGATION ==="+ investigateConstrained et theta_opt stdErrs totalParams+ putStrLn ""++-- | Fit with multiple random restarts+fitMultipleRestarts :: Distribution -> Maybe (VU.Vector Double) -> [VU.Vector Double] -> VU.Vector Double+ -> Fix SRTree -> Int -> Int -> IO [(Double, VU.Vector Double)]+fitMultipleRestarts dist mYErr xTr yTr tree nParams nRep = do+ let funAndGrad = compileLossAndGrad MultiThread (NLL dist) mYErr xTr yTr tree+ runRestart = do+ theta0 <- VU.replicateM nParams (randomRIO (-2, 2))+ let (theta, lossVal, _) = minimizeNLLWith funAndGrad TNEWTON 200 theta0+ pure (negate lossVal, theta)+ results <- sequence [ runRestart | _ <- [1..nRep] ]+ -- Also try from zeros+ let theta0_zero = VU.replicate nParams 0.0+ (theta_zero, loss_zero, _) = minimizeNLLWith funAndGrad TNEWTON 200 theta0_zero+ pure $ (negate loss_zero, theta_zero) : results++-- | Show a list of CIs with parameter names+showCIList :: [(String, CI)] -> String+showCIList = intercalate "\n " . map (\(name, CI est lo hi) ->+ name ++ ": " ++ showF lo ++ " <= " ++ showF est ++ " <= " ++ showF hi)+ where showF x+ | isNaN x = "NaN"+ | isInfinite x = if x > 0 then "+Inf" else "-Inf"+ | otherwise = show (fromIntegral (round (x * 1e4) :: Int) / 1e4 :: Double)++-- | Detailed investigation of the Constrained backend+investigateConstrained :: EvalTree -> VU.Vector Double -> VU.Vector Double -> Int -> IO ()+investigateConstrained et theta_opt stdErrs nParams = do+ let nll_opt = ctNLL et theta_opt+ n = ctRows et+ k = VU.length theta_opt+ chi2_1 = 3.841 -- chi2 quantile for 1 df at 0.95++ putStrLn $ " nll_opt = " ++ show nll_opt+ putStrLn $ " n = " ++ show n ++ ", k = " ++ show k+ putStrLn $ " chi2_1(0.95) = " ++ show chi2_1+ putStrLn ""++ -- Corrected tau_max' calculation+ let tau_max' = chi2_1 / 2+ tau_max_old = nll_opt * chi2_1 / fromIntegral n -- OLD (buggy)+ putStrLn $ " Corrected tau_max' (chi2_1/2) = " ++ show tau_max'+ putStrLn $ " OLD tau_max' (nll_opt * chi2_1 / n) = " ++ show tau_max_old ++ " (WRONG)"+ putStrLn ""++ -- Test each parameter+ forM_ [0 .. nParams - 1] $ \ix -> do+ putStrLn $ " Parameter t" ++ show ix ++ ":"+ putStrLn $ " MLE = " ++ show (theta_opt VU.! ix)+ putStrLn $ " StdErr = " ++ show (stdErrs VU.! ix)++ -- Test getEndPoint directly+ let getPoint isLeft = getEndPoint et theta_opt tau_max' (stdErrs VU.! ix) ix isLeft+ catch (do+ let leftPt = getPoint True+ rightPt = getPoint False+ putStrLn $ " Left endpoint = " ++ show leftPt+ putStrLn $ " Right endpoint = " ++ show rightPt+ when (isNaN leftPt || isNaN rightPt) $+ putStrLn $ " *** NaN detected! ***"+ when (leftPt > rightPt) $+ putStrLn $ " *** Left > Right: reversed interval! ***"+ ) (\(e :: SomeException) -> putStrLn $ " ERROR in getEndPoint: " ++ show e)++ -- Test the profiling function (fix ix, re-optimize others)+ putStrLn $ " Testing ctOptimizerFixed..."+ catch (do+ let delta = stdErrs VU.! ix * 0.5+ theta_left = VU.generate nParams (\j -> if j == ix then (theta_opt VU.! ix) - delta else theta_opt VU.! j)+ theta_right = VU.generate nParams (\j -> if j == ix then (theta_opt VU.! ix) + delta else theta_opt VU.! j)+ reopt_left = ctOptimizerFixed et ix theta_left+ reopt_right = ctOptimizerFixed et ix theta_right+ nll_left = ctNLL et reopt_left+ nll_right = ctNLL et reopt_right+ putStrLn $ " theta_left (fixed at " ++ show (theta_opt VU.! ix - delta) ++ ") -> reopt NLL = " ++ show nll_left+ putStrLn $ " theta_right (fixed at " ++ show (theta_opt VU.! ix + delta) ++ ") -> reopt NLL = " ++ show nll_right+ putStrLn $ " NLL increase left: " ++ show (nll_left - nll_opt)+ putStrLn $ " NLL increase right: " ++ show (nll_right - nll_opt)+ when (isNaN nll_left || isNaN nll_right) $+ putStrLn $ " *** NaN in re-optimized NLL! ***"+ ) (\(e :: SomeException) -> putStrLn $ " ERROR in ctOptimizerFixed: " ++ show e)+ putStrLn ""
src/Algorithm/EqSat/SearchSR.hs view
@@ -213,10 +213,9 @@ printExprFun 0 bec Nothing -> pure () ---paretoFront :: Int -> (Int -> EClassId -> RndEGraph ()) -> RndEGraph ()+paretoFront :: (Fix SRTree -> RndEGraph (Double, [Target])) -> Int -> (Int -> EClassId -> RndEGraph b) -> RndEGraph [b] paretoFront fitFun maxSize printExprFun = go 1 0 (-(1.0/0.0)) where- go :: Int -> Int -> Double -> RndEGraph [[String]] go n ix f | n > maxSize = pure [] | otherwise = do
src/Algorithm/EqSat/Simplify.hs view
@@ -262,6 +262,7 @@ myCost (Var _) = 1 myCost (Const _) = 3 myCost (Param _) = 3+myCost (Y _) = 1 myCost (Bin op l r) = 2 + l + r myCost (Uni _ t) = 3 + t
src/Algorithm/SRTree/AD/Unboxed.hs view
@@ -48,7 +48,6 @@ import qualified Data.Vector.Unboxed.Mutable as VUM import qualified Data.Vector as VB import qualified Data.Vector.Mutable as VMB-import Debug.Trace (trace, traceShow) import qualified Data.IntMap.Strict as IntMap import Data.List ( foldl', foldl1' ) import Data.Maybe (isJust, fromMaybe)
src/Algorithm/SRTree/ConfidenceIntervals.hs view
@@ -1,4 +1,4 @@-{-# language ViewPatterns, ScopedTypeVariables, MultiWayIf, FlexibleContexts #-}+{-# language ViewPatterns, ScopedTypeVariables, MultiWayIf, FlexibleContexts, BangPatterns #-} ------------------------------------------------------------------------------- -- | -- Module : Algorithm.SRTree.ConfidenceIntervals@@ -29,7 +29,6 @@ import Numeric.Optimization.NLOPT import System.IO.Unsafe ( unsafePerformIO ) import Control.Monad.Catch ( catch, SomeException )- import Debug.Trace ( trace ) -- | profile likelihood algorithms: Bates (classical), ODE (faster), Constrained (fastest)@@ -87,9 +86,10 @@ paramCI (Profile stats profiles) nSamples _ alpha = zipWith3 CI theta lows highs where- -- for the profile likelihood we use the square root of the F-distribution with (1-alpha)+ -- for the profile likelihood we use the square root of the F-distribution+ -- with 1 numerator df (each parameter is profiled individually) k = length theta- t = sqrt $ quantile (fDistribution k (fromIntegral $ nSamples - k)) (1 - alpha)+ t = sqrt $ quantile (fDistribution 1 (fromIntegral $ nSamples - k)) (1 - alpha) stdErr = _stdErr stats lows = map (`_tau2theta` (-t)) profiles highs = map (`_tau2theta` t) profiles@@ -182,43 +182,97 @@ Left g -> Left $ g . invleft op vl Right vr -> Right $ evalOp op vl vr +-- | Recompute standard errors from the Hessian at a given theta.+-- Used when a profile walk restarts from a new optimum.+recomputeStdErr :: EvalTree -> Target -> Target+recomputeStdErr et t = stdErr+ where+ 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)+ 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]-getAllProfiles ptype et theta stdErr estCIs alpha = getAll 0 []+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 = []+ | otherwise = go 0 et theta stdErr estCIs where- k = U.length theta- n = ctRows et- tau_max = sqrt $ quantile (fDistribution k (n - k)) (1 - 0.01)- tau_max' = sqrt $ quantile (fDistribution k (n - k)) (1 - alpha)+ restartLimit = 5 :: Int - 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+ go restarts et' theta' stdErr' estCIs'+ | restarts >= restartLimit = profileAll restarts et' theta' stdErr' estCIs'+ | otherwise = profileAll restarts et' theta' stdErr' estCIs' - getAll ix acc | ix == k = acc- | ix == k-1 && ptype == Constrained && ctDist et == Gaussian = case getProfileODE et theta (stdErr U.! ix) (estCIs !! ix) tau_max ix of- Left t -> getAllProfiles ptype et t stdErr estCIs alpha- Right p -> getAll (ix + 1) (acc <> [p])- | otherwise = case profFun ix of- Left t -> getAllProfiles ptype et t stdErr estCIs alpha- Right p -> getAll (ix + 1) (acc <> [p])+ profileAll restarts et' theta' stdErr' estCIs' = go' 0 []+ where+ k = U.length theta'+ n = ctRows et'+ -- For profiling a single parameter, the threshold is chi2_1 (1 df),+ -- not chi2_k (k df). The profile likelihood ratio for ONE parameter+ -- follows chi2_1 under H0.+ tau_max = sqrt $ quantile (fDistribution 1 (n - k)) (1 - 0.01)+ nll_opt = ctNLL et' (ctOptimizer et' theta')+ chi2_1 = quantile (fDistribution 1 (n - k)) (1 - alpha)+ -- Profile likelihood CI: 2*(L(theta_hat) - L(theta)) <= chi2_1+ -- => ctNLL(theta) <= ctNLL(theta_hat) + chi2_1/2+ -- So tau_max for the constrained method = chi2_1/2+ tau_max' = chi2_1 / 2 + -- If estCIs is empty, compute Laplace CIs as initial estimates+ -- (needed by ODE fallback for the last Gaussian parameter)+ 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))+ 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)+ in map (\ix -> CI (theta' U.! ix) ((theta' U.! ix) - tVal * (se U.! ix)) ((theta' U.! ix) + tVal * (se U.! ix))) [0..k-1]+ 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++ go' ix acc | ix == k = 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+ 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 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+ (tau2theta, theta2tau) = createSplines taus thetas stdErr_i tau_max ix optTh pure $ ProfileT taus thetas optTh tau2theta theta2tau- where+ where p0 = ([0], [theta_opt])- kmax = 300+ kmax = 500 nll_opt = ctNLL et theta_opt theta_opt = ctOptimizer et theta optTh = theta_opt U.! ix@@ -227,7 +281,7 @@ go 0 delta _ _ acc = Right acc go k delta t inv_slope acc@(taus, thetas) | isNaN inv_slope = Right acc- | nll_cond < nll_opt = Left theta_t+ | 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'@@ -237,8 +291,17 @@ theta_t = minimizer theta_delta (nll_cond, grad) = ctGradNLL et theta_t zv = grad U.! ix- inv_slope' = min 4.0 . max 0.0625 . abs $ (tau / (stdErr_i * zv))- tau = signum delta * sqrt (2*nll_cond - 2*nll_opt)+ -- 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)@@ -253,32 +316,41 @@ taus = U.fromList [-tau_max, tau_max] thetas = [theta, theta] theta_i = theta U.! ix- getPoint = getEndPoint et theta tau_max ix+ getPoint = getEndPoint et theta tau_max stdErr_i ix leftPt = getPoint True rightPt = getPoint False tau2theta tau = if tau < 0 then leftPt else rightPt -getEndPoint :: EvalTree -> Target -> Double -> Int -> Bool -> Double-getEndPoint et theta tau_max ix isLeft =- case minimizeAugLag problem (G.convert theta_opt) of- Right sol -> solutionParams sol VS.! ix- Left _ -> theta_opt U.! ix+getEndPoint :: EvalTree -> Target -> Double -> Double -> Int -> Bool -> Double+getEndPoint et theta tau_max stdErr_i ix isLeft+ | isNaN mle = 0/0 -- NaN: MLE itself is NaN+ | f mle >= 0 = 0/0 -- NaN: MLE violates constraint+ | isLeft && f lo <= 0 = 0/0 -- NaN: constraint satisfied at left bound+ | not isLeft && f hi <= 0 = 0/0 -- NaN: constraint satisfied at right bound+ | isLeft = bisect lo mle 0+ | otherwise = bisect mle hi 0 where n = U.length theta- theta_opt = ctOptimizer et theta nll_opt = ctNLL et theta_opt loss_crit = nll_opt + tau_max-- loss = subtract loss_crit . ctNLL et . G.convert- obj = (if isLeft then id else negate) . (VS.! ix)+ mle = theta_opt U.! ix+ -- Use a wide search range: 50x the standard error, with a minimum of 50x |mle|+ -- This ensures we don't miss the CI boundary for parameters near zero+ searchScale = max (abs mle * 50) (stdErr_i * 50)+ lo = mle - searchScale+ hi = mle + searchScale - stop = ObjectiveRelativeTolerance 1e-4 :| [MaximumEvaluations 1000]- localAlg = NELDERMEAD obj [] Nothing- local = LocalProblem (fromIntegral n) stop localAlg- constraint = InequalityConstraint (Scalar loss) 1e-6+ -- Profiled NLL: fix theta[ix]=t, re-optimize all other params+ f t = let x = U.generate n (\j -> if j == ix then t else theta_opt U.! j)+ reopt = ctOptimizerFixed et ix (G.convert x)+ in ctNLL et reopt - loss_crit - problem = AugLagProblem [] [] (AUGLAG_LOCAL local [constraint] [])+ bisect a b k+ | k >= 60 || abs (b - a) < 1e-12 = (a + b) / 2+ | f mid <= 0 = if isLeft then bisect a mid (k+1) else bisect mid b (k+1)+ | otherwise = if isLeft then bisect mid b (k+1) else bisect a mid (k+1)+ where mid = (a + b) / 2 {-# INLINE getEndPoint #-} -- Based on@@ -290,7 +362,7 @@ | otherwise = let (taus', thetas') = solLeft <> ([0], [theta_opt]) <> solRight taus = U.fromList taus' thetas = thetas'- (tau2theta, theta2tau) = createSplines taus thetas stdErr_i tau_max ix+ (tau2theta, theta2tau) = createSplines taus thetas stdErr_i tau_max ix optTh in pure $ ProfileT taus thetas optTh tau2theta theta2tau where dflt = ProfileT (U.fromList [-tau_max, tau_max]) [theta, theta] (theta U.! ix) (const (theta U.! ix)) (const tau_max)@@ -349,12 +421,26 @@ hess = hessianNLL dist mYerr xss ys tree theta fexcept :: SomeException -> IO Columns- fexcept e = trace ("cov NegDef" <> show (toRowMajor hess)) $ pure ident+ fexcept _ = pure ident - cov = unsafePerformIO $ catch (invChol hess) fexcept+ covRaw = unsafePerformIO $ catch (invChol hess) 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.+ -- So cov_code = inv(X^T X) and cov_correct = MSE * inv(X^T X) = MSE * cov_code.+ sigma2 = case dist of+ LeastSquares -> let mse = compileLoss xss (buildLoss (NLL LeastSquares) (fromIntegral n) tree) ys mYerr theta+ in max mse 1e-10 -- avoid division by zero+ _ -> 1.0 -- no scaling needed for NLL-based losses++ scaleFactor = case dist of+ LeastSquares -> sigma2+ _ -> 1.0++ cov = fromRowMajor k k $ U.map (* scaleFactor) (toRowMajor covRaw)+ covMat = toRowMajor cov- stdErr = U.generate k (\ix -> sqrt $ covMat U.! (ix * k + ix))+ stdErr = U.generate k (\ix -> sqrt $ max 0 (covMat U.! (ix * k + ix))) stdErrSq = case outer stdErr stdErr of Right v -> v@@ -364,16 +450,63 @@ corr = fromRowMajor k k $ U.generate (k * k) (\ix -> covMat U.! ix / stdErrSqMat U.! ix) -- Create splines for profile-t-createSplines :: Target -> Columns -> Double -> Double -> Int -> (Double -> Double, Double -> Double)-createSplines taus thetas se tau_max ix+-- We enforce monotonicity of theta w.r.t. tau: if the profile walk produced+-- non-monotonic pairs (theta[i] < theta[i-1] for positive tau direction or vice versa),+-- we keep only the outermost monotonic subsequence to prevent spline extrapolation garbage.+createSplines :: Target -> Columns -> Double -> Double -> Int -> Double -> (Double -> Double, Double -> Double)+createSplines taus thetas se tau_max ix optTh | n < 2 = (genSplineFun [(-tau_max, -se), (tau_max, se)], genSplineFun [(-se, 0), (se, 1)]) | otherwise = (tau2theta, theta2tau) where n = U.length taus cols = getCol ix thetas- nubOnFirst = nubBy (\x y -> fst x == fst y)- tau2theta = genSplineFun $ nubOnFirst $ sortOnFirst taus cols- theta2tau = genSplineFun $ nubOnFirst $ sortOnFirst cols taus+ rawPairs = sortOnFirst taus cols+ monoPairs = enforceMonotonicTau rawPairs+ _ = trace ("createSplines: raw=" ++ show (length rawPairs) ++ " mono=" ++ show (length monoPairs) ++ " head=" ++ show (take 3 monoPairs) ++ " last=" ++ show (reverse $ take 3 $ reverse monoPairs)) ()+ tau2theta = genSplineFun monoPairs+ theta2tau = genSplineFun $ enforceMonotonicTheta optTh $ sortOnFirst cols taus++-- | Enforce monotonicity for (tau, theta) pairs sorted by tau.+-- Split at tau=0; both halves keep theta non-decreasing:+-- negative half: as tau increases from -tau_max toward 0, theta increases+-- positive half: as tau increases from 0 toward tau_max, theta increases+enforceMonotonicTau :: [(Double, Double)] -> [(Double, Double)]+enforceMonotonicTau [] = []+enforceMonotonicTau [p] = [p]+enforceMonotonicTau pts = negMono ++ posMono+ where+ (neg, pos) = span (\(t, _) -> t <= 0) pts+ negMono = monotoneInc neg+ posMono = monotoneInc pos++-- | Enforce monotonicity for (theta, tau) pairs sorted by theta.+-- Split at theta=optTh; both halves keep tau non-decreasing:+-- left half: as theta increases toward optTh, tau increases toward 0+-- right half: as theta increases from optTh, tau increases from 0+enforceMonotonicTheta :: Double -> [(Double, Double)] -> [(Double, Double)]+enforceMonotonicTheta _ [] = []+enforceMonotonicTheta _ [p] = [p]+enforceMonotonicTheta optTh pts = negMono ++ posMono+ where+ (neg, pos) = span (\(t, _) -> t <= optTh) pts+ negMono = monotoneInc neg+ posMono = monotoneInc pos++-- | Keep longest prefix of non-decreasing second elements.+monotoneInc :: [(Double, Double)] -> [(Double, Double)]+monotoneInc [] = []+monotoneInc [x] = [x]+monotoneInc ((t0,th0):(t1,th1):rest)+ | th1 >= th0 = (t0,th0) : monotoneInc ((t1,th1):rest)+ | otherwise = monotoneInc ((t0,th0):rest)++-- | Keep longest prefix of non-increasing second elements.+monotoneDec :: [(Double, Double)] -> [(Double, Double)]+monotoneDec [] = []+monotoneDec [x] = [x]+monotoneDec ((t0,th0):(t1,th1):rest)+ | th1 <= th0 = (t0,th0) : monotoneDec ((t1,th1):rest)+ | otherwise = monotoneDec ((t0,th0):rest) getCol :: Int -> Columns -> Target getCol ix mtx = U.generate (length mtx) (\j -> (mtx !! j) U.! ix)
src/Algorithm/SRTree/Likelihoods.hs view
@@ -47,7 +47,6 @@ import Control.Concurrent (getNumCapabilities) import Control.Concurrent.Async (forConcurrently) -import Debug.Trace import Data.SRTree.Print import Control.Monad.State.Strict import Control.Monad.Identity
src/Algorithm/SRTree/ModelSelection.hs view
@@ -38,8 +38,6 @@ import qualified Data.Vector.Unboxed as U import Algorithm.SRTree.Compile -import Debug.Trace- -- | Bayesian information criterion bic :: EvaluatedTree -> Double bic et = valParams et * log (valRows et) + 2 * valLoss et
src/Algorithm/SRTree/NonlinearOpt.hs view
@@ -12,7 +12,13 @@ -- ----------------------------------------------------------------------------- module Algorithm.SRTree.NonlinearOpt- where+ ( minimizeNLLWith+ , minimizeNLL'+ , minimizeNLL+ , minimizeNLLWithFixedParam'+ , minimizeNLLWithFixedParam+ , compileLossAndGrad+ ) where import Algorithm.SRTree.Likelihoods import Numeric.Optimization.NLOPT@@ -31,8 +37,6 @@ import Control.Monad.State.Strict import Control.Monad.Identity -import Debug.Trace- minimizeNLLWith :: (VS.Vector Double -> (Double, VS.Vector Double)) -> (ObjectiveD -> (Maybe VectorStorage) -> LocalAlgorithm) -> Int -> Target -> (Target, Double, Int) minimizeNLLWith funAndGrad alg niter t0 | niter == 0 = (t0, f, 0)@@ -53,13 +57,20 @@ t_opt' = G.convert t_opt {-# INLINE minimizeNLLWith #-} +-- | Compile the loss function and gradient for a tree, returning a reusable+-- closure. Use this when you need to optimize the same expression with+-- multiple random restarts — compile once, call the closure many times.+compileLossAndGrad :: ADBackEnd -> Loss -> Maybe Target -> Columns -> Target -> Fix SRTree -> VS.Vector Double -> (Double, VS.Vector Double)+compileLossAndGrad backend dist mYerr xss ys tree =+ let m = V.length ys+ tree' = buildLoss dist (fromIntegral m) tree+ in compileFunAndGrad backend xss ys mYerr tree'+ -- | minimizes the negative log-likelihood of the expression minimizeNLL' :: (ObjectiveD -> (Maybe VectorStorage) -> LocalAlgorithm) -> ADBackEnd -> Loss -> Maybe Target -> Int -> Columns -> Target -> Fix SRTree -> Target -> (Target, Double, Int) minimizeNLL' alg backend dist mYerr niter xss ys tree t0 = minimizeNLLWith funAndGrad alg niter t0 where- m = V.length ys- tree' = buildLoss dist (fromIntegral m) tree- funAndGrad = compileFunAndGrad backend xss ys mYerr tree'+ funAndGrad = compileLossAndGrad backend dist mYerr xss ys tree minimizeNLL :: ADBackEnd -> Loss -> Maybe Target -> Int -> Columns -> Target -> Fix SRTree -> Target -> (Target, Double, Int)
src/Algorithm/SRTree/Utils.hs view
@@ -15,7 +15,6 @@ import Data.List (unfoldr) import Data.SRTree.Eval-import Debug.Trace (traceShow) -- | Internal helper to get dimensions (rows, columns) matSize :: Columns -> (Int, Int)
src/Text/ParseSR.hs view
@@ -26,8 +26,6 @@ import qualified Data.Map.Strict as Map import Data.List.Split ( splitOn ) -import Debug.Trace (trace, traceShow)- -- * Data types -- | Parser of a symbolic regression tree with `Int` variable index and@@ -68,9 +66,6 @@ --parsePat :: B.ByteString -> Either String Pattern --parsePat = eitherResult . (`feed` "") . parse parsePatExpr . putEOL . B.strip--eitherResult' :: Show r => Result r -> Either String r-eitherResult' res = trace (show res) $ eitherResult res -- * Parsers
srtree.cabal view
@@ -5,7 +5,7 @@ -- see: https://github.com/sol/hpack name: srtree-version: 3.0.0.2+version: 3.0.0.3 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@@ -196,10 +196,30 @@ , zlib >=0.6.3 && <0.8 default-language: Haskell2010 +executable test-ci+ main-is: Main.hs+ other-modules:+ Paths_srtree+ hs-source-dirs:+ apps/TestCI+ ghc-options: -threaded -rtsopts -with-rtsopts=-N -O2+ build-depends:+ base >=4.19 && <5+ , srtree+ , bytestring >=0.11 && <0.13+ , containers >=0.6.7 && <0.9+ , vector >=0.12 && <0.14+ , exceptions >=0.10 && <0.11+ , directory >=1.3 && <1.4+ , filepath >=1.4.0.0 && <1.6+ , random >=1.2 && <1.4+ default-language: Haskell2010+ test-suite srtree-test type: exitcode-stdio-1.0 main-is: Spec.hs other-modules:+ CITests EqSatTests StoreTests Paths_srtree
+ test/CITests.hs view
@@ -0,0 +1,149 @@+module CITests (tests) where++import Test.HUnit+import Algorithm.SRTree.ConfidenceIntervals+ ( monotoneInc, monotoneDec, enforceMonotonicTau, enforceMonotonicTheta, createSplines )+import Algorithm.SRTree.Utils ( genSplineFun )+import qualified Data.Vector.Unboxed as U++eps :: Double+eps = 1e-9++approxEq :: Double -> Double -> Bool+approxEq a b = abs (a - b) < eps++-- | monotoneInc keeps longest non-decreasing prefix of second elements+test_monotoneInc :: Test+test_monotoneInc = TestLabel "monotoneInc" $ TestCase $ do+ -- already non-decreasing: keep all+ assertEqual "all kept" [(1,10),(2,20),(3,30)] (monotoneInc [(1,10),(2,20),(3,30)])+ -- drop trailing decrease+ assertEqual "drop tail" [(1,10),(2,20)] (monotoneInc [(1,10),(2,20),(3,15)])+ -- single element+ assertEqual "single" [(1,5)] (monotoneInc [(1,5)])+ -- empty+ assertEqual "empty" [] (monotoneInc [])+ -- flat is ok (non-decreasing)+ assertEqual "flat ok" [(1,5),(2,5),(3,5)] (monotoneInc [(1,5),(2,5),(3,5)])+ -- decrease at start: keeps first, skips the decrease, then keeps later increase+ assertEqual "early decrease" [(1,10),(3,15)] (monotoneInc [(1,10),(2,5),(3,15)])++-- | monotoneDec keeps longest non-increasing prefix of second elements+test_monotoneDec :: Test+test_monotoneDec = TestLabel "monotoneDec" $ TestCase $ do+ -- already non-increasing: keep all+ assertEqual "all kept" [(1,30),(2,20),(3,10)] (monotoneDec [(1,30),(2,20),(3,10)])+ -- drop trailing increase+ assertEqual "drop tail" [(1,30),(2,20)] (monotoneDec [(1,30),(2,20),(3,25)])+ -- single element+ assertEqual "single" [(1,5)] (monotoneDec [(1,5)])+ -- empty+ assertEqual "empty" [] (monotoneDec [])+ -- flat is ok (non-increasing)+ assertEqual "flat ok" [(1,5),(2,5),(3,5)] (monotoneDec [(1,5),(2,5),(3,5)])++-- | enforceMonotonicTau: split at tau=0, negative half non-increasing theta,+-- positive half non-decreasing theta+test_enforceMonotonicTau :: Test+test_enforceMonotonicTau = TestLabel "enforceMonotonicTau" $ TestCase $ do+ -- well-formed data (monotonic in both halves)+ let wellFormed = [(-2.0, 0.2), (-1.0, 0.5), (0.0, 1.0), (1.0, 1.5), (2.0, 1.8)]+ assertEqual "well-formed" wellFormed (enforceMonotonicTau wellFormed)++ -- non-monotonic negative half (bump): (-1.0, 0.5) then (-0.5, 0.7) is increasing+ let nonMonNeg = [(-2.0, 0.2), (-1.5, 0.4), (-1.0, 0.5), (-0.5, 0.7), (0.0, 1.0), (0.5, 1.3), (1.0, 1.5)]+ let result = enforceMonotonicTau nonMonNeg+ -- negative half: theta non-decreasing from -tau_max to 0, so all kept+ assertBool "negative half preserved" (length result >= 5)++ -- non-monotonic positive half: theta decreases at tau=1.5+ let nonMonPos = [(-1.0, 0.5), (0.0, 1.0), (0.5, 1.3), (1.0, 1.5), (1.5, 1.4), (2.0, 1.8)]+ let result2 = enforceMonotonicTau nonMonPos+ -- should drop (1.5, 1.4) since it breaks non-decreasing+ let posPart = filter (\(t,_) -> t > 0) result2+ pairs = zip posPart (tail posPart)+ assertBool "positive monotonic" (all (\((_,a), (_,b)) -> b >= a) pairs)+ where++-- | enforceMonotonicTheta: split at theta=optTh, left half non-increasing tau,+-- right half non-decreasing tau+test_enforceMonotonicTheta :: Test+test_enforceMonotonicTheta = TestLabel "enforceMonotonicTheta" $ TestCase $ do+ -- optTh = 1.0, data sorted by theta+ let optTh = 1.0+ wellFormed = [(0.2, -2.0), (0.5, -1.0), (1.0, 0.0), (1.5, 1.0), (2.0, 2.0)]+ assertEqual "well-formed" wellFormed (enforceMonotonicTheta optTh wellFormed)++ -- split at theta=1.0 (optTh), not theta=0+ -- data: theta < 1.0 should have negative tau, theta > 1.0 should have positive tau+ let mixedThetas = [(0.5, -1.0), (0.8, -0.5), (1.0, 0.0), (1.2, 0.5), (1.5, 1.0)]+ let result = enforceMonotonicTheta optTh mixedThetas+ assertEqual "all kept" mixedThetas result++ -- non-monotonic: tau jumps back at theta=1.2+ let nonMon = [(0.5, -1.0), (1.0, 0.0), (1.2, 0.8), (1.5, 0.6), (2.0, 2.0)]+ let result2 = enforceMonotonicTheta optTh nonMon+ -- right half (theta > 1.0): tau non-decreasing, so (1.5, 0.6) after (1.2, 0.8) is dropped+ assertBool "right half monotonic" (length result2 < length nonMon)++-- | createSplines: basic spline creation and evaluation+test_createSplines :: Test+test_createSplines = TestLabel "createSplines" $ TestCase $ do+ -- Create a simple linear profile: tau = theta - 1.0 (optTh = 1.0)+ let n = 20+ optTh = 1.0+ se = 0.5+ tau_max = 3.0+ taus = U.fromList [ -tau_max + 2*tau_max * fromIntegral i / fromIntegral (n-1) | i <- [0..n-1] ]+ -- theta = 1.0 + tau/3 (linear relationship)+ thetas = [ U.fromList [ optTh + (taus U.! i) / 3.0 | _ <- [0] ] | i <- [0..n-1] ]+ (tau2theta, _theta2tau) = createSplines taus thetas se tau_max 0 optTh++ -- at tau=0, should return approximately optTh+ let atZero = tau2theta 0.0+ assertBool ("tau2theta(0) ~ optTh: " ++ show atZero) (approxEq atZero optTh)++ -- at tau=tau_max, should be approximately optTh + tau_max/3+ let atMax = tau2theta tau_max+ expected_atMax = optTh + tau_max / 3.0+ assertBool ("tau2theta(tau_max) ~ expected: " ++ show atMax ++ " vs " ++ show expected_atMax)+ (abs (atMax - expected_atMax) < 0.5) -- generous tolerance for spline overshoot++ -- at tau=-tau_max, should be approximately optTh - tau_max/3+ let atMin = tau2theta (-tau_max)+ expected_atMin = optTh - tau_max / 3.0+ assertBool ("tau2theta(-tau_max) ~ expected: " ++ show atMin ++ " vs " ++ show expected_atMin)+ (abs (atMin - expected_atMin) < 0.5)++-- | Regression: negative-tau data must not be dropped+test_negative_tau_preserved :: Test+test_negative_tau_preserved = TestLabel "negative_tau_preserved" $ TestCase $ do+ let optTh = 1.0+ se = 0.5+ tau_max = 3.0+ -- Monotonically decreasing theta for negative tau+ negTaus = [-2.5, -2.0, -1.5, -1.0, -0.5]+ posTaus = [0.5, 1.0, 1.5, 2.0, 2.5]+ taus = U.fromList (negTaus ++ [0.0] ++ posTaus)+ thetas = [ U.fromList [ optTh + (taus U.! i) / 3.0 ] | i <- [0 .. U.length taus - 1] ]+ (tau2theta, _) = createSplines taus thetas se tau_max 0 optTh++ -- The spline MUST give a value below optTh for negative tau+ let atNeg = tau2theta (-2.0)+ assertBool ("tau2theta(-2.0) < optTh: got " ++ show atNeg)+ (atNeg < optTh - 0.1)++ -- The spline MUST give a value above optTh for positive tau+ let atPos = tau2theta 2.0+ assertBool ("tau2theta(2.0) > optTh: got " ++ show atPos)+ (atPos > optTh + 0.1)++tests :: Test+tests = TestLabel "ConfidenceIntervals" $ TestList+ [ test_monotoneInc+ , test_monotoneDec+ , test_enforceMonotonicTau+ , test_enforceMonotonicTheta+ , test_createSplines+ , test_negative_tau_preserved+ ]
test/Spec.hs view
@@ -7,6 +7,7 @@ import Algorithm.SRTree.AD.Unboxed (CompiledTree, compileTree, compileTreeMulti, evalGrad, evalGradVec, evalGradMulti) import qualified EqSatTests import qualified StoreTests+import qualified CITests import Data.SRTree.Random (randomTree, tossBiased, randomFrom) import System.Random (mkStdGen) import Control.Monad.State.Strict (evalStateT)@@ -109,6 +110,7 @@ , TestLabel "benchgrad" test_benchgrad , TestLabel "eqsat" EqSatTests.tests , TestLabel "store" StoreTests.tests+ , CITests.tests ] if failures counts /= 0 || errors counts /= 0 then error "Some tests failed"