packages feed

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