- 人工智能
- 机器学习
- 深度学习
- 概率编程
【免费下载链接】pyro
Deep universal probabilistic programming with Python and PyTorch
本篇技术指南围绕 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 参考)明确指出其包含三大创新点:
- 带隐藏混杂因子的因果效应推断生成模型;
- 模型与指导使用"孪生神经网络"(twin neural nets),使
t=0与t=1两组条件分布参数不共享,从而支持高度不平衡的处理分配(imbalanced treatment); - 自定义训练损失,在标准 ELBO 之外加入额外项,使指导网络能够回答反事实(counterfactual)查询。
对外的主要接口是CEVAE类,同时暴露可定制的组件:Model、Guide、TraceCausalEffect_ELBO以及各类工具(FullyConnected、DistributionNet及其子类、PreWhitener等)。
CEVAE构造参数
| 参数 | 类型 | 默认值 | 含义 |
|---|---|---|---|
feature_dim | int | 必填 | 特征空间x的维度 |
outcome_dist | str | "bernoulli" | 结果分布类型,可选"bernoulli"、"exponential"、"laplace"、"normal"、"studentt" |
latent_dim | int | 20 | 潜变量z的维度 |
hidden_dim | int | 200 | 全连接网络隐藏层维度 |
num_layers | int | 3 | 全连接网络隐藏层层数 |
num_samples | int | 100 | ite()方法默认蒙特卡洛采样数 |
从源码可见,构造函数会逐一校验上述尺寸参数必须为正整数,否则抛出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 | -n | 50 | 训练轮数 |
--batch-size | -b | 100 | 批大小 |
--learning-rate | -lr | 1e-3 | 初始学习率 |
--learning-rate-decay | -lrd | 0.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)执行逻辑可以拆解为:
- 先用
guide(x)采样一批潜变量z(block隐藏掉y、t站点,仅保留z); - 对每个
z,用poutine.do将处理变量强制钉死为t=0或t=1,再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 | 单个logits(clamp 到 [-10, 10]) | 二元结果 |
exponential | ExponentialNet | rate(softplus 约束,reciprocal 得到 scale) | 非负连续结果 |
laplace | LaplaceNet | loc, scale | 拉普拉斯结果 |
normal | NormalNet | loc, scale | 高斯结果 |
studentt | StudentTNet | df, 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
相关推荐
Pyro因果推断实战指南:使用do-calculus精准分析干预效果
Pyro因果推断实战指南:使用do calculus精准分析干预效果 在数据科学和机器学习领域, 因果推断 正成为解决复杂问题的关键工具。Pyro作为基于PyT
人工智能机器学习深度学习概率编程AI_Tutorial因果推断应用:从理论到工业实践完整解析
AI_Tutorial因果推断应用:从理论到工业实践完整解析 因果推断作为人工智能领域的核心技术,正在工业界掀起一场革命。AI_Tutorial项目汇集了来自各
如何用AI让老旧视频重获新生?Video2X的3个神奇应用场景
如何用AI让老旧视频重获新生?Video2X的3个神奇应用场景 你正在寻找解决老旧视频画质模糊、帧率低下的方法吗?Video2X或许就是你需要的答案。这款开源工
音视频视频处理图像处理深度学习
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考