人工智能机器学习深度学习概率编程【免费下载链接】pyroDeep universal probabilistic programming with Python and PyTorch项目地址https://gitcode.com/gh_mirrors/py/pyro点击查看免费下载本教程以 Pyro 仓库中的官方示例 examples/toy_mixture_model_discrete_enumeration.py对应教程页 tutorial/source/toy_mixture_model_discrete_enumeration.rst为主线完整演示如何在变分推断SVI中通过**离散枚举discrete enumeration**将隐离散变量精确边际化。读完本文你将掌握三个核心步骤用pyro.infer.config_enumerate装饰模型、在目标 sample 站点标注infer{enumerate: parallel}、并改用TraceEnum_ELBO作为损失函数同时理解Vindex、max_plate_nesting等关键机制在源码层面的工作原理。1. 示例模型A → [B] → C 的三变量混合模型示例的目标结构为(A) - [B] - (C)其中A是已观测的 Bernoulli 变量带 Beta 先验B是隐藏变量方括号表示不观测是两个 Bernoulli 分布的混合具体取哪个分布由 A 的取值True/False决定其参数本身也有 Beta 先验C是已观测变量同 B 一样是两个 Bernoulli 分布的混合由 B 的取值决定三者之上有一个plate表示n条独立观测数据。因为 B 是隐藏的离散变量我们希望把它从模型中**精确边际化marginalize**掉而不是用采样估计。示例文档明确给出了三条操作步骤在模型方法上标记pyro.infer.config_enumerate在模型中的 B 采样站点标注infer{enumerate: parallel}给pyro.infer.SVI传入pyro.infer.TraceEnum_ELBO损失函数。这三步是本文后续展开的主线也是 Pyro 中模型内离散枚举的标准范式。2. 数据生成用真实 CPD 构造观测generate_data(num_obs)先从先验超参数中采样出真实的条件概率表CPD再用这些 CPD 生成观测数据prior { A: torch.tensor([1.0, 10.0]), B: torch.tensor([[10.0, 1.0], [1.0, 10.0]]), C: torch.tensor([[10.0, 1.0], [1.0, 10.0]]), } CPDs { p_A: Beta(prior[A][0], prior[A][1]).sample(), p_B: Beta(prior[B][:, 0], prior[B][:, 1]).sample(), p_C: Beta(prior[C][:, 0], prior[C][:, 1]).sample(), } data {A: Bernoulli(torch.ones(num_obs) * CPDs[p_A]).sample()} data[B] Bernoulli(torch.gather(CPDs[p_B], 0, data[A].type(torch.long))).sample() data[C] Bernoulli(torch.gather(CPDs[p_C], 0, data[B].type(torch.long))).sample()要点解读p_A是标量概率p_B、p_C是长度为 2 的向量索引 0/1 分别对应当前父节点取 False/True 时的成功概率torch.gather负责按父节点取值查表例如data[B]的成功概率由data[A]的真实取值决定从而构造出条件依赖返回的prior稍后直接用作 guide 中变分参数的初始值CPDs用作评估后验质量的基准真值。3. 模型与 Guide三个关键步骤的落点3.1 模型枚举标注与Vindexpyro.infer.config_enumerate def model(prior, obs, num_obs): p_A pyro.sample(p_A, dist.Beta(1, 1)) p_B pyro.sample(p_B, dist.Beta(torch.ones(2), torch.ones(2)).to_event(1)) p_C pyro.sample(p_C, dist.Beta(torch.ones(2), torch.ones(2)).to_event(1)) with pyro.plate(data_plate, num_obs): A pyro.sample(A, dist.Bernoulli(p_A.expand(num_obs)), obsobs[A]) # Vindex used to ensure proper indexing into the enumerated sample sites B pyro.sample( B, dist.Bernoulli(Vindex(p_B)[A.type(torch.long)]), infer{enumerate: parallel}, ) pyro.sample(C, dist.Bernoulli(Vindex(p_C)[B.type(torch.long)]), obsobs[C])三个关键点pyro.infer.config_enumerate来自 pyro/infer/enum.py。它本质上是poutine.infer_config的封装会给所有支持枚举has_enumerate_support True的离散 sample 站点注入枚举配置。注意源码注释强调它不会覆盖站点上已有的infer{enumerate: ...}注解——这正是本例中 B 站点显式写infer{enumerate: parallel}依然生效的原因。它同时支持作为函数guide config_enumerate(guide)或装饰器config_enumerate两种用法参数包括defaultparallel、expandFalse、num_samplesNone、tmcdiagonal且会对非法参数如defaultsequential配合num_samples直接抛ValueError。infer{enumerate: parallel}标注 B 为并行枚举站点。B 不出现在 guide 中因此枚举发生在模型侧。TraceEnum_ELBO的类文档明确指出要在模型中对某站点做枚举就标记infer{enumerate: parallel}并确保该站点不出现在 guide 中而 guide 侧的枚举可以选sequential或parallel。从源码看真正执行枚举的是 pyro/poutine/enum_messenger.py 中的EnumMessenger它拦截_pyro_sample对满足条件的站点调用enumerate_site取出全部支撑取值再通过全局_ENUM_ALLOCATOR分配一个负索引的枚举维度_enumerate_dim使枚举以张量维度的形式向量化展开。TraceEnum_ELBO._get_traces中正是先后用EnumMessenger(first_available_dim-1 - self.max_plate_nesting)包裹 guide、再用不带参数的EnumMessenger()包裹 model从而保证模型枚举维度分配在 guide 枚举维度左侧更全局。Vindex来自 pyro/ops/indexing.py 的向量化高级索引工具。枚举会让p_B、p_C以及 A/B 的取值都带上额外的批次维度枚举维度普通p_B[A]这种索引会破坏广播语义或把枚举维度当成普通维度误处理。Vindex(p_B)[A.type(torch.long)]会按广播语义对批次维度做高级索引确保按父节点取值选择混合成分的操作在枚举场景下依然正确。源码注释给出了精确的行为约定Vindex(x)[i, :, j]中i, j可以携带批次维和枚举维但不能有事件维...表示未知的批次维度且只能出现在最左侧。3.2 Guide为连续参数匹配变分族def guide(prior, obs, num_obs): a pyro.param(a, prior[A], constraintconstraints.positive) pyro.sample(p_A, dist.Beta(a[0], a[1])) b pyro.param(b, prior[B], constraintconstraints.positive) pyro.sample(p_B, dist.Beta(b[:, 0], b[:, 1]).to_event(1)) c pyro.param(c, prior[C], constraintconstraints.positive) pyro.sample(p_C, dist.Beta(c[:, 0], c[:, 1]).to_event(1))guide 只处理连续参数p_A / p_B / p_C用 Beta 分布不包含隐变量 B 的近似分布——这正是精确枚举相对于用 guide 近似隐变量的差别离散隐变量在模型侧被穷举求和掉了。constraints.positive保证 Beta 的形状参数恒为正初始值直接取自先验超参数prior加速收敛。guide 与 model 的签名保持一致prior, obs, num_obsSVI 会以相同参数调用两者。4. 训练TraceEnum_ELBO与max_plate_nestingdef train(prior, data, num_steps, num_obs): pyro.clear_param_store() # max_plate_nesting 1 because there is a single plate in the model loss_func pyro.infer.TraceEnum_ELBO(max_plate_nesting1) svi pyro.infer.SVI(model, guide, pyro.optim.Adam({lr: 0.01}), lossloss_func) losses [] for _ in tqdm(range(num_steps)): loss svi.step(prior, data, num_obs) losses.append(loss) ...关键点逐一展开pyro.clear_param_store()清空参数库避免多次运行示例时残留上一次的参数值。max_plate_nesting1本模型只有一层data_plate因此嵌套深度为 1。ELBO基类文档pyro/infer/elbo.py说明max_plate_nesting是并行枚举 sample 站点时必需的上界若省略ELBO 会通过运行一次 (model, guide) 对来猜测有效值但当模型或 guide 结构是动态的时这种猜测可能不正确。在本示例的静态结构中显式传入 1 是推荐做法。从TraceEnum_ELBO._get_traces的源码可以看出其作用它决定EnumMessenger的first_available_dim -1 - max_plate_nesting即枚举维度必须分配在 plate 维度更左侧的位置从而保证形状规则plate 内变量不能依赖 plate 外变量。pyro.optim.Adam({lr: 0.01})Pyro 的优化器封装学习率 0.01是示例默认值。TraceEnum_ELBO定义在 pyro/infer/traceenum_elbo.py。它支持两类能力对离散站点的穷举枚举本示例用法与对 guide 中任意站点的局部并行采样。其 ELBO 计算核心是_compute_dice_elbo先把模型因子按是否依赖枚举维度分类用contract_tensor_tree对依赖枚举维度的因子做张量消息传递基于opt_einsum的shared_intermediates缓存与SampleRing把隐变量求和掉再由Dice算子计算带log_factors的期望loss_and_grads对每个粒子的 ELBO 反向传播返回-elbo作为损失。它额外提供了一个便捷能力compute_marginals()可返回每个模型侧枚举站点的边际分布要求num_particles1且无 guide 侧枚举。形状校验启用验证模式时_get_trace会调用check_traceenum_requirements检查模型与 guide 的站点匹配与形状合法性若strict_enumeration_warning开启且没有任何站点配置枚举会警告应改用Trace_ELBO。训练完成后示例把参数库导出为 numpy 数组用于评估其中对a做了a[None, :]重塑使其与b、c的二维形状对齐b天然是(2, 2)的 CPD 形状。5. 评估对比真实 CPD 与预测 CPDdef evaluate(CPDs, posterior_params): true_p_A, pred_p_A get_true_pred_CPDs(CPDs[p_A], posterior_params[a]) true_p_B, pred_p_B get_true_pred_CPDs(CPDs[p_B], posterior_params[b]) true_p_C, pred_p_C get_true_pred_CPDs(CPDs[p_C], posterior_params[c]) ... def get_true_pred_CPDs(CPD, posterior_param): true_p CPD.numpy() pred_p posterior_param[:, 0] / np.sum(posterior_param, axis1) return true_p, pred_p预测概率由 Beta 分布形状参数的均值alpha / (alpha beta)给出posterior_param[:, 0] / posterior_param.sum(axis1)正是这一公式。示例依次打印p_A、p_B True | A False/True、p_C True | B False/True三组真实 vs 预测对比用于直观验证枚举变分推断确实恢复了真实的 CPD。你可以将其与 tests/infer/test_enum.py 中大量参数化的枚举测试对照理解——该测试文件系统覆盖了config_enumerate的 sequential/parallel 配置、iter_discrete_traces的轨迹展开顺序如深度为 5 时恰好产生2**5条轨迹、guide 侧枚举与模型侧枚举的合法性约束以及TraceEnum_ELBO的梯度正确性。6. 运行方式与命令行参数示例通过argparse暴露两个可选参数见if __name__ __main__部分参数短选项默认值含义--num-steps-n4000变分推断迭代步数--num-obs-o10000独立观测数据条数运行命令在仓库根目录下python examples/toy_mixture_model_discrete_enumeration.py -n 4000 -o 10000脚本开头还带有一个版本断言assert pyro.__version__.startswith(1.9.1)说明该示例面向 Pyro 1.9.1 系列编写若你使用的是其他版本可依据本仓库源码尤其 pyro/infer/enum.py 与 pyro/infer/traceenum_elbo.py核对 API 差异。运行过程中会先打印带tqdm进度条的迭代过程结束时弹出损失曲线图并在终端输出三组 CPD 的对比结果。7. 从示例到实战离散枚举的适用场景与注意事项结合本例与源码实现可以总结出在 Pyro 中应用离散枚举的通用要点三步范式是固定的模型上加config_enumerate或对 guide 调用config_enumerate、目标离散站点标注infer{enumerate: parallel}、SVI 改用TraceEnum_ELBO。三者缺一不可没有前者站点不会获得枚举配置没有标注则EnumMessenger不会展开该站点不用TraceEnum_ELBO则普通 ELBO 无法解释枚举维度。枚举发生在模型侧还是 guide 侧若离散站点不出现在 guide 中则在模型侧枚举精确求和若出现在 guide 中则在 guide 侧枚举。TraceEnum_ELBO文档明确假设plate 外的变量永远不能依赖 plate 内的变量违反这一依赖结构会导致错误结果这也是max_plate_nesting必须正确的深层原因。离散支撑大小决定可行性枚举是对dist.enumerate_support()的穷举支撑集合呈指数增长iter_discrete_traces在深度为depth的模型上会产出2**depth条轨迹见测试 tests/infer/test_enum.py。因此它适合支撑较小的离散变量支撑很大时config_enumerate的num_samples参数可切换到 Tensor Monte CarloTMC局部采样tmcmixture或diagonal作为替代。Vindex是枚举场景的必备索引工具一旦采样值带上枚举批次维度普通tensor[idx]索引极易出错务必改用Vindex(tensor)[idx]。通过本示例你已经掌握了 Pyro 离散枚举的最小完整闭环数据生成 → 模型/guide 编写 →TraceEnum_ELBO训练 → CPD 评估。后续可将同一范式迁移到 HMM、主题模型、混合模型等更复杂的含隐变量场景相关扩展实现可继续参考仓库中 pyro/infer/enum.py、pyro/infer/traceenum_elbo.py 与 pyro/poutine/enum_messenger.py 的源码注释。赞分享人工智能机器学习深度学习概率编程【免费下载链接】pyroDeep universal probabilistic programming with Python and PyTorch项目地址https://gitcode.com/gh_mirrors/py/pyro点击查看免费下载相关推荐Allegro5性能优化秘籍内存管理、线程安全和渲染效率Allegro5性能优化秘籍内存管理、线程安全和渲染效率 Allegro5是一款功能强大的跨平台游戏编程库广泛应用于2D游戏开发。本文将分享Allegro5人工智能机器学习深度学习概率编程zls枚举类型完整的枚举和联合支持zls枚举类型完整的枚举和联合支持 还在为Zig语言中复杂的枚举类型编写而头疼zlsZig Language Server为你提供了完整的枚举和联合类型开发工具Pyro Funsor 后端下的隐马尔可夫模型精确枚举推断从 collapsed SVI 到 MAP Baum-Welch 实战指南Pyro Funsor 后端下的隐马尔可夫模型精确枚举推断从 collapsed SVI 到 MAP Baum Welch 实战指南 本篇技术指南围绕 Pyr人工智能机器学习深度学习概率编程上一篇ComfyUI-Manager终极指南5步解决InvalidChannel错误让AI绘画扩展管理回归正轨下一篇Windows Cleaner终极指南免费开源工具彻底解决C盘空间不足问题创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
