0 | module Control.Monad.Sample.Instances
2 | import Control.Monad.Identity
6 | import Control.Monad.Distribution
7 | import Control.Monad.Sample.Definition
11 | [pickFirst] MonadSample Identity where
12 | sample {i = (S k)} (MkDist xs) = Id FZ
16 | [pickMax] MonadSample Identity where
17 | sample {i = (S k)} d = Id (argmax d.logits)
21 | [pickMin] MonadSample Identity where
22 | sample {i = (S k)} d = Id (argmin d.logits)
27 | MonadSample IO where
28 | sample {i = S j} (MkDist xs) = do
29 | let dist = softargmaxImpl xs
30 | cumSum = cumulativeSum dist
31 | r <- randomRIO (0.0, 1.0)
32 | case findBin (#> cumSum) r of
38 | let logits : Dist "coin" 2
39 | logits = MkDist (># [-(1.099), 1.099])
40 | is <- sequence (replicate 1000 (sample logits))
42 | printLn (count (== 0) is)
43 | printLn (count (== 1) is)
49 | let logits = diracDelta {name="dirac"} {i=10} index
50 | inds <- sequence (replicate 1000 (sample logits))
51 | printLn (take 10 inds)
52 | printLn (count (== index) inds)