You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何用递归模式在Haskell中基于monad-bayes表示概率分布

Implementing 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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.05.25 04:21:38