☰
Pyro Poutine 深度指南:用可组合效应处理器(Effect Handlers)构建概率编程与自定义推断算法
2026/9/25 8:47:50 网站建设 项目流程
  • 人工智能
  • 机器学习
  • 深度学习
  • 概率编程

【免费下载链接】pyro

Deep universal probabilistic programming with Python and PyTorch

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

Poutine 是 Pyro 内置推断算法之下的核心基础设施——一组可组合的效应处理器(effect handlers),用于记录、拦截和修改概率程序中每一次pyro.sample、pyro.param等原语调用的行为。Pyro 的几乎所有推断算法(SVI、MCMC、枚举推断等)都是把这些 handler 叠加到随机函数(stochastic function)上构建出来的。读完本文,你将掌握 Poutine 的三种调用方式、完整 handler 清单与参数语义、Trace与Runtime的底层数据结构、消息处理管线,以及如何编写自定义 Messenger 来创造新的推断算法。文档原文位于 docs/source/poutine.rst,全部源码位于 pyro/poutine/ 目录。

一、什么是 Poutine:为什么推断算法需要效应处理器

概率编程的关键难题是:模型是一个普通的 Python 函数,但推断算法需要"看到"并"改写"函数内部每一个采样语句的行为。例如变分推断需要把采样替换成从 guide 中采样、需要记录所有采样点的对数概率、需要把某些采样点固定为观测值。

Pyro 借鉴编程语言中的**代数效应(Algebraic Effects)**思想解决这一问题。Poutine 是一组将"效应"(effect)施加在随机函数上的工具:每个 handler 都接收一个随机函数,返回一个行为被修改的新随机函数。在深入 Poutine 之前,建议先阅读 Matija Pretnar 的《An Introduction to Algebraic Effects and Handlers》了解效应处理器解决的一般性问题(文中外部教程链接见原文,此处不重复给出)。

核心结论是:所有内置推断算法都可以用这几行代码的范式表达——先用trace记录执行过程得到Trace,再用replay等 handler 改写执行,最后从Trace中提取对数概率组装损失:

guide_tr = poutine.trace(guide).get_trace(...) model_tr = poutine.trace(poutine.replay(conditioned_model, trace=guide_tr)).get_trace(...) monte_carlo_elbo = model_tr.log_prob_sum() - guide_tr.log_prob_sum()

这段代码即 pyro/poutine/handlers.py 模块 docstring 中给出的"用几行代码实现推断算法"的经典示例,它正是Trace_ELBO的核心思想。

二、三种调用方式与自由组合

Poutine handler 可以当作高阶函数、装饰器或上下文管理器使用,且可以任意嵌套组合。以如下模型为例:

def model(x): s = pyro.param("s", torch.tensor(0.5)) z = pyro.sample("z", dist.Normal(x, s)) return z ** 2

方式一:高阶函数。condition把采样点z标记为观测,返回与model输入输出签名完全一致的函数:

conditioned_model = poutine.condition(model, data={"z": 1.0})

方式二:装饰器:

@pyro.condition(data={"z": 1.0}) def model(x): s = pyro.param("s", torch.tensor(0.5)) z = pyro.sample("z", dist.Normal(x, s)) return z ** 2

方式三:上下文管理器,作用于一段代码块而非整个函数:

with pyro.condition(data={"z": 1.0}): s = pyro.param("s", torch.tensor(0.5)) z = pyro.sample("z", dist.Normal(0., s)) y = z ** 2

自由组合:handler 是纯函数式的,可以无限制叠加,例如trace(condition(model))会同时记录执行并应用条件化。

从实现看(pyro/poutine/handlers.py),每个 handler 都由_make_handler工厂生成:它把 Messenger 类包装成一个handler(fn=None, *args, **kwargs)函数——fn非空时返回msngr(fn)包装后的函数,fn为空时返回 Messenger 实例本身(供上下文管理器或装饰器场景使用)。注意如果第一个参数既不可调用也不是可迭代对象,会抛出ValueError提示"你是否想把它作为关键字参数传入",这是把data等参数误放在首位时最常见的报错。

三、Trace:执行轨迹的图数据结构

Trace(pyro/poutine/trace_struct.py)是 Pyro 程序的执行记录:单次执行中对每个pyro.sample()和pyro.param()调用的完整记录。它是一个有向图,节点代表原语调用或输入输出,边代表节点间的条件依赖关系。

3.1 获取 Trace 与节点元数据

trace = pyro.poutine.trace(model).get_trace(0.0) logp = trace.log_prob_sum() params = [trace.nodes[name]["value"].unconstrained() for name in trace.param_nodes]

trace.nodes是一个collections.OrderedDict,按执行顺序包含_INPUT、s、z、_RETURN等键。每个节点的值是一个消息字典,以采样点z为例:

trace.nodes["z"] # {'type': 'sample', 'name': 'z', 'is_observed': False, # 'fn': Normal(), 'value': tensor(0.6480), 'args': (), 'kwargs': {}, # 'infer': {}, 'scale': 1.0, 'cond_indep_stack': (), # 'done': True, 'stop': False, 'continuation': None}

各字段含义:

字段含义
type消息类型,常见有sample、param、plate、markov,也可自定义
infer用户或算法附加的推断元数据字典(枚举、观测、辅助变量等)
args/kwargspyro.sample传给fn.__call__或fn.log_prob的参数
scale计算联合对数概率时对该站点对数概率的缩放因子
cond_indep_stack对应pyro.plate上下文的不可变条件独立性栈
done/stop/continuationPyro 内部消息处理控制字段,用户一般不直接触碰

infer字典的类型定义见 pyro/poutine/runtime.py,常用键包括:enumerate(取值"sequential"或"parallel",启用离散枚举)、expand(枚举时是否展开分布)、is_auxiliary(是否为辅助变量)、is_observed、obs(观测值)、num_samples、tmc(TraceTMC_ELBO 的"diagonal"/"mixture"近似)等。

3.2 Trace 的关键方法与属性

  • log_prob_sum():遍历所有 sample 站点,用fn.log_prob(value, *args, **kwargs)计算对数概率,经scale_and_mask(缩放与掩码)后求和得到标量联合对数概率;结果按站点记忆化(memoized),且支持site_filter只统计特定站点。开启验证时会对 NaN/Inf 发出告警(pyro/poutine/trace_struct.py)。
  • compute_log_prob()/compute_score_parts():批量计算每个站点的log_prob、log_prob_sum、unscaled_log_prob、score_parts(ScoreParts 三件套:log_prob、score_function、entropy_term),全部记忆化。compute_score_parts是Trace_ELBO处理不可重参数化站点时梯度估计的基础。
  • 属性:observation_nodes(观测站点)、param_nodes(参数站点)、stochastic_nodes(未观测的采样站点)、reparameterized_nodes(has_rsample=True的可重参数化站点)、nonreparam_stochastic_nodes(不可重参数化采样站点)。
  • 图操作:add_node(重复节点默认报错)、add_edge、remove_node、predecessors/successors、topological_sort(拓扑排序,reverse=True反向)、copy(浅拷贝)。
  • detach_():原地把所有 sample 值.detach(),用于切断梯度。
  • format_shapes():生成站点形状表格的字符串,TraceHandler在模型抛出ValueError/RuntimeError时会自动把形状表追加到异常信息中,这是 Pyro 报错信息中"Trace Shapes"表格的来源。
  • symbolize_dims()/pack_tensors():为 plate 维与枚举维分配唯一符号(偶数为 plate 维、奇数为枚举维),并计算打包张量,是并行枚举与 Tensor 计算内部优化环节。

graph_type参数支持"flat"与"dense"两种图:flat 只记录执行序与观测依赖;dense 会在TraceMessenger.__exit__时调用identify_dense_edges(pyro/poutine/trace_messenger.py),根据cond_indep_stack上相同名字、不同counter的 frame 判断条件独立,补全所有条件依赖边,供 tracegraph 类 ELBO 使用。

四、Runtime:消息与执行栈

Runtime(pyro/poutine/runtime.py)是 Poutine 的执行引擎,核心是全局效应栈与消息处理管线。

4.1 效应栈与 Message

  • 全局栈_PYRO_STACK是一个List[Messenger],所有激活中的 handler 按with嵌套顺序压栈。
  • Message是 Pyro 内部的消息类型(TypedDict),即 Trace 中每个节点的结构,字段定义见 pyro/poutine/runtime.py。
  • am_i_wrapped()返回当前是否处于任一 poutine 包裹中(len(_PYRO_STACK) > 0)。

4.2 apply_stack:消息的四阶段管线

当程序在 handler 包裹下执行pyro.sample等原语时,effectful装饰器会构造一条初始消息并调用apply_stack(pyro/poutine/runtime.py):

  1. 下行处理:从栈底到栈顶,对每个 Messenger 调用_process_message(msg);若消息stop字段变为 True 则提前终止;
  2. 默认行为:调用default_process_message——若消息已完成、已观测或已有值则标记done=True;否则执行msg"fn"采样并标记完成;
  3. 上行后处理:从栈顶到栈底,对每个 Messenger 调用_postprocess_message(msg),把执行结果写回消息并更新 messenger 内部状态(如TraceMessenger在此把节点加入 trace);
  4. 延续:若消息携带continuation回调,则调用它。

effectful(pyro/poutine/runtime.py)是这一机制的入口:它要求每个操作必须有type标签(如"sample"),并自动把name、infer、obs参数转为消息字段,obs is not None时is_observed=True。pyro.sample、pyro.param、pyro.factor、pyro.plate等原语本质上都是effectful装饰的底层函数。未包裹时effectful直接透传原始调用,因此裸模型在无 handler 环境下运行就是普通的 Python 函数。

4.3 辅助运行时设施

  • NonlocalExit异常(pyro/poutine/runtime.py):由EscapeMessenger抛出,携带当前站点消息,用于从程序内部"非局部跳出";reset_stack()在多次重执行前重置栈中 frame 状态(poutine.queue中反复重入依赖它)。
  • get_mask():返回_inspect()["mask"],记录外层poutine.mask的掩码效果,可用于跳过昂贵的pyro.factor计算:
def model(): if poutine.get_mask() is not False: log_density = my_expensive_computation() pyro.factor("foo", log_density)
  • get_plates():返回当前cond_indep_stack中的 plate frame 元组。
  • 维度分配器:_DimAllocator(为 plate 从右向左分配维度,维冲突会给出"Try moving the dim of one plate to the left"的提示)与_EnumAllocator(为并行枚举分配可回收维度,要求first_available_dim < 0)。

五、Handler 全览:签名、参数与用途

所有 handler 在 pyro/poutine/handlers.py 中定义(含完整类型签名与默认值),并通过 pyro/poutine/init.py 导出为poutine.xxx,同时通过pyro.xxx顶层命名空间可直接使用。

Handler核心参数作用
tracegraph_type("flat"/"dense")、param_only记录执行轨迹,param_only=True时只记录参数;退出时若为 dense 图补全依赖边
conditiondata(字典或 Trace)把字典中名字对应的 sample 站点改为观测站点
dodata(Dict[str, Tensor/Number])干预(do-calculus):把名字对应站点的采样替换为固定值并切断该站点的对数概率贡献
substitutedata替换站点值为指定值,但不改变其观测/采样性质
replaytrace、params用给定 Trace 中已记录的值重放采样,匹配不到名字的站点正常采样;params指定只重放部分参数
blockhide/expose/hide_types/expose_types/hide_all/expose_all、hide_fn/expose_fn对外屏蔽部分站点,默认屏蔽全部;隐藏判定规则见下文
seedrng_seed(int)在执行前设置全局随机种子,保证可复现
scalescale(float 或 Tensor)缩放站点对数概率,等价于对 ELBO 项加权
maskmask(bool 或 BoolTensor)掩码(屏蔽)部分站点对数概率,等价于对站点施加 0/1 权重
liftprior(分布、字典或可调用)把采样点改为从指定先验分布采样(lift 效应,用于 guide 先验注入)
broadcast无自动广播采样值/观测值到 plate 形状,消除手动广播样板
enumfirst_available_dim离散枚举采样支持(配合config_enumerate使用)
markovhistory(默认 1)、keep、dim、name声明马尔可夫依赖;history=0时类似pyro.plate;keep=True时 frame 可重放(同层相邻分支可互相依赖);dim/name为接口桩,行为尚未实现
escapeescape_fn满足谓词(如discrete_escape)时抛出NonlocalExit跳出执行
collapse无折叠(collapsing out)条件独立结构
equalizesites、type、keep_dist平衡多个站点的分布
infer_configconfig_fn按站点动态改写infer推断配置字典
reparamconfig(字典或可调用,值为Reparam)按配置重参数化采样站点(如loc_scale、neutra、haar等,见 pyro/infer/reparam/)
uncondition无取消条件化,把观测站点恢复为潜在变量
queuequeue、max_tries(默认1e6)、extend_fn(默认enum_extend)、escape_fn(默认discrete_escape)、num_samples(默认 -1)顺序枚举离散变量的复合操作:从队列取部分 trace 执行,遇NonlocalExit时扩展部分 trace 并放回队列,直到得到完整 trace

5.1 block 的隐藏判定规则

BlockMessenger(pyro/poutine/block_messenger.py)默认行为是屏蔽一切。一个站点被隐藏当且仅当以下条件之一成立:

  1. hide_fn(msg) is True或(not expose_fn(msg)) is True;
  2. msg["name"] in hide;
  3. msg["type"] in hide_types(注意观测站点会被当作"observe"类型处理);
  4. msg["name"] not in expose且msg["type"] not in expose_types;
  5. hide、hide_types、expose_types均为None。

_make_default_hide_fn会做一致性校验:hide_all与expose_all不能同时为真,hide与expose不能有交集,hide_types与expose_types同理;显式给出expose/expose_types时hide_all会被置为 True(即"只放行显式暴露的站点")。典型用法如poutine.block(fn_inner, hide=["a"])——内层 trace 能看到站点a和b,外层任何效应都看不到a。

5.2 与推断枚举相关的 handler

  • config_enumerate(见 docs/source/poutine.rst 中autofunction:: pyro.infer.enum.config_enumerate,实现在 pyro/infer/enum.py)用于为采样站点批量配置infer={"enumerate": ...}策略。
  • enum_extend(pyro/poutine/util.py)通过fn.enumerate_support()枚举站点支撑集生成多个扩展 trace;discrete_escape(同文件 L111-L128)判断站点是否为离散、未观测、未入 trace 且具有has_enumerate_support。两者配合queue即可实现精确的顺序枚举推断;mc_extend(L83-L108)则以num_samples次蒙特卡洛采样扩展 trace,用于对个别站点做 MC 边缘化。

六、Messenger:效应的底层实现与自定义扩展

文档明确指出:Messenger 对象是 handler 所暴露效应的底层实现。高级用户可以直接修改已有 handler 背后的 messenger,或编写新 messenger 实现新效应,并保证与库其余部分正确组合。

6.1 Messenger 基类契约

Messenger(pyro/poutine/messenger.py)本身是一个上下文管理器,基类对所有 Pyro 原语实现默认行为——因此Messenger()(fn)生成的联合分布与原函数完全相同。

  • __call__(fn):返回一个包装函数,执行时在with self:下运行fn(*args, **kwargs);
  • __enter__():把自身压入_PYRO_STACK栈底(同一实例不能安装两次,否则抛ValueError),必须返回self;
  • __exit__():正常退出时弹栈;若包裹代码抛异常,则从栈中找到自身位置并移除自身及以下所有 frame;
  • _process_message(msg)/_postprocess_message(msg):按msg["type"]动态分派到_pyro_{type}或_pyro_post_{type}方法(如_pyro_sample、_pyro_post_sample),消息原地更新;
  • register(fn, type, post)/unregister(fn, type):动态为效应添加/移除操作(post=True注册后处理),可用于为第三方库生成包装器:
@SomeMessengerClass.register def some_function(msg): ...do_something... return msg

6.2 关键 Messenger 的实现要点

  • TraceMessenger(pyro/poutine/trace_messenger.py):在_pyro_post_sample/_pyro_post_param中把消息加入 trace;_pyro_post_sample会跳过infer["_do_not_trace"]的辅助站点(须同时满足is_auxiliary且未观测)。TraceHandler.__call__负责加入_INPUT/_RETURN节点,并在异常时附加format_shapes()表格。
  • ConditionMessenger(pyro/poutine/condition_messenger.py):_pyro_sample中若站点名在data中,则把msg["value"]设为字典值(或 Trace 中该节点值),is_observed置为value is not None——即把 sample 变 observe,等价于在pyro.sample中加obs=value。

6.3 实验性工具:block_messengers

block_messengers(predicate)(pyro/poutine/messenger.py)是一个实验性上下文管理器:把满足谓词的 messenger 暂时从_PYRO_STACK中替换为平凡 messenger(不调用其__enter__/__exit__),用于选择性屏蔽外层 handler,并 yield 被屏蔽的 messenger 列表。

七、Utilities:推断工具函数

pyro.poutine.util(pyro/poutine/util.py)提供推断辅助函数:

  • enable_validation(is_validate)/is_validation_enabled():全局开关 poutine 验证(默认与__debug__一致),开启时log_prob_sum等计算会对 NaN/Inf 告警;同时注册为 Pyro 全局设置validate_poutine。
  • site_is_subsample(site):判断站点是否来自plate内的 subsample 语句(分布类型名为_Subsample);site_is_factor判断是否来自pyro.factor(分布类型名为Unit)。
  • prune_subsample_sites(trace):复制并移除所有 subsample 站点。
  • enum_extend/mc_extend/discrete_escape/all_escape:见 5.2 节,是顺序枚举与 variance reduction 的子程序。

八、实战:用 trace + replay 组装最小推断算法

把前述知识连起来,一个最小可用的"guide + model" ELBO 组装流程如下:

import pyro import pyro.distributions as dist import pyro.poutine as poutine def model(data): z = pyro.sample("z", dist.Normal(0, 1)) pyro.sample("x", dist.Normal(z, 1), obs=data) def guide(data): loc = pyro.param("loc", torch.tensor(0.0)) scale = pyro.param("scale", torch.tensor(1.0), constraint=constraints.positive) pyro.sample("z", dist.Normal(loc, scale)) # 1) 记录 guide 执行,得到参考 trace guide_tr = poutine.trace(guide).get_trace(data) # 2) 用 guide 的采样值重放 model,同时记录 model trace model_tr = poutine.trace(poutine.replay(model, trace=guide_tr)).get_trace(data) # 3) 由两段 trace 计算 ELBO 损失 elbo = guide_tr.log_prob_sum() - model_tr.log_prob_sum()

这里replay保证 model 中站点z使用 guide 采样出的同一批值(这是 SVI 中"引导重参数化"的基础),trace负责两侧的完整记录,log_prob_sum提供标量损失。这正是 pyro/poutine/handlers.py 文档所演示的范式,也是Trace_ELBO(pyro/infer/trace_elbo.py)等实现的雏形。

九、测试与更多资源

Pyro 为 Poutine 提供了完整的测试覆盖(tests/poutine/),包括:

  • test_poutines.py:各 handler 的功能测试;
  • test_nesting.py:handler 嵌套/组合正确性;
  • test_trace_struct.py:Trace 图结构(节点、边、拓扑排序)与log_prob_sum等数值计算;
  • test_runtime.py、test_mapdata.py、test_properties.py:运行时栈、plate/map-data 交互与性质测试;
  • test_counterfactual.py:condition/do反事实推断测试。

更多实践示例可参考 docs/source/poutine.rst、effect_handlers.ipynb 教程,以及 pyro/infer/ 下各 ELBO、枚举与 MCMC 算法对 handler 的组合运用。从源码结构看,pyro.infer、pyro.contrib中的大量高级功能(config_enumerate、TraceTMC_ELBO、DiscreteHMM等)都以本文介绍的Trace、Runtime与 Messenger 体系为底层依赖——理解 Poutine 就等于拿到了阅读 Pyro 全部推断代码的钥匙。

  • 人工智能
  • 机器学习
  • 深度学习
  • 概率编程

【免费下载链接】pyro

Deep universal probabilistic programming with Python and PyTorch

项目地址:https://gitcode.com/gh_mirrors/py/pyro
点击查看免费下载
上一篇:jemalloc mallctl 内存监控完整指南:3 个函数、1 张症状表、8 项上线检查清单
下一篇:CANN/asc-devkit SIMD API文档

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

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

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

立即咨询