- 人工智能
- 机器学习
- 深度学习
- 概率编程
【免费下载链接】pyro
Deep universal probabilistic programming with Python and PyTorch
本篇技术指南聚焦 Pyro 仓库 examples/rsa 目录下的 Rational Speech Acts(RSA,理性言语行为)示例集,它们演示如何用概率程序实现“关于推理的推理”(reasoning about reasoning)——即一个智能体不仅对自己相信什么建模,还对其他智能体的信念与推理过程进行递归建模。读者将掌握基于Search/BestFirstSearch的精确与启发式枚举推理工具、HashingMarginal边际化分布的正确用法,以及从谢林协调博弈、虚假信念博弈到泛型语句、夸张语、组合语义语法解析等五类经典 RSA 模型的完整实现与运行方式。
目录与文件构成
本示例集位于 examples/rsa,包含 6 个 Python 脚本,均改编自 Noah Goodman 及其合作者的公开工作:
| 文件 | 内容 | 原始来源(见 examples/rsa/README.md) |
|---|---|---|
generics.py | 泛型语句(generic statements)的 RSA 语用模型 | Probabilistic Language Understanding 第 07 章 |
hyperbole.py | 夸张语(hyperbole)的 RSA 模型 | Probabilistic Language Understanding 第 03 章(非字面语言) |
schelling.py | 谢林协调博弈:两位间谍递归推理约定会面地点 | ForestDB 的 schelling 模型 |
schelling_false.py | 带虚假信念的谢林博弈:Alice 实际想避开 Bob | ForestDB 的 schelling-falsebelief 模型 |
search_inference.py | 全部示例共用的推理工具:Search、BestFirstSearch、HashingMarginal、memoize | Design and Implementation of Probabilistic Programming Languages(dippl)第 03 章枚举 |
semantic_parsing.py | 将 RSA 语用学与 CCG 组合语义语法结合的“语义-语用杂糅”模型 | dippl 的 zSemanticPragmaticMashup 示例 |
这些脚本同时被 tests/test_examples.py 收录为冒烟测试用例(rsa/generics.py --num-samples=10等),意味着它们不仅可用于学习,还可在仓库 CI 中作为可执行示例运行。
注意:所有脚本入口处均带有
assert pyro.__version__.startswith("1.9.1")版本断言,因此请使用匹配 1.9.x 系列的 Pyro 环境运行。
核心推理工具:search_inference.py
search_inference.py 是整个示例集的基石,提供了四类基础设施:
memoize:基于functools.lru_cache的记忆化装饰器,用于缓存Marginal计算结果(下文详解)。HashingMarginal:把TracePosterior对象转换成可采样、可求对数概率、可枚举支撑集的Distribution。Search:基于队列的精确枚举推理。BestFirstSearch:按概率优先的启发式枚举推理。
其中Search与BestFirstSearch都继承自 pyro/infer/abstract_infer.py 中的TracePosterior,只需实现_traces()方法逐条产出(trace, log_weight)即可复用 Pyro 的迹后验基础设施。
Search:精确枚举推理
Search的完整实现思路如下(见 examples/rsa/search_inference.py):
class Search(TracePosterior): """Exact inference by enumerating over all possible executions""" def __init__(self, model, max_tries=int(1e6), **kwargs): self.model = model self.max_tries = max_tries super().__init__(**kwargs) def _traces(self, *args, **kwargs): q = queue.Queue() q.put(poutine.Trace()) p = poutine.trace(poutine.queue(self.model, queue=q, max_tries=self.max_tries)) while not q.empty(): tr = p.get_trace(*args, **kwargs) yield tr, tr.log_prob_sum()其工作原理是:维护一个“部分迹”(partial trace)队列,配合 Pyro 的poutine.queue消息处理机制,每遇到一个未取值的离散sample站点就将其所有可能取值分支展开入队,直到得到完整执行迹;随后以tr.log_prob_sum()(该迹全部采样点对数概率之和)作为权重,穷举出所有可能的执行路径。max_tries控制尝试次数上限(默认1e6),防止组合爆炸时无限循环。
从源码结构看,poutine.queue的实现(见 pyro/poutine/handlers.py)印证了这一机制:它依次使用trace(escape(replay(wrapped, next_trace)))续跑部分迹,当escape_fn命中未执行的离散采样点时抛出NonlocalExit,再调用默认扩展函数util.enum_extend把该站点各取值分支的新迹压入队列,循环直到得到一条完整迹。这正对应了 dippl 教程中“枚举式概率程序执行”的经典算法。
BestFirstSearch:按概率优先的枚举
当状态空间较大、精确枚举不可行时,BestFirstSearch 改为用PriorityQueue按迹的log_prob_sum()排序,优先扩展高概率分支;还引入了一个微小的随机扰动- torch.rand(1).item() * 1e-2来打破并列优先级。num_samples默认取 100,即默认枚举前 100 条最高概率的执行迹;若队列提前耗尽则提前退出。注释中说明:当所有执行都被枚举完时,其结果与Search精确等价,否则是高概率近似。
HashingMarginal:把迹后验变成分布
HashingMarginal 继承dist.Distribution,将一个TracePosterior(蒙特卡洛或枚举后验)转换为可当普通分布使用的对象:
- 将每个迹的
_RETURN返回值(或显式指定的sites列表对应的站点值)作为“结果”; - 对结果做哈希去重(张量用
value.cpu().contiguous().numpy().tobytes()哈希,字典递归转成 key-value 元组再哈希); - 相同结果的多个迹权重用
logsumexp在对数域累加,最后归一化成一个Categorical; - 对外暴露
sample()、log_prob()、enumerate_support(),以及mean/variance属性(后者用加权平均实现)。
由于它真正实现了Distribution接口,示例模型可以把“某智能体的推理结果”当作一个分布,直接pyro.sample(..., obs=...)嵌套进更高层的推理,这正是 RSA 递归建模的关键技术。注释同时坦承“整个对象目前非常低效”,因此源码中普遍用@memoize(maxsize=10)缓存_dist_and_values()结果来摊薄重复计算。
RSA 建模骨架:Marginal 装饰器
每个示例文件顶部都定义了本地Marginal装饰器,把“用某个推理算法运行模型并求边际”这一操作包装为可记忆化函数。以 generics.py 为例:
def Marginal(fn): return memoize(lambda *args: HashingMarginal(Search(fn).run(*args)))含义是:Marginal(fn)(*args)用Search精确枚举运行fn,把返回的迹后验封装成HashingMarginal分布,并按参数记忆化缓存。由于Marginal是装饰器,源码中直接写作:
@Marginal def listener0(utterance, threshold, prior): ...即可得到“给定参数后的边际分布”函数。这样设计使得 RSA 各层智能体(字面听众 → 说话者 → 语用听众 → 更高层说话者)可以像普通函数一样互相调用、互相采样,形成递归嵌套的概率程序。
谢林协调博弈:递归推理的最小范例
schelling.py 演示了两个间谍 Alice 与 Bob 在无法通信的情况下、仅靠递归推理选择同一会面地点的博弈,是理解“推理的推理”最直观的入口。
模型结构(见 examples/rsa/schelling.py):
def location(preference): # 两人共享的先验偏好:抛一枚偏置硬币决定去哪个地点 return pyro.sample("loc", Bernoulli(preference)) def alice(preference, depth): # Alice 通过推理 Bob 的选择来决定去向 alice_prior = location(preference) with poutine.block(): bob_marginal = HashingMarginal(Search(bob).run(preference, depth - 1)) return pyro.sample("bob_choice", bob_marginal, obs=alice_prior) def bob(preference, depth): bob_prior = location(preference) if depth > 0: with poutine.block(): alice_marginal = HashingMarginal(Search(alice).run(preference, depth)) return pyro.sample("alice_choice", alice_marginal, obs=bob_prior) else: return bob_prior关键点有三个:
- 深度递归:
bob在depth > 0时推理 Alice,alice推理bob(depth-1),形成alice → bob → alice → …的递归链;depth控制推理层级。 poutine.block()屏蔽内层采样:当 Alice 把 Bob 的决策过程作为“子程序”调用时,block()保证内层Search枚举产生的采样点不会泄漏到外层迹中,避免命名冲突与错误嵌套。obs=条件化:pyro.sample("bob_choice", bob_marginal, obs=alice_prior)表示“Alice 相信 Bob 会选与我先验一致的地点”,即 Alice 的条件化推理。
运行方式(CLI 参数见 examples/rsa/schelling.py):
python examples/rsa/schelling.py --num-samples=10 --depth=2 --preference=0.6程序先打印 Bob 决策过程的边际概率分布,再对bob_decision蒙特卡洛采样num_samples次,估计 Bob 选择其偏好地点的经验频率。可以尝试把depth从 0 逐级调大,观察递归推理层数对协调概率的收敛影响。
虚假信念博弈:心智理论(Theory of Mind)
schelling_false.py 在协调博弈之上加入了“虚假信念”:表面上两位间谍都想会面,实际Alice 想要避开 Bob。它额外定义了alice_fb(见 examples/rsa/schelling_false.py),在推理出 Bob 的去向后故意选择相反地点:
def alice_fb(preference, depth): alice_prior = location(preference) with poutine.block(): bob_marginal = HashingMarginal(Search(bob).run(preference, depth - 1)) pyro.sample("bob_choice", bob_marginal, obs=alice_prior) return 1 - alice_prior # 反向选择注意这里alice_fb依然先采样"bob_choice"让 Bob 的决策过程被“观察到”,但在返回值上取反,从而建模“Alice 知晓 Bob 的推理并反其道而行”。而alice(普通版本)仍保留原逻辑,bob推理的是普通alice——这正是“虚假信念”的来源:Bob 以为 Alice 想会面,Alice 却实际在逃避。示例最终估计的是Alice 实际选择偏好地点的经验频率:
python examples/rsa/schelling_false.py --num-samples=10 --depth=3 --preference=0.55默认递归深度为 3,比基本谢林博弈深一层,因为虚假信念建模需要更长的推理链才能体现“Bob 的误解被 Alice 利用”。
泛型语句(Generics)的语用推理
generics.py 建模“泛型语句”的语义:例如“蚊子传播疟疾”“鸟会下蛋”这类不依赖全称量化的概括性表述。它构建了一个完整的 RSA 递归栈(模型定义见 examples/rsa/generics.py):
- 结构化先验
structured_prior_model:用Bernoulli(theta)决定属性是否存在,存在时再用离散化的 Beta 密度(discretize_beta_pdf,bins 取[0.01, 0.1, …, 0.99])枚举属性流行度。示例用 4 组(theta, gamma, delta)参数对应“有翅膀、下蛋、传播疟疾、是雌性”四种属性(见 generics.py),其中疟疾用theta=0.1, gamma=0.01, delta=2.0表示“罕见但一旦出现几乎必然传播”的属性。 - 真值函数
meaning(utterance, state, threshold):定义了"generic is true"(state > threshold)、"mu"(恒真)、"some"(state > 0)、"most"(state >= 0.5)、"all"(state >= 0.99)等话语的语义。 - 推理层级:
listener0(字面听众,用pyro.factor施加 −99999 的硬性真值约束)→speaker1(带s1Optimality=5.0的说话者最优性缩放)→listener1(语用听众,同时推理状态与阈值)→speaker2(在给定流行度下选择话语)。
其中“说话者最优性”通过poutine.scale实现(见 generics.py):
with poutine.scale(scale=torch.tensor(s1Optimality)): pyro.sample("L0_score", L0, obs=state)scale相当于把该采样点的对数概率乘以 5.0,数值上等价于 RSA 中说话者选择公式的 softmax 温度参数(optimality / rationality parameter α):α 越大,说话者越倾向于选择效用最高的那个话语。
脚本会对四种属性的听众解释输出支撑集概率,并对“传播疟疾、下蛋、是雌性、狮子(用下蛋属性的低流行度 0.01 模拟)”四个说话者场景输出话语概率,复现文献中“罕见且强烈属性更容易被泛型概括”的经典结论。运行:
python examples/rsa/generics.py --num-samples=10夸张语(Hyperbole)的语用推理
hyperbole.py 建模“非字面语言”——例如用精确数字表达夸张语义。模型要素包括:
- 状态空间:
State = namedtuple("State", ["price", "valence"]),价格取自 10 个离散值(50 到 10001,见 hyperbole.py),valence(正/负评价)由条件于价格的 Bernoulli 先验决定。 - 问答维度(Question Under Discussion, QUD):
qud_fns定义了price、valence、priceValence、approxPrice、approxPriceValence五种“讨论问题”,其中approxPrice会把价格舍入到 10 的整数倍(approx()函数)。 - 话语成本:
utterance_cost给“精确数字”额外加上preciseNumberCost = 1.0的成本(取负后作为 Categorical logits),建模“精确表达比近似表达更费力”。 - 推理层级:
literal_listener(字面听众,pyro.factor硬性约束价格匹配)→speaker(以alpha = 1.0的最优性参数选择话语)→pragmatic_listener(联合推理价格、valence、QUD 并条件化听到的话语,见 hyperbole.py)。
运行(见 hyperbole.py):
python examples/rsa/hyperbole.py --price=10000程序打印语用听众在听到--price所指话语后,对全部 20 个(price, valence)状态的后验概率。源码中保留的test_truth()(hyperbole.py)还内置了一组 20 个状态的手工计算期望值,可与 Pyro 输出逐项对照,验证模型正确性。
组合语义 × RSA 语用:semantic_parsing.py
semantic_parsing.py 是最复杂的示例:它把CCG 风格的组合语义(词汇意义 + 句法类型 + 函数复合)与RSA 语用推理拼合在一起。与前几个示例不同,它使用BestFirstSearch(默认num_samples=100)而非Search(见 semantic_parsing.py),因为组合枚举的搜索空间更大,需要按概率优先截断。
核心部件:
- 词汇语义:
Meaning抽象类及其子类(BlondMeaning、NiceMeaning、TallMeaning、BobMeaning、SomeMeaning、AllMeaning、NoneMeaning、UndefinedMeaning),每个意义同时携带sem(world)(在给定世界上求值的语义函数)与syn()(CCG 方向性句法类型,如{"dir": "L", "int": "NP", "out": "S"}表示“左侧取 NP 返回 S”)。 - 组合过程:
can_apply检查相邻词义句法是否可复合,combine_meaning随机(pyro.sample("ix_c", Categorical(...)))选择一条可复合规则,apply_world_passing实现语义函数的世界传递复合,combine_meanings递归直到只剩一个意义。 - 世界先验:
world_prior逐个生成对象(每个对象有三条Bernoulli(0.5)属性),并用累加的pyro.factor施加“意义为真”的软约束(heuristic对真返回 0、假返回 −100)。 - RSA 层级:
literal_listener→speaker→rsa_listener,其中speaker条件化于字面听众、rsa_listener条件化于说话者,形成完整嵌套。
运行:
python examples/rsa/semantic_parsing.py --num-samples=10主程序演示两个查询(见 semantic_parsing.py):字面听众对"all blond people are nice"求“是否有任何人是 nice”的 QUD;RSA 语用听众对"some of the blond people are nice"求“所有金发者是否都 nice”的 QUD。后者展示语用推理如何让“some”隐含地传递“并非所有”的会话含义。
运行环境与注意事项
- 版本匹配:所有脚本入口断言
pyro.__version__.startswith("1.9.1"),需在 Pyro 1.9.1 系列环境中运行;仓库 setup.py 与 pyproject.toml 中的依赖(PyTorch、pyro-ppl的对应版本)须一并满足。 - 数值精度:各脚本均调用
torch.set_default_dtype(torch.float64),以双精度保证离散枚举和对数权重累加的数值稳定性(例如HashingMarginal中的logsumexp累加与log_prob(...).exp()后验输出)。 - 枚举开销:
Search会穷举全部离散分支,模型状态空间(尤其是semantic_parsing.py的组合枚举)可能呈指数增长;空间过大时可改用BestFirstSearch并调小num_samples,或先从小--depth/ 小状态空间起步验证。 - 嵌套与屏蔽:递归推理必须配合
poutine.block()使用(各示例中with poutine.block():包裹内层HashingMarginal(...)),否则内层枚举采样点会污染外层迹;这也是 Pyro 中嵌套推理(nested inference)的标准姿势。 - 可直接运行验证:五个示例均已列入 tests/test_examples.py 的冒烟测试列表,可用
pytest tests/test_examples.py或逐个执行上述命令验证环境。
总结与延伸阅读
本示例集以不到 1500 行代码覆盖了 RSA 语用学建模的完整技术栈:以Search/BestFirstSearch实现精确与近似枚举推理,以HashingMarginal将任意迹后验封装为可嵌套采样的分布,再以Marginal装饰器把“推理求边际”变成记忆化的一等公民,最终通过pyro.sample(..., obs=...)与poutine.scale逐层构建说话者–听众的递归信念嵌套。从最简单的谢林博弈(协调推理)、虚假信念博弈(心智理论),到泛型语句、夸张语(语用推理)与 CCG 组合语义拼合(语义-语用融合),五个模型依次递进,是学习概率编程中“嵌套推理”模式的绝佳教材。
若希望进一步深入,可以:
- 阅读 Pyro 官方推理文档 docs/source/inference.rst 与 docs/source/inference_algos.rst,对比枚举推理与 SVI、MCMC 等近似推断的适用场景;
- 结合 pyro/infer/abstract_infer.py 中
TracePosterior的接口,仿照Search实现自定义的迹后验推理算法; - 参考 pyro/poutine/handlers.py 中
queue、escape、replay等消息处理器的组合方式,理解枚举式执行的底层机制。
- 人工智能
- 机器学习
- 深度学习
- 概率编程
【免费下载链接】pyro
Deep universal probabilistic programming with Python and PyTorch
相关推荐
5个实用技巧:用Chrome.ahk实现浏览器自动化控制的终极指南
5个实用技巧:用Chrome.ahk实现浏览器自动化控制的终极指南 你是否厌倦了重复性的网页操作?是否希望用脚本语言直接控制Chrome浏览器?Chrome.a
浏览器控制RPA戴森球计划工厂蓝图解决方案:3000+优化设计提升建造效率
戴森球计划工厂蓝图解决方案:3000+优化设计提升建造效率 FactoryBluePrints项目为戴森球计划玩家提供了系统性的工厂布局解决方案,通过超过300
游戏开发Qwen大语言模型微调:从理论到实践的完整指南
Qwen大语言模型微调:从理论到实践的完整指南 你是否曾经遇到过这样的困境:想要微调一个强大的语言模型,却发现显存不足、训练时间长、效果不理想?这些问题在传统全
大模型人工智能微调模型量化模型评测本地部署模型推理服务Qwen
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考