- 人工智能
- 机器学习
- 深度学习
- 概率编程
【免费下载链接】pyro
Deep universal probabilistic programming with Python and PyTorch
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/kwargs | pyro.sample传给fn.__call__或fn.log_prob的参数 |
scale | 计算联合对数概率时对该站点对数概率的缩放因子 |
cond_indep_stack | 对应pyro.plate上下文的不可变条件独立性栈 |
done/stop/continuation | Pyro 内部消息处理控制字段,用户一般不直接触碰 |
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):
- 下行处理:从栈底到栈顶,对每个 Messenger 调用
_process_message(msg);若消息stop字段变为 True 则提前终止; - 默认行为:调用
default_process_message——若消息已完成、已观测或已有值则标记done=True;否则执行msg"fn"采样并标记完成; - 上行后处理:从栈顶到栈底,对每个 Messenger 调用
_postprocess_message(msg),把执行结果写回消息并更新 messenger 内部状态(如TraceMessenger在此把节点加入 trace); - 延续:若消息携带
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 | 核心参数 | 作用 |
|---|---|---|
trace | graph_type("flat"/"dense")、param_only | 记录执行轨迹,param_only=True时只记录参数;退出时若为 dense 图补全依赖边 |
condition | data(字典或 Trace) | 把字典中名字对应的 sample 站点改为观测站点 |
do | data(Dict[str, Tensor/Number]) | 干预(do-calculus):把名字对应站点的采样替换为固定值并切断该站点的对数概率贡献 |
substitute | data | 替换站点值为指定值,但不改变其观测/采样性质 |
replay | trace、params | 用给定 Trace 中已记录的值重放采样,匹配不到名字的站点正常采样;params指定只重放部分参数 |
block | hide/expose/hide_types/expose_types/hide_all/expose_all、hide_fn/expose_fn | 对外屏蔽部分站点,默认屏蔽全部;隐藏判定规则见下文 |
seed | rng_seed(int) | 在执行前设置全局随机种子,保证可复现 |
scale | scale(float 或 Tensor) | 缩放站点对数概率,等价于对 ELBO 项加权 |
mask | mask(bool 或 BoolTensor) | 掩码(屏蔽)部分站点对数概率,等价于对站点施加 0/1 权重 |
lift | prior(分布、字典或可调用) | 把采样点改为从指定先验分布采样(lift 效应,用于 guide 先验注入) |
broadcast | 无 | 自动广播采样值/观测值到 plate 形状,消除手动广播样板 |
enum | first_available_dim | 离散枚举采样支持(配合config_enumerate使用) |
markov | history(默认 1)、keep、dim、name | 声明马尔可夫依赖;history=0时类似pyro.plate;keep=True时 frame 可重放(同层相邻分支可互相依赖);dim/name为接口桩,行为尚未实现 |
escape | escape_fn | 满足谓词(如discrete_escape)时抛出NonlocalExit跳出执行 |
collapse | 无 | 折叠(collapsing out)条件独立结构 |
equalize | sites、type、keep_dist | 平衡多个站点的分布 |
infer_config | config_fn | 按站点动态改写infer推断配置字典 |
reparam | config(字典或可调用,值为Reparam) | 按配置重参数化采样站点(如loc_scale、neutra、haar等,见 pyro/infer/reparam/) |
uncondition | 无 | 取消条件化,把观测站点恢复为潜在变量 |
queue | queue、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)默认行为是屏蔽一切。一个站点被隐藏当且仅当以下条件之一成立:
hide_fn(msg) is True或(not expose_fn(msg)) is True;msg["name"] in hide;msg["type"] in hide_types(注意观测站点会被当作"observe"类型处理);msg["name"] not in expose且msg["type"] not in expose_types;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 msg6.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
相关推荐
Pyro概率编程终极指南:MCMC与SVI推理算法深度对比解析
Pyro概率编程终极指南:MCMC与SVI推理算法深度对比解析 Pyro作为基于PyTorch构建的深度概率编程库,提供了强大的推理引擎功能,其中 MCMC (
人工智能机器学习深度学习概率编程OpenAEV终极指南:如何构建企业级安全验证平台的完整教程
OpenAEV终极指南:如何构建企业级安全验证平台的完整教程 OpenAEV(Open Adversarial Exposure Validation Plat
Metasploit in Termux核心组件解析:从数据库配置到模块加载
Metasploit in Termux核心组件解析:从数据库配置到模块加载 Metasploit Framework是一款强大的渗透测试工具,而在Termux
网络安全
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考