人工智能机器学习深度学习概率编程【免费下载链接】pyroDeep universal probabilistic programming with Python and PyTorch项目地址https://gitcode.com/gh_mirrors/py/pyro点击查看免费下载本篇技术指南围绕 Pyro 官方的 CEVAE 示例教程tutorial/source/cevae.rst 及其内嵌的 synthetic.py 完整示例展开讲解如何在 Pyro 中使用pyro.contrib.cevae.CEVAE进行存在隐藏混杂因子hidden confounder时的因果效应推断包括个体处理效应ITE与平均处理效应ATE的估计、反事实counterfactual查询的实现原理以及完整的训练、评估与 JIT 加速流程。读完本文你将掌握 CEVAE 的模型-指导Model/Guide架构、do算子驱动的反事实推断写法、TraceCausalEffect_ELBO目标函数并能直接复现仓库中的端到端示例。背景为什么因果效应推断需要深度潜变量模型在经典的随机对照试验RCT中处理变量t与潜在特征相互独立直接比较处理组与对照组的平均结果即可得到无偏的 ATE。但在观察性研究中处理分配往往受未观测的混杂因子影响——例如某个病人的病情严重程度未观测同时决定了其是否接受治疗以及预后结果此时简单的分组均值之差naive ATE是有偏的。CEVAECausal Effect Variational Autoencoder正是为解决这一问题设计的生成式模型。它假定存在一个隐藏的混杂因子Z并假设数据由如下图模型生成Z → X X 是 Z 的带噪声部分观测 Z → t 处理分配受 Z 影响即存在混杂 Z → y Z 直接影响结果 t → y 处理直接影响结果其中t是二元处理变量如用药与否y是结果如康复与否Z是未观测混杂因子X是Z的带噪函数如病历中的可观测特征。该图模型直接对应 pyro/contrib/cevae/init.py 中CEVAE类的文档字符串其核心思想是利用变分推断学出Z的后验分布再通过 do-操作在潜变量空间中人为指定处理变量取值从而剥离混杂得到因果效应。认识pyro.contrib.cevae模块CEVAE 的实现位于 pyro/contrib/cevae/init.py模块文档也可在 docs/source/contrib.cevae.rst 查看 API 参考明确指出其包含三大创新点带隐藏混杂因子的因果效应推断生成模型模型与指导使用孪生神经网络twin neural nets使t0与t1两组条件分布参数不共享从而支持高度不平衡的处理分配imbalanced treatment自定义训练损失在标准 ELBO 之外加入额外项使指导网络能够回答反事实counterfactual查询。对外的主要接口是CEVAE类同时暴露可定制的组件Model、Guide、TraceCausalEffect_ELBO以及各类工具FullyConnected、DistributionNet及其子类、PreWhitener等。CEVAE构造参数参数类型默认值含义feature_dimint必填特征空间x的维度outcome_diststrbernoulli结果分布类型可选bernoulli、exponential、laplace、normal、studenttlatent_dimint20潜变量z的维度hidden_dimint200全连接网络隐藏层维度num_layersint3全连接网络隐藏层层数num_samplesint100ite()方法默认蒙特卡洛采样数从源码可见构造函数会逐一校验上述尺寸参数必须为正整数否则抛出ValueError随后构造Model(config)与Guide(config)两个PyroModule并持有。三步走的使用范式源码 docstring 给出了最精炼的使用范式cevae CEVAE(feature_dim5) cevae.fit(x_train, t_train, y_train) ite cevae.ite(x_test) # individual treatment effect ate ite.mean() # average treatment effect即构造 → 训练 → 推断。ite()返回长度为len(x_test)的个体效应向量对其求均值即得到 ATE。示例全景synthetic.py 的完整工作流教程页 tutorial/source/cevae.rst 的正文即完整内嵌了示例脚本 examples/contrib/cevae/synthetic.py该脚本同时被 tutorial/source/index.rst 收录在 Deep Generative Models 教程目录下。脚本参考了 Louizos 等 2017 年的论文Causal Effect Inference with Deep Latent-Variable Models但将原论文假设的feature_dim1、latent_dim5扩大为更一般的规模。整条流水线分为四个阶段数据生成、训练、评估、JIT 加速。1. 命令行参数一览脚本通过argparse暴露全部超参数默认值如下参数简写默认值含义--num-data—1000样本数量--feature-dim—5特征维度--latent-dim—20潜变量维度--hidden-dim—200隐藏层维度--num-layers—3隐藏层层数--num-epochs-n50训练轮数--batch-size-b100批大小--learning-rate-lr1e-3初始学习率--learning-rate-decay-lrd0.1学习率衰减系数末期学习率 初始学习率 × 该值--weight-decay—1e-4权重衰减--seed—1234567890随机种子--jit—False训练后用 TorchScript 编译--cuda—False使用 CUDA等价于torch.set_default_device(cuda)在仓库根目录下按如下方式运行示例脚本入口为 examples/contrib/cevae/synthetic.pypython examples/contrib/cevae/synthetic.py # 默认配置 python examples/contrib/cevae/synthetic.py -n 100 -b 200 -lr 5e-4 # 自定义训练超参 python examples/contrib/cevae/synthetic.py --jit --cuda # JIT 编译 GPU脚本开头还会断言pyro.__version__以1.9.1开头并对pyro的 logger 开启DEBUG级输出便于观察每个 minibatch 的 loss。2. 合成数据生成制造有混杂的数据generate_data(args)复现了论文 [1] 的生成过程用 Pyro 概率编程原语直接采样z dist.Bernoulli(0.5).sample([args.num_data]) # 隐藏混杂因子二元 x dist.Normal(z, 5 * z 3 * (1 - z)).sample([args.feature_dim]).t() t dist.Bernoulli(0.75 * z 0.25 * (1 - z)).sample() # 处理分配与 z 相关 → 混杂 y dist.Bernoulli(logits3 * (z 2 * (2 * t - 2))).sample()这里的关键设计是t的伯努利概率0.75*z 0.25*(1-z)依赖隐藏因子z意味着z同时驱动了处理分配与结果生成这正是隐藏混杂的数据学体现——仅凭x,t,y无法直接给出无偏的效应估计。此外样本张量x的形状为(num_data, feature_dim)每个特征维度独立采样后再转置拼接。3. 真值 ITE 的蒙特卡洛近似由于是合成数据可以对照真实因果效应评估模型。源码利用z的真实取值直接计算反事实期望之差t0_t1 torch.tensor([[0.0], [1.0]]) y_t0, y_t1 dist.Bernoulli(logits3 * (z 2 * (2 * t0_t1 - 2))).mean true_ite y_t1 - y_t0即对每个个体分别用t0与t1代入结果分布p(y|z,t)求期望二者之差即为该个体的真实 ITE对全体样本求均值得到真实 ATE。后续输出中true ATE就是以此为基准的。4. 训练训练前先固定随机种子并清空参数存储pyro.set_rng_seed(args.seed) pyro.clear_param_store() cevae CEVAE( feature_dimargs.feature_dim, latent_dimargs.latent_dim, hidden_dimargs.hidden_dim, num_layersargs.num_layers, num_samples10, # 注意示例中 ite 采样数取 10而非默认 100 ) cevae.fit( x_train, t_train, y_train, num_epochsargs.num_epochs, batch_sizeargs.batch_size, learning_rateargs.learning_rate, learning_rate_decayargs.learning_rate_decay, weight_decayargs.weight_decay, )示例特意将num_samples10在保证评估精度的同时控制反事实推断的蒙特卡洛开销。5. 评估三种 ATE 对比评估阶段重新生成一批测试数据并同时打印三条基准线true_ate true_ite.mean() # 真实 ATE用 z 计算 naive_ate y_test[t_test 1].mean() - y_test[t_test 0].mean() # 朴素 ATE分组均值差 est_ite cevae.ite(x_test) # CEVAE 估计的 ITE est_ate est_ite.mean() # CEVAE 估计的 ATE其中naive ATE是忽略混杂、直接比较处理组/对照组均值的结果通常与真实值存在偏差estimated ATE是 CEVAE 通过潜变量反事实推断得到的结果。若--jit开启则先用cevae.to_script_module()编译再调用ite()。深入源码一生成模型与指导网络的架构CEVAE由两个 PyroModule 组成Model生成模型与Guide推断模型。Model因果生成过程Model.forward严格按图模型采样并全部置于pyro.plate(data, size, subsamplex)中以便小批量训练z pyro.sample(z, self.z_dist()) # z ~ N(0,I)对角标准正态 x pyro.sample(x, self.x_dist(z), obsx) # x ~ p(x|z)神经网络输出对角高斯 t pyro.sample(t, self.t_dist(z), obst) # t ~ Bernoulli(logitsf(z)) y pyro.sample(y, self.y_dist(t, z), obsy) # y ~ p(y|t,z)三个条件分布都由神经网络参数化x_dist(z)x_nn是一个DiagNormalNet其网络维度为[latent_dim] [hidden_dim]*num_layers [feature_dim]输出loc, scale构造to_event(1)的对角高斯t_dist(z)BernoulliNet将z映射为单个logitsy_dist(t, z)核心设计——结果网络被拆成y0_nn与y1_nn两个独立网络分别建模p(y|t0,z)与p(y|t1,z)前向时用torch.where(t, p1, p0)按处理取值拼接参数。源码注释明确说明Parameters are not shared among t values这一孪生网络结构正是为支持高度不平衡的处理分配而设计。Guide反事实推断的变分近似Guide.forward定义了与生成过程对应的推断网络采样顺序为t pyro.sample(t, self.t_dist(x), obst, infer{is_auxiliary: True}) y pyro.sample(y, self.y_dist(t, x), obsy, infer{is_auxiliary: True}) pyro.sample(z, self.z_dist(y, t, x)) # z ~ q(z|y,t,x)作为嵌入这里t、y两个站点被标记为is_auxiliary辅助站点——源码注释指出它们仅用于预测并参与 CEVAE 的辅助损失不参与标准 ELBO 中潜变量的推断只有z站点走常规 ELBO。Guide 同样采用共享前几层 按 t 分裂最后一层的孪生结构y_nn/z_nn先提取共享隐藏表示再由y0_nn/y1_nn、z0_nn/z1_nn分别输出两组参数最终z_dist构造dist.Normal(loc, scale).to_event(1)的对角高斯后验。深入源码二TraceCausalEffect_ELBO特殊目标函数CEVAE 的训练不使用标准Trace_ELBO而是其子类TraceCausalEffect_ELBO。源码 docstring 给出了目标函数最大化形式-loss ELBO log q(t|x) log q(y|t,x)实现上_differentiable_loss_particle首先构造标准-ELBO找出 Guide 轨迹中所有被观测的站点即辅助站点t、y将它们从复制后的 guide trace 中剔除后再调用父类计算loss, surrogate_loss随后把被剔除站点的log_prob_sum以负号追加进损失即加上log q(t|x) log q(y|t,x)两项。loss()方法再用torch_item去掉梯度信息返回标量。换言之Guide 不仅要像普通变分推断那样逼近p(z|·)还要学会直接预测t与y这正是后续反事实查询能力的基础。反事实推断的底层机制do 算子 轨迹重放ite(x)方法在 pyro/contrib/cevae/init.py 中按如下公式估计个体处理效应ITE(x) E[ y | Xx, do(t1) ] − E[ y | Xx, do(t0) ]其实现用到了 Pyro 的poutine.do、poutine.replay与poutine.trace三件套with pyro.plate(num_particles, num_samples, dim-2): with poutine.trace() as tr, poutine.block(hide[y, t]): self.guide(x) # 采样 z ~ q(z|y,t,x)但隐藏 y、t 站点 with poutine.do(datadict(ttorch.zeros(()))): y0 poutine.replay(self.model.y_mean, tr.trace)(x) # do(t0) 下的期望结果 with poutine.do(datadict(ttorch.ones(()))): y1 poutine.replay(self.model.y_mean, tr.trace)(x) # do(t1) 下的期望结果 ite (y1 - y0).mean(0)执行逻辑可以拆解为先用guide(x)采样一批潜变量zblock隐藏掉y、t站点仅保留z对每个z用poutine.do将处理变量强制钉死为t0或t1再replay到model.y_mean得到反事实期望E[y | z, do(t·)]在num_samples个粒子维度上取均值得到每个个体的 ITE。由于对每个样本都要做num_samples次采样、且结果期望按num_samples²的组合方式求平均源码注释标明其复杂度为O(len(x) * num_samples ** 2)。此外ite()内部会先做数据白化PreWhitener按训练集的均值/标准差做标准化num_samples与batch_size均可通过参数覆盖默认值。结果分布扩展outcome_dist与 DistributionNet 家族CEVAE 并不局限于伯努利结果。pyro.contrib.cevae通过DistributionNet抽象出输出某类分布参数 构造分布的统一接口Model/Guide在初始化时按config[outcome_dist]字符串动态查表选择子类DistributionNet.get_class。目前已支持五种结果分布outcome_dist网络类输出的分布参数说明bernoulli默认BernoulliNet单个logitsclamp 到 [-10, 10]二元结果exponentialExponentialNetratesoftplus 约束reciprocal 得到 scale非负连续结果laplaceLaplaceNetloc, scale拉普拉斯结果normalNormalNetloc, scale高斯结果studenttStudentTNetdf, loc, scale共享df 1厚尾结果所有网络的参数层都由FullyConnected带 ELU 激活的多层感知机搭建并对loc、scale做保守的 clamp例如NormalNet将scale约束在[1e-3, 1e6]。需要说明的是ExponentialNet命名沿袭自实现中对尺度参数的 softplus 处理实际返回的是rate 1/scale并最终以dist.Exponential(rate)构造分布。测试 tests/contrib/cevae/test_cevae.py 中会遍历DistributionNet.__subclasses__()自动覆盖全部五种分布做冒烟测试其中exponential结果在喂入模型前会clamp_(min1e-20)以保证正值。fit()训练接口详解CEVAE.fit(x, t, y, ...)是端到端的训练入口签名如下含默认值fit(x, t, y, num_epochs100, batch_size100, learning_rate1e-3, learning_rate_decay0.1, weight_decay1e-4, log_every100)其内部流程与几个值得注意的实现细节输入校验断言x为 2D 且x.size(-1) feature_dim、t.shape x.shape[:1]、y形状与其自身第一维一致数据白化用PreWhitener(x)基于训练集统计量构建标准化器训练与ite()推断都经过它DataLoader以TensorDataset(x, t, y)shuffleTrue构造批次generator与x.device对齐以支持 GPU学习率调度优化器使用ClippedAdamPyro 提供的梯度裁剪版 Adam学习率衰减按lrd learning_rate_decay ** (1 / num_steps)计算——源码注释保证初始学习率为learning_rate、末期学习率收敛到learning_rate * learning_rate_decay衰减粒度取决于批数与轮数的乘积num_stepsSVI 驱动SVI(self.model, self.guide, optim, TraceCausalEffect_ELBO())每个 step 的 loss 除以全量样本数log_every控制每多少步输出一次 DEBUG 日志并断言 loss 无 NaN返回值返回每个 epoch 的 loss 列表。JIT 编译与模型序列化--jit选项走的是to_script_module()方法先将模块切到eval模式用torch.randn(2, feature_dim)伪造输入通过torch.jit.trace_module(self, {ite: (fake_x,)}, check_traceFalse)将ite方法编译为 TorchScript。注意两处关键处理一是关闭pyro.validation_enabled(False)二是check_traceFalse——源码注释明确解释这是因为 CEVAE 内部存在非确定性节点蒙特卡洛采样无法通过严格的 trace 一致性检查。序列化能力在 tests/contrib/cevae/test_cevae.py 的test_serialization中得到验证分别对纯 Python 版本torch.save/torch.load与 JIT 版本torch.jit.save/torch.jit.load保存再加载固定随机种子后比较ite(x)输出断言与原始结果在atol0.1内一致。测试中还标注了已知问题torch 2.x下 JIT 路径存在上游 issue会xfail。正确性验证测试套件如何保障仓库用两类测试为 CEVAE 背书test_smoke对num_data ∈ {1, 100, 200}、feature_dim ∈ {1, 2}、全部五种outcome_dist的组合做冒烟测试验证fit(x, t, y, num_epochs2)后ite(x)形状为(num_data,)——即使是单样本也要求反事实推断链路完整可跑test_serialization如前所述验证 Python/JIT 两条序列化路径下推断结果的一致性。这两个测试tests/contrib/cevae/test_cevae.py与 API 文档页 docs/source/contrib.cevae.rst 一起构成了除示例脚本之外的完整参考闭环。运行环境与版本注意事项版本断言示例要求pyro.__version__以1.9.1开头运行前请确认环境中的 Pyro 版本匹配GPU 支持--cuda通过torch.set_default_device(cuda)设置默认设备DataLoader 的generator亦与设备对齐测试test_cuda.py体系对其它分布类目有覆盖CEVAE 的 CUDA 路径同样依赖 PyTorch 默认设备机制随机性训练与推断分别设置pyro.set_rng_seed(args.seed)评估真实 ITE 时用的是测试集新采样的z与训练数据无关依赖核心依赖为 PyTorch 与 Pyro含pyro.contrib自动注册的DistributionNet子类体系示例仅需标准库argparse、logging与torch。小结CEVAE 展示了概率编程语言在因果推断上的独特优势模型即代码、干预即变换。借助 Pyro 的poutine.do与poutine.replay原本需要专门实现的反事实推断被压缩为几十行可读代码而孪生神经网络与TraceCausalEffect_ELBO的配合则让模型在隐藏混杂存在时依然能输出接近真实值的 ATE。建议读者沿着本文脉络依次阅读 示例脚本 → 模块实现 → 测试用例并在自己的数据上从默认参数出发逐步调整latent_dim、hidden_dim、num_samples与outcome_dist以获得与业务场景匹配的因果效应估计。参考文献C. Louizos, U. Shalit, J. Mooij, D. Sontag, R. Zemel, M. Welling (2017).Causal Effect Inference with Deep Latent-Variable Models.该论文即示例与模块 docstring 中引用的 [1]其开源参考实现也是本模块的设计来源。赞分享人工智能机器学习深度学习概率编程【免费下载链接】pyroDeep universal probabilistic programming with Python and PyTorch项目地址https://gitcode.com/gh_mirrors/py/pyro点击查看免费下载相关推荐Pyro因果推断实战指南使用do-calculus精准分析干预效果Pyro因果推断实战指南使用do calculus精准分析干预效果 在数据科学和机器学习领域 因果推断 正成为解决复杂问题的关键工具。Pyro作为基于PyT人工智能机器学习深度学习概率编程AI_Tutorial因果推断应用从理论到工业实践完整解析AI_Tutorial因果推断应用从理论到工业实践完整解析 因果推断作为人工智能领域的核心技术正在工业界掀起一场革命。AI_Tutorial项目汇集了来自各如何用AI让老旧视频重获新生Video2X的3个神奇应用场景如何用AI让老旧视频重获新生Video2X的3个神奇应用场景 你正在寻找解决老旧视频画质模糊、帧率低下的方法吗Video2X或许就是你需要的答案。这款开源工音视频视频处理图像处理深度学习上一篇MCP协议标准化进程Awesome MCP Servers在行业中的影响力下一篇Extism运行时完整指南解锁WebAssembly执行引擎的强大功能创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考