☰
Pyro 中的理性言语行为(RSA)嵌套推理示例:从博弈论协调到语用学建模
2026/9/25 3:34:17 网站建设 项目流程
  • 人工智能
  • 机器学习
  • 深度学习
  • 概率编程

【免费下载链接】pyro

Deep universal probabilistic programming with Python and PyTorch

项目地址:https://gitcode.com/gh_mirrors/py/pyro
点击查看免费下载

本篇技术指南聚焦 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 实际想避开 BobForestDB 的 schelling-falsebelief 模型
search_inference.py全部示例共用的推理工具:Search、BestFirstSearch、HashingMarginal、memoizeDesign 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

关键点有三个:

  1. 深度递归:bob在depth > 0时推理 Alice,alice推理bob(depth-1),形成alice → bob → alice → …的递归链;depth控制推理层级。
  2. poutine.block()屏蔽内层采样:当 Alice 把 Bob 的决策过程作为“子程序”调用时,block()保证内层Search枚举产生的采样点不会泄漏到外层迹中,避免命名冲突与错误嵌套。
  3. 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”隐含地传递“并非所有”的会话含义。

运行环境与注意事项

  1. 版本匹配:所有脚本入口断言pyro.__version__.startswith("1.9.1"),需在 Pyro 1.9.1 系列环境中运行;仓库 setup.py 与 pyproject.toml 中的依赖(PyTorch、pyro-ppl的对应版本)须一并满足。
  2. 数值精度:各脚本均调用torch.set_default_dtype(torch.float64),以双精度保证离散枚举和对数权重累加的数值稳定性(例如HashingMarginal中的logsumexp累加与log_prob(...).exp()后验输出)。
  3. 枚举开销:Search会穷举全部离散分支,模型状态空间(尤其是semantic_parsing.py的组合枚举)可能呈指数增长;空间过大时可改用BestFirstSearch并调小num_samples,或先从小--depth/ 小状态空间起步验证。
  4. 嵌套与屏蔽:递归推理必须配合poutine.block()使用(各示例中with poutine.block():包裹内层HashingMarginal(...)),否则内层枚举采样点会污染外层迹;这也是 Pyro 中嵌套推理(nested inference)的标准姿势。
  5. 可直接运行验证:五个示例均已列入 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

项目地址:https://gitcode.com/gh_mirrors/py/pyro
点击查看免费下载

相关推荐

上一篇:18节点EP144架构实战:DeepSeek Open Infra Index分布式推理性能提升545%的终极指南
下一篇:终极Rofi主题开发指南:从RASI语法到自定义样式的完整教程

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询