如何用递归模式在Haskell中基于monad-bayes表示概率分布
sUnrolled in monad-bayes for Unrolling Conditional Distributions 嘿,我刚好在monad-bayes里折腾过类似的序列生成需求,来给你唠唠怎么实现这个sUnrolled——先补全下你没写完的数学定义(毕竟这是实现的核心):通常我们说的sUnrolled,应该是把单步的条件分布s展开成一个生成字符序列的分布:给定初始状态st :: [Char],生成长度为n的序列chars,其中每一步的下一个字符都由s根据当前的完整历史(初始状态+已生成的字符)采样得到。用数学语言写就是:
对于序列
c₁, c₂, ..., cₙ:
c₁ ~ s(st)c₂ ~ s(st ++ [c₁])- ...
cₖ ~ s(st ++ [c₁, ..., cₖ₋₁])
最终sUnrolled n st就是这个序列的概率分布。
下面是几种不同场景下的实现方案:
1. 基础固定长度实现
先从最直观的递归版本开始,利用monad的特性一步步采样并拼接序列:
import Control.Monad.Bayes.Class (MonadDist) sUnrolled :: MonadDist m => Int -> [Char] -> m [Char] sUnrolled 0 _ = return [] -- 长度为0时直接返回空序列 sUnrolled n st = do nextChar <- s st -- 用当前状态采样下一个字符 -- 递归生成剩下的n-1个字符,更新状态为原状态+新字符 restOfSequence <- sUnrolled (n-1) (st ++ [nextChar]) return (nextChar : restOfSequence) -- 拼接当前字符和剩余序列
这个版本逻辑直白,很容易理解,但有个小问题:每次用st ++ [nextChar]更新状态时,列表拼接是O(k)复杂度(k是当前列表长度),如果生成很长的序列,性能会打折扣。
2. 性能优化版(避免重复拼接)
我们可以用反向累积的方式优化,最后再把序列反转回来,把整体复杂度从O(n²)降到O(n):
sUnrolledOpt :: MonadDist m => Int -> [Char] -> m [Char] sUnrolledOpt n st = reverse <$> go n st [] where -- 辅助递归函数:k是剩余要生成的字符数,currentSt是当前状态,acc是累积的反向序列 go 0 _ acc = return acc go k currentSt acc = do nextChar <- s currentSt -- 把新字符加到累积列表头部(O(1)操作),同时更新状态 go (k-1) (currentSt ++ [nextChar]) (nextChar : acc)
这个版本和基础版逻辑完全一致,只是用了更高效的累积方式,长序列场景下体验会好很多。
3. 可变长度:生成直到满足终止条件
如果你的需求不是固定长度,而是要生成序列直到某个字符出现(比如生成到'x'就停止),可以改成带终止条件的递归:
sUnrolledUntil :: MonadDist m => (Char -> Bool) -> [Char] -> m [Char] sUnrolledUntil stopCondition st = do nextChar <- s st if stopCondition nextChar then return [nextChar] else (nextChar :) <$> sUnrolledUntil stopCondition (st ++ [nextChar])
比如你想生成直到出现'z'的序列,就可以调用sUnrolledUntil (== 'z') initialState。
测试验证
用monad-bayes的Sampler实例来跑个测试看看效果:
import Control.Monad.Bayes.Sampler (Sampler, sampleIO) -- 先写个示例的s:不管输入是什么,都从a/b/c里均匀采样 exampleS :: MonadDist m => [Char] -> m Char exampleS _ = uniform ['a', 'b', 'c'] -- 测试生成长度为5的序列 main :: IO () main = sampleIO (sUnrolled 5 "") >>= print
运行后会输出类似"bacab"的随机序列,完全符合预期。而且因为我们用的是MonadDist类型类,这个实现可以无缝切换到其他概率monad(比如Weighted做重要性采样,或者Trace做MCMC)。
内容的提问来源于stack exchange,提问作者Reuben

