news 2026/9/25 2:43:46

Pyro 因果效应推断实战:CEVAE(Causal Effect VAE)从理论到代码

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Pyro 因果效应推断实战:CEVAE(Causal Effect VAE)从理论到代码
  • 人工智能
  • 机器学习
  • 深度学习
  • 概率编程

【免费下载链接】pyro

Deep 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)是有偏的。

CEVAE(Causal 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 参考)明确指出其包含三大创新点:

  1. 带隐藏混杂因子的因果效应推断生成模型;
  2. 模型与指导使用"孪生神经网络"(twin neural nets),使t=0与t=1两组条件分布参数不共享,从而支持高度不平衡的处理分配(imbalanced treatment);
  3. 自定义训练损失,在标准 ELBO 之外加入额外项,使指导网络能够回答反事实(counterfactual)查询。

对外的主要接口是CEVAE类,同时暴露可定制的组件:Model、Guide、TraceCausalEffect_ELBO以及各类工具(FullyConnected、DistributionNet及其子类、PreWhitener等)。

CEVAE构造参数

参数类型默认值含义
feature_dimint必填特征空间x的维度
outcome_diststr"bernoulli"结果分布类型,可选"bernoulli"、"exponential"、"laplace"、"normal"、"studentt"
latent_dimint20潜变量z的维度
hidden_dimint200全连接网络隐藏层维度
num_layersint3全连接网络隐藏层层数
num_samplesint100ite()方法默认蒙特卡洛采样数

从源码可见,构造函数会逐一校验上述尺寸参数必须为正整数,否则抛出ValueError;随后构造Model(config)与Guide(config)两个PyroModule并持有。

三步走的使用范式

源码 docstring 给出了最精炼的使用范式:

cevae = CEVAE(feature_dim=5) 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_dim=1、latent_dim=5扩大为更一般的规模。整条流水线分为四个阶段:数据生成、训练、评估、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.py):

python 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(logits=3 * (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(logits=3 * (z + 2 * (2 * t0_t1 - 2))).mean true_ite = y_t1 - y_t0

即对每个个体,分别用t=0与t=1代入结果分布p(y|z,t)求期望,二者之差即为该个体的真实 ITE;对全体样本求均值得到真实 ATE。后续输出中,true ATE就是以此为基准的。

4. 训练

训练前先固定随机种子并清空参数存储:

pyro.set_rng_seed(args.seed) pyro.clear_param_store() cevae = CEVAE( feature_dim=args.feature_dim, latent_dim=args.latent_dim, hidden_dim=args.hidden_dim, num_layers=args.num_layers, num_samples=10, # 注意:示例中 ite 采样数取 10(而非默认 100) ) cevae.fit( x_train, t_train, y_train, num_epochs=args.num_epochs, batch_size=args.batch_size, learning_rate=args.learning_rate, learning_rate_decay=args.learning_rate_decay, weight_decay=args.weight_decay, )

示例特意将num_samples=10,在保证评估精度的同时控制反事实推断的蒙特卡洛开销。

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, subsample=x)中以便小批量训练:

z = pyro.sample("z", self.z_dist()) # z ~ N(0,I),对角标准正态 x = pyro.sample("x", self.x_dist(z), obs=x) # x ~ p(x|z),神经网络输出对角高斯 t = pyro.sample("t", self.t_dist(z), obs=t) # t ~ Bernoulli(logits=f(z)) y = pyro.sample("y", self.y_dist(t, z), obs=y) # 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映射为单个logits;
  • y_dist(t, z):核心设计——结果网络被拆成y0_nn与y1_nn两个独立网络,分别建模p(y|t=0,z)与p(y|t=1,z),前向时用torch.where(t, p1, p0)按处理取值拼接参数。源码注释明确说明:"Parameters are not shared among t values",这一"孪生网络"结构正是为支持高度不平衡的处理分配而设计。

Guide:反事实推断的变分近似

Guide.forward定义了与生成过程对应的推断网络,采样顺序为:

t = pyro.sample("t", self.t_dist(x), obs=t, infer={"is_auxiliary": True}) y = pyro.sample("y", self.y_dist(t, x), obs=y, 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 | X=x, do(t=1) ] − E[ y | X=x, do(t=0) ]

其实现用到了 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(data=dict(t=torch.zeros(()))): y0 = poutine.replay(self.model.y_mean, tr.trace)(x) # do(t=0) 下的期望结果 with poutine.do(data=dict(t=torch.ones(()))): y1 = poutine.replay(self.model.y_mean, tr.trace)(x) # do(t=1) 下的期望结果 ite = (y1 - y0).mean(0)

执行逻辑可以拆解为:

  1. 先用guide(x)采样一批潜变量z(block隐藏掉y、t站点,仅保留z);
  2. 对每个z,用poutine.do将处理变量强制钉死为t=0或t=1,再replay到model.y_mean,得到反事实期望E[y | z, do(t=·)];
  3. 在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单个logits(clamp 到 [-10, 10])二元结果
exponentialExponentialNetrate(softplus 约束,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_(min=1e-20)以保证正值。

fit()训练接口详解

CEVAE.fit(x, t, y, ...)是端到端的训练入口,签名如下(含默认值):

fit(x, t, y, num_epochs=100, batch_size=100, learning_rate=1e-3, learning_rate_decay=0.1, weight_decay=1e-4, log_every=100)

其内部流程与几个值得注意的实现细节:

  • 输入校验:断言x为 2D 且x.size(-1) == feature_dim、t.shape == x.shape[:1]、y形状与其自身第一维一致;
  • 数据白化:用PreWhitener(x)基于训练集统计量构建标准化器,训练与ite()推断都经过它;
  • DataLoader:以TensorDataset(x, t, y)+shuffle=True构造批次,generator与x.device对齐以支持 GPU;
  • 学习率调度:优化器使用ClippedAdam(Pyro 提供的梯度裁剪版 Adam),学习率衰减按lrd = learning_rate_decay ** (1 / num_steps)计算——源码注释保证"初始学习率为learning_rate、末期学习率收敛到learning_rate * learning_rate_decay",衰减粒度取决于批数与轮数的乘积num_steps;
  • SVI 驱动: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_trace=False)将ite方法编译为 TorchScript。注意两处关键处理:一是关闭pyro.validation_enabled(False),二是check_trace=False——源码注释明确解释这是因为 CEVAE 内部存在非确定性节点(蒙特卡洛采样),无法通过严格的 trace 一致性检查。

序列化能力在 tests/contrib/cevae/test_cevae.py 的test_serialization中得到验证:分别对纯 Python 版本(torch.save/torch.load)与 JIT 版本(torch.jit.save/torch.jit.load)保存再加载,固定随机种子后比较ite(x)输出,断言与原始结果在atol=0.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_epochs=2)后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],其开源参考实现也是本模块的设计来源。)

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

【免费下载链接】pyro

Deep universal probabilistic programming with Python and PyTorch

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

相关推荐

上一篇:MCP协议标准化进程:Awesome MCP Servers在行业中的影响力
下一篇:Extism运行时完整指南:解锁WebAssembly执行引擎的强大功能

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/25 2:37:40

MySQL后台注入靶场实战:从环境搭建到提权完整链路

简介:这份资源是一套存在SQL注入漏洞的网站源码,面向正在学习Web安全、需要动手复现注入攻击的初学者与进阶者,可用于本地或空间搭建靶场环境,练习后台注入的探测与利用思路。压缩包共844个文件,约4.95MB,以…

作者头像 李华
网站建设 2026/9/25 2:36:11

英伟达暑期实习笔试样题解析:GPU体系结构与深度学习考点

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/25 2:35:32

华为没有像 SAP 那样公开“MetaERP 采购模块白皮书”,所以下面这套分析是基于华为 MetaERP 公开架构表述 + 高端 ERP 采购到付款(P2P / Procure-to-Pay)通用范

华为没有像 SAP 那样公开“MetaERP 采购模块白皮书”,所以下面这套分析是基于华为 MetaERP 公开架构表述 高端 ERP 采购到付款(P2P / Procure-to-Pay)通用范式 华为“阳光采购 / 业财一体 / 元数据驱动”实践反推出来的工程化解读&#xff…

作者头像 李华