Pyro Poutine 深度指南:用可组合效应处理器(Effect Handlers)构建概率编程与自定义推断算法
发布时间:2026/9/25 8:47:51来源:尧图网络
人工智能机器学习深度学习概率编程【免费下载链接】pyroDeep 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, traceguide_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(fnNone, *args, **kwargs)函数——fn非空时返回msngr(fn)包装后的函数fn为空时返回 Messenger 实例本身供上下文管理器或装饰器场景使用。注意如果第一个参数既不可调用也不是可迭代对象会抛出ValueError提示你是否想把它作为关键字参数传入这是把data等参数误放在首位时最常见的报错。三、Trace执行轨迹的图数据结构Tracepyro/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、tmcTraceTMC_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_partsScoreParts 三件套log_prob、score_function、entropy_term全部记忆化。compute_score_parts是Trace_ELBO处理不可重参数化站点时梯度估计的基础。属性observation_nodes观测站点、param_nodes参数站点、stochastic_nodes未观测的采样站点、reparameterized_nodeshas_rsampleTrue的可重参数化站点、nonreparam_stochastic_nodes不可重参数化采样站点。图操作add_node重复节点默认报错、add_edge、remove_node、predecessors/successors、topological_sort拓扑排序reverseTrue反向、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_edgespyro/poutine/trace_messenger.py根据cond_indep_stack上相同名字、不同counter的 frame 判断条件独立补全所有条件依赖边供 tracegraph 类 ELBO 使用。四、Runtime消息与执行栈Runtimepyro/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_stackpyro/poutine/runtime.py下行处理从栈底到栈顶对每个 Messenger 调用_process_message(msg)若消息stop字段变为 True 则提前终止默认行为调用default_process_message——若消息已完成、已观测或已有值则标记doneTrue否则执行msgfn采样并标记完成上行后处理从栈顶到栈底对每个 Messenger 调用_postprocess_message(msg)把执行结果写回消息并更新 messenger 内部状态如TraceMessenger在此把节点加入 trace延续若消息携带continuation回调则调用它。effectfulpyro/poutine/runtime.py是这一机制的入口它要求每个操作必须有type标签如sample并自动把name、infer、obs参数转为消息字段obs is not None时is_observedTrue。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_typeflat/dense、param_only记录执行轨迹param_onlyTrue时只记录参数退出时若为 dense 图补全依赖边conditiondata字典或 Trace把字典中名字对应的 sample 站点改为观测站点dodataDict[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_seedint在执行前设置全局随机种子保证可复现scalescalefloat 或 Tensor缩放站点对数概率等价于对 ELBO 项加权maskmaskbool 或 BoolTensor掩码屏蔽部分站点对数概率等价于对站点施加 0/1 权重liftprior分布、字典或可调用把采样点改为从指定先验分布采样lift 效应用于 guide 先验注入broadcast无自动广播采样值/观测值到 plate 形状消除手动广播样板enumfirst_available_dim离散枚举采样支持配合config_enumerate使用markovhistory默认 1、keep、dim、name声明马尔可夫依赖history0时类似pyro.platekeepTrue时 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 并放回队列直到得到完整 trace5.1 block 的隐藏判定规则BlockMessengerpyro/poutine/block_messenger.py默认行为是屏蔽一切。一个站点被隐藏当且仅当以下条件之一成立hide_fn(msg) is True或(not expose_fn(msg)) is Truemsg[name] in hidemsg[type] in hide_types注意观测站点会被当作observe类型处理msg[name] not in expose且msg[type] not in expose_typeshide、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 与推断枚举相关的 handlerconfig_enumerate见 docs/source/poutine.rst 中autofunction:: pyro.infer.enum.config_enumerate实现在 pyro/infer/enum.py用于为采样站点批量配置infer{enumerate: ...}策略。enum_extendpyro/poutine/util.py通过fn.enumerate_support()枚举站点支撑集生成多个扩展 tracediscrete_escape同文件 L111-L128判断站点是否为离散、未观测、未入 trace 且具有has_enumerate_support。两者配合queue即可实现精确的顺序枚举推断mc_extendL83-L108则以num_samples次蒙特卡洛采样扩展 trace用于对个别站点做 MC 边缘化。六、Messenger效应的底层实现与自定义扩展文档明确指出Messenger 对象是 handler 所暴露效应的底层实现。高级用户可以直接修改已有 handler 背后的 messenger或编写新 messenger 实现新效应并保证与库其余部分正确组合。6.1 Messenger 基类契约Messengerpyro/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)动态为效应添加/移除操作postTrue注册后处理可用于为第三方库生成包装器SomeMessengerClass.register def some_function(msg): ...do_something... return msg6.2 关键 Messenger 的实现要点TraceMessengerpyro/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()表格。ConditionMessengerpyro/poutine/condition_messenger.py_pyro_sample中若站点名在data中则把msg[value]设为字典值或 Trace 中该节点值is_observed置为value is not None——即把 sample 变 observe等价于在pyro.sample中加obsvalue。6.3 实验性工具block_messengersblock_messengers(predicate)pyro/poutine/messenger.py是一个实验性上下文管理器把满足谓词的 messenger 暂时从_PYRO_STACK中替换为平凡 messenger不调用其__enter__/__exit__用于选择性屏蔽外层 handler并 yield 被屏蔽的 messenger 列表。七、Utilities推断工具函数pyro.poutine.utilpyro/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 语句分布类型名为_Subsamplesite_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), obsdata) def guide(data): loc pyro.param(loc, torch.tensor(0.0)) scale pyro.param(scale, torch.tensor(1.0), constraintconstraints.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, traceguide_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_ELBOpyro/infer/trace_elbo.py等实现的雏形。九、测试与更多资源Pyro 为 Poutine 提供了完整的测试覆盖tests/poutine/包括test_poutines.py各 handler 的功能测试test_nesting.pyhandler 嵌套/组合正确性test_trace_struct.pyTrace 图结构节点、边、拓扑排序与log_prob_sum等数值计算test_runtime.py、test_mapdata.py、test_properties.py运行时栈、plate/map-data 交互与性质测试test_counterfactual.pycondition/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 全部推断代码的钥匙。赞分享人工智能机器学习深度学习概率编程【免费下载链接】pyroDeep universal probabilistic programming with Python and PyTorch项目地址https://gitcode.com/gh_mirrors/py/pyro点击查看免费下载相关推荐Pyro概率编程终极指南MCMC与SVI推理算法深度对比解析Pyro概率编程终极指南MCMC与SVI推理算法深度对比解析 Pyro作为基于PyTorch构建的深度概率编程库提供了强大的推理引擎功能其中 MCMC 人工智能机器学习深度学习概率编程OpenAEV终极指南如何构建企业级安全验证平台的完整教程OpenAEV终极指南如何构建企业级安全验证平台的完整教程 OpenAEVOpen Adversarial Exposure Validation PlatMetasploit in Termux核心组件解析从数据库配置到模块加载Metasploit in Termux核心组件解析从数据库配置到模块加载 Metasploit Framework是一款强大的渗透测试工具而在Termux网络安全上一篇jemalloc mallctl 内存监控完整指南3 个函数、1 张症状表、8 项上线检查清单下一篇CANN/asc-devkit SIMD API文档创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
网站建设高端定制企业官网