imp-ppl (empty) → 0.1.0.0
raw patch · 41 files changed
+3998/−0 lines, 41 filesdep +basedep +containersdep +imp-ppl
Dependencies added: base, containers, imp-ppl, mtl, tasty, tasty-hunit, time, vector
Files
- CHANGELOG.md +6/−0
- LICENSE +21/−0
- README.md +212/−0
- bench/Bench.hs +144/−0
- bench/Ellsberg.hs +72/−0
- bench/Robot.hs +47/−0
- imp-ppl.cabal +147/−0
- src/Imp.hs +57/−0
- src/Imp/BDD.hs +30/−0
- src/Imp/BDD/Builder.hs +166/−0
- src/Imp/BDD/Compile.hs +107/−0
- src/Imp/BDD/WMC.hs +50/−0
- src/Imp/DSL.hs +107/−0
- src/Imp/DSL/Combinators.hs +181/−0
- src/Imp/DSL/Grade.hs +40/−0
- src/Imp/Examples/Basic.hs +32/−0
- src/Imp/Examples/Ellsberg.hs +20/−0
- src/Imp/Examples/IMDP.hs +78/−0
- src/Imp/Examples/Iteration.hs +36/−0
- src/Imp/Examples/Knightian.hs +37/−0
- src/Imp/Examples/MontyHall.hs +29/−0
- src/Imp/Examples/Polytope.hs +28/−0
- src/Imp/Examples/TwoChild.hs +15/−0
- src/Imp/Inference.hs +12/−0
- src/Imp/Inference/Approx.hs +77/−0
- src/Imp/Inference/Enumerate.hs +86/−0
- src/Imp/Inference/Optimize.hs +79/−0
- src/Imp/Inference/Symbolic.hs +88/−0
- src/Imp/Prelude.hs +10/−0
- src/Imp/Semiring.hs +81/−0
- test/Main.hs +19/−0
- test/Test/BDD.hs +205/−0
- test/Test/Combinators.hs +188/−0
- test/Test/DSL.hs +125/−0
- test/Test/Examples.hs +212/−0
- test/Test/Inference.hs +322/−0
- test/Test/Semiring.hs +71/−0
- test/Test/Util.hs +70/−0
- viz/Main.hs +286/−0
- viz/Viz.hs +369/−0
- viz/Viz/ConvexHull.hs +36/−0
+ CHANGELOG.md view
@@ -0,0 +1,6 @@+# Revision history for imp-ppl++## 0.1.0.0 — 2026-08-06++* First release: the graded monad DSL, BDD compilation, and four credal+ inference backends (enumeration, interval, gradient, symbolic).
+ LICENSE view
@@ -0,0 +1,21 @@+MIT License++Copyright (c) 2026 Jack Liell-Cock++Permission is hereby granted, free of charge, to any person obtaining a copy+of this software and associated documentation files (the "Software"), to deal+in the Software without restriction, including without limitation the rights+to use, copy, modify, merge, publish, distribute, sublicense, and/or sell+copies of the Software, and to permit persons to whom the Software is+furnished to do so, subject to the following conditions:++The above copyright notice and this permission notice shall be included in all+copies or substantial portions of the Software.++THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR+IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,+FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE+AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER+LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,+OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE+SOFTWARE.
+ README.md view
@@ -0,0 +1,212 @@+# Imp++Imprecise probabilistic programming via BDDs in Haskell.++**Imp** is a domain-specific language for discrete probabilistic programs with *Knightian uncertainty*, where some probabilities are not precisely known.+Programs compile to binary decision diagrams (BDDs) for inference via weighted model counting.++Unlike standard probabilistic programming, **Imp** computes *credal sets*.+These are convex sets of distributions that capture all possibilities consistent with the specified uncertainty,+enabling reasoning under ambiguity in settings like decision-making, planning, and robust inference.++## Key ideas++| Concept | Description |+|---------|-------------|+| **Graded Monad DSL** | Programs have type `Imp (g :: [Symbol]) a`, where `g` is a type-level set of *Knightian names* tracking sources of adversarial uncertainty. Uses `QualifiedDo` for ergonomic do-notation. |+| **BDD Compilation** | Both probabilistic flips and Knightian choices compile to BDD variables. Probabilistic variables get weights for WMC whereas Knightian variables are left free. |+| **Semiring-parametric Inference** | Inference is performed via semiring-parametric WMC enabling different inference methods based on the instantiation. |+| **Enumeration via Probabilities** | Computes the full credal set by enumerating Knightian valuations, giving exact lower and upper probabilities. |+| **Optimization via Dual Numbers** | Performs gradient ascent over the Knightian weights to search for probability bounds. |+| **Approximation via Intervals** | Uses interval arithmetic to give a sound outer approximation of interval bounds in one WMC pass. |+| **Symbolic via Polynomials** | Recovers polynomial representation of the credal set for later optimization (i.e. by a corner search). |++## Quick start++```haskell+{-# LANGUAGE DataKinds, QualifiedDo, RebindableSyntax, TypeApplications #-}+import Imp++data Ball = Red | Black | Yellow deriving (Eq, Ord, Show)++-- Ellsberg's urn: 30 red, 60 either black or yellow+ellsberg :: Imp '["split"] Ball+ellsberg = Imp.do+ isRed <- flip (1/3)+ isBlack <- interval @"split" 0.0 1.0+ Imp.return $ if isRed then Red+ else if isBlack then Black+ else Yellow+```++### Running inference++```haskell+-- Per-value marginal bounds (exact enumeration)+marginal ellsberg+-- fromList [(Red,(0.333,0.333)),(Black,(0.0,0.667)),(Yellow,(0.0,0.667))]++-- Lower/upper probability of an event+intervalProbability ellsberg (== Red)+-- (0.333, 0.333)++intervalProbability ellsberg (\b -> b == Red || b == Yellow)+-- (0.333, 1.0)+```++Precise marginals can be calculated for programs with no Knightian uncertainty:++```haskell+fairCoin :: Imp '[] Bool+fairCoin = flip 0.5++preciseMarginal fairCoin+-- fromList [(False,0.5),(True,0.5)]+```++## Building and running++Requires GHC 9.10+ and Cabal. Add `imp-ppl` to build-depends, then `import Imp`.++```bash+# Build the library+cabal build++# Run the test suite+cabal test++# Dev tooling below (requires the `dev` flag)++# Benchmarks (each emits CSV to stdout: N,method,time_s):+cabal run -f dev bench-ellsberg # n-way Ellsberg, N = 2..15+cabal run -f dev bench-robot # n-step robot IMDP, N = 1..8++# Visualization tool (generates index.html with credal set plots and BDD diagrams)+cabal run -f dev viz+open index.html+```++## Modules++| Module | Description |+|--------|-------------|+| `Imp` | Single-import entry point: custom prelude, DSL combinators, and the inference API |+| `Imp.Prelude` | The standard `Prelude` minus the six names the DSL replaces |+| `Imp.DSL` | Graded monad GADT, `flip`, `knight`, `interval`, `observe`, `tag` |+| `Imp.DSL.Grade` | Type-level grade algebra: `Merge`, `Union`, `TagAll` |+| `Imp.DSL.Combinators` | Iteration combinators: `mapName`, `intervalMap`, `foldN`, `scanN`, `foldMN` |+| `Imp.Semiring` | `Semiring` class, `ProbS`, `DualS`, `IntervalS`, `PolyS` |+| `Imp.BDD` | Core BDD types |+| `Imp.BDD.Builder` | Hash-consed BDD manager, ITE/And/Or operations |+| `Imp.BDD.Compile` | Compilation from `Imp` programs to BDDs |+| `Imp.BDD.WMC` | Semiring-parametric weighted model counting |+| `Imp.Inference` | Re-exports the four inference backends |+| `Imp.Inference.Enumerate` | Exact inference by Knightian valuation enumeration |+| `Imp.Inference.Approx` | Sound outer bounds via interval WMC |+| `Imp.Inference.Optimize` | Gradient ascent via dual number WMC |+| `Imp.Inference.Symbolic` | Exact inference via polynomial WMC and corner search |+| `Imp.Examples.Basic` | Simple coin-flip programs with no Knightian uncertainty |+| `Imp.Examples.Ellsberg` | The Ellsberg paradox: 30 Red balls, 60 Black or Yellow in unknown proportion |+| `Imp.Examples.IMDP` | Interval MDP: robot navigation on a line |+| `Imp.Examples.Iteration` | Random walks with imprecise step probability |+| `Imp.Examples.Knightian` | Knightian names controlling correlation between choices |+| `Imp.Examples.MontyHall` | Monty Hall problem with imprecise host behavior |+| `Imp.Examples.Polytope` | Polytope credal sets from composed intervals |+| `Imp.Examples.TwoChild` | The imprecise two-child problem |++## GHC extensions++Users need only four pragmas plus one import:++```haskell+{-# LANGUAGE DataKinds, QualifiedDo, RebindableSyntax, TypeApplications #-}+import Imp+```++The single import provides both unqualified access (DSL combinators, `ifThenElse` for `RebindableSyntax`, and the inference API) and the `Imp.` qualifier that `QualifiedDo` desugaring uses (`Imp.do`, `Imp.return`).+`RebindableSyntax` also rebinds ordinary monadic do, so put driver code in a separate module importing `Imp.Inference` (not `Imp`, which shadows `return`, `flip`, and `fmap`), or keep one module and use `import qualified Prelude as P` with `P.do` for `IO`.++Internally, the library uses `DataKinds`, `DerivingStrategies`, `GADTs`, and `TypeFamilies` as cabal default-extensions, plus per-file `UndecidableInstances` and `AllowAmbiguousTypes` where needed.++## Main features++### Probabilistic and Knightian choices++`flip p` creates a coin flip that returns `True` with probability `p`.+`knight @"name"` creates a Knightian binary choice returning `True` with unknown probability.+`interval @"name" lo hi` is shorthand for a Bernoulli with unknown probability in `[lo, hi]`.++### Conditioning++`observe` conditions on a Boolean being `True`:++```haskell+conditioned :: Imp '["bias"] Bool+conditioned = Imp.do+ biased <- interval @"bias" 0.3 0.7+ observe biased+ Imp.return biased+```++### Tagging and iteration++`tag @"name"` scopes all Knightian choices in a subprogram under a tag, enabling reuse.+Combinators like `foldN`, `scanN`, and `foldMN` iterate over numbered tags for multi-step models:++```haskell+-- Robot taking 5 steps with interval transition probabilities+trajectory :: Imp '["step1.d", "step2.d", "step3.d", "step4.d", "step5.d"] [Position]+trajectory = scanN @5 @"step" dynamics startPos step+```++### Inference modes++| Function | Description | Complexity |+|----------|-------------|------------|+| `preciseMarginal` | Exact marginals for programs with no Knightian uncertainty | O(\|BDD\|) |+| `marginal` | Exact per-value lower/upper bounds via enumeration | O(2^k x \|BDD\|) |+| `intervalProbability` | Exact lower/upper P(event) via enumeration | O(2^k x \|BDD\|) |+| `intervalExpectation` | Exact lower/upper E[f] via enumeration | O(2^k x \|BDD\|) |+| `marginalApprox` | Sound outer approximation via interval WMC | O(\|BDD\|) |+| `intervalProbabilityApprox` | Sound outer approximation of P(event) | O(\|BDD\|) |+| `intervalExpectationApprox` | Sound outer approximation of E[f] | O(\|BDD\|) |+| `marginalSymbolic` | Exact per-value bounds via symbolic WMC | O(2^k x \|BDD\|) |+| `intervalProbabilitySymbolic` | Exact lower/upper P(event) via symbolic WMC | O(2^k x \|BDD\|) |+| `intervalExpectationSymbolic` | Exact lower/upper E[f] via symbolic WMC | O(2^k x \|BDD\|) |+| `credalVertices` | The distributions per feasible Knightian valuation | O(2^k x \|BDD\|) |+| `optimizeProbability` | Gradient-based optimization over the credal set | O(steps x k x \|BDD\|) |+| `optimizeExpectation` | Gradient-based optimization of expectations | O(steps x k x \|BDD\|) |++Where *k* = number of Knightian variables.++## Examples++See `Imp.Examples.*` for worked examples:++- **Basic:** simple coin flips+- **Ellsberg:** the classic ambiguity-aversion paradox (30 red, 60 black/yellow unknown split)+- **Knightian:** `dependent`/`independent` showing how Knightian names control correlation+- **IMDP:** robot navigation on a line with interval transition probabilities, including compositional reuse via `tag`+- **Iteration:** random walks using `intervalMap`+- **MontyHall:** Monty Hall problem with Knightian host bias+- **Polytope:** composed interval choices forming higher-dimensional credal sets+- **TwoChild:** an imprecise variant of the two-child problem, demonstrating `observe`++## Accompanying paper++The library accompanies the paper [Imprecise Probabilistic Programming, Precisely (Functional Pearl)](https://doi.org/10.1145/3828698) (to appear).+The corresponding code is tagged [`icfp2026`](https://github.com/jacklc3/imp/tree/icfp2026).+The current version differs from the paper as follows.++| Name | Change |+|------|--------|+| `compile` | Returns a tuple of BDD manager, variable weights, Knightian variable indices, and worlds as a map rather than an association list |+| `wmc`, `wmcBatch` | Take a `Weight -> (s, s)` interpretation of the probability/Knightian weights into the semiring, rather than `WMCParams`; added `wmcBatch` for `Traversable`s |+| `ProbS` | Renamed from `RealS` |+| `foldMN` | Renamed from `foldmN` |+| `preciseMarginal`, `credalVertices` | Return maps instead of association lists |+| `marginal`, `marginalApprox`, `marginalSymbolic` | Return maps of lower/upper pairs instead of lists of triples |+| `optimizeExpectation`, `optimizeProbability` | Renamed from `credalOptimize*`, take the program as the first argument, and return maps of Knightian weights rather than association lists |+| `(>>=)`, `(>>)` | Carry an `Ord` constraint, permitting intermediate grouping of worlds in `compile`, which reduces the branching factor |+| `Imp.Inference.Symbolic` | Implemented the polynomial semiring and corner search that the paper left for future work |+| Empty credal set | Now uniformly throws an error when inferring bounds |
+ bench/Bench.hs view
@@ -0,0 +1,144 @@+{-# OPTIONS_GHC -fno-full-laziness #-}+{-# LANGUAGE GADTs #-}+module Bench+ ( BenchCase(..)+ , BenchConfig(..)+ , defaultConfig+ , runAll+ ) where++import Control.Monad (replicateM, when)+import Data.List (sort)+import Data.Time.Clock (getCurrentTime, diffUTCTime)+import System.IO (hFlush, stdout)+import System.Mem (performGC)++import Imp.DSL (Imp)+import Imp.Inference (intervalProbability, intervalProbabilityApprox)+import Imp.Inference.Optimize (optimizeProbability)+import Imp.Inference.Symbolic (intervalProbabilitySymbolic)++-- | A single benchmark case.+data BenchCase where+ BenchCase :: Ord a => Int -> Imp g a -> (a -> Bool) -> BenchCase++-- | Median timing across the four inference methods, in seconds.+data BenchResult = BenchResult+ { brN :: Int+ , brExact :: Maybe Double+ , brSymbolic :: Maybe Double+ , brApprox :: Double+ , brGrad :: Double+ }++data BenchConfig = BenchConfig+ { bcReps :: Int -- ^ median over this many reps+ , bcGradSteps :: Int -- ^ gradient-descent step count+ , bcGradLR :: Double -- ^ gradient-descent learning rate+ , bcMaxExactN :: Maybe Int -- ^ skip exact enumeration when @N > this@+ , bcMaxSymbolicN :: Maybe Int -- ^ skip symbolic when @N > this@+ , bcGradBounds :: Bool -- ^ run gradient in both directions+ , bcGcBetween :: Bool -- ^ 'performGC' before each timed rep+ }++-- | Defaults tuned for the Ellsberg-style benches (fast individual reps,+-- high-dimensional gradient landscapes). Slower benches (e.g. Robot)+-- override 'bcGradSteps', 'bcGradLR', and 'bcGradBounds' to suit.+defaultConfig :: BenchConfig+defaultConfig = BenchConfig+ { bcReps = 10+ , bcGradSteps = 30+ , bcGradLR = 0.3+ , bcMaxExactN = Just 20+ , bcMaxSymbolicN = Just 20+ , bcGradBounds = True+ , bcGcBetween = True+ }++-- ---------------------------------------------------------------------------+-- Pipeline+-- ---------------------------------------------------------------------------++-- | Run every case, emitting CSV (header + per-case rows) to stdout in real-time.+runAll :: BenchConfig -> [BenchCase] -> IO ()+runAll cfg cs = do+ putStrLn csvHeader+ hFlush stdout+ mapM_ (\c -> runBench cfg c >>= emitCsvRow) cs++-- | Run a single benchmark case, timing each of the four inference+-- methods and returning their medians.+runBench :: BenchConfig -> BenchCase -> IO BenchResult+runBench cfg (BenchCase n prog predicate) = do+ let skipExact = maybe False (n >) (bcMaxExactN cfg)+ skipSymbolic = maybe False (n >) (bcMaxSymbolicN cfg)+ steps = bcGradSteps cfg+ lr = bcGradLR cfg+ tExact <- if skipExact+ then return Nothing+ else Just <$> timeFresh cfg (\() -> intervalProbability prog predicate)+ tSymbolic <- if skipSymbolic+ then return Nothing+ else Just <$> timeFresh cfg (\() -> intervalProbabilitySymbolic prog predicate)+ tApprox <- timeFresh cfg (\() -> intervalProbabilityApprox prog predicate)+ tGrad <- timeFresh cfg $ \() ->+ if bcGradBounds cfg+ then gradientBounds steps lr prog predicate+ else let (_, v) = optimizeProbability prog predicate steps lr+ in (v, v)+ return (BenchResult n tExact tSymbolic tApprox tGrad)++-- | Gradient-descent bounds.+gradientBounds :: Ord a => Int -> Double -> Imp g a -> (a -> Bool) -> (Double, Double)+gradientBounds nSteps lr prog predicate =+ let (_, pMax) = optimizeProbability prog predicate nSteps lr+ (_, pMin) = optimizeProbability prog predicate nSteps (-lr)+ in (pMin, pMax)++-- ---------------------------------------------------------------------------+-- Timing+-- ---------------------------------------------------------------------------++-- | Median wall-clock time (in seconds) of @bcReps@ fresh evaluations of+-- the supplied nullary function.+timeFresh :: BenchConfig -> (() -> (Double, Double)) -> IO Double+timeFresh cfg compute = do+ ts <- replicateM (bcReps cfg) $ do+ when (bcGcBetween cfg) performGC+ (_, t) <- timeIt $+ let (a, b) = compute () in a `seq` b `seq` return (a, b)+ return t+ return (median ts)++timeIt :: IO a -> IO (a, Double)+timeIt act = do+ t0 <- getCurrentTime+ x <- act+ t1 <- getCurrentTime+ return (x, realToFrac (diffUTCTime t1 t0))++median :: [Double] -> Double+median xs =+ let s = sort xs+ m = length s+ in if odd m+ then s !! (m `div` 2)+ else 0.5 * (s !! (m `div` 2 - 1) + s !! (m `div` 2))++-- ---------------------------------------------------------------------------+-- CSV output+-- ---------------------------------------------------------------------------++-- | CSV column header: @N,method,time_s@.+csvHeader :: String+csvHeader = "N,method,time_s"++-- | Four CSV rows (one per method) for a single result.+emitCsvRow :: BenchResult -> IO ()+emitCsvRow r = do+ let n = brN r+ putStrLn (show n ++ ",exact," ++ maybe "NaN" show (brExact r))+ putStrLn (show n ++ ",symbolic," ++ maybe "NaN" show (brSymbolic r))+ putStrLn (show n ++ ",interval," ++ show (brApprox r))+ putStrLn (show n ++ ",gradient," ++ show (brGrad r))+ hFlush stdout
+ bench/Ellsberg.hs view
@@ -0,0 +1,72 @@+{-# LANGUAGE AllowAmbiguousTypes #-}+{-# LANGUAGE ExplicitNamespaces #-}+{-# LANGUAGE QualifiedDo #-}+{-# LANGUAGE RebindableSyntax #-}+{-# LANGUAGE UndecidableInstances #-}++-- | An n-way Ellsberg stick-break-chain benchmark.++module Main where++import qualified Imp+import Imp.DSL (Imp, knight, IfThenElse(..))+import Imp.DSL.Combinators (GenNames)+import Imp.DSL.Grade (Merge)++import Prelude+import Data.Proxy (Proxy(..))+import GHC.TypeLits (KnownNat, KnownSymbol, Nat, Symbol, natVal, type (-))++import Bench+ ( BenchCase(..)+ , defaultConfig+ , runAll+ )++type family StickGrade (names :: [Symbol]) :: [Symbol] where+ StickGrade '[] = '[]+ StickGrade (n ': ns) = Merge '[n] (StickGrade ns)++class StickBreak (names :: [Symbol]) where+ stickChain :: Int -> Imp (StickGrade names) Int++instance StickBreak '[] where+ stickChain = Imp.return++instance (KnownSymbol n, StickBreak ns) => StickBreak (n ': ns) where+ stickChain base = Imp.do+ stop <- knight @n+ if stop+ then Imp.return base+ else stickChain @ns (base + 1)++mkEllsberg+ :: forall (n :: Nat).+ (KnownNat n, StickBreak (GenNames (n - 2) "c"))+ => Imp (StickGrade (GenNames (n - 2) "c")) Int+mkEllsberg = Imp.do+ isZero <- Imp.flip (1 / fromIntegral (natVal (Proxy @n)))+ if isZero+ then Imp.return (0 :: Int)+ else stickChain @(GenNames (n - 2) "c") 1++benchCases :: [BenchCase]+benchCases =+ [ BenchCase 2 (mkEllsberg @2) (== 1)+ , BenchCase 3 (mkEllsberg @3) (== 2)+ , BenchCase 4 (mkEllsberg @4) (== 3)+ , BenchCase 5 (mkEllsberg @5) (== 4)+ , BenchCase 6 (mkEllsberg @6) (== 5)+ , BenchCase 7 (mkEllsberg @7) (== 6)+ , BenchCase 8 (mkEllsberg @8) (== 7)+ , BenchCase 9 (mkEllsberg @9) (== 8)+ , BenchCase 10 (mkEllsberg @10) (== 9)+ , BenchCase 11 (mkEllsberg @11) (== 10)+ , BenchCase 12 (mkEllsberg @12) (== 11)+ , BenchCase 13 (mkEllsberg @13) (== 12)+ , BenchCase 14 (mkEllsberg @14) (== 13)+ , BenchCase 15 (mkEllsberg @15) (== 14)+ ]++main :: IO ()+main = runAll defaultConfig benchCases
+ bench/Robot.hs view
@@ -0,0 +1,47 @@+{-# LANGUAGE AllowAmbiguousTypes #-}++-- | An n-step robot IMDP benchmark.++module Main where++import GHC.TypeLits (Nat)++import Imp.DSL (Imp, interval)+import Imp.DSL.Combinators (ConcatMapTag, GenNames, TagFoldM, foldN)+import Imp.Examples.IMDP (Position(..), step)++import Bench+ ( BenchCase(..)+ , BenchConfig(..)+ , defaultConfig+ , runAll+ )++robotStep :: Imp '["d"] Bool+robotStep = interval @"d" 0.6 0.9++mkRobot :: forall (n :: Nat). TagFoldM (GenNames n "t") '["d"]+ => Imp (ConcatMapTag (GenNames n "t") '["d"]) Position+mkRobot = foldN @n @"t" robotStep P0 step++benchCases :: [BenchCase]+benchCases =+ [ BenchCase 1 (mkRobot @1) (== P2)+ , BenchCase 2 (mkRobot @2) (== P2)+ , BenchCase 3 (mkRobot @3) (== P2)+ , BenchCase 4 (mkRobot @4) (== P2)+ , BenchCase 5 (mkRobot @5) (== P2)+ , BenchCase 6 (mkRobot @6) (== P2)+ , BenchCase 7 (mkRobot @7) (== P2)+ , BenchCase 8 (mkRobot @8) (== P2)+ ]++robotConfig :: BenchConfig+robotConfig = defaultConfig+ { bcGradSteps = 10+ , bcGradLR = 0.1+ , bcGradBounds = False+ }++main :: IO ()+main = runAll robotConfig benchCases
+ imp-ppl.cabal view
@@ -0,0 +1,147 @@+cabal-version: 3.0+name: imp-ppl+version: 0.1.0.0+synopsis: Imprecise probabilistic programming via BDDs+description:+ A DSL for discrete probabilistic programs with Knightian uncertainty,+ where some probabilities are not precisely known. Programs compile to+ binary decision diagrams, and inference computes credal sets,+ and lower and upper probabilities via semiring-parametric weighted+ model counting.+license: MIT+license-file: LICENSE+author: Jack Liell-Cock <jackliellcock@gmail.com>+maintainer: Jack Liell-Cock <jackliellcock@gmail.com>+copyright: (c) 2026 Jack Liell-Cock+category: Statistics+homepage: https://github.com/jacklc3/imp+bug-reports: https://github.com/jacklc3/imp/issues+build-type: Simple+tested-with: GHC == 9.10.1+ , GHC == 9.14.1++extra-doc-files:+ README.md+ CHANGELOG.md++source-repository head+ type: git+ location: https://github.com/jacklc3/imp.git++source-repository this+ type: git+ location: https://github.com/jacklc3/imp.git+ tag: v0.1.0.0++-- Development-only visualisation and benchmark executables.+flag dev+ description: Build the viz and benchmark executables.+ default: False+ manual: True++library+ hs-source-dirs: src+ exposed-modules:+ Imp+ Imp.Prelude+ Imp.BDD+ Imp.BDD.Builder+ Imp.BDD.Compile+ Imp.BDD.WMC+ Imp.Semiring+ Imp.DSL+ Imp.DSL.Combinators+ Imp.DSL.Grade+ Imp.Inference+ Imp.Inference.Approx+ Imp.Inference.Enumerate+ Imp.Inference.Optimize+ Imp.Inference.Symbolic+ Imp.Examples.Basic+ Imp.Examples.Ellsberg+ Imp.Examples.IMDP+ Imp.Examples.Iteration+ Imp.Examples.Knightian+ Imp.Examples.MontyHall+ Imp.Examples.Polytope+ Imp.Examples.TwoChild+ build-depends:+ base >= 4.20 && < 5,+ containers >= 0.6 && < 0.9,+ mtl >= 2.3 && < 2.4,+ vector >= 0.13 && < 0.14+ default-language: GHC2021+ default-extensions:+ DataKinds+ DerivingStrategies+ GADTs+ TypeFamilies+ ghc-options: -Wall -Wcompat -Wmissing-deriving-strategies++-- Development aid: renders credal sets and BDDs to index.html.+executable viz+ if !flag(dev)+ buildable: False+ hs-source-dirs: viz+ main-is: Main.hs+ other-modules:+ Viz+ Viz.ConvexHull+ build-depends:+ base,+ imp-ppl,+ containers+ default-language: GHC2021+ ghc-options: -Wall++-- Shared benchmark pipeline.+common bench-common+ if !flag(dev)+ buildable: False+ hs-source-dirs: bench+ build-depends:+ base,+ imp-ppl,+ time >= 1.12 && < 1.15+ default-language: GHC2021+ default-extensions:+ DataKinds+ TypeFamilies+ ghc-options: -Wall+ if flag(dev)+ ghc-options: -O2++executable bench-robot+ import: bench-common+ main-is: Robot.hs+ other-modules: Bench++executable bench-ellsberg+ import: bench-common+ main-is: Ellsberg.hs+ other-modules: Bench++test-suite test+ type: exitcode-stdio-1.0+ hs-source-dirs: test+ main-is: Main.hs+ other-modules:+ Test.BDD+ Test.Combinators+ Test.DSL+ Test.Examples+ Test.Inference+ Test.Semiring+ Test.Util+ build-depends:+ base,+ imp-ppl,+ containers,+ mtl,+ vector >= 0.13 && < 0.14,+ tasty >= 1.4 && < 1.6,+ tasty-hunit >= 0.10 && < 0.11+ default-language: GHC2021+ default-extensions:+ DataKinds+ ghc-options: -Wall -Wcompat -Wmissing-deriving-strategies
+ src/Imp.hs view
@@ -0,0 +1,57 @@+-- | Single-import entry point for Imp.+--+-- @+-- {-\# LANGUAGE DataKinds, QualifiedDo, RebindableSyntax, TypeApplications \#-}+-- import Imp+-- @+--+-- The single import provides both:+--+-- * Unqualified access to the Prelude, the DSL combinators, and the inference API.+-- * The @Imp.@ qualifier used by @QualifiedDo@ desugaring (@Imp.do@, @Imp.return@, @Imp.>>=@).+--+-- Driver code that only runs inference should import "Imp.Inference" directly.+module Imp+ ( module Imp.DSL+ , module Imp.DSL.Combinators+ , module Imp.Inference+ , module Imp.Prelude+ ) where++import Imp.Prelude+import Imp.DSL+ ( Imp+ , flip+ , knight+ , interval+ , observe+ , tag+ , IfThenElse(..)+ , IfR+ , Merge+ , Union+ , TagAll+ , return+ , (>>=)+ , (>>)+ , fmap+ , (<$>)+ )+import Imp.DSL.Combinators+ ( GenNames+ , ConcatMapTag+ , MapName(..)+ , intervalMap+ , knightMap+ , intervalN+ , knightN+ , TagFoldM(..)+ , foldMN+ , tagFold+ , tagMap+ , tagScan+ , tagN+ , foldN+ , scanN+ )+import Imp.Inference
+ src/Imp/BDD.hs view
@@ -0,0 +1,30 @@+-- | Core BDD types.+module Imp.BDD+ ( VarLabel(..)+ , NodeId(..)+ , BDD(..)+ , BDDNode(..)+ ) where++-- | A BDD variable label (index into the variable ordering).+newtype VarLabel = VarLabel { unVarLabel :: Int }+ deriving newtype (Eq, Ord, Show)++-- | A BDD node ID (index into the node table).+newtype NodeId = NodeId { unNodeId :: Int }+ deriving newtype (Eq, Ord, Show)++-- | A BDD pointer with complemented edges for O(1) negation.+data BDD+ = BDDTrue+ | BDDFalse+ | BDDRef !NodeId -- ^ positive reference+ | BDDComp !NodeId -- ^ complemented reference+ deriving stock (Eq, Ord, Show)++-- | Internal BDD node: variable, low child, and high child.+data BDDNode = BDDNode+ { bddVar :: !VarLabel+ , bddLow :: !BDD+ , bddHigh :: !BDD+ } deriving stock (Eq, Ord, Show)
+ src/Imp/BDD/Builder.hs view
@@ -0,0 +1,166 @@+-- | Hash-consed BDD manager and the core BDD operations.+module Imp.BDD.Builder+ ( BDDManager+ , nodeTable+ , BDDM+ -- Manager operations.+ , emptyManager+ , newVar+ , lookupNode+ -- BDD operations.+ , bddNot+ , bddAnd+ , bddOr+ , bddAny+ , bddIte+ , bddRestrict+ ) where++import Control.Monad.State.Strict+import Data.IntMap.Strict (IntMap)+import qualified Data.IntMap.Strict as IntMap+import Data.Map.Strict (Map)+import qualified Data.Map.Strict as Map++import Imp.BDD++-- | Hash-consed node store with memo caches.+data BDDManager = BDDManager+ { nextNodeId :: !Int+ -- | The hash-consed node store, keyed by t'NodeId'.+ , nodeTable :: !(IntMap BDDNode)+ , uniqueTable :: !(Map BDDNode NodeId)+ , iteCache :: !(Map (BDD, BDD, BDD) BDD)+ , restrictCache :: !(Map (BDD, VarLabel, Bool) BDD)+ , nextVarId :: !Int+ } deriving stock (Show)++-- | A manager with no nodes or variables.+emptyManager :: BDDManager+emptyManager = BDDManager 0 IntMap.empty Map.empty Map.empty Map.empty 0++-- | BDD-manager state monad.+type BDDM = State BDDManager++-- | Allocate a fresh BDD variable and return its positive literal+-- together with its t'VarLabel'.+newVar :: BDDM (BDD, VarLabel)+newVar = do+ v <- VarLabel <$> gets nextVarId+ modify' (\mgr -> mgr { nextVarId = nextVarId mgr + 1 })+ bdd <- mkNode (BDDNode v BDDFalse BDDTrue)+ return (bdd, v)++-- | Negate a BDD in O(1) via complemented edges.+bddNot :: BDD -> BDD+bddNot BDDTrue = BDDFalse+bddNot BDDFalse = BDDTrue+bddNot (BDDRef n) = BDDComp n+bddNot (BDDComp n) = BDDRef n++-- | Conjunction of two BDDs.+bddAnd :: BDD -> BDD -> BDDM BDD+bddAnd a b = bddIte a b BDDFalse++-- | Disjunction of two BDDs.+bddOr :: BDD -> BDD -> BDDM BDD+bddOr a b = bddIte a BDDTrue b++-- | OR together a list of BDDs using a balanced tree fold.+bddAny :: [BDD] -> BDDM BDD+bddAny [] = return BDDFalse+bddAny [g] = return g+bddAny xs = do+ let (l, r) = splitAt (length xs `div` 2) xs+ lv <- bddAny l+ rv <- bddAny r+ bddOr lv rv++-- | BDD if-then-else terminal cases and memoisation.+bddIte :: BDD -> BDD -> BDD -> BDDM BDD+bddIte f g h+ | f == BDDTrue = return g+ | f == BDDFalse = return h+ | g == BDDTrue && h == BDDFalse = return f+ | g == BDDFalse && h == BDDTrue = return (bddNot f)+ | g == h = return g+ | otherwise = do+ cached <- gets (Map.lookup (f, g, h) . iteCache)+ case cached of+ Just result -> return result+ Nothing -> do+ result <- bddIteExpand f g h+ modify' $ \s -> s { iteCache = Map.insert (f, g, h) result (iteCache s) }+ return result++-- | Shannon expansion on the topmost variable.+bddIteExpand :: BDD -> BDD -> BDD -> BDDM BDD+bddIteExpand f g h = do+ nf <- lookupNode f+ ng <- lookupNode g+ nh <- lookupNode h+ -- f is never constant here, so at least one node exists.+ let topV = minimum [bddVar n | Just n <- [nf, ng, nh]]+ cofactor branch bdd node = case node of+ Just n | bddVar n == topV -> if branch then bddHigh n else bddLow n+ _ -> bdd+ lo <- bddIte (cofactor False f nf) (cofactor False g ng) (cofactor False h nh)+ hi <- bddIte (cofactor True f nf) (cofactor True g ng) (cofactor True h nh)+ mkNode (BDDNode topV lo hi)++-- | Restrict a BDD by fixing a variable to True or False, with memoisation.+bddRestrict :: BDD -> VarLabel -> Bool -> BDDM BDD+bddRestrict BDDTrue _ _ = return BDDTrue+bddRestrict BDDFalse _ _ = return BDDFalse+bddRestrict bdd var val = do+ cached <- gets (Map.lookup (bdd, var, val) . restrictCache)+ case cached of+ Just result -> return result+ Nothing -> do+ result <- bddRestrictExpand bdd var val+ modify' $ \s ->+ s { restrictCache = Map.insert (bdd, var, val) result (restrictCache s) }+ return result++-- | Restrict one level and recurse.+bddRestrictExpand :: BDD -> VarLabel -> Bool -> BDDM BDD+bddRestrictExpand bdd var val = do+ mnode <- lookupNode bdd+ case mnode of+ Nothing -> return bdd+ Just (BDDNode v lo hi)+ | v == var -> return $ if val then hi else lo+ | v > var -> return bdd -- var not in this sub-BDD+ | otherwise -> do+ lo' <- bddRestrict lo var val+ hi' <- bddRestrict hi var val+ mkNode (BDDNode v lo' hi')++-- | Look up a BDD reference.+lookupNode :: BDD -> BDDM (Maybe BDDNode)+lookupNode BDDTrue = return Nothing+lookupNode BDDFalse = return Nothing+lookupNode (BDDRef (NodeId n)) = IntMap.lookup n <$> gets nodeTable+lookupNode (BDDComp (NodeId n)) = do+ node <- IntMap.lookup n <$> gets nodeTable+ return $ (\(BDDNode v lo hi) -> BDDNode v (bddNot lo) (bddNot hi)) <$> node++-- | Hash-consed node creation with complemented-edge normalization.+mkNode :: BDDNode -> BDDM BDD+mkNode node@(BDDNode v lo hi)+ | lo == hi = return lo+ -- Normalize nodes to never have a complemented low edge.+ | BDDTrue <- lo = bddNot <$> mkNode (BDDNode v (bddNot lo) (bddNot hi))+ | BDDComp _ <- lo = bddNot <$> mkNode (BDDNode v (bddNot lo) (bddNot hi))+ | otherwise = do+ mgr <- get+ case Map.lookup node (uniqueTable mgr) of+ Just nid -> return (BDDRef nid)+ Nothing -> do+ let nidInt = nextNodeId mgr+ nid = NodeId nidInt+ put mgr { nextNodeId = nidInt + 1+ , nodeTable = IntMap.insert nidInt node (nodeTable mgr)+ , uniqueTable = Map.insert node nid (uniqueTable mgr)+ }+ return (BDDRef nid)
+ src/Imp/BDD/Compile.hs view
@@ -0,0 +1,107 @@+{-# LANGUAGE LambdaCase #-}+{-# LANGUAGE AllowAmbiguousTypes #-}+-- | Compilation from 'Imp' programs to BDDs.+module Imp.BDD.Compile+ ( Compiled+ , compile+ ) where++import Control.Monad.Reader+import Control.Monad.State.Strict+import Data.IntMap.Strict (IntMap)+import qualified Data.IntMap.Strict as IntMap+import Data.Map.Strict (Map)+import qualified Data.Map.Strict as Map+import Data.Proxy (Proxy(..))+import GHC.TypeLits (KnownSymbol, symbolVal)++import Imp.BDD+import Imp.BDD.Builder+import Imp.BDD.WMC (Weight(..))+import Imp.DSL (Imp(..))++-- | Compiled result: BDD manager, variable weights, Knightian variables and index,+-- and the worlds: each return value with its BDD guard.+type Compiled a = (BDDManager, IntMap Weight, Map String Int, Map a BDD)++-- | State accumulated during compilation.+data CompileState = CompileState+ { csWeights :: !(IntMap Weight)+ , csKnights :: !(Map String (BDD, Int))+ }++-- | Read-only environment for compilation.+newtype CompileEnv = CompileEnv+ { ceTag :: String+ }++initState :: CompileState+initState = CompileState IntMap.empty Map.empty++type CompileM = ReaderT CompileEnv (StateT CompileState BDDM)++-- | Compile an Imp program to its worlds, one guard per return value.+compile :: Ord a => Imp g a -> Compiled a+compile prog =+ let comp = runStateT (runReaderT (compileM prog) (CompileEnv "")) initState+ ((worlds, st), mgr) = runState comp emptyManager+ in (mgr, csWeights st, snd <$> csKnights st, worlds)++-- | Lift a BDD-manager action into the compile monad.+liftBDDM :: BDDM a -> CompileM a+liftBDDM = lift . lift++-- | Resolve a (possibly tag-prefixed) Knightian name.+resolveName :: String -> CompileM String+resolveName baseName = do+ currentTag <- asks ceTag+ return $ if null currentTag then baseName else currentTag ++ "." ++ baseName++-- | Allocate a fresh probabilistic BDD variable with the given Bernoulli weight.+flipVar :: Double -> CompileM BDD+flipVar p = do+ (var, varLabel) <- liftBDDM newVar+ modify' $ \s -> s+ { csWeights = IntMap.insert (unVarLabel varLabel) (Prob p) (csWeights s) }+ return var++-- | Return the BDD variable for a Knightian name, allocating one if it+-- hasn't been seen before.+knightVar :: forall n. KnownSymbol n => CompileM BDD+knightVar = do+ name <- resolveName (symbolVal (Proxy :: Proxy n))+ knights <- gets csKnights+ case Map.lookup name knights of+ Just (var, _) -> return var+ Nothing -> do+ (var, varLabel) <- liftBDDM newVar+ let i = Map.size knights+ modify' $ \s -> s+ { csWeights = IntMap.insert (unVarLabel varLabel) (Knight i) (csWeights s)+ , csKnights = Map.insert name (var, i) (csKnights s)+ }+ return var++-- | Compile a program to worlds.+compileM :: Ord a => Imp g a -> CompileM (Map a BDD)+compileM = \case+ ImpReturn a -> return (Map.singleton a BDDTrue)+ ImpFlip p -> do+ v <- flipVar p+ return (Map.fromList [(True, v), (False, bddNot v)])+ ImpKnight (_ :: Proxy n) -> do+ v <- knightVar @n+ return (Map.fromList [(True, v), (False, bddNot v)])+ ImpObserve b ->+ return (Map.singleton () (if b then BDDTrue else BDDFalse))+ ImpBind m f -> do+ worlds <- compileM m+ conjoined <- mapM (\(a, g) -> traverse (liftBDDM . bddAnd g) =<< compileM (f a))+ (Map.toList worlds)+ liftBDDM (traverse bddAny (Map.unionsWith (++) [ (: []) <$> c | c <- conjoined ]))+ ImpTag (_ :: Proxy t) inner -> do+ let baseTag = symbolVal (Proxy :: Proxy t)+ resolvedTag <- if null baseTag then asks ceTag else resolveName baseTag+ local (\env -> env { ceTag = resolvedTag }) (compileM inner)+ ImpBranch True t _ -> compileM t+ ImpBranch False _ f -> compileM f
+ src/Imp/BDD/WMC.hs view
@@ -0,0 +1,50 @@+-- | Semiring-parametric weighted model counting over BDDs.+module Imp.BDD.WMC+ ( Weight(..)+ , wmc+ , wmcBatch+ ) where++import Control.Monad.State.Strict (StateT, evalStateT, gets, modify', lift, runState)+import Data.IntMap.Strict (IntMap, (!))+import qualified Data.IntMap.Lazy as Lazy+import Data.Map.Strict (Map)+import qualified Data.Map.Strict as Map++import Imp.BDD+import Imp.BDD.Builder (BDDManager, BDDM, lookupNode)+import Imp.Semiring++-- | A BDD variable weight.+data Weight = Prob !Double | Knight !Int+ deriving stock (Show, Eq)++-- | Weighted model count of a single BDD.+wmc :: Semiring s => (Weight -> (s, s)) -> BDDManager -> IntMap Weight -> BDD -> s+wmc f mgr weights bdd =+ fst (runState (evalStateT (wmcM (Lazy.map f weights) bdd) Map.empty) mgr)++-- | Weighted model count of a batch of BDDs, sharing memo across the batch.+wmcBatch :: (Semiring s, Traversable t)+ => (Weight -> (s, s)) -> BDDManager -> IntMap Weight -> t BDD -> t s+wmcBatch f mgr weights bdds =+ fst (runState (evalStateT (traverse (wmcM (Lazy.map f weights)) bdds) Map.empty) mgr)++wmcM :: Semiring s => IntMap (s, s) -> BDD -> StateT (Map BDD s) BDDM s+wmcM _ BDDTrue = return one+wmcM _ BDDFalse = return zero+wmcM weights bdd = do+ cached <- gets (Map.lookup bdd)+ case cached of+ Just val -> return val+ Nothing -> do+ mnode <- lift (lookupNode bdd)+ !result <- case mnode of+ Nothing -> return one+ Just node -> do+ let (wLo, wHi) = weights ! unVarLabel (bddVar node)+ loVal <- wmcM weights (bddLow node)+ hiVal <- wmcM weights (bddHigh node)+ return ((wLo .*. loVal) .+. (wHi .*. hiVal))+ modify' (Map.insert bdd result)+ return result
+ src/Imp/DSL.hs view
@@ -0,0 +1,107 @@+{-# LANGUAGE AllowAmbiguousTypes #-}+-- | The graded monad DSL: programs graded by their Knightian choice names.+module Imp.DSL+ ( Imp(..)+ , flip+ , knight+ , interval+ , observe+ , tag+ , IfThenElse(..)+ , IfR+ , Merge+ , Union+ , TagAll+ , Imp.DSL.return+ , (Imp.DSL.>>=)+ , (Imp.DSL.>>)+ , Imp.DSL.fmap+ , (Imp.DSL.<$>)+ ) where++import Prelude hiding (return, (>>=), (>>), flip, fmap, (<$>))+import Data.Proxy (Proxy(..))+import GHC.TypeLits (KnownSymbol, Symbol)+import Imp.DSL.Grade (Merge, Union, TagAll)++-- | A probabilistic program graded by @g@,+-- a type-level list of Symbols tracking which Knightian names this program uses.+data Imp (g :: [Symbol]) a where+ ImpReturn :: a -> Imp '[] a+ ImpBind :: Ord a => Imp g1 a -> (a -> Imp g2 b) -> Imp (Merge g1 g2) b+ ImpFlip :: !Double -> Imp '[] Bool+ ImpKnight :: KnownSymbol n => Proxy n -> Imp '[n] Bool+ ImpObserve :: Bool -> Imp '[] ()+ ImpTag :: KnownSymbol t => Proxy t -> Imp g a -> Imp (TagAll t g) a+ ImpBranch :: Bool -> Imp g1 a -> Imp g2 a -> Imp (Union g1 g2) a++-- | Graded monad return.+return :: a -> Imp '[] a+return = ImpReturn++-- | Graded monad bind.+infixl 1 >>=+(>>=) :: Ord a => Imp g1 a -> (a -> Imp g2 b) -> Imp (Merge g1 g2) b+(>>=) = ImpBind++-- | Graded monad sequence.+infixl 1 >>+(>>) :: Ord a => Imp g1 a -> Imp g2 b -> Imp (Merge g1 g2) b+m >> n = ImpBind m (const n)++-- | Graded functor map.+fmap :: Ord a => (a -> b) -> Imp g a -> Imp g b+fmap f m = ImpBind m (ImpReturn . f)++-- | Operator form of the graded functor map.+infixl 4 <$>+(<$>) :: Ord a => (a -> b) -> Imp g a -> Imp g b+(<$>) = fmap++-- | Probabilistic coin flip with probability @p@ of True; @p@ must lie in @[0,1]@.+flip :: Double -> Imp '[] Bool+flip p | p >= 0 && p <= 1 = ImpFlip p+ | otherwise = error ("flip: probability " ++ show p ++ " outside [0,1]")++-- | Knightian (adversarial) binary choice, named at the type level.+-- Usage: @knight \@\"x\"@+knight :: forall n. KnownSymbol n => Imp '[n] Bool+knight = ImpKnight (Proxy @n)++-- | Interval-valued probability: @interval \@\"n\" lo hi@ returns a @Bool@+-- with @P(True) ∈ [lo, hi]@. The bounds may be given in either order.+interval :: forall n. KnownSymbol n => Double -> Double -> Imp '[n] Bool+interval lo hi = ImpBind (ImpKnight (Proxy @n)) $ \x ->+ flip (if x then hi else lo)++-- | Condition on a Boolean predicate being True.+observe :: Bool -> Imp '[] ()+observe = ImpObserve++-- | Tag all Knightian choices in a subprogram enabling reusage.+tag :: forall t g a. KnownSymbol t => Imp g a -> Imp (TagAll t g) a+tag = ImpTag (Proxy @t)++-- ---------------------------------------------------------------------------+-- IfThenElse: grade-aware branching for RebindableSyntax+-- ---------------------------------------------------------------------------++-- | Compute the result type of an @if-then-else@ expression.+-- For @Imp@ branches with potentially different grades, the result grade is their 'Union'.+type family IfR t f where+ IfR (Imp g1 a) (Imp g2 a) = Imp (Union g1 g2) a+ IfR t f = t++-- | Overloaded @if-then-else@ for use with @RebindableSyntax@.+--+-- The pure instance is marked @INCOHERENT@ so that GHC can commit to+-- it when the branch types are not yet determined.+class IfThenElse t f where+ ifThenElse :: Bool -> t -> f -> IfR t f++instance {-# INCOHERENT #-} (a ~ b, IfR a b ~ a) => IfThenElse a b where+ ifThenElse True t _ = t+ ifThenElse False _ f = f++instance IfThenElse (Imp g1 a) (Imp g2 a) where+ ifThenElse = ImpBranch
+ src/Imp/DSL/Combinators.hs view
@@ -0,0 +1,181 @@+{-# LANGUAGE AllowAmbiguousTypes, UndecidableInstances #-}+-- | Iteration combinators over generated Knightian names.+module Imp.DSL.Combinators+ ( GenNames+ , ConcatMapTag+ , MapName(..)+ , intervalMap+ , knightMap+ , intervalN+ , knightN+ , TagFoldM(..)+ , foldMN+ , tagFold+ , tagMap+ , tagScan+ , tagN+ , foldN+ , scanN+ ) where++import GHC.TypeLits+ ( Natural, Symbol, AppendSymbol, ConsSymbol, NatToChar, KnownSymbol+ , Div, Mod, type (+), type (-) )++import Prelude hiding (return, (>>=), (>>), flip, fmap)+import Data.Proxy (Proxy(..))+import Imp.DSL.Grade (Merge, TagAll)+import Imp.DSL (Imp(..), knight, interval)++-- | Convert a type-level 'Natural' to its decimal 'Symbol' representation.+type family NatToSymbol (n :: Natural) :: Symbol where+ NatToSymbol n = NatToSymbolGo (Div n 10) (Mod n 10)++type family NatToSymbolGo (q :: Natural) (r :: Natural) :: Symbol where+ NatToSymbolGo 0 r = ConsSymbol (NatToChar (r + 48)) ""+ NatToSymbolGo q r = AppendSymbol (NatToSymbol q) (ConsSymbol (NatToChar (r + 48)) "")++-- | Number of decimal digits in a 'Natural'.+type family Digits (n :: Natural) :: Natural where+ Digits n = DigitsH (Div n 10)++type family DigitsH (q :: Natural) :: Natural where+ DigitsH 0 = 1+ DigitsH q = 1 + DigitsH (Div q 10)++-- | @k@ zeros: @Zeros 3 = \"000\"@.+type family Zeros (k :: Natural) :: Symbol where+ Zeros 0 = ""+ Zeros k = AppendSymbol "0" (Zeros (k - 1))++-- | Decimal representation of @n@, left-padded with zeros to @width@ digits.+type PadNat (width :: Natural) (n :: Natural) =+ AppendSymbol (Zeros (width - Digits n)) (NatToSymbol n)++-- | Generate a list of 1-indexed numbered names.+-- Padding matters because grades are kept sorted lexicographically.+type family GenNames (count :: Natural) (base :: Symbol) :: [Symbol] where+ GenNames 0 _ = '[]+ GenNames n base = GenNamesGo n (Digits n) base++type family GenNamesGo (n :: Natural) (width :: Natural) (base :: Symbol) :: [Symbol] where+ GenNamesGo 0 _ _ = '[]+ GenNamesGo n width base =+ Merge (GenNamesGo (n - 1) width base) '[AppendSymbol base (PadNat width n)]++-- | Map 'TagAll' over a list of tags and concatenate the results.+type family ConcatMapTag (tags :: [Symbol]) (k :: [Symbol]) :: [Symbol] where+ ConcatMapTag '[] _ = '[]+ ConcatMapTag (t ': ts) k = Merge (TagAll t k) (ConcatMapTag ts k)++-- | Graded @map@: traverse a type-level list of Knightian names,+-- applying a single-name computation to each.+--+-- @+-- mapName \@'[\"x\", \"y\"] knight -- two independent Knightian choices+-- @+class MapName (names :: [Symbol]) where+ mapName :: Ord a => (forall n. KnownSymbol n => Imp '[n] a) -> Imp names [a]++instance MapName '[] where+ mapName _ = ImpReturn []++instance (KnownSymbol n, MapName ns, Merge '[n] ns ~ (n ': ns)) => MapName (n ': ns) where+ mapName f = ImpBind (f @n) $ \x ->+ ImpBind (mapName @ns f) (ImpReturn . (x :))++-- | Independent Knightian choice per name.+knightMap :: MapName names => Imp names [Bool]+knightMap = mapName knight++-- | Independent interval per name.+intervalMap :: MapName names => Double -> Double -> Imp names [Bool]+intervalMap lo hi = mapName (interval lo hi)++-- | Independent numbered Knightian choices.+knightN :: forall n base.+ MapName (GenNames n base) =>+ Imp (GenNames n base) [Bool]+knightN = knightMap @(GenNames n base)++-- | Independent numbered intervals.+intervalN :: forall n base.+ MapName (GenNames n base) =>+ Double -> Double -> Imp (GenNames n base) [Bool]+intervalN lo hi = intervalMap @(GenNames n base) lo hi++-- | Graded left fold where the step function is monadic with unit grade.+-- Enables conditioning/flips at each time step.+--+-- @+-- tagFoldM \@'[\"move1\", \"move2\"] robotDynamics P1 $ \\pos move -> Imp.do+-- let pos' = step3 pos move+-- observe (pos' /= P0)+-- Imp.return pos'+-- @+class TagFoldM (tags :: [Symbol]) (g :: [Symbol]) where+ tagFoldM :: (Ord a, Ord b)+ => Imp g a -> b -> (b -> a -> Imp '[] b) -> Imp (ConcatMapTag tags g) b++instance TagFoldM '[] g where+ tagFoldM _ acc _ = ImpReturn acc++instance ( KnownSymbol t+ , TagFoldM ts g+ ) => TagFoldM (t ': ts) g where+ tagFoldM prog acc f = ImpBind (ImpTag (Proxy @t) prog) $ \x ->+ ImpBind (f acc x) $ \acc' ->+ tagFoldM @ts @g prog acc' f++-- | @foldMN \@n \@base prog acc f@: numbered monadic fold.+foldMN :: forall n base g a b.+ (TagFoldM (GenNames n base) g, Ord a, Ord b) =>+ Imp g a -> b -> (b -> a -> Imp '[] b) -> Imp (ConcatMapTag (GenNames n base) g) b+foldMN = tagFoldM @(GenNames n base) @g++-- | Graded left fold over tagged subprograms.+--+-- @+-- tagFold \@'[\"move1\", \"move2\"] robotDynamics P0 step3+-- @+tagFold :: forall tags g a b. (TagFoldM tags g, Ord a, Ord b) =>+ Imp g a -> b -> (b -> a -> b) -> Imp (ConcatMapTag tags g) b+tagFold prog acc f = tagFoldM @tags prog acc (\b a -> ImpReturn (f b a))++-- | Graded @map@ over tagged subprograms.+--+-- @+-- tagMap \@'[\"move1\", \"move2\", \"move3\"] robotDynamics+-- @+tagMap :: forall tags g a. (TagFoldM tags g, Ord a)+ => Imp g a -> Imp (ConcatMapTag tags g) [a]+tagMap prog = ImpBind (tagFold @tags prog [] (\as a -> a : as)) (ImpReturn . reverse)++-- | Graded left scan over tagged subprograms.+--+-- @+-- tagScan \@'[\"move1\", \"move2\", \"move3\"] robotDynamics P1 step3+-- @+tagScan :: forall tags g a b. (TagFoldM tags g, Ord a, Ord b) =>+ Imp g a -> b -> (b -> a -> b) -> Imp (ConcatMapTag tags g) [b]+tagScan prog acc f =+ ImpBind (tagFold @tags prog (acc, []) (\(b, bs) a -> let b' = f b a in (b', b' : bs)))+ (ImpReturn . reverse . snd)++-- | @tagN \@n \@base prog@: numbered tag iteration.+tagN :: forall n base g a.+ (TagFoldM (GenNames n base) g, Ord a) =>+ Imp g a -> Imp (ConcatMapTag (GenNames n base) g) [a]+tagN = tagMap @(GenNames n base)++-- | @foldN \@n \@base prog acc f@: numbered fold over tagged subprograms.+foldN :: forall n base g a b.+ (TagFoldM (GenNames n base) g, Ord a, Ord b) =>+ Imp g a -> b -> (b -> a -> b) -> Imp (ConcatMapTag (GenNames n base) g) b+foldN = tagFold @(GenNames n base) @g++-- | @scanN \@n \@base prog acc f@: numbered scan over tagged subprograms.+scanN :: forall n base g a b.+ (TagFoldM (GenNames n base) g, Ord a, Ord b) =>+ Imp g a -> b -> (b -> a -> b) -> Imp (ConcatMapTag (GenNames n base) g) [b]+scanN = tagScan @(GenNames n base) @g
+ src/Imp/DSL/Grade.hs view
@@ -0,0 +1,40 @@+{-# LANGUAGE UndecidableInstances #-}+-- | Type-level grade operations.+module Imp.DSL.Grade+ ( Merge+ , Union+ , TagAll+ ) where++import GHC.TypeLits+ ( Symbol, AppendSymbol, CmpSymbol , TypeError, ErrorMessage(..) )++-- | Sorted merge of two sorted @[Symbol]@ lists, erroring on overlap.+type family Merge (xs :: [Symbol]) (ys :: [Symbol]) :: [Symbol] where+ Merge '[] ys = ys+ Merge xs '[] = xs+ Merge (x ': xs) (y ': ys) = MergeH (CmpSymbol x y) x xs y ys++type family MergeH (o :: Ordering) (x :: Symbol) (xs :: [Symbol])+ (y :: Symbol) (ys :: [Symbol]) :: [Symbol] where+ MergeH 'LT x xs y ys = x ': Merge xs (y ': ys)+ MergeH 'EQ x _ _ _ = TypeError ('Text "Duplicate Knightian name: " ':<>: 'ShowType x)+ MergeH 'GT x xs y ys = y ': Merge (x ': xs) ys++-- | Sorted union of two sorted @[Symbol]@ lists, allowing overlap.+type family Union (xs :: [Symbol]) (ys :: [Symbol]) :: [Symbol] where+ Union '[] ys = ys+ Union xs '[] = xs+ Union (x ': xs) (y ': ys) = UnionH (CmpSymbol x y) x xs y ys++type family UnionH (o :: Ordering) (x :: Symbol) (xs :: [Symbol])+ (y :: Symbol) (ys :: [Symbol]) :: [Symbol] where+ UnionH 'LT x xs y ys = x ': Union xs (y ': ys)+ UnionH 'EQ x xs _ ys = x ': Union xs ys+ UnionH 'GT x xs y ys = y ': Union (x ': xs) ys++-- | Prepend every symbol in a list with a tag and dot separator.+type family TagAll (tag :: Symbol) (xs :: [Symbol]) :: [Symbol] where+ TagAll _ '[] = '[]+ TagAll "" xs = xs+ TagAll t (x ': xs) = AppendSymbol t (AppendSymbol "." x) ': TagAll t xs
+ src/Imp/Examples/Basic.hs view
@@ -0,0 +1,32 @@+{-# LANGUAGE QualifiedDo, RebindableSyntax #-}+-- | Simple coin-flip programs with no Knightian uncertainty.+module Imp.Examples.Basic+ ( fairCoin+ , twoCoins+ , biasedCoin+ , coinOr+ ) where++import Imp++-- | Fair coin flip.+fairCoin :: Imp '[] Bool+fairCoin = flip 0.5++-- | Two independent coin flips.+twoCoins :: Imp '[] (Bool, Bool)+twoCoins = Imp.do+ a <- flip 0.5+ b <- flip 0.5+ Imp.return (a, b)++-- | Biased coin.+biasedCoin :: Imp '[] Bool+biasedCoin = flip 0.7++-- | OR of two fair coins.+coinOr :: Imp '[] Bool+coinOr = Imp.do+ a <- flip 0.5+ b <- flip 0.5+ Imp.return (a || b)
+ src/Imp/Examples/Ellsberg.hs view
@@ -0,0 +1,20 @@+{-# LANGUAGE QualifiedDo, RebindableSyntax #-}+-- | The Ellsberg paradox: 30 Red balls, 60 Black or Yellow in unknown proportion.+module Imp.Examples.Ellsberg+ ( Ball(..)+ , ellsberg+ ) where++import Imp++-- | Balls in the urn.+data Ball = Red | Black | Yellow+ deriving stock (Eq, Ord, Show)++-- | Ellsberg's urn.+ellsberg :: Imp '["split"] Ball+ellsberg = Imp.do+ isRed <- flip (1/3)+ isBlack <- interval @"split" 0.0 1.0+ Imp.return $ if isRed then Red+ else (if isBlack then Black else Yellow)
+ src/Imp/Examples/IMDP.hs view
@@ -0,0 +1,78 @@+{-# LANGUAGE QualifiedDo, RebindableSyntax #-}+-- | Interval MDP: robot navigation on a line.+module Imp.Examples.IMDP+ ( Position(..)+ , step+ , simpleRobot+ , simpleRobot3+ , Move(..)+ , step3+ , robotDynamics+ , complexRobot+ ) where++import Imp++-- | Robot position.+data Position = P0 | P1 | P2 deriving stock (Eq, Ord, Show)++-- | Move right on True with 'P2' absorbing.+step :: Position -> Bool -> Position+step P2 _ = P2+step P1 True = P2+step P1 False = P1+step P0 True = P1+step P0 False = P0++-- | 2-step IMDP from position 0.+simpleRobot :: Imp '["move1", "move2"] Position+simpleRobot = Imp.do+ move1 <- interval @"move1" 0.6 0.9+ let pos1 = step P0 move1+ move2 <- interval @"move2" 0.6 0.9+ Imp.return (step pos1 move2)++-- | 3-step IMDP from position 0.+simpleRobot3 :: Imp '["move1", "move2", "move3"] Position+simpleRobot3 = Imp.do+ move1 <- interval @"move1" 0.6 0.9+ let pos1 = step P0 move1+ move2 <- interval @"move2" 0.6 0.9+ let pos2 = step pos1 move2+ move3 <- interval @"move3" 0.6 0.9+ Imp.return (step pos2 move3)++-- ---------------------------------------------------------------------------+-- Compositional robot: reusable dynamics with 'tag'+-- ---------------------------------------------------------------------------++-- | Three-way movement outcome.+data Move = Backwards | Stationary | Forwards+ deriving stock (Eq, Ord, Show)++-- | Step function with three-way movement: the robot can now move backwards.+step3 :: Position -> Move -> Position+step3 P0 Backwards = P0+step3 P0 Stationary = P0+step3 P0 Forwards = P1+step3 P1 Backwards = P0+step3 P1 Stationary = P1+step3 P1 Forwards = P2+step3 P2 _ = P2 -- goal is absorbing++-- | Imprecise robot dynamics as a reusable subprogram.+robotDynamics :: Imp '["b", "f"] Move+robotDynamics = Imp.do+ goForward <- interval @"f" 0.5 0.8+ goBack <- interval @"b" 0.0 0.2+ Imp.return $ if goForward then Forwards+ else if goBack then Backwards+ else Stationary++-- | Compositional 2-step robot using 'tag' to reuse 'robotDynamics'.+complexRobot :: Imp '["move1.b", "move1.f", "move2.b", "move2.f"] Position+complexRobot = Imp.do+ move1 <- tag @"move1" robotDynamics+ let pos1 = step3 P0 move1+ move2 <- tag @"move2" robotDynamics+ Imp.return (step3 pos1 move2)
+ src/Imp/Examples/Iteration.hs view
@@ -0,0 +1,36 @@+{-# LANGUAGE QualifiedDo, RebindableSyntax #-}+-- | Random walks with imprecise step probability.+module Imp.Examples.Iteration+ ( walk1+ , walk2+ , walk3+ , walk3Asym+ ) where++import Imp++-- | 1-step walk.+walk1 :: Imp '["s1"] Int+walk1 = Imp.do+ steps <- intervalMap @'["s1"] 0.3 0.7+ Imp.return $ length (filter id steps)++-- | 2-step walk.+walk2 :: Imp '["s1", "s2"] Int+walk2 = Imp.do+ steps <- intervalMap @'["s1", "s2"] 0.3 0.7+ Imp.return $ length (filter id steps)++-- | 3-step walk.+walk3 :: Imp '["s1", "s2", "s3"] Int+walk3 = Imp.do+ steps <- intervalMap @'["s1", "s2", "s3"] 0.3 0.7+ Imp.return $ length (filter id steps)++-- | Asymmetric intervals break the swap-symmetry.+walk3Asym :: Imp '["s1", "s2", "s3"] Int+walk3Asym = Imp.do+ s1 <- interval @"s1" 0.1 0.9+ s2 <- interval @"s2" 0.3 0.7+ s3 <- interval @"s3" 0.45 0.55+ Imp.return $ length (filter id [s1, s2, s3])
+ src/Imp/Examples/Knightian.hs view
@@ -0,0 +1,37 @@+{-# LANGUAGE QualifiedDo, RebindableSyntax #-}+-- | Knightian names controlling correlation between choices.+module Imp.Examples.Knightian+ ( Three(..)+ , dependent+ , independent+ ) where++import Imp++-- | Three-valued outcome.+data Three = Red | Green | Blue+ deriving stock (Eq, Ord, Show)++-- | Dependent Knightian choices. Both branches share the same+-- Knightian variable, so the outcomes are correlated.+dependent :: Imp '["a1"] Three+dependent = Imp.do+ x <- flip 0.5+ if x then Imp.do+ y <- knight @"a1"+ Imp.return (if y then Red else Green)+ else Imp.do+ y <- knight @"a1"+ Imp.return (if y then Red else Blue)++-- | Independent Knightian choices. Both branches have different+-- Knightian variables, so more outcomes are possible.+independent :: Imp '["a1", "a2"] Three+independent = Imp.do+ x <- flip 0.5+ if x then Imp.do+ y <- knight @"a1"+ Imp.return (if y then Red else Green)+ else Imp.do+ y <- knight @"a2"+ Imp.return (if y then Red else Blue)
+ src/Imp/Examples/MontyHall.hs view
@@ -0,0 +1,29 @@+{-# LANGUAGE QualifiedDo, RebindableSyntax #-}+-- | Monty Hall problem with imprecise host behavior.+module Imp.Examples.MontyHall+ ( Door(..)+ , montyHall+ ) where++import Imp++-- | The three doors.+data Door = Door1 | Door2 | Door3+ deriving stock (Eq, Ord, Show)++-- | The Monty Hall encoding.+montyHall :: Imp '["host_bias"] Bool+montyHall = Imp.do+ c1 <- flip (1/3)+ c2 <- flip (1/2)+ hostBias <- knight @"host_bias"+ let car = if c1 then Door1 else if c2 then Door2 else Door3+ host = case car of+ Door1 -> if hostBias then Door2 else Door3+ Door2 -> Door3+ Door3 -> Door2+ switchTo = case host of+ Door2 -> Door3+ Door3 -> Door2+ Door1 -> error "Impossible..."+ Imp.return (switchTo == car)
+ src/Imp/Examples/Polytope.hs view
@@ -0,0 +1,28 @@+{-# LANGUAGE QualifiedDo, RebindableSyntax #-}+-- | Polytope credal sets from composed intervals.+module Imp.Examples.Polytope+ ( Three(..)+ , polytope+ , polytope2+ ) where++import Imp++-- | Three-valued outcome.+data Three = Red | Green | Blue+ deriving stock (Eq, Ord, Show)++-- | Two independent intervals composed to give an imprecise distribution.+polytope :: Imp '["p", "q"] Three+polytope = Imp.do+ p <- interval @"p" 0.2 0.8+ q <- interval @"q" 0.3 0.7+ Imp.return $ if p then Red else (if q then Green else Blue)++-- | Three intervals composed to give a finer polytope.+polytope2 :: Imp '["b", "g", "r"] Three+polytope2 = Imp.do+ r <- interval @"r" 0.2 0.5+ g <- interval @"g" 0.2 0.5+ b <- interval @"b" 0.2 0.5+ Imp.return $ if r then Red else (if g then Green else (if b then Blue else Red))
+ src/Imp/Examples/TwoChild.hs view
@@ -0,0 +1,15 @@+{-# LANGUAGE QualifiedDo, RebindableSyntax #-}+-- | The two-child problem.+module Imp.Examples.TwoChild+ ( twoChild+ ) where++import Imp++-- | Imprecise variant of the classical two-child problem.+twoChild :: Imp '["alien"] Bool+twoChild = Imp.do+ humanBoy <- flip 0.5+ alienBoy <- knight @"alien"+ observe (humanBoy || alienBoy)+ Imp.return (humanBoy && alienBoy)
+ src/Imp/Inference.hs view
@@ -0,0 +1,12 @@+-- | The full inference API: re-exports every inference backend.+module Imp.Inference+ ( module Imp.Inference.Enumerate+ , module Imp.Inference.Approx+ , module Imp.Inference.Optimize+ , module Imp.Inference.Symbolic+ ) where++import Imp.Inference.Enumerate+import Imp.Inference.Approx+import Imp.Inference.Optimize+import Imp.Inference.Symbolic
+ src/Imp/Inference/Approx.hs view
@@ -0,0 +1,77 @@+-- | Approximate credal inference via one-pass interval WMC. Bounds are+-- sound outer approximations, but are not tight in general.+--+-- Throws an error if the evidence is unsatisfiable, i.e. the credal set is empty.+module Imp.Inference.Approx+ ( marginalApprox+ , intervalProbabilityApprox+ , intervalExpectationApprox+ ) where++import Control.Monad.State.Strict (runState)+import Data.Map.Strict (Map)+import qualified Data.Map.Strict as Map++import Imp.BDD.Builder (bddAny)+import Imp.BDD.Compile (compile)+import Imp.BDD.WMC (Weight(..), wmc, wmcBatch)+import Imp.DSL (Imp)+import Imp.Semiring++-- | Flips get a point interval, Knights the full unit interval.+intervalWeights :: Weight -> (IntervalS, IntervalS)+intervalWeights w = case w of+ Prob p -> (IntervalS (1 - p) (1 - p), IntervalS p p)+ Knight _ -> (IntervalS 0 1, IntervalS 0 1)++-- | Group the return values by @key@, returning the total-mass interval+-- and one interval per group+compileApprox :: (Ord a, Ord k) => Imp g a -> (a -> k) -> (IntervalS, Map k IntervalS)+compileApprox prog key =+ let (mgr, weights, _, worlds) = compile prog+ blocks = Map.fromListWith (++) [ (key v, [g]) | (v, g) <- Map.toList worlds ]+ (bdds, mgr') = runState (traverse bddAny blocks) mgr+ (evidence, mgr'') = runState (bddAny (Map.elems bdds)) mgr'+ in ( wmc intervalWeights mgr'' weights evidence+ , wmcBatch intervalWeights mgr'' weights bdds )++-- | Interval division where the quotient stays in @[0, 1]@.+intervalDiv :: IntervalS -> IntervalS -> IntervalS+intervalDiv (IntervalS nLo nHi) (IntervalS dLo dHi) =+ let div' a b = if b > 0 then a / b else 0+ in IntervalS (div' nLo dHi) (div' nHi (max nHi dLo))++-- | Divide a count box by the total box, tightened by the complementary count.+condition :: IntervalS -> IntervalS -> IntervalS -> (Double, Double)+condition total@(IntervalS _ tHi) count remainder =+ let IntervalS lo hi = intervalDiv count total+ IntervalS rLo rHi = intervalDiv remainder total+ lo' = max lo (1 - rHi)+ hi' = min hi (1 - rLo)+ in if tHi <= 0 then error "No feasible probabilities"+ -- clamp to min in case of rounding errors+ else (min lo' hi', hi')++-- | Approximate per-value marginal bounds via interval WMC.+marginalApprox :: Ord a => Imp g a -> Map a (Double, Double)+marginalApprox prog =+ let (total, counts) = compileApprox prog id+ remainder v = sumS (Map.elems (Map.delete v counts))+ in Map.mapWithKey (\v c -> condition total c (remainder v)) counts++-- | Approximate interval probability via interval WMC.+intervalProbabilityApprox :: Ord a => Imp g a -> (a -> Bool) -> (Double, Double)+intervalProbabilityApprox prog predicate =+ let (total, counts) = compileApprox prog predicate+ getInterval b = Map.findWithDefault zero b counts+ in condition total (getInterval True) (getInterval False)++-- | Approximate lower and upper expectations via interval WMC, clamped to+-- the range of scores the posterior can still reach.+intervalExpectationApprox :: Ord a => Imp g a -> (a -> Double) -> (Double, Double)+intervalExpectationApprox prog score =+ let reachable = [ (score v, b) | (v, b) <- Map.toList (marginalApprox prog), snd b > 0 ]+ scores = map fst reachable+ (los, his) = unzip [ if s >= 0 then (s * lo, s * hi) else (s * hi, s * lo)+ | (s, (lo, hi)) <- reachable ]+ in (max (minimum scores) (sum los), min (maximum scores) (sum his))
+ src/Imp/Inference/Enumerate.hs view
@@ -0,0 +1,86 @@+-- | Exact credal inference by enumerating Knightian valuations.+--+-- The functions returning bounds throw an error if the evidence is+-- unsatisfiable, i.e. the credal set is empty.+module Imp.Inference.Enumerate+ ( preciseMarginal+ , credalVertices+ , intervalProbability+ , intervalExpectation+ , marginal+ ) where++import Control.Monad (foldM)+import Control.Monad.State.Strict (runState)+import Data.Bool (bool)+import Data.Functor.Compose (Compose(..))+import Data.Maybe (mapMaybe)+import qualified Data.IntMap.Strict as IntMap+import Data.Map.Strict (Map)+import qualified Data.Map.Strict as Map++import Imp.BDD (VarLabel(..))+import Imp.BDD.Builder (bddRestrict)+import Imp.BDD.Compile (compile)+import Imp.BDD.WMC (Weight(..), wmcBatch)+import Imp.DSL (Imp)+import Imp.Semiring (ProbS(..))++-- | Flips get Bernoulli branch pair. Knights have been conditioned away.+probWeights :: Weight -> (ProbS, ProbS)+probWeights w = case w of+ Prob p -> (ProbS (1 - p), ProbS p)+ Knight _ -> error "probWeights: Knights should have been conditioned away"++-- | Compute the min and max of a non-empty list in a single strict pass.+bounds :: Ord a => [a] -> (a, a)+bounds [] = error "No feasible probabilities"+bounds (x:xs) = foldl' step (x, x) xs+ where step (!mn, !mx) y = (min mn y, max mx y)++-- | Precise marginal distribution. Only for programs with no Knightian choices.+-- Outcomes with probability zero are omitted, unlike 'marginal', which reports+-- every return value.+preciseMarginal :: Ord a => Imp '[] a -> Map a Double+preciseMarginal prog =+ let (mgr, weights, _, worlds) = compile prog+ counts = unProb <$> wmcBatch probWeights mgr weights worlds+ total = sum counts+ in if total > 0+ then Map.filter (> 0) ((/ total) <$> counts)+ else error "No feasible probabilities"++-- | Compute the vertices of the credal set, one distribution per+-- feasible Knightian valuation. The extreme points are a subset+-- of these.+credalVertices :: Ord a => Imp g a -> [Map a Double]+credalVertices prog =+ let (mgr, weights, _, worlds) = compile prog+ valuations = sequence [ [(VarLabel k, False), (VarLabel k, True)]+ | (k, Knight _) <- IntMap.toList weights ]+ condition val =+ traverse (\bdd -> foldM (\g (vl, b) -> bddRestrict g vl b) bdd val) worlds+ (valEvents, mgr') = runState (mapM condition valuations) mgr+ counts = unProb <$> wmcBatch probWeights mgr' weights (Compose valEvents)+ normalize row =+ let !total = sum row+ in if total > 0 then Just ((/ total) <$> row) else Nothing+ in mapMaybe normalize (getCompose counts)++-- | Exact lower and upper probability of an event, by valuation enumeration.+intervalProbability :: Ord a => Imp g a -> (a -> Bool) -> (Double, Double)+intervalProbability prog predicate =+ intervalExpectation prog (bool 0 1 . predicate)++-- | Exact lower and upper expectation of a real-valued function, by+-- valuation enumeration.+intervalExpectation :: Ord a => Imp g a -> (a -> Double) -> (Double, Double)+intervalExpectation prog score =+ bounds [ sum [score v * p | (v, p) <- Map.toList d]+ | d <- credalVertices prog ]++-- | Exact per-value marginal bounds, by valuation enumeration.+marginal :: Ord a => Imp g a -> Map a (Double, Double)+marginal prog = case credalVertices prog of+ [] -> error "No feasible probabilities"+ dists -> bounds <$> Map.unionsWith (++) [ (\x -> [x]) <$> d | d <- dists ]
+ src/Imp/Inference/Optimize.hs view
@@ -0,0 +1,79 @@+-- | Approximate credal inference by gradient ascent over the Knightian+-- weights. Values are inner approximations: every iterate is a point+-- inside the credal set, so they never overshoot the exact bounds.+--+-- Throws an error if the evidence is unsatisfiable, i.e. the credal set is empty.+module Imp.Inference.Optimize+ ( optimizeExpectation+ , optimizeProbability+ ) where++import Control.Applicative ((<|>))+import Control.Monad (guard)+import Data.Bool (bool)+import Data.Maybe (fromMaybe)+import Data.Ord (clamp)+import Data.Map.Strict (Map)+import qualified Data.Map.Strict as Map+import qualified Data.Vector as V++import Imp.BDD.Compile (compile)+import Imp.BDD.WMC (Weight(..), wmcBatch)+import Imp.DSL (Imp)+import Imp.Semiring++-- | A constant dual number.+dualConst :: Double -> DualS+dualConst x = DualS x V.empty++-- | Dual number division; undefined for a non-positive denominator.+dualDiv :: DualS -> DualS -> Maybe DualS+dualDiv (DualS a da) (DualS b db)+ | b <= 0 = Nothing+ | otherwise = Just (DualS (a / b) (V.zipWith (\a' b' -> (a' * b - a * b') / (b * b)) da db))++-- | The dual-number reading of the variable weights at the given parameters.+dualWeights :: V.Vector Double -> Weight -> (DualS, DualS)+dualWeights params w = case w of+ Prob p -> (dualConst (1 - p), dualConst p)+ Knight i -> let p = params V.! i+ k = length params+ in ( DualS (1 - p) (V.generate k (\j -> if j == i then -1 else 0))+ , DualS p (V.generate k (\j -> if j == i then 1 else 0)) )++-- | Gradient ascent over the Knightian variables to maximize expected score over+-- the credal set. Ascent starts deterministically at weight @0.5@.+--+-- Returns @(weights, expectation)@ where @weights@ maps each Knightian+-- variable name to its optimized weight.+--+-- Stops early when the gradient magnitude falls below @1e-6@.+--+-- Use a negative learning rate for gradient /descent/ (minimization).+optimizeExpectation :: Ord a+ => Imp g a+ -> (a -> Double)+ -> Int -- ^ maximum steps+ -> Double -- ^ learning rate+ -> (Map String Double, Double)+optimizeExpectation prog score steps lr =+ let (mgr, weights, knights, worlds) = compile prog+ params0 = V.replicate (Map.size knights) 0.5+ go n params = do+ let counts = wmcBatch (dualWeights params) mgr weights worlds+ aggr = sumS [ dualConst (score v) .*. count | (v, count) <- Map.toList counts ]+ DualS e de <- dualDiv aggr (sumS (Map.elems counts))+ let params' = V.zipWith (\p dp -> clamp (0, 1) (p + lr * dp)) params de+ continue = n > 0 && V.sum (V.map (\x -> x * x) de) > 1e-12+ (guard continue >> go (n - 1) params') <|> Just ((params V.!) <$> knights, e)+ in fromMaybe (error "No feasible probabilities") (go steps params0)++-- | Gradient ascent over the Knightian variables to maximize P(event).+optimizeProbability :: Ord a+ => Imp g a+ -> (a -> Bool)+ -> Int -- ^ maximum steps+ -> Double -- ^ learning rate+ -> (Map String Double, Double)+optimizeProbability prog predicate steps lr =+ optimizeExpectation prog (bool 0 1 . predicate) steps lr
+ src/Imp/Inference/Symbolic.hs view
@@ -0,0 +1,88 @@+-- | Exact credal inference via a symbolic (multilinear-polynomial) semiring.+--+-- Bounds are extracted by optimizing the objective over the parameter box corners.+--+-- Throws an error if the evidence is unsatisfiable, i.e. the credal set is empty.+module Imp.Inference.Symbolic+ ( marginalSymbolic+ , intervalProbabilitySymbolic+ , intervalExpectationSymbolic+ ) where++import Data.Bits (bit, testBit, clearBit, (.|.))+import Data.Bool (bool)+import Data.Map.Strict (Map)+import qualified Data.Map.Strict as Map++import Imp.BDD.Compile (compile)+import Imp.BDD.WMC (Weight(..), wmcBatch)+import Imp.DSL (Imp)+import Imp.Semiring++-- | A constant polynomial.+polyConst :: Double -> PolyS+polyConst p = PolyS (if p == 0 then Map.empty else Map.singleton 0 p)++-- | The polynomial reading of the variable weights.+polyWeights :: Weight -> (PolyS, PolyS)+polyWeights w = case w of+ Prob p -> (polyConst (1 - p), polyConst p)+ -- Two indicators per Knightian variable, one for @0@ and one for @1@.+ Knight i -> ( PolyS (Map.singleton (bit (2 * i)) 1)+ , PolyS (Map.singleton (bit (2 * i + 1)) 1) )++-- | Each value's count polynomial together with the total-mass polynomial.+compilePolys :: Ord a => Imp g a -> (Map a PolyS, PolyS)+compilePolys prog =+ let (mgr, weights, _, worlds) = compile prog+ counts = wmcBatch polyWeights mgr weights worlds+ in (counts, sumS (Map.elems counts))++-- | Exact per-value marginal bounds, computed symbolically.+marginalSymbolic :: Ord a => Imp g a -> Map a (Double, Double)+marginalSymbolic prog =+ let (counts, total) = compilePolys prog+ in Map.map (optimizeRatio total) counts++-- | Exact lower and upper probability of an event, computed symbolically.+intervalProbabilitySymbolic :: Ord a => Imp g a -> (a -> Bool) -> (Double, Double)+intervalProbabilitySymbolic prog predicate =+ intervalExpectationSymbolic prog (bool 0 1 . predicate)++-- | Exact lower and upper expectation of a real-valued function, computed symbolically.+intervalExpectationSymbolic :: Ord a => Imp g a -> (a -> Double) -> (Double, Double)+intervalExpectationSymbolic prog score =+ let (counts, total) = compilePolys prog+ aggr = sumS [polyConst (score v) .*. count | (v, count) <- Map.toList counts]+ in optimizeRatio total aggr++-- | Bounds over the feasible corners of the parameter box.+optimizeRatio :: PolyS -> PolyS -> (Double, Double)+optimizeRatio den num =+ let mask = foldl' (.|.) 0 (Map.keys (unPoly num) ++ Map.keys (unPoly den))+ free = [ i | i <- takeWhile (\i -> bit (2 * i) <= mask) [0 ..]+ , testBit mask (2 * i) || testBit mask (2 * i + 1) ]+ in case cornerRatios free den num of+ [] -> error "No feasible probabilities"+ ratios -> (minimum ratios, maximum ratios)++-- | Ratio at every feasible corner reachable by fixing the Knightian variables to @0@ and @1@.+cornerRatios :: [Int] -> PolyS -> PolyS -> [Double]+cornerRatios [] den num =+ let constDen = constTerm den in [ constTerm num / constDen | constDen > 0 ]+cornerRatios (i : rest) den num =+ cornerRatios rest (fixKnight i False den) (fixKnight i False num)+ ++ cornerRatios rest (fixKnight i True den) (fixKnight i True num)++-- | Fix Knightian parameter @i@ to @0@ or @1@.+fixKnight :: Int -> Bool -> PolyS -> PolyS+fixKnight i b (PolyS m) =+ PolyS $ Map.filter (/= 0) $ Map.fromListWith (+)+ [ (clearBit mask thisBit, c) | (mask, c) <- Map.toList m, not (testBit mask thatBit) ]+ where+ thisBit = if b then 2 * i + 1 else 2 * i+ thatBit = if b then 2 * i else 2 * i + 1++-- | The constant term (value once every parameter is fixed).+constTerm :: PolyS -> Double+constTerm (PolyS m) = Map.findWithDefault 0 0 m
+ src/Imp/Prelude.hs view
@@ -0,0 +1,10 @@+-- | The standard "Prelude", minus the six names the graded DSL replaces.+-- @return@, @(>>=)@, @(>>)@, @flip@, @fmap@ and @(\<$\>)@.+--+-- @RebindableSyntax@, which the DSL requires, implies @NoImplicitPrelude@,+-- so this module supplies the ordinary Prelude that would otherwise be missing.+module Imp.Prelude+ ( module Prelude+ ) where++import Prelude hiding (return, (>>=), (>>), flip, fmap, (<$>))
+ src/Imp/Semiring.hs view
@@ -0,0 +1,81 @@+-- | The 'Semiring' class and the four WMC semirings.+module Imp.Semiring+ ( Semiring(..)+ , ProbS(..)+ , DualS(..)+ , IntervalS(..)+ , PolyS(..)+ , sumS+ ) where++import Data.Bits ((.|.))+import qualified Data.Map.Strict as Map+import qualified Data.Vector as V++-- | The algebra WMC is parameterised over.+class Semiring a where+ zero :: a+ one :: a+ (.+.) :: a -> a -> a+ (.*.) :: a -> a -> a++infixl 6 .+.+infixl 7 .*.++-- | Sum a list in the semiring.+sumS :: Semiring s => [s] -> s+sumS = foldl' (.+.) zero++-- | Probability semiring for precise WMC.+newtype ProbS = ProbS { unProb :: Double }+ deriving newtype (Show, Eq)++instance Semiring ProbS where+ zero = ProbS 0+ one = ProbS 1+ ProbS a .+. ProbS b = ProbS (a + b)+ ProbS a .*. ProbS b = ProbS (a * b)++-- | Dual number semiring for forward-mode AD through WMC.+data DualS = DualS !Double !(V.Vector Double)+ deriving stock (Show, Eq)++instance Semiring DualS where+ zero = DualS 0 V.empty+ one = DualS 1 V.empty+ DualS a da .+. DualS b db = DualS (a + b) (vplus da db)+ DualS a da .*. DualS b db =+ DualS (a * b) (vplus (vscale a db) (vscale b da))++vplus :: V.Vector Double -> V.Vector Double -> V.Vector Double+vplus a b+ | V.null a = b+ | V.null b = a+ | otherwise = V.zipWith (+) a b++vscale :: Double -> V.Vector Double -> V.Vector Double+vscale s v = V.map (* s) v++-- | Interval semiring for one-pass, sound WMC.+data IntervalS = IntervalS !Double !Double+ deriving stock (Show, Eq)++instance Semiring IntervalS where+ zero = IntervalS 0 0+ one = IntervalS 1 1+ IntervalS a b .+. IntervalS c d = IntervalS (a + c) (b + d)+ IntervalS a b .*. IntervalS c d = IntervalS (a * c) (b * d)++-- | Polynomial semiring for symbolic WMC.+newtype PolyS = PolyS { unPoly :: Map.Map Integer Double }+ deriving stock (Show, Eq)++instance Semiring PolyS where+ zero = PolyS Map.empty+ one = PolyS (Map.singleton 0 1)+ PolyS a .+. PolyS b = PolyS $ Map.filter (/= 0) $ Map.unionWith (+) a b+ PolyS a .*. PolyS b = PolyS $ Map.filter (/= 0) $ Map.fromListWith (+)+ [ (ma .|. mb, ca * cb)+ | (ma, ca) <- Map.toList a+ , (mb, cb) <- Map.toList b+ ]
+ test/Main.hs view
@@ -0,0 +1,19 @@+module Main where++import Test.Tasty+import qualified Test.Semiring as Semiring+import qualified Test.BDD as BDD+import qualified Test.DSL as DSL+import qualified Test.Combinators as Combinators+import qualified Test.Inference as Inference+import qualified Test.Examples as Examples++main :: IO ()+main = defaultMain $ testGroup "imp"+ [ Semiring.tests+ , BDD.tests+ , DSL.tests+ , Combinators.tests+ , Inference.tests+ , Examples.tests+ ]
+ test/Test/BDD.hs view
@@ -0,0 +1,205 @@+-- | Tests the BDD manager and WMC algorithm.+module Test.BDD (tests) where++import Control.Monad.State.Strict (gets, runState)+import Data.IntMap.Strict (IntMap)+import qualified Data.IntMap.Strict as IntMap+import Test.Tasty+import Test.Tasty.HUnit++import Imp.BDD+import Imp.BDD.Builder+import Imp.BDD.WMC+import Imp.Semiring++probsOf :: [(VarLabel, Double)] -> IntMap Weight+probsOf probs = IntMap.fromList [ (unVarLabel vl, Prob p) | (vl, p) <- probs ]++probF :: Weight -> (ProbS, ProbS)+probF w = case w of+ Prob p -> (ProbS (1 - p), ProbS p)+ Knight _ -> error "probF: not used in these tests"++-- | Run a manager action from an empty manager.+runBDD :: BDDM a -> a+runBDD m = fst (runState m emptyManager)++nodeCount :: BDDM Int+nodeCount = gets (IntMap.size . nodeTable)++tests :: TestTree+tests = testGroup "BDD"+ [ testGroup "Variables and negation"+ [ testCase "newVar creates distinct variables" $+ runBDD (do+ (v1, l1) <- newVar+ (v2, l2) <- newVar+ return (v1 /= v2 && l1 /= l2)) @? "variables should be distinct"++ , testCase "bddNot on terminals" $ do+ bddNot BDDTrue @?= BDDFalse+ bddNot BDDFalse @?= BDDTrue++ , testCase "bddNot is an involution" $+ runBDD (do+ (v, _) <- newVar+ return (bddNot (bddNot v) == v)) @? "not . not = id"+ ]++ , testGroup "ITE terminal cases"+ [ testCase "bddAnd/bddOr with constants" $+ runBDD (do+ (v, _) <- newVar+ at <- bddAnd v BDDTrue+ af <- bddAnd v BDDFalse+ of' <- bddOr v BDDFalse+ ot <- bddOr v BDDTrue+ return [at == v, af == BDDFalse, of' == v, ot == BDDTrue])+ @?= [True, True, True, True]++ , testCase "bddIte f False True = bddNot f" $+ runBDD (do+ (v, _) <- newVar+ r <- bddIte v BDDFalse BDDTrue+ return (r == bddNot v)) @? "ite f 0 1 should be the complemented edge"++ , testCase "bddIte f g g = g" $+ runBDD (do+ (a, _) <- newVar+ (b, _) <- newVar+ r <- bddIte a b b+ return (r == b)) @? "both branches equal"++ , testCase "bddAny of no disjuncts is False" $+ runBDD (bddAny []) @?= BDDFalse++ , testCase "bddAny of one disjunct is itself" $+ runBDD (do+ (v, _) <- newVar+ r <- bddAny [v]+ return (r == v)) @? "singleton disjunction"++ , testCase "bddAny agrees with a right fold of bddOr" $+ runBDD (do+ (a, _) <- newVar+ (b, _) <- newVar+ (c, _) <- newVar+ balanced <- bddAny [a, b, c]+ folded <- bddOr a =<< bddOr b c+ return (balanced == folded)) @? "balanced fold = linear fold"+ ]++ , testGroup "Canonicalisation and sharing"+ [ testCase "De Morgan: a AND b and NOT (NOT a OR NOT b) are the same node" $+ runBDD (do+ (a, _) <- newVar+ (b, _) <- newVar+ conj <- bddAnd a b+ disj <- bddOr (bddNot a) (bddNot b)+ return (conj == bddNot disj)) @? "complement pair must not be interned twice"++ , testCase "De Morgan interns 3 nodes, not 4" $+ runBDD (do+ (a, _) <- newVar+ (b, _) <- newVar+ _ <- bddAnd a b+ _ <- bddOr (bddNot a) (bddNot b)+ nodeCount) @?= 3++ , testCase "bddAnd is commutative and shares the node" $+ runBDD (do+ (a, _) <- newVar+ (b, _) <- newVar+ ab <- bddAnd a b+ ba <- bddAnd b a+ n <- nodeCount+ return (ab == ba, n)) @?= (True, 3)+ ]++ , testGroup "Restrict"+ [ testCase "restrict a AND b on the top variable" $+ runBDD (do+ (a, la) <- newVar+ (b, _) <- newVar+ ab <- bddAnd a b+ t <- bddRestrict ab la True+ f <- bddRestrict ab la False+ return (t == b, f)) @?= (True, BDDFalse)++ , testCase "restrict a AND b on the lower variable" $+ runBDD (do+ (a, _) <- newVar+ (b, lb) <- newVar+ ab <- bddAnd a b+ t <- bddRestrict ab lb True+ f <- bddRestrict ab lb False+ return (t == a, f)) @?= (True, BDDFalse)++ , testCase "restrict on a variable above the root is a no-op" $+ runBDD (do+ (_, la) <- newVar+ (b, _) <- newVar+ r <- bddRestrict b la True+ return (r == b)) @? "variable not in this sub-BDD"++ , testCase "restrict terminals" $+ runBDD (do+ (_, la) <- newVar+ t <- bddRestrict BDDTrue la True+ f <- bddRestrict BDDFalse la True+ return (t, f)) @?= (BDDTrue, BDDFalse)+ ]++ , testGroup "lookupNode"+ [ testCase "terminals have no node" $+ runBDD (do+ t <- lookupNode BDDTrue+ f <- lookupNode BDDFalse+ return (t, f)) @?= (Nothing, Nothing)++ , testCase "a complemented reference negates both children" $+ runBDD (do+ (v, _) <- newVar+ pos <- lookupNode v+ neg <- lookupNode (bddNot v)+ return (pos, neg))+ @?= ( Just (BDDNode (VarLabel 0) BDDFalse BDDTrue)+ , Just (BDDNode (VarLabel 0) BDDTrue BDDFalse) )+ ]++ , testGroup "WMC"+ [ testCase "constants" $ do+ unProb (wmc probF emptyManager IntMap.empty BDDTrue) @?= 1.0+ unProb (wmc probF emptyManager IntMap.empty BDDFalse) @?= 0.0++ , testCase "single variable" $ do+ let ((var, vl), mgr) = runState newVar emptyManager+ unProb (wmc probF mgr (probsOf [(vl, 0.7)]) var) @?= 0.7++ , testCase "AND of two variables" $ do+ let ((r, l1, l2), mgr) = runState (do+ (a, la) <- newVar+ (b, lb) <- newVar+ ab <- bddAnd a b+ return (ab, la, lb)) emptyManager+ unProb (wmc probF mgr (probsOf [(l1, 0.5), (l2, 0.5)]) r) @?= 0.25++ , testCase "OR of two variables (complemented edge under a weight)" $ do+ let ((r, l1, l2), mgr) = runState (do+ (a, la) <- newVar+ (b, lb) <- newVar+ ab <- bddOr a b+ return (ab, la, lb)) emptyManager+ unProb (wmc probF mgr (probsOf [(l1, 0.5), (l2, 0.5)]) r) @?= 0.75++ , testCase "wmcBatch over a shared manager" $ do+ let ((bdds, l1, l2), mgr) = runState (do+ (a, la) <- newVar+ (b, lb) <- newVar+ conj <- bddAnd a b+ disj <- bddOr a b+ return ([a, b, conj, disj], la, lb)) emptyManager+ map unProb (wmcBatch probF mgr (probsOf [(l1, 0.5), (l2, 0.25)]) bdds)+ @?= [0.5, 0.25, 0.125, 0.625]+ ]+ ]
+ test/Test/Combinators.hs view
@@ -0,0 +1,188 @@+-- | Tests the iteration combinators. The type signatures below pin+-- 'GenNames' and 'ConcatMapTag' at compile time.+{-# LANGUAGE QualifiedDo #-}+module Test.Combinators (tests) where++import Prelude hiding (return, (>>=), (>>), flip)+import qualified Data.Map.Strict as Map+import Test.Tasty+import Test.Tasty.HUnit+import Test.Util (assertApprox, assertBounds, assertMap, knightNames)++import qualified Imp.DSL as Imp+import Imp.DSL (Imp, interval, observe)+import Imp.DSL.Combinators+import Imp.Inference+import Imp.Examples.IMDP as IMDP++-- | Must reproduce the hand-written 'IMDP.complexRobot', grade included.+robotFold2 :: Imp '["move1.b", "move1.f", "move2.b", "move2.f"] Position+robotFold2 = foldN @2 @"move" robotDynamics P0 step3++-- | A third step, which is the only way to reach @step3 P2 _@.+robotFold3 :: Imp '["move1.b", "move1.f", "move2.b", "move2.f"+ , "move3.b", "move3.f"] Position+robotFold3 = foldN @3 @"move" robotDynamics P0 step3++-- | The trajectory rather than the endpoint.+robotScan2 :: Imp '["move1.b", "move1.f", "move2.b", "move2.f"] [Position]+robotScan2 = scanN @2 @"move" robotDynamics P0 step3++-- | The moves themselves, with no accumulator.+robotTag2 :: Imp '["move1.b", "move1.f", "move2.b", "move2.f"] [Move]+robotTag2 = tagN @2 @"move" robotDynamics++-- | Explicit-tag equivalents of the three above.+robotTagFold, robotTagScan :: Imp '["move1.b", "move1.f", "move2.b", "move2.f"] Position+robotTagFold = tagFold @'["move1", "move2"] robotDynamics P0 step3+robotTagScan = Imp.fmap last (tagScan @'["move1", "move2"] robotDynamics P0 step3)++robotTagMap :: Imp '["move1.b", "move1.f", "move2.b", "move2.f"] [Move]+robotTagMap = tagMap @'["move1", "move2"] robotDynamics++-- | A single imprecise step, reused by the numbered folds below.+oneStep :: Imp '["d"] Bool+oneStep = interval @"d" 0.4 0.6++-- | Count the @True@ steps.+countTrue :: Int -> Bool -> Int+countTrue acc x = if x then acc + 1 else acc++trueFold :: Imp '["t1.d", "t2.d", "t3.d"] Int+trueFold = foldN @3 @"t" oneStep 0 countTrue++trueScan :: Imp '["t1.d", "t2.d", "t3.d"] [Int]+trueScan = scanN @3 @"t" oneStep 0 countTrue++-- | 'foldMN' can condition at each step; here at most two steps may succeed.+cappedFold :: Imp '["t1.d", "t2.d", "t3.d"] Int+cappedFold = foldMN @3 @"t" oneStep 0 $ \acc x -> Imp.do+ let acc' = countTrue acc x+ observe (acc' < 3)+ Imp.return acc'++-- | 'GenNames' pads to the width of the largest index, keeping grades sorted.+names0 :: Imp '[] [Bool]+names0 = knightN @0 @"k"++names9 :: Imp '["k1","k2","k3","k4","k5","k6","k7","k8","k9"] [Bool]+names9 = knightN @9 @"k"++names10 :: Imp '["k01","k02","k03","k04","k05","k06","k07","k08","k09","k10"] [Bool]+names10 = knightN @10 @"k"++tests :: TestTree+tests = testGroup "Combinators"+ [ testGroup "mapName family"+ [ testCase "knightMap: one free choice per name" $+ assertMap "knightMap"+ [ ([False, False], (0, 1)), ([False, True], (0, 1))+ , ([True, False], (0, 1)), ([True, True], (0, 1)) ]+ (marginal (knightMap @'["x", "y"]))++ , testCase "intervalMap: two independent [0.3, 0.7] choices" $+ assertMap "intervalMap"+ [ ([False, False], (0.09, 0.49)), ([False, True], (0.09, 0.49))+ , ([True, False], (0.09, 0.49)), ([True, True], (0.09, 0.49)) ]+ (marginal (intervalMap @'["x", "y"] 0.3 0.7))++ , testCase "intervalN @3: P(all three) = [0.3^3, 0.7^3]" $ do+ assertBounds "all" (0.027, 0.343) (intervalProbability (intervalN @3 @"s" 0.3 0.7) and)+ assertBounds "any" (0.657, 0.973) (intervalProbability (intervalN @3 @"s" 0.3 0.7) or)++ , testCase "knightN @0 is the empty list at unit grade" $+ assertMap "knightN @0" [([], (1, 1))] (marginal names0)++ , testCase "knightN @3: 2^3 valuations, all unconstrained" $ do+ knightNames (knightN @3 @"k" :: Imp '["k1", "k2", "k3"] [Bool])+ @?= ["k1", "k2", "k3"]+ length (credalVertices (knightN @3 @"k" :: Imp '["k1", "k2", "k3"] [Bool])) @?= 8+ ]++ , testGroup "Tagged iteration over robotDynamics"+ [ testCase "foldN @2 reproduces complexRobot" $ do+ marginal robotFold2 @?= marginal IMDP.complexRobot+ length (credalVertices robotFold2) @?= 16+ assertBounds "P(P2)" (0.25, 0.64) (intervalProbability robotFold2 (== P2))++ , testCase "foldN @3: three steps, P2 absorbing" $+ assertMap "foldN @3"+ [ (P0, (0.008, 0.195)), (P1, (0.0832, 0.375)), (P2, (0.475, 0.896)) ]+ (marginal robotFold3)++ , testCase "scanN @2: trajectories, and P0 cannot jump to P2" $ do+ assertMap "scanN @2"+ [ ([P0, P0], (0.04, 0.25)), ([P0, P1], (0.1, 0.4))+ , ([P1, P0], (0.0, 0.08)), ([P1, P1], (0.08, 0.4))+ , ([P1, P2], (0.25, 0.64)) ]+ (marginal robotScan2)+ assertBounds "last = P2" (0.25, 0.64)+ (intervalProbability robotScan2 ((== P2) . last))++ , testCase "tagN @2: the move pairs" $+ assertMap "tagN @2"+ [ ([Backwards, Backwards], (0.0, 0.01))+ , ([Backwards, Stationary], (0.0, 0.05))+ , ([Backwards, Forwards], (0.0, 0.08))+ , ([Stationary, Backwards], (0.0, 0.05))+ , ([Stationary, Stationary],(0.0256, 0.25))+ , ([Stationary, Forwards], (0.08, 0.4))+ , ([Forwards, Backwards], (0.0, 0.08))+ , ([Forwards, Stationary], (0.08, 0.4))+ , ([Forwards, Forwards], (0.25, 0.64)) ]+ (marginal robotTag2)++ , testCase "tagN @2: the list is in tag order, not reversed" $ do+ -- Identical steps make every bound symmetric under reversal,+ -- so using per-tag weights.+ let fwdThenBack ms = case ms of [Forwards, Backwards] -> 1.0; _ -> 0.0+ (weights, val) = optimizeExpectation robotTag2 fwdThenBack 500 0.5+ weights @?= Map.fromList+ [("move1.b", 0.5), ("move1.f", 1.0), ("move2.b", 1.0), ("move2.f", 0.0)]+ assertApprox "P(forwards then backwards)" 0.08 val++ , testCase "the numbered wrappers agree with the explicit-tag versions" $ do+ marginal robotTagFold @?= marginal robotFold2+ marginal robotTagMap @?= marginal robotTag2+ marginal robotTagScan @?= marginal robotFold2+ ]++ , testGroup "Numbered iteration over a single interval"+ [ testCase "foldN @3: count of successes" $+ assertMap "foldN @3"+ [ (0, (0.064, 0.216)), (1, (0.288, 0.432))+ , (2, (0.288, 0.432)), (3, (0.064, 0.216)) ]+ (marginal trueFold)++ , testCase "scanN @3: every path has the same bounds" $+ assertMap "scanN @3"+ [ ([0,0,0], (0.064, 0.216)), ([0,0,1], (0.064, 0.216))+ , ([0,1,1], (0.064, 0.216)), ([0,1,2], (0.064, 0.216))+ , ([1,1,1], (0.064, 0.216)), ([1,1,2], (0.064, 0.216))+ , ([1,2,2], (0.064, 0.216)), ([1,2,3], (0.064, 0.216)) ]+ (marginal trueScan)++ , testCase "foldMN @3: the per-step observe removes the all-successes path" $ do+ assertMap "foldMN @3"+ [ (0, (4/49, 3/13)), (1, (18/49, 6/13))+ , (2, (4/13, 27/49)), (3, (0, 0)) ]+ (marginal cappedFold)+ assertMap "foldMN @3 symbolic"+ [ (0, (4/49, 3/13)), (1, (18/49, 6/13))+ , (2, (4/13, 27/49)), (3, (0, 0)) ]+ (marginalSymbolic cappedFold)++ , testCase "foldMN @3: tag names join with a dot, one per step" $ do+ let (weights, prob) = optimizeProbability cappedFold (== 2) 200 0.1+ weights @?= Map.fromList [("t1.d", 1.0), ("t2.d", 1.0), ("t3.d", 1.0)]+ assertApprox "ascent reaches the exact upper bound" (27/49) prob+ ]++ , testGroup "GenNames padding"+ [ testCase "one digit up to 9, two digits from 10" $ do+ knightNames names0 @?= []+ knightNames names9 @?= ["k1","k2","k3","k4","k5","k6","k7","k8","k9"]+ knightNames names10+ @?= ["k01","k02","k03","k04","k05","k06","k07","k08","k09","k10"]+ ]+ ]
+ test/Test/DSL.hs view
@@ -0,0 +1,125 @@+-- | Tests the graded monad primitives and the grade algebra. The type+-- signatures below only compile if 'Merge', 'Union' and 'TagAll' agree.+{-# LANGUAGE QualifiedDo #-}+module Test.DSL (tests) where++import Prelude hiding (return, (>>=), (>>), flip)+import Test.Tasty+import Test.Tasty.HUnit+import Test.Util (assertApprox, assertBounds, assertDist, assertMap, knightNames)++import qualified Imp.DSL as Imp+import Imp.DSL (Imp, flip, ifThenElse, interval, knight, observe, tag)+import Imp.Inference++-- | 'Merge' sorts, so binding out of alphabetical order still gives a sorted grade.+sorted :: Imp '["a", "b", "c"] Bool+sorted = Imp.do+ c <- knight @"c"+ a <- knight @"a"+ b <- knight @"b"+ Imp.return (a && b && c)++-- | The grade of a branch is the 'Union' of its arms. Explicit because+-- this module does not enable @RebindableSyntax@.+branched :: Imp '["x", "y"] Bool+branched = Imp.do+ c <- flip 0.5+ ifThenElse c+ (Imp.do { a <- knight @"x"; Imp.return a })+ (Imp.do { b <- knight @"y"; Imp.return (not b) })++-- | 'TagAll' prefixes every name in the subprogram.+tagged :: Imp '["t.k"] Bool+tagged = tag @"t" (knight @"k")++-- | Nested tags compose left to right.+nested :: Imp '["a.b.k"] Bool+nested = tag @"a" (tag @"b" (knight @"k"))++-- | An 'interval' keeps its bounds under a tag.+taggedInterval :: Imp '["m.i"] Bool+taggedInterval = tag @"m" (interval @"i" 0.25 0.75)++-- | The README's conditioning example.+conditioned :: Imp '["bias"] Bool+conditioned = Imp.do+ biased <- interval @"bias" 0.3 0.7+ observe biased+ Imp.return biased++tests :: TestTree+tests = testGroup "DSL"+ [ testGroup "Primitives"+ [ testCase "flip p is a Bernoulli" $+ assertDist "flip 0.7" [(False, 0.3), (True, 0.7)] (preciseMarginal (flip 0.7))++ , testCase "knight is completely unconstrained" $+ assertMap "knight" [(False, (0, 1)), (True, (0, 1))] (marginal (knight @"k"))++ , testCase "interval lo hi bounds both outcomes" $ do+ assertMap "interval" [(False, (0.25, 0.75)), (True, (0.25, 0.75))]+ (marginal (interval @"i" 0.25 0.75))+ length (credalVertices (interval @"i" 0.25 0.75)) @?= 2++ , testCase "observe renormalises away the rejected world" $ do+ let prog = Imp.do { h <- flip 0.5; observe h; Imp.return h }+ assertDist "observed" [(True, 1.0)] (preciseMarginal prog)++ , testCase "(>>) discards the first result but keeps its evidence" $+ assertDist "sequenced" [(False, 0.75), (True, 0.25)]+ (preciseMarginal (observe True Imp.>> flip 0.25))++ , testCase "graded fmap maps the value, not the weight" $ do+ assertDist "fmap" [(False, 0.7), (True, 0.3)]+ (preciseMarginal (Imp.fmap not (flip 0.7)))+ assertDist "<$>" [(False, 0.7), (True, 0.3)]+ (preciseMarginal (not Imp.<$> flip 0.7))+ ]++ , testGroup "Grades"+ [ testCase "Merge sorts names regardless of bind order" $+ knightNames sorted @?= ["a", "b", "c"]++ , testCase "Union: a probabilistic condition reaches both branches" $ do+ knightNames branched @?= ["x", "y"]+ length (credalVertices branched) @?= 4+ assertMap "branched" [(False, (0.0, 1.0)), (True, (0.0, 1.0))] (marginal branched)++ , testCase "ifThenElse on non-Imp branches is the ordinary conditional" $ do+ ifThenElse True "yes" "no" @?= "yes"+ ifThenElse False "yes" "no" @?= "no"++ , testCase "tag prefixes the Knightian name" $+ knightNames tagged @?= ["t.k"]++ , testCase "nested tags join with dots, outermost first" $+ knightNames nested @?= ["a.b.k"]++ , testCase "the empty tag leaves the name alone" $+ knightNames (tag @"t" (tag @"" (knight @"k")) :: Imp '["t.k"] Bool) @?= ["t.k"]++ , testCase "tag preserves the interval bounds it wraps" $ do+ knightNames taggedInterval @?= ["m.i"]+ assertMap "tagged interval" [(False, (0.25, 0.75)), (True, (0.25, 0.75))]+ (marginal taggedInterval)+ ]++ , testGroup "Conditioning"+ [ testCase "conditioned: every backend gives P(True) = 1" $ do+ assertMap "marginal" [(False, (0, 0)), (True, (1, 1))] (marginal conditioned)+ assertMap "marginalSymbolic" [(False, (0, 0)), (True, (1, 1))]+ (marginalSymbolic conditioned)+ assertMap "marginalApprox" [(False, (0, 0)), (True, (1, 1))]+ (marginalApprox conditioned)+ assertBounds "intervalProbability" (1, 1) (intervalProbability conditioned id)+ assertBounds "intervalProbabilityApprox" (1, 1)+ (intervalProbabilityApprox conditioned id)+ assertBounds "intervalProbabilitySymbolic" (1, 1)+ (intervalProbabilitySymbolic conditioned id)++ , testCase "conditioned: optimization cannot leave the credal set" $ do+ assertApprox "ascent" 1.0 (snd (optimizeProbability conditioned id 200 0.1))+ assertApprox "descent" 1.0 (snd (optimizeProbability conditioned id 200 (-0.1)))+ ]+ ]
+ test/Test/Examples.hs view
@@ -0,0 +1,212 @@+-- | Facts about the example programs in Imp.Examples.*+module Test.Examples (tests) where++import Data.List (nub, sort)+import qualified Data.Map.Strict as Map+import Test.Tasty+import Test.Tasty.HUnit+import Test.Util (assertApprox, assertBounds, assertDist, assertMap, roundDist)+import Imp.Inference+import Imp.Examples.Basic as Basic+import Imp.Examples.Ellsberg as E+import Imp.Examples.IMDP as IMDP+import Imp.Examples.Iteration as Iter+import Imp.Examples.Knightian as K+import Imp.Examples.MontyHall as MH+import Imp.Examples.Polytope as P+import Imp.Examples.TwoChild as TC++tests :: TestTree+tests = testGroup "Examples"+ [ testGroup "Basic"+ [ testCase "fairCoin" $+ assertDist "fairCoin" [(False, 0.5), (True, 0.5)] (preciseMarginal Basic.fairCoin)++ , testCase "biasedCoin" $+ assertDist "biasedCoin" [(False, 0.3), (True, 0.7)]+ (preciseMarginal Basic.biasedCoin)++ , testCase "twoCoins: uniform over the four pairs" $+ assertDist "twoCoins"+ [ ((False, False), 0.25), ((False, True), 0.25)+ , ((True, False), 0.25), ((True, True), 0.25) ]+ (preciseMarginal Basic.twoCoins)++ , testCase "coinOr: P(True) = 0.75" $+ assertDist "coinOr" [(False, 0.25), (True, 0.75)] (preciseMarginal Basic.coinOr)+ ]++ , testGroup "Knightian names"+ [ testCase "dependent: line segment, 2 vertices" $ do+ let verts = credalVertices K.dependent+ length verts @?= 2+ sort (map roundDist verts) @?= sort+ [ Map.fromList [(K.Blue, 0.0), (K.Green, 0.0), (K.Red, 1.0)]+ , Map.fromList [(K.Blue, 0.5), (K.Green, 0.5), (K.Red, 0.0)]+ ]++ , testCase "independent: quadrilateral, 4 vertices" $ do+ let verts = credalVertices K.independent+ length verts @?= 4+ sort (map roundDist verts) @?= sort+ [ Map.fromList [(K.Blue, 0.0), (K.Green, 0.0), (K.Red, 1.0)]+ , Map.fromList [(K.Blue, 0.5), (K.Green, 0.0), (K.Red, 0.5)]+ , Map.fromList [(K.Blue, 0.0), (K.Green, 0.5), (K.Red, 0.5)]+ , Map.fromList [(K.Blue, 0.5), (K.Green, 0.5), (K.Red, 0.0)]+ ]++ , testCase "independent is strictly larger than dependent" $ do+ let dep = map roundDist (credalVertices K.dependent)+ ind = map roundDist (credalVertices K.independent)+ all (`elem` ind) dep @? "every dependent vertex is an independent one"+ length (nub ind) > length (nub dep) @? "independent has strictly more vertices"++ , testCase "sharing a name does not change the marginals, only the joint" $ do+ let bounds = [ (K.Red, (0.0, 1.0)), (K.Green, (0.0, 0.5)), (K.Blue, (0.0, 0.5)) ]+ assertMap "dependent" bounds (marginal K.dependent)+ assertMap "independent" bounds (marginal K.independent)++ , testCase "dependent: E[score]" $+ assertBounds "E" (0.25, 1.0)+ (intervalExpectation K.dependent+ (\v -> case v of K.Red -> 1.0; K.Green -> 0.5; K.Blue -> 0.0))+ ]++ , testGroup "MontyHall"+ [ testCase "P(switch wins) = 2/3 regardless of host bias" $+ assertMap "montyHall" [(False, (1/3, 1/3)), (True, (2/3, 2/3))]+ (marginal MH.montyHall)+ ]++ , testGroup "TwoChild"+ [ testCase "P(both boys | at least one boy)" $+ assertMap "twoChild" [(False, (0.5, 1.0)), (True, (0.0, 0.5))]+ (marginal TC.twoChild)++ , testCase "credal vertices: P(True) takes values 0 and 1/2" $ do+ let verts = credalVertices TC.twoChild+ sort (map (Map.findWithDefault 0 True) verts) @?= [0.0, 0.5]+ ]++ , testGroup "Polytope"+ [ testCase "polytope: 4 vertices, all summing to 1" $ do+ let verts = credalVertices P.polytope+ length verts @?= 4+ mapM_ (\dist -> assertApprox "sum" 1.0 (sum dist)) verts++ , testCase "polytope: marginal bounds" $+ assertMap "polytope"+ [ (P.Red, (0.2, 0.8)), (P.Green, (0.06, 0.56)), (P.Blue, (0.06, 0.56)) ]+ (marginal P.polytope)++ , testCase "polytope2: 8 vertices, all distinct" $ do+ let verts = credalVertices P.polytope2+ length verts @?= 8+ length (nub (map roundDist verts)) @?= 8++ , testCase "polytope2: Red picks up the fall-through branch" $+ assertMap "polytope2"+ [ (P.Red, (0.4, 0.82)), (P.Green, (0.1, 0.4)), (P.Blue, (0.05, 0.32)) ]+ (marginal P.polytope2)+ ]++ , testGroup "Iteration"+ [ testCase "walk1: 2 valuations, 2 distinct vertices" $ do+ let verts = credalVertices Iter.walk1+ length verts @?= 2+ length (nub (map roundDist verts)) @?= 2+ assertMap "walk1" [(0, (0.3, 0.7)), (1, (0.3, 0.7))] (marginal Iter.walk1)++ , testCase "walk2: 4 valuations, 3 distinct" $ do+ let verts = credalVertices Iter.walk2+ length verts @?= 4+ length (nub (map roundDist verts)) @?= 3+ assertMap "walk2" [(0, (0.09, 0.49)), (1, (0.42, 0.58)), (2, (0.09, 0.49))]+ (marginal Iter.walk2)++ , testCase "walk3: 8 valuations, 4 distinct" $ do+ let verts = credalVertices Iter.walk3+ length verts @?= 8+ length (nub (map roundDist verts)) @?= 4+ assertMap "walk3"+ [ (0, (0.027, 0.343)), (1, (0.189, 0.469))+ , (2, (0.189, 0.469)), (3, (0.027, 0.343)) ]+ (marginal Iter.walk3)++ , testCase "walk3Asym: 8 valuations, 8 distinct" $ do+ let verts = credalVertices Iter.walk3Asym+ length verts @?= 8+ length (nub (map roundDist verts)) @?= 8+ assertMap "walk3Asym"+ [ (0, (0.0135, 0.3465)), (1, (0.1695, 0.4995))+ , (2, (0.1695, 0.4995)), (3, (0.0135, 0.3465)) ]+ (marginal Iter.walk3Asym)+ ]++ , testGroup "Ellsberg paradox"+ [ testCase "credal set: 2 vertices" $+ length (credalVertices E.ellsberg) @?= 2++ , testCase "Red is precise, Black and Yellow are maximally imprecise" $+ assertMap "ellsberg"+ [ (E.Red, (1/3, 1/3)), (E.Black, (0.0, 2/3)), (E.Yellow, (0.0, 2/3)) ]+ (marginal E.ellsberg)++ , testCase "P(Black or Yellow) = [2/3, 2/3]" $+ assertBounds "P" (2/3, 2/3)+ (intervalProbability E.ellsberg (\b -> b == E.Black || b == E.Yellow))++ , testCase "ambiguity aversion: known gambles are precise" $ do+ assertBounds "I (bet Red)" (1/3, 1/3)+ (intervalExpectation E.ellsberg (\b -> if b == E.Red then 1.0 else 0.0))+ assertBounds "II (bet Black)" (0.0, 2/3)+ (intervalExpectation E.ellsberg (\b -> if b == E.Black then 1.0 else 0.0))+ assertBounds "IV (bet Black or Yellow)" (2/3, 2/3)+ (intervalExpectation E.ellsberg+ (\b -> if b == E.Black || b == E.Yellow then 1.0 else 0.0))+ ]++ , testGroup "Interval MDP"+ [ testCase "simpleRobot: 4 valuations, 3 distinct" $ do+ let verts = credalVertices IMDP.simpleRobot+ length verts @?= 4+ length (nub (map roundDist verts)) @?= 3++ , testCase "simpleRobot: marginal bounds" $+ assertMap "simpleRobot"+ [ (IMDP.P0, (0.01, 0.16)), (IMDP.P1, (0.18, 0.48)), (IMDP.P2, (0.36, 0.81)) ]+ (marginal IMDP.simpleRobot)++ , testCase "simpleRobot3: 8 valuations, 4 distinct" $ do+ let verts = credalVertices IMDP.simpleRobot3+ length verts @?= 8+ length (nub (map roundDist verts)) @?= 4++ , testCase "simpleRobot3: E[position]" $ do+ assertMap "simpleRobot3"+ [ (IMDP.P0, (0.001, 0.064)), (IMDP.P1, (0.027, 0.288))+ , (IMDP.P2, (0.648, 0.972)) ]+ (marginal IMDP.simpleRobot3)+ assertBounds "E[position]" (1.584, 1.971)+ (intervalExpectation IMDP.simpleRobot3+ (\v -> case v of IMDP.P0 -> 0; IMDP.P1 -> 1; IMDP.P2 -> 2))+ ]++ , testGroup "Compositional robot"+ [ testCase "complexRobot: 16 valuations" $ do+ let verts = credalVertices IMDP.complexRobot+ length verts @?= 16+ mapM_ (\dist -> assertApprox "sum" 1.0 (sum dist)) verts++ , testCase "complexRobot: exact marginal bounds" $+ assertMap "complexRobot"+ [ (IMDP.P0, (0.04, 0.3)), (IMDP.P1, (0.288, 0.5)), (IMDP.P2, (0.25, 0.64)) ]+ (marginal IMDP.complexRobot)++ , testCase "complexRobot: P(reach goal) = [0.25, 0.64]" $ do+ assertBounds "enumeration" (0.25, 0.64)+ (intervalProbability IMDP.complexRobot (== IMDP.P2))+ assertBounds "symbolic" (0.25, 0.64)+ (intervalProbabilitySymbolic IMDP.complexRobot (== IMDP.P2))+ ]+ ]
+ test/Test/Inference.hs view
@@ -0,0 +1,322 @@+-- | Behaviour of the inference backends on minimal inline programs, plus+-- cross-backend agreement. Example-program facts live in Test.Examples.+{-# LANGUAGE QualifiedDo #-}+module Test.Inference (tests) where++import Prelude hiding (return, (>>=), (>>), flip)+import qualified Data.Map.Strict as Map+import Test.Tasty+import Test.Tasty.HUnit+import Test.Util (assertApprox, assertBounds, assertContains, assertDist, assertMap,+ assertNoFeasible)+import qualified Imp.DSL as Imp+import Imp.DSL (Imp, flip, knight, observe)+import Imp.DSL.Combinators (GenNames, knightN)+import Imp.Inference+import Imp.Examples.Ellsberg as E+import Imp.Examples.IMDP as IMDP+import Imp.Examples.Knightian as K+import Imp.Examples.MontyHall as MH+import Imp.Examples.Polytope as P+import Imp.Examples.TwoChild as TC++andCoins :: Imp '[] Bool+andCoins = Imp.do+ a <- flip 0.5+ b <- flip 0.5+ Imp.return (a && b)++-- | A flip widened by a Knightian choice.+orKnight :: Imp '["k"] Bool+orKnight = Imp.do+ a <- flip 0.5+ k <- knight @"k"+ Imp.return (a || k)++-- | Conditioning on a Knightian choice.+observedKnight :: Imp '["k"] Bool+observedKnight = Imp.do+ k <- knight @"k"+ observe k+ Imp.return k++-- | Conditioning and widening by a Knightian choice.+observedOr :: Imp '["k"] Bool+observedOr = Imp.do+ k <- knight @"k"+ h <- flip 0.5+ observe (k || h)+ Imp.return h++-- | Evidence satisfiable only via a rare fault. A feasibility threshold+-- scaled too loosely discards a genuinely feasible corner.+rareFault :: Imp (GenNames 4 "s") Bool+rareFault = Imp.do+ ss <- knightN @4 @"s"+ f <- flip 1e-18+ observe (not (or ss) || f)+ Imp.return (or ss)++-- | Empty credal set.+infeasible :: Imp '[] Bool+infeasible = Imp.do+ h <- flip 0.5+ observe (h && not h)+ Imp.return h++-- | Mixed-sign reward.+reward :: Position -> Double+reward v = case v of IMDP.P2 -> 10; IMDP.P1 -> -1; _ -> 0++-- | Mixed-sign score over the three-way polytope.+mixedScore :: P.Three -> Double+mixedScore v = case v of P.Red -> 2.0; P.Green -> -1.0; P.Blue -> 0.5++-- | Enumerate and Symbolic must agree on both marginals and expectations.+agree :: (Ord a, Show a) => String -> Imp g a -> (a -> Double) -> Assertion+agree label prog score = do+ let enum = marginal prog+ sym = marginalSymbolic prog+ Map.keys sym @?= Map.keys enum+ mapM_ (\(v, (e, s)) -> assertBounds (label ++ " marginal " ++ show v) e s)+ (Map.toList (Map.intersectionWith (,) enum sym))+ assertBounds (label ++ " expectation")+ (intervalExpectation prog score) (intervalExpectationSymbolic prog score)++tests :: TestTree+tests = testGroup "Inference"+ [ testGroup "Precise"+ [ testCase "fair coin" $+ assertDist "flip 0.5" [(False, 0.5), (True, 0.5)] (preciseMarginal (flip 0.5))++ , testCase "biased coin" $+ assertDist "flip 0.7" [(False, 0.3), (True, 0.7)] (preciseMarginal (flip 0.7))++ , testCase "AND of two fair coins" $+ assertDist "and" [(False, 0.75), (True, 0.25)] (preciseMarginal andCoins)+ ]++ , testGroup "Exact enumeration"+ [ testCase "no Knightian vars: bounds collapse to a point" $ do+ assertBounds "P(flip 0.5)" (0.5, 0.5) (intervalProbability (flip 0.5) id)+ assertBounds "P(flip 0.7)" (0.7, 0.7) (intervalProbability (flip 0.7) id)+ assertBounds "P(and)" (0.25, 0.25) (intervalProbability andCoins id)++ , testCase "flip OR knight: P(True) = [0.5, 1]" $+ assertBounds "P" (0.5, 1.0) (intervalProbability orKnight id)++ , testCase "expectation of a scaled indicator" $+ assertBounds "E" (6.0, 6.0)+ (intervalExpectation (flip 0.6) (\b -> if b then 10.0 else 0.0))++ , testCase "conditioning can make a Knightian valuation infeasible" $ do+ assertMap "marginal" [(False, (0, 0)), (True, (1, 1))] (marginal observedKnight)+ length (credalVertices observedKnight) @?= 1++ , testCase "observedOr: P(True) = [0.5, 1]" $+ assertBounds "P" (0.5, 1.0) (intervalProbability observedOr id)+ ]++ , testGroup "Interval approximation"+ [ testCase "exact when each Knightian var is read once" $+ assertMap "marginalApprox" [(False, (1/3, 1/3)), (True, (2/3, 2/3))]+ (marginalApprox MH.montyHall)++ , testCase "exact when the Knightian name is shared" $+ assertMap "marginalApprox"+ [ (K.Red, (0.0, 1.0)), (K.Green, (0.0, 0.5)), (K.Blue, (0.0, 0.5)) ]+ (marginalApprox K.dependent)++ , testCase "loose but sound on simpleRobot" $ do+ assertMap "marginalApprox"+ [ (IMDP.P0, (0.0, 0.25)), (IMDP.P1, (0.0, 1.0)), (IMDP.P2, (0.0, 1.0)) ]+ (marginalApprox IMDP.simpleRobot)+ containsExact "simpleRobot" (marginal IMDP.simpleRobot)+ (marginalApprox IMDP.simpleRobot)++ , testCase "contains the exact marginal" $ do+ containsExact "complexRobot" (marginal IMDP.complexRobot)+ (marginalApprox IMDP.complexRobot)+ containsExact "polytope2" (marginal P.polytope2) (marginalApprox P.polytope2)+ containsExact "twoChild" (marginal TC.twoChild) (marginalApprox TC.twoChild)++ , testCase "contains the exact probability under conditioning" $ do+ assertContains "observedKnight" (intervalProbability observedKnight id)+ (intervalProbabilityApprox observedKnight id)+ assertContains "observedOr" (intervalProbability observedOr id)+ (intervalProbabilityApprox observedOr id)+ assertContains "twoChild" (intervalProbability TC.twoChild id)+ (intervalProbabilityApprox TC.twoChild id)++ , testCase "contains the exact expectation" $ do+ assertContains "complexRobot" (intervalExpectation IMDP.complexRobot reward)+ (intervalExpectationApprox IMDP.complexRobot reward)+ assertContains "observedOr"+ (intervalExpectation observedOr (\b -> if b then 3.0 else -1.0))+ (intervalExpectationApprox observedOr (\b -> if b then 3.0 else -1.0))++ , testCase "a value with no remaining mass drops out of the score range" $ do+ assertBounds "positive on True" (3.0, 3.0)+ (intervalExpectationApprox observedKnight (\b -> if b then 3.0 else -1.0))+ assertBounds "negative on True" (-1.0, -1.0)+ (intervalExpectationApprox observedKnight (\b -> if b then -1.0 else 3.0))++ , testCase "the score range clamps a sum of loose per-value boxes" $ do+ -- The P0/P1 boxes sum past the extreme score; the clamp pulls it back.+ assertBounds "two negatives" (-1.0, 0.0)+ (intervalExpectationApprox IMDP.simpleRobot+ (\v -> case v of IMDP.P2 -> 0; _ -> -1))+ assertBounds "two positives" (0.0, 1.0)+ (intervalExpectationApprox IMDP.simpleRobot+ (\v -> case v of IMDP.P2 -> 0; _ -> 1))++ , testCase "no Knightian vars: expectation is exact" $+ assertBounds "E" (6.0, 6.0)+ (intervalExpectationApprox (flip 0.6) (\b -> if b then 10.0 else 0.0))++ , testCase "a predicate no world satisfies" $ do+ assertBounds "const True" (1.0, 1.0)+ (intervalProbabilityApprox (flip 0.5) (const True))+ assertBounds "const False" (0.0, 0.0)+ (intervalProbabilityApprox (flip 0.5) (const False))+ ]++ , testGroup "Symbolic"+ [ testCase "no Knightian: P(True) = [0.7, 0.7]" $+ assertBounds "P" (0.7, 0.7) (intervalProbabilitySymbolic (flip 0.7) id)++ , testCase "dependent P(Green) = [0, 0.5]" $+ assertBounds "P" (0.0, 0.5) (intervalProbabilitySymbolic K.dependent (== K.Green))++ , testCase "ellsberg P(Red) = [1/3, 1/3], P(Black) = [0, 2/3]" $ do+ assertBounds "Red" (1/3, 1/3) (intervalProbabilitySymbolic E.ellsberg (== E.Red))+ assertBounds "Black" (0, 2/3) (intervalProbabilitySymbolic E.ellsberg (== E.Black))++ , testCase "montyHall: P(switch wins) = [2/3, 2/3]" $+ assertBounds "P" (2/3, 2/3) (intervalProbabilitySymbolic MH.montyHall id)++ , testCase "polytope P(R) = [0.2, 0.8]" $+ assertBounds "P" (0.2, 0.8) (intervalProbabilitySymbolic P.polytope (== P.Red))++ , testCase "expectation: ellsberg E[bet Red] = [1/3, 1/3]" $+ assertBounds "E" (1/3, 1/3)+ (intervalExpectationSymbolic E.ellsberg (\b -> if b == E.Red then 1.0 else 0.0))++ , testCase "expectation: no Knightian is exact" $+ assertBounds "E" (6.0, 6.0)+ (intervalExpectationSymbolic (flip 0.6) (\b -> if b then 10.0 else 0.0))++ , testCase "at least as tight as the interval approximation" $ do+ assertContains "ellsberg Black" (intervalProbabilitySymbolic E.ellsberg (== E.Black))+ (intervalProbabilityApprox E.ellsberg (== E.Black))+ assertContains "polytope2 Red" (intervalProbabilitySymbolic P.polytope2 (== P.Red))+ (intervalProbabilityApprox P.polytope2 (== P.Red))+ ]++ , testGroup "Enumeration and Symbolic agree"+ [ testCase "dependent" $ agree "dependent" K.dependent+ (\v -> case v of K.Red -> 1.0; K.Green -> 0.5; K.Blue -> 0.0)+ , testCase "complexRobot (mixed-sign reward)" $+ agree "complexRobot" IMDP.complexRobot reward+ , testCase "polytope2 (mixed-sign score)" $+ agree "polytope2" P.polytope2 mixedScore+ , testCase "montyHall" $+ agree "montyHall" MH.montyHall (\b -> if b then 1.0 else 0.0)+ , testCase "twoChild (observe)" $+ agree "twoChild" TC.twoChild (\b -> if b then 3.0 else -1.0)+ , testCase "observedKnight (an infeasible corner)" $+ agree "observedKnight" observedKnight (\b -> if b then 1.0 else 0.0)+ , testCase "observedOr (non-constant denominator)" $+ agree "observedOr" observedOr (\b -> if b then 3.0 else -1.0)+ , testCase "rareFault (a feasible corner with tiny evidence)" $+ agree "rareFault" rareFault (\b -> if b then 1.0 else 0.0)+ ]++ , testGroup "Gradient optimization"+ [ testCase "complexRobot: ascent on P(P2) reaches the exact upper bound" $ do+ let (weights, prob) = optimizeProbability IMDP.complexRobot (== IMDP.P2) 200 0.1+ -- move*.b has zero gradient once move*.f is 1, so it keeps its 0.5 init.+ weights @?= Map.fromList+ [("move1.b", 0.5), ("move1.f", 1.0), ("move2.b", 0.5), ("move2.f", 1.0)]+ assertApprox "P(P2)" 0.64 prob+ assertApprox "= exact upper" prob+ (snd (intervalProbability IMDP.complexRobot (== IMDP.P2)))++ , testCase "complexRobot: ascent on the reward prefers backing off P1" $ do+ let (weights, val) = optimizeExpectation IMDP.complexRobot reward 200 0.1+ weights @?= Map.fromList+ [("move1.b", 0.5), ("move1.f", 1.0), ("move2.b", 1.0), ("move2.f", 1.0)]+ assertApprox "E[reward]" 6.112 val+ assertApprox "= exact upper" (snd (intervalExpectation IMDP.complexRobot reward)) val++ , testCase "complexRobot: descent is not a mirror of ascent" $ do+ let (weights, val) = optimizeExpectation IMDP.complexRobot reward 200 (-0.1)+ weights @?= Map.fromList+ [("move1.b", 0.5), ("move1.f", 0.0), ("move2.b", 0.0), ("move2.f", 0.0)]+ assertApprox "E[reward]" 2.0 val+ assertApprox "= exact lower" (fst (intervalExpectation IMDP.complexRobot reward)) val++ , testCase "ellsberg: P(Red) is precise, so ascent cannot move it" $ do+ let (weights, prob) = optimizeProbability E.ellsberg (== E.Red) 200 0.1+ weights @?= Map.fromList [("split", 0.5)]+ assertApprox "P(Red)" (1/3) prob++ , testCase "ellsberg: E[bet Red] is precise too" $+ assertApprox "E" (1/3)+ (snd (optimizeExpectation E.ellsberg+ (\b -> if b == E.Red then 1 else 0) 200 0.1))++ , testCase "conditioning: iterates stay inside the credal set" $ do+ assertApprox "ascent" 1.0 (snd (optimizeProbability observedOr id 200 0.1))+ assertApprox "descent" 0.5 (snd (optimizeProbability observedOr id 200 (-0.1)))+ let (lo, hi) = intervalProbability observedOr id+ assertBounds "exact" (0.5, 1.0) (lo, hi)++ , testCase "one ascent step moves by the exact gradient" $ do+ -- P(h | k or h) = 1/(1 + p_k), so the derivative at 0.5 is -4/9.+ let (weights, prob) = optimizeProbability observedOr id 1 1.0+ assertApprox "weight" (0.5 - 4/9) (weights Map.! "k")+ assertApprox "P" (18/19) prob++ , testCase "no Knightian vars: empty weights and the exact value" $ do+ let (weights, prob) = optimizeProbability (flip 0.6) id 100 0.1+ weights @?= Map.empty+ assertApprox "prob" 0.6 prob+ let (weights', val) = optimizeExpectation (flip 0.6) (\b -> if b then 10 else 0) 100 0.1+ weights' @?= Map.empty+ assertApprox "E[f]" 6.0 val+ ]++ , testGroup "Empty credal set"+ [ testCase "credalVertices is empty rather than an error" $+ credalVertices infeasible @?= []++ , testCase "every bound-producing entry point throws" $ do+ assertNoFeasible "preciseMarginal" (preciseMarginal infeasible)+ assertNoFeasible "marginal" (marginal infeasible)+ assertNoFeasible "intervalProbability" (intervalProbability infeasible id)+ assertNoFeasible "intervalExpectation"+ (intervalExpectation infeasible (\b -> if b then 1 else 0))+ assertNoFeasible "marginalApprox" (marginalApprox infeasible)+ assertNoFeasible "intervalProbabilityApprox" (intervalProbabilityApprox infeasible id)+ assertNoFeasible "intervalExpectationApprox"+ (intervalExpectationApprox infeasible (\b -> if b then 1 else 0))+ assertNoFeasible "marginalSymbolic" (marginalSymbolic infeasible)+ assertNoFeasible "intervalProbabilitySymbolic"+ (intervalProbabilitySymbolic infeasible id)+ assertNoFeasible "intervalExpectationSymbolic"+ (intervalExpectationSymbolic infeasible (\b -> if b then 1 else 0))+ assertNoFeasible "optimizeProbability" (optimizeProbability infeasible id 100 0.1)+ assertNoFeasible "optimizeExpectation"+ (optimizeExpectation infeasible (\b -> if b then 1 else 0) 100 0.1)+ ]+ ]++-- | Assert that the approximate bounds contain the exact ones.+containsExact :: (Ord a, Show a)+ => String -> Map.Map a (Double, Double) -> Map.Map a (Double, Double)+ -> Assertion+containsExact label exact approx = do+ Map.keys approx @?= Map.keys exact+ mapM_ (\(v, (e, a)) -> assertContains (label ++ " " ++ show v) e a)+ (Map.toList (Map.intersectionWith (,) exact approx))
+ test/Test/Semiring.hs view
@@ -0,0 +1,71 @@+-- | Tests the four WMC semirings directly.+module Test.Semiring (tests) where++import qualified Data.Map.Strict as Map+import qualified Data.Vector as V+import Test.Tasty+import Test.Tasty.HUnit++import Imp.Semiring++-- | Monomial keys are bitmasks, so @x0@ is @bit 0@.+x0, x1, notX0 :: PolyS+x0 = PolyS (Map.singleton 1 1)+x1 = PolyS (Map.singleton 2 1)+notX0 = PolyS (Map.fromList [(0, 1), (1, -1)])++tests :: TestTree+tests = testGroup "Semiring"+ [ testGroup "ProbS"+ [ testCase "sumS adds" $ unProb (sumS [ProbS 0.1, ProbS 0.2, ProbS 0.3]) @?= 0.6000000000000001++ , testCase "sumS of nothing is zero" $ unProb (sumS []) @?= 0.0++ , testCase "identities" $ do+ (zero .+. ProbS 0.25) @?= ProbS 0.25+ (one .*. ProbS 0.25) @?= ProbS 0.25+ (zero .*. ProbS 0.25) @?= ProbS 0.0+ ]++ , testGroup "IntervalS"+ [ testCase "componentwise product and sum" $ do+ (IntervalS 0.2 0.4 .*. IntervalS 0.5 0.6) @?= IntervalS 0.1 0.24+ (IntervalS 0.2 0.4 .+. IntervalS 0.5 0.6) @?= IntervalS 0.7 1.0++ , testCase "identities" $ do+ (zero .+. IntervalS 0.2 0.4) @?= IntervalS 0.2 0.4+ (one .*. IntervalS 0.2 0.4) @?= IntervalS 0.2 0.4+ ]++ , testGroup "DualS"+ [ testCase "product rule" $+ (DualS 2 (V.fromList [1, 0]) .*. DualS 3 (V.fromList [0, 1]))+ @?= DualS 6 (V.fromList [3, 2])++ , testCase "sum adds gradients" $+ (DualS 2 (V.fromList [1, 0]) .+. DualS 3 (V.fromList [0, 1]))+ @?= DualS 5 (V.fromList [1, 1])++ , testCase "empty gradient is absorbed, not zipped away" $ do+ (zero .+. DualS 5 (V.fromList [1, 2])) @?= DualS 5 (V.fromList [1, 2])+ (one .*. DualS 5 (V.fromList [1, 2])) @?= DualS 5 (V.fromList [1, 2])+ ]++ , testGroup "PolyS"+ [ testCase "monomials are idempotent" $+ (x0 .*. x0) @?= x0++ , testCase "distinct monomials union their masks" $+ unPoly (x0 .*. x1) @?= Map.singleton 3 1.0++ , testCase "addition drops cancelling terms" $+ (PolyS (Map.singleton 0 1) .+. PolyS (Map.singleton 0 (-1))) @?= zero++ , testCase "identities" $ do+ unPoly (one :: PolyS) @?= Map.singleton 0 1.0+ unPoly (zero :: PolyS) @?= Map.empty+ (one .*. notX0) @?= notX0+ (zero .+. notX0) @?= notX0+ (notX0 .*. x0) @?= zero+ ]+ ]
+ test/Test/Util.hs view
@@ -0,0 +1,70 @@+-- | Assertion and inspection helpers shared by the test modules.+module Test.Util+ ( assertApprox+ , assertBounds+ , assertMap+ , assertDist+ , assertContains+ , assertNoFeasible+ , knightNames+ , roundDist+ ) where++import Control.Exception (ErrorCall, evaluate, try)+import Data.List (isPrefixOf)+import Data.Map.Strict (Map)+import qualified Data.Map.Strict as Map+import Test.Tasty.HUnit++import Imp.DSL (Imp)+import Imp.Inference (optimizeProbability)++-- | The Knightian variable names a program allocates, as the compiler sees them.+knightNames :: Ord a => Imp g a -> [String]+knightNames prog = Map.keys (fst (optimizeProbability prog (const True) 1 0.1))++-- | Round a distribution's probabilities to 3 decimal places for comparison.+roundDist :: Map a Double -> Map a Double+roundDist = fmap (\v -> fromIntegral (round (v * 1000) :: Int) / 1000)++-- | Assert that two Doubles are approximately equal (within 1e-6).+assertApprox :: String -> Double -> Double -> Assertion+assertApprox label expected actual =+ abs (actual - expected) < 1e-6 @?+ (label ++ ": expected " ++ show expected ++ ", got " ++ show actual)++-- | Assert a lower/upper bound pair.+assertBounds :: String -> (Double, Double) -> (Double, Double) -> Assertion+assertBounds label (elo, ehi) (lo, hi) = do+ assertApprox (label ++ " lower") elo lo+ assertApprox (label ++ " upper") ehi hi++-- | Assert the exact key list and pair of a bounds map.+assertMap :: (Ord k, Show k)+ => String -> [(k, (Double, Double))] -> Map k (Double, Double) -> Assertion+assertMap label expected actual = do+ Map.keys actual @?= map fst expected+ mapM_ (\(k, b) -> assertBounds (label ++ " " ++ show k) b (actual Map.! k)) expected++-- | Assert the exact key list and probability of a distribution map.+assertDist :: (Ord k, Show k) => String -> [(k, Double)] -> Map k Double -> Assertion+assertDist label expected actual = do+ Map.keys actual @?= map fst expected+ mapM_ (\(k, p) -> assertApprox (label ++ " " ++ show k) p (actual Map.! k)) expected++-- | Assert that the second pair of bounds contains the first.+assertContains :: String -> (Double, Double) -> (Double, Double) -> Assertion+assertContains label (ilo, ihi) (olo, ohi) = do+ olo <= ilo + 1e-9 @?+ (label ++ ": outer lower " ++ show olo ++ " > inner lower " ++ show ilo)+ ohi >= ihi - 1e-9 @?+ (label ++ ": outer upper " ++ show ohi ++ " < inner upper " ++ show ihi)++-- | Assert that forcing the value throws the empty-credal-set error.+assertNoFeasible :: Show a => String -> a -> Assertion+assertNoFeasible label x = do+ result <- try (evaluate (length (show x)))+ case result of+ Left e -> "No feasible probabilities" `isPrefixOf` show (e :: ErrorCall) @?+ (label ++ ": unexpected error: " ++ show e)+ Right _ -> assertFailure (label ++ ": expected an error, got " ++ show x)
+ viz/Main.hs view
@@ -0,0 +1,286 @@+module Main where++import Data.List (nub)+import qualified Data.IntMap.Strict as IntMap+import qualified Data.Map.Strict as Map++import Imp.BDD (VarLabel(..))+import Imp.BDD.Compile (compile)+import Imp.BDD.WMC (Weight(..))+import Imp.DSL (Imp)+import Imp.Inference (credalVertices, intervalExpectation, intervalProbability)+import Imp.Examples.Ellsberg as E+import Imp.Examples.IMDP as IMDP+import Imp.Examples.Iteration as Iter+import Imp.Examples.Knightian as K+import Imp.Examples.Polytope as P+import Viz++main :: IO ()+main = do+ writeFile "index.html" page+ putStrLn "Wrote index.html — open in a browser."++page :: String+page = htmlPage overviewSimplex figureEntries+ [iterSection, ellsbergSection, imdpSection]++-- | The Figure 1 examples overlaid on one simplex.+overviewSimplex :: String+overviewSimplex = simplexSVG ("R", "G", "B")+ [ CredalLayer "polytope (interval)" "#2E7D32" (bary3 (P.Red, P.Green, P.Blue) P.polytope)+ , CredalLayer "independent (quadrilateral)" "#1565C0" (bary3 (K.Red, K.Green, K.Blue) K.independent)+ , CredalLayer "dependent (line segment)" "#C62828" (bary3 (K.Red, K.Green, K.Blue) K.dependent)+ ]++figureEntries :: [VizEntry]+figureEntries =+ [ VizEntry "dependent — 1 Knightian choice (line segment)"+ srcDependent (mkBddSVG K.dependent)+ , VizEntry "independent — 2 Knightian choices (quadrilateral)"+ srcIndependent (mkBddSVG K.independent)+ , VizEntry "polytope — composed intervals"+ (srcPolytope ++ expectNote) (mkBddSVG P.polytope)+ ]++-- | Lower/upper expectation shown under the polytope snippet.+expectNote :: String+expectNote =+ let (lo, hi) = intervalExpectation P.polytope+ (\v -> case v of P.Red -> 1.0; P.Green -> 0.5; P.Blue -> 0.0)+ in "\n-- E[f] for f(R)=1, f(G)=0.5, f(B)=0:\n-- E[f] \x2208 [" ++ showRound3 lo ++ ", " ++ showRound3 hi ++ "]"++iterSection :: String+iterSection =+ let layer :: String -> String -> Imp g Int -> CredalLayer+ layer name color walk =+ let verts = credalVertices walk+ label = name ++ ": " ++ show (length verts) ++ "\x2192"+ ++ show (nDistinct verts) ++ " distinct"+ in CredalLayer label color (map (toBary3 (0, 1, 2) . capAt2) verts)+ labels = ("0 (R)", "1 (G)", "\x2265 2 (B)")+ simplexSym = simplexSVG labels+ [ layer "walk3" "#6A1B9A" Iter.walk3+ , layer "walk2" "#1565C0" Iter.walk2+ , layer "walk1" "#C62828" Iter.walk1+ ]+ simplexAsym = simplexSVG labels [layer "walk3Asym" "#E65100" Iter.walk3Asym]+ in unlines+ [ ""+ , "<h2>Iteration — building up Knightian choices</h2>"+ , "<p class=\"legend\">Random walk: at each step, move right with P ∈ [0.3, 0.7]."+ , " Positions: 0 → R, 1 → G, ≥2 → B."+ , " Each iteration adds a fresh Knightian name to the grade."+ , " <code>intervalMap</code> is the graded <code>mapM</code>: it traverses"+ , " a type-level list of names, giving each an independent interval choice.</p>"+ , "<div class=\"row\">"+ , "<div class=\"card\">"+ , " <h3>Symmetric intervals — vertex collapse</h3>"+ , " <p class=\"legend\">With identical intervals, swapping valuations (hi,lo) ↔ (lo,hi)"+ , " gives the same distribution. 2<sup>n</sup> valuations → (n+1) distinct vertices.</p>"+ , simplexSym+ , " " ++ srcBlock srcWalk+ , "</div>"+ , "<div class=\"card\">"+ , " <h3>Asymmetric intervals — no collapse</h3>"+ , " <p class=\"legend\">Different interval widths break the symmetry."+ , " All 8 valuations give distinct distributions."+ , " Dots inside the shaded polygon are interior to the convex hull.</p>"+ , simplexAsym+ , " " ++ srcBlock srcWalkAsym+ , "</div>"+ , "</div>"+ ]++ellsbergSection :: String+ellsbergSection =+ let simplex = simplexSVG ("Red", "Black", "Yellow")+ [CredalLayer "ellsberg (line segment)" "#1565C0" (bary3 (E.Red, E.Black, E.Yellow) ellsberg)]+ gambleRow (name, bet, event) =+ let (lo, hi) = intervalProbability ellsberg event+ precise = if hi - lo < 1e-9 then "<strong>Yes</strong>" else "No"+ in " <tr><td>" ++ name ++ "</td><td>" ++ bet ++ "</td><td>["+ ++ showRound3 lo ++ ", " ++ showRound3 hi ++ "]</td><td>" ++ precise ++ "</td></tr>"+ gambles =+ [ ("I", "Red", (== E.Red))+ , ("II", "Black", (== E.Black))+ , ("III", "Red ∨ Yellow", \b -> b == E.Red || b == E.Yellow)+ , ("IV", "Black ∨ Yellow", \b -> b == E.Black || b == E.Yellow)+ ]+ in unlines $+ [ ""+ , "<h2>Textbook: Ellsberg paradox (1961)</h2>"+ , "<p class=\"legend\">Urn with 90 balls: 30 Red (known), 60 Black or Yellow (unknown split)."+ , " P(Red) = ⅓ is precise; P(Black) + P(Yellow) = ⅔ is precise;"+ , " but individually P(Black) ∈ [0, ⅔] and P(Yellow) ∈ [0, ⅔].</p>"+ , "<div class=\"row\">"+ , "<div class=\"card\">"+ , " <h3>Credal set on the simplex</h3>"+ , simplex+ , " " ++ srcBlock srcEllsberg+ , "</div>"+ , "<div class=\"card\">"+ , " <h3>Gamble analysis</h3>"+ , " <p class=\"legend\">Lower/upper expectations of the four Ellsberg gambles."+ , " People prefer gambles with <em>precise</em> expectations (ambiguity aversion).</p>"+ , " <table class=\"api\">"+ , " <tr><th>Gamble</th><th>Bet on</th><th>E[win]</th><th>Precise?</th></tr>"+ ] ++ map gambleRow gambles +++ [ " </table>"+ , " <p class=\"legend\">Preferring I over II <em>and</em> IV over III is inconsistent"+ , " with any single prior, but consistent with the credal set.</p>"+ , "</div>"+ , "</div>"+ ]++imdpSection :: String+imdpSection =+ let layer :: String -> String -> Imp g Position -> CredalLayer+ layer name color robot = CredalLayer+ (name ++ ": " ++ show (nDistinct (credalVertices robot)) ++ " distinct")+ color (bary3 (P0, P1, P2) robot)+ simplex = simplexSVG ("Start (0)", "Mid (1)", "Goal (2)")+ [ layer "simpleRobot3" "#2E7D32" IMDP.simpleRobot3+ , layer "simpleRobot" "#C62828" IMDP.simpleRobot+ ]+ reachRow :: Int -> Imp g Position -> String -> String+ reachRow n robot choices =+ let (lo, hi) = intervalProbability robot (== P2)+ in " <tr><td>" ++ show n ++ "</td><td>[" ++ showRound3 lo ++ ", "+ ++ showRound3 hi ++ "]</td><td>" ++ choices ++ "</td></tr>"+ in unlines+ [ ""+ , "<h2>Interval MDP — robot navigation</h2>"+ , "<p class=\"legend\">Robot at position 0 on {0, 1, 2}. At each step, moves right"+ , " with P ∈ [0.6, 0.9]. Position 2 is absorbing (goal).</p>"+ , "<div class=\"row\">"+ , "<div class=\"card\">"+ , " <h3>Credal sets (2 and 3 steps)</h3>"+ , simplex+ , " " ++ srcBlock srcRobot+ , "</div>"+ , "<div class=\"card\">"+ , " <h3>Reachability analysis</h3>"+ , " <p class=\"legend\">Lower/upper probability of reaching the goal (position 2).</p>"+ , " <table class=\"api\">"+ , " <tr><th>Steps</th><th>P(reach goal)</th><th>Knightian choices</th></tr>"+ , reachRow 2 IMDP.simpleRobot "move1, move2"+ , reachRow 3 IMDP.simpleRobot3 "move1, move2, move3"+ , " </table>"+ , " <p class=\"legend\">Symmetric intervals cause vertex collapse:"+ , " with identical step intervals, (hi,lo) and (lo,hi) give the same distribution.</p>"+ , "</div>"+ , "</div>"+ ]++-- | Compile to one event BDD per value, render as SVG.+mkBddSVG :: (Ord a, Show a) => Imp g a -> String+mkBddSVG prog =+ let (mgr, weights, knights, events) = compile prog+ knightName = IntMap.fromList [ (i, name) | (name, i) <- Map.toList knights ]+ names = Map.fromList $+ [ (VarLabel k, knightName IntMap.! i) | (k, Knight i) <- IntMap.toList weights ] +++ [ (VarLabel k, "flip") | (k, Prob _) <- IntMap.toList weights ]+ in bddSVG mgr names [ (show v, bdd) | (v, bdd) <- Map.toList events ]++-- | A program's credal-set vertices in barycentric coordinates.+bary3 :: Ord a => (a, a, a) -> Imp g a -> [(Double, Double, Double)]+bary3 keys prog = map (toBary3 keys) (credalVertices prog)++-- | Project a distribution onto barycentric coordinates for the+-- three given outcomes (missing outcomes get probability 0).+toBary3 :: Ord a => (a, a, a) -> Map.Map a Double -> (Double, Double, Double)+toBary3 (k1, k2, k3) dist =+ ( Map.findWithDefault 0 k1 dist+ , Map.findWithDefault 0 k2 dist+ , Map.findWithDefault 0 k3 dist )++-- | Merge all walk positions >= 2 so the distribution fits a 2-simplex.+capAt2 :: Map.Map Int Double -> Map.Map Int Double+capAt2 = Map.mapKeysWith (+) (min 2)++-- | Round a distribution's probabilities to 3 decimal places for comparison.+roundDist :: Map.Map a Double -> Map.Map a Double+roundDist = fmap (\v -> fromIntegral (round (v * 1000) :: Int) / 1000)++-- | Number of distinct distributions, up to rounding.+nDistinct :: Ord a => [Map.Map a Double] -> Int+nDistinct = length . nub . map roundDist++showRound3 :: Double -> String+showRound3 x = show (fromIntegral (round (x * 1000) :: Int) / 1000 :: Double)++-- Source snippets shown on the page, verbatim from Imp.Examples.*.++srcDependent :: String+srcDependent = unlines+ [ "dependent :: Imp '[\"a1\"] Three"+ , "dependent = Imp.do"+ , " x <- flip 0.5"+ , " if x then Imp.do"+ , " y <- knight @\"a1\""+ , " Imp.return (if y then Red else Green)"+ , " else Imp.do"+ , " y <- knight @\"a1\""+ , " Imp.return (if y then Red else Blue)"+ ]++srcIndependent :: String+srcIndependent = unlines+ [ "independent :: Imp '[\"a1\", \"a2\"] Three"+ , "independent = Imp.do"+ , " x <- flip 0.5"+ , " if x then Imp.do"+ , " y <- knight @\"a1\""+ , " Imp.return (if y then Red else Green)"+ , " else Imp.do"+ , " y <- knight @\"a2\""+ , " Imp.return (if y then Red else Blue)"+ ]++srcPolytope :: String+srcPolytope = unlines+ [ "polytope :: Imp '[\"p\", \"q\"] Three"+ , "polytope = Imp.do"+ , " p <- interval @\"p\" 0.2 0.8"+ , " q <- interval @\"q\" 0.3 0.7"+ , " Imp.return $ if p then Red else (if q then Green else Blue)"+ ]++srcWalk :: String+srcWalk = unlines+ [ "walk3 :: Imp '[\"s1\", \"s2\", \"s3\"] Int"+ , "walk3 = Imp.do"+ , " steps <- intervalMap @'[\"s1\", \"s2\", \"s3\"] 0.3 0.7"+ , " Imp.return $ length (filter id steps)"+ ]++srcWalkAsym :: String+srcWalkAsym = unlines+ [ "walk3Asym :: Imp '[\"s1\", \"s2\", \"s3\"] Int"+ , "walk3Asym = Imp.do"+ , " s1 <- interval @\"s1\" 0.1 0.9"+ , " s2 <- interval @\"s2\" 0.3 0.7"+ , " s3 <- interval @\"s3\" 0.45 0.55"+ , " Imp.return $ length (filter id [s1, s2, s3])"+ ]++srcEllsberg :: String+srcEllsberg = unlines+ [ "ellsberg :: Imp '[\"split\"] Ball"+ , "ellsberg = Imp.do"+ , " isRed <- flip (1/3)"+ , " isBlack <- interval @\"split\" 0.0 1.0"+ , " Imp.return $ if isRed then Red"+ , " else (if isBlack then Black else Yellow)"+ ]++srcRobot :: String+srcRobot = unlines+ [ "simpleRobot :: Imp '[\"move1\", \"move2\"] Position"+ , "simpleRobot = Imp.do"+ , " move1 <- interval @\"move1\" 0.6 0.9"+ , " let pos1 = step P0 move1"+ , " move2 <- interval @\"move2\" 0.6 0.9"+ , " Imp.return (step pos1 move2)"+ ]
+ viz/Viz.hs view
@@ -0,0 +1,369 @@+module Viz+ ( -- * Simplex visualization+ CredalLayer(..)+ , simplexSVG+ -- * BDD visualization+ , bddSVG+ -- * HTML page+ , VizEntry(..)+ , htmlPage+ , srcBlock+ ) where++import Data.List (sort)+import qualified Data.IntMap.Strict as IntMap+import qualified Data.Map.Strict as Map+import qualified Data.Set as Set++import Imp.BDD (BDD(..), BDDNode(..), VarLabel(..), NodeId(..))+import Imp.BDD.Builder (BDDManager, nodeTable)++import Viz.ConvexHull++lookupNodeTable :: NodeId -> BDDManager -> Maybe BDDNode+lookupNodeTable (NodeId n) mgr = IntMap.lookup n (nodeTable mgr)++-- | A layer on the simplex: legend name, colour, and barycentric points.+data CredalLayer = CredalLayer !String !String ![(Double, Double, Double)]++-- | Generate SVG of credal sets on the probability 2-simplex.+-- The three labels name the vertices (bottom-left, bottom-right, top).+simplexSVG :: (String, String, String) -> [CredalLayer] -> String+simplexSVG (label1, label2, label3) layers = unlines $+ [ "<svg xmlns=\"http://www.w3.org/2000/svg\" viewBox=\"0 0 450 420\">"+ , " <defs><style>"+ , " .vtx { font: bold 16px sans-serif; }"+ , " .leg { font: 13px sans-serif; }"+ , " </style></defs>"+ , " <rect width=\"100%\" height=\"100%\" fill=\"white\"/>"+ -- simplex outline+ , " <polygon points=\"" ++ triPts ++ "\""+ , " fill=\"#f5f5f5\" stroke=\"#bbb\" stroke-width=\"1\"/>"+ ] +++ -- grid lines at 0.25, 0.5, 0.75+ gridLines +++ -- credal-set layers (first in list drawn first = behind)+ concatMap drawLayer layers +++ -- vertex labels+ [ " <text x=\"" ++ showCoord (fst v1 - 5) ++ "\" y=\"" ++ showCoord (snd v1 + 22)+ ++ "\" text-anchor=\"middle\" class=\"vtx\">" ++ label1 ++ "</text>"+ , " <text x=\"" ++ showCoord (fst v2 + 5) ++ "\" y=\"" ++ showCoord (snd v2 + 22)+ ++ "\" text-anchor=\"middle\" class=\"vtx\">" ++ label2 ++ "</text>"+ , " <text x=\"" ++ showCoord (fst v3) ++ "\" y=\"" ++ showCoord (snd v3 - 10)+ ++ "\" text-anchor=\"middle\" class=\"vtx\">" ++ label3 ++ "</text>"+ ] +++ -- legend+ [ " <g transform=\"translate(15,15)\">" ] +++ [ " <rect x=\"0\" y=\"" ++ show (i*24) ++ "\" width=\"16\" height=\"16\""+ ++ " fill=\"" ++ color ++ "\" fill-opacity=\"0.4\""+ ++ " stroke=\"" ++ color ++ "\" stroke-width=\"1.5\"/>"+ ++ "<text x=\"22\" y=\"" ++ show (i*24+13) ++ "\" class=\"leg\">"+ ++ name ++ "</text>"+ | (i, CredalLayer name color _) <- zip [0::Int ..] layers ] +++ [ " </g>"+ , "</svg>"+ ]+ where+ side = 320.0+ cx = 225.0+ base = 380.0+ top' = base - side * sqrt 3 / 2++ v1 = (cx - side / 2, base) -- bottom-left+ v2 = (cx + side / 2, base) -- bottom-right+ v3 = (cx, top') -- top++ triPts = showCoord (fst v1) ++ "," ++ showCoord (snd v1) ++ " "+ ++ showCoord (fst v2) ++ "," ++ showCoord (snd v2) ++ " "+ ++ showCoord (fst v3) ++ "," ++ showCoord (snd v3)++ bary (p1, p2, p3) =+ ( p1 * fst v1 + p2 * fst v2 + p3 * fst v3+ , p1 * snd v1 + p2 * snd v2 + p3 * snd v3 )++ gridLines = concatMap (\t ->+ [ gridLine t v1 v2 v3+ , gridLine t v2 v3 v1+ , gridLine t v3 v1 v2+ ]) [0.25, 0.5, 0.75]++ gridLine t (ax,ay) (bx,by) (ex,ey) =+ let x1 = t*ax + (1-t)*bx; y1 = t*ay + (1-t)*by+ x2 = t*ax + (1-t)*ex; y2 = t*ay + (1-t)*ey+ in " <line x1=\"" ++ showCoord x1 ++ "\" y1=\"" ++ showCoord y1+ ++ "\" x2=\"" ++ showCoord x2 ++ "\" y2=\"" ++ showCoord y2+ ++ "\" stroke=\"#ddd\" stroke-width=\"0.5\"/>"++ drawLayer (CredalLayer _ color pts) =+ let xys = map bary pts+ outline = convexHull2D xys+ polyPts = unwords [showCoord x ++ "," ++ showCoord y | (x,y) <- outline]+ in [ " <polygon points=\"" ++ polyPts ++ "\""+ ++ " fill=\"" ++ color ++ "\" fill-opacity=\"0.25\""+ ++ " stroke=\"" ++ color ++ "\" stroke-width=\"2.5\"/>" ]+ ++ [ " <circle cx=\"" ++ showCoord x ++ "\" cy=\"" ++ showCoord y+ ++ "\" r=\"4\" fill=\"" ++ color ++ "\"/>"+ | (x,y) <- xys ]++-- | Generate an inline SVG of a BDD graph, laid out top-to-bottom.+bddSVG :: BDDManager+ -> Map.Map VarLabel String -- ^ human-readable variable names+ -> [(String, BDD)] -- ^ named root BDDs+ -> String+bddSVG mgr varNames roots =+ let -- Collect reachable nodes+ allNids = collectNids mgr (map snd roots)+ nodeList = [ (nid, node) | nid <- allNids+ , Just node <- [lookupNodeTable nid mgr] ]++ -- Group by variable, sorted ascending+ byVar = Map.fromListWith (flip (++))+ [(bddVar node, [nid]) | (nid, node) <- nodeList]+ layers = Map.toAscList byVar -- [(VarLabel, [NodeId])]+ nLayers = length layers++ -- Layout constants+ nodeR = 18.0+ layerH = 80.0+ colW = 90.0+ topPad = 45.0+ rootH = 30.0++ -- Determine SVG width from widest row+ maxCols = maximum $ [length nids | (_, nids) <- layers]+ ++ [length roots, 2]+ svgW = max 260 (fromIntegral maxCols * colW + 60)++ -- Center a row of items horizontally, return list of (item, x)+ centerRow :: [a] -> [(a, Double)]+ centerRow [] = []+ centerRow [a] = [(a, svgW / 2)]+ centerRow items =+ let n = length items+ tw = fromIntegral (n - 1) * colW+ x0 = (svgW - tw) / 2+ in [(item, x0 + fromIntegral i * colW) | (i, item) <- zip [0 :: Int ..] items]++ -- Root label positions+ rootRow = centerRow roots++ -- Node positions (by layer)+ nodePos :: Map.Map NodeId (Double, Double)+ nodePos = Map.fromList $ concat+ [ [(nid, (x, y))+ | (nid, x) <- centerRow (sort nids)]+ | (li, (_, nids)) <- zip [0::Int ..] layers+ , let y = topPad + rootH + fromIntegral li * layerH+ ]++ -- Terminal row+ termY = topPad + rootH + fromIntegral nLayers * layerH+ trueX = svgW / 2 - colW / 2+ falseX = svgW / 2 + colW / 2+ svgH = termY + 55++ -- Resolve a BDD to its (x,y) position+ posOf :: BDD -> (Double, Double)+ posOf BDDTrue = (trueX, termY)+ posOf BDDFalse = (falseX, termY)+ posOf (BDDRef nid) = Map.findWithDefault (svgW/2, 0) nid nodePos+ posOf (BDDComp nid) = Map.findWithDefault (svgW/2, 0) nid nodePos++ isComp (BDDComp _) = True+ isComp _ = False++ mkLine x1 y1 x2 y2 dashed comp =+ " <line x1=\"" ++ showCoord x1 ++ "\" y1=\"" ++ showCoord y1+ ++ "\" x2=\"" ++ showCoord x2 ++ "\" y2=\"" ++ showCoord y2+ ++ "\" stroke=\"" ++ (if comp then "#d32f2f" else if dashed then "#999" else "#444")+ ++ "\" stroke-width=\"" ++ (if dashed then "1" else "1.5")+ ++ "\"" ++ (if dashed then " stroke-dasharray=\"5,3\"" else "") ++ "/>"++ mkCompDot x y =+ "<circle cx=\"" ++ showCoord x ++ "\" cy=\"" ++ showCoord y+ ++ "\" r=\"4\" fill=\"white\" stroke=\"#d32f2f\" stroke-width=\"1.5\"/>"++ mkCircle x y label =+ " <circle cx=\"" ++ showCoord x ++ "\" cy=\"" ++ showCoord y ++ "\" r=\"" ++ showCoord nodeR+ ++ "\" fill=\"white\" stroke=\"#333\" stroke-width=\"1.5\"/>"+ ++ "<text x=\"" ++ showCoord x ++ "\" y=\"" ++ showCoord (y + 5)+ ++ "\" text-anchor=\"middle\" font-family=\"sans-serif\" font-size=\"12\">"+ ++ label ++ "</text>"++ mkSquare x y label color =+ " <rect x=\"" ++ showCoord (x - 14) ++ "\" y=\"" ++ showCoord (y - 14)+ ++ "\" width=\"28\" height=\"28\" rx=\"3\" fill=\""+ ++ color ++ "\" stroke=\"#333\" stroke-width=\"1\"/>"+ ++ "<text x=\"" ++ showCoord x ++ "\" y=\"" ++ showCoord (y + 5)+ ++ "\" text-anchor=\"middle\" font-family=\"sans-serif\""+ ++ " font-size=\"13\" font-weight=\"bold\">"+ ++ label ++ "</text>"++ -- Draw an edge from a source point to a child BDD: the connecting+ -- line plus, when the child is a complemented reference, the little+ -- red complemented-edge dot at the child end.+ edgeFrom x1 y1 child dashed =+ let (x2, y2) = posOf child+ r2 = if child == BDDTrue || child == BDDFalse then 14 else nodeR+ in mkLine x1 y1 x2 (y2 - r2) dashed (isComp child)+ ++ if isComp child then mkCompDot x2 (y2 - r2) else ""++ header = "<svg xmlns=\"http://www.w3.org/2000/svg\""+ ++ " viewBox=\"0 0 " ++ showCoord svgW ++ " " ++ showCoord svgH ++ "\""+ ++ " width=\"" ++ showCoord (min svgW 400) ++ "\">"+ ++ "<rect width=\"100%\" height=\"100%\" fill=\"white\"/>"+ edgeElems = concat+ [ case lookupNodeTable nid mgr of+ Just (BDDNode _ lo hi) ->+ let (x, y) = Map.findWithDefault (0,0) nid nodePos+ in [ edgeFrom x (y + nodeR) lo True -- lo = dashed (0-branch)+ , edgeFrom x (y + nodeR) hi False -- hi = solid (1-branch)+ ]+ Nothing -> []+ | nid <- allNids+ ]+ rootEdges =+ [ edgeFrom rx (rootH + 8) bdd False+ | ((_, bdd), (_, rx)) <- zip roots rootRow+ ]+ nodeElems =+ [ mkCircle x y (vName varNames (bddVar node))+ | (nid, node) <- nodeList+ , let (x, y) = Map.findWithDefault (0,0) nid nodePos+ ]+ termElems =+ [ mkSquare trueX termY "T" "#c8e6c9"+ , mkSquare falseX termY "F" "#ffcdd2"+ ]+ rootLabels =+ [ " <text x=\"" ++ showCoord rx ++ "\" y=\"" ++ showCoord rootH+ ++ "\" text-anchor=\"middle\" font-family=\"sans-serif\""+ ++ " font-size=\"14\" font-weight=\"bold\">" ++ name ++ "</text>"+ | ((name, _), (_, rx)) <- zip roots rootRow+ ]+ edgeAnnotations = concat+ [ case lookupNodeTable nid mgr of+ Just (BDDNode _ lo hi) ->+ let (x, y) = Map.findWithDefault (0,0) nid nodePos+ (lx, _) = posOf lo+ (hx, _) = posOf hi+ labelY = y + nodeR + 12+ loLabelX = x + (lx - x) * 0.3 - 8+ hiLabelX = x + (hx - x) * 0.3 + 8+ mkLbl lx' ly lbl clr =+ "<text x=\"" ++ showCoord lx' ++ "\" y=\"" ++ showCoord ly+ ++ "\" font-family=\"sans-serif\" font-size=\"10\""+ ++ " fill=\"" ++ clr ++ "\">" ++ lbl ++ "</text>"+ in [ mkLbl loLabelX labelY "0" "#999"+ , mkLbl hiLabelX labelY "1" "#444"+ ]+ Nothing -> []+ | nid <- allNids+ ]++ in unlines $+ [header]+ ++ edgeElems ++ rootEdges+ ++ nodeElems ++ termElems+ ++ rootLabels ++ edgeAnnotations+ ++ ["</svg>"]++-- | An entry for the HTML page: title, source code, and BDD SVG.+data VizEntry = VizEntry+ { veTitle :: !String+ , veSource :: !String -- ^ Haskell source snippet+ , veBddSvg :: !String -- ^ inline SVG+ }++-- | Generate a self-contained HTML page with embedded SVG visualizations.+htmlPage :: String -- ^ simplex SVG (inline)+ -> [VizEntry] -- ^ example entries+ -> [String] -- ^ extra HTML sections appended after examples+ -> String+htmlPage simplex entries extras = unlines $+ [ "<!DOCTYPE html>"+ , "<html lang=\"en\"><head><meta charset=\"utf-8\">"+ , "<title>imp — credal set visualization</title>"+ , "<style>"+ , " body { font-family: system-ui, sans-serif; max-width: 960px;"+ , " margin: 0 auto; padding: 24px; background: #fafafa; color: #222; }"+ , " h1 { font-size: 1.6rem; }"+ , " h2 { font-size: 1.2rem; margin-top: 2rem; color: #555; }"+ , " .row { display: flex; gap: 24px; flex-wrap: wrap; justify-content: center;"+ , " align-items: flex-start; }"+ , " .card { background: white; border: 1px solid #e0e0e0; border-radius: 8px;"+ , " padding: 16px; text-align: center; }"+ , " .card h3 { margin: 0 0 8px; font-size: 1rem; color: #333; }"+ , " svg { max-width: 100%; height: auto; }"+ , " .legend { font-size: 0.85rem; color: #666; margin-top: 12px; }"+ , " pre.src { background: #1e1e2e; color: #cdd6f4; padding: 12px 16px;"+ , " border-radius: 6px; text-align: left; font-size: 0.82rem;"+ , " line-height: 1.45; overflow-x: auto; margin: 10px 0 0; }"+ , " table.api { border-collapse: collapse; width: 100%; margin: 12px 0; font-size: 0.9rem; }"+ , " table.api th, table.api td { border: 1px solid #ddd; padding: 6px 10px; text-align: left; }"+ , " table.api th { background: #f0f0f0; }"+ , " table.api code { background: #eee; padding: 1px 4px; border-radius: 3px;"+ , " font-size: 0.85em; }"+ , "</style>"+ , "</head><body>"+ , "<h1>imp — imprecise probabilistic programming</h1>"+ , "<p>Credal sets and BDD compilation for examples from"+ , " Liell-Cock & Staton (POPL 2025), Figure 1.</p>"+ , ""+ , "<h2>Credal sets on the probability simplex</h2>"+ , "<div class=\"row\"><div class=\"card\">"+ , simplex+ , "<div class=\"legend\">Dashed grid at 0.25 / 0.5 / 0.75 probability levels.</div>"+ , "</div></div>"+ , ""+ , "<h2>Source & BDD compilation</h2>"+ , "<p class=\"legend\">Solid lines = hi (1) branch Dashed = lo (0) branch"+ , " <span style=\"color:#d32f2f\">Red ○</span> = complemented edge</p>"+ , "<div class=\"row\">"+ ] +++ concatMap (\(VizEntry title src svg) ->+ [ "<div class=\"card\">"+ , " <h3>" ++ title ++ "</h3>"+ , " " ++ srcBlock src+ , svg+ , "</div>"+ ]) entries +++ [ "</div>" ] +++ extras +++ [ "</body></html>" ]++-- | A source snippet as a styled, HTML-escaped @<pre>@ block.+srcBlock :: String -> String+srcBlock src = "<pre class=\"src\">" ++ escapeHtml src ++ "</pre>"++escapeHtml :: String -> String+escapeHtml = concatMap esc+ where+ esc '&' = "&"+ esc '<' = "<"+ esc '>' = ">"+ esc c = [c]++-- | Render an SVG coordinate: rounded to two decimals.+showCoord :: Double -> String+showCoord x = show (fromIntegral (round (x * 100) :: Int) / 100 :: Double)++vName :: Map.Map VarLabel String -> VarLabel -> String+vName varNames vl = case Map.lookup vl varNames of+ Just n -> n+ Nothing -> "v" ++ show (unVarLabel vl)++-- | All internal node ids reachable from the given roots, each visited once.+collectNids :: BDDManager -> [BDD] -> [NodeId]+collectNids mgr = Set.toList . foldl' go Set.empty+ where+ go seen BDDTrue = seen+ go seen BDDFalse = seen+ go seen (BDDRef nid) = visit seen nid+ go seen (BDDComp nid) = visit seen nid++ visit seen nid+ | nid `Set.member` seen = seen+ | otherwise = case lookupNodeTable nid mgr of+ Nothing -> Set.insert nid seen+ Just (BDDNode _ lo hi) -> go (go (Set.insert nid seen) lo) hi
+ viz/Viz/ConvexHull.hs view
@@ -0,0 +1,36 @@+-- | Convex hull algorithm for credal set display.+module Viz.ConvexHull+ ( convexHull2D+ ) where++import Data.List (sort, nubBy)++-- | Compute the convex hull of a set of 2D points.+-- Returns the hull vertices in counterclockwise order.+-- Uses Andrew's monotone chain algorithm: O(n log n).+convexHull2D :: [(Double, Double)] -> [(Double, Double)]+convexHull2D pts+ | length unique <= 2 = unique+ | otherwise =+ let sorted = sort unique+ lower = reverse (buildHull sorted)+ upper = reverse (buildHull (reverse sorted))+ in init lower ++ init upper+ where+ unique = dedup2D pts+ buildHull = foldl' add []+ add stack p = p : dropRightTurns stack+ where+ dropRightTurns (b : a : rest)+ | cross2D a b p <= 0 = dropRightTurns (a : rest)+ dropRightTurns s = s++-- | Cross product of vectors (b - a) and (c - a).+cross2D :: (Double, Double) -> (Double, Double) -> (Double, Double) -> Double+cross2D (ax, ay) (bx, by) (cx, cy) =+ (bx - ax) * (cy - ay) - (by - ay) * (cx - ax)++-- | Remove near-duplicate 2D points.+dedup2D :: [(Double, Double)] -> [(Double, Double)]+dedup2D = nubBy (\(x1, y1) (x2, y2) -> abs (x1 - x2) < epsSamePoint && abs (y1 - y2) < epsSamePoint)+ where epsSamePoint = 1e-10