imp-ppl-0.1.0.0: src/Imp/BDD/WMC.hs
-- | 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