HyperAIHyperAI

Command Palette

Search for a command to run...

因果基础模型

Christopher Stith Hossein Rahmani Jesse C. Cresswell

摘要

因果推断是从数据中估计治疗或干预效果的一种实践。传统上,每个新问题都需要一个定制的流程:首先提出因果机制,选择兼容的估计器,最后进行训练。与此同时,在多样化的环境和模态中,机器学习的大部分已转向基础模型的范式:即在大规模上预训练一次,然后无需微调即可应用于新任务的网络。因果基础模型(CFMs)将这一范式引入因果推断。CFMs 是预训练的神经网络,能够通过上下文学习在全新数据集上估计因果量(如平均治疗效果),而无需更新模型。这项工作为该新兴领域提供了实用介绍。在讨论 CFMs 之前,我们总结了因果推断和机器学习所需的背景知识。全文包含示例代码和 Jupyter 笔记本,可通过点击图标访问。完整代码库可在 github.com/layer6ai-labs/cfms 获取。

一句话总结

Layer 6 AI 与 TD Bank Group 的研究人员提出了因果基础模型(CFMs\text{CFMs}CFMs),这是一种预训练神经网络,能够通过上下文学习在全新数据集上估计因果量(如平均处理效应\text{平均处理效应}平均处理效应),而无需更新模型。文章还提供了该新兴领域的实用入门介绍,包含示例代码和 Jupyter notebooks。

核心贡献

  • 本文提供了因果基础模型(CFMs)的实用、动手式入门介绍,包含示例代码和 Jupyter notebooks,完整代码库已发布于 github.com/layer6ai-labs/cfms。
  • 工作将 CFMs 形式化为预训练的基于 transformer 的模型,通过在新数据集上的上下文学习估计因果量(如条件平均处理效应、条件干预分布),无需微调,并将这一定义与 CaML、BBCI 和 CInA 等前期工作进行了对比定位。
  • 实验结果表明,当前的 CFMs 在因果基准测试上达到了具有竞争力的性能,同时与传统需要逐问题训练和调参的估计器相比,大幅缩短了部署时间。

引言

因果推断是回答经济学、医学和政策领域中干预性问题的核心,其目标是估计处理或政策的效果,而非单纯的关联。传统估计方法,如贝叶斯加性回归树、双重机器学习、因果森林等,对每个新问题都需要一套劳动密集型的流程:提出因果机制、选择估计器、调整超参数并进行训练,且任务之间无法复用。这限制了可扩展性和速度,尤其在数据规模大或多样性高的实际场景中更为突出。

为解决这些局限,作者引入了因果基础模型(CFMs),即预训练的基于 transformer 的神经网络,可直接应用于未见过的因果推断任务,无需进一步训练或微调。CFMs 在数据生成过程和因果机制的先验上进行训练,使其能够在推理时通过上下文学习执行摊销贝叶斯推断。作者提供了 CFMs 的实用、动手式入门介绍,包括示例代码和 Jupyter notebooks,并将首批公开可用的 CFMs 与传统估计器进行了基准对比。其主要贡献在于让 CFMs 能够被更广泛的受众使用,展示了这些模型在推理速度上的显著提升以及具有竞争力或更优的性能,同时回顾了该领域的最新发展和应用。

方法

因果基础模型(CFMs)利用先验数据拟合网络来摊销对因果量(如条件期望潜在结果(CEPO)或条件干预分布(CID))的贝叶斯推断。与传统因果估计器需要对每个新数据集进行迭代模型选择、超参数调优和训练不同,CFMs 将这一流程压缩为单次前向传播。如下图所示,CFM 工作流保持预训练参数固定,并将观测数据集作为上下文用于上下文学习,从而在推理时无需更新权重。

为训练 CFMs,作者依赖结构因果模型(SCMs)来定义数据生成过程。SCM 在贝叶斯网络上增加了显式的结构方程,将各节点联系起来。例如,考虑内生变量 {X,T,Y}\{X, T, Y\}{X,T,Y} 和外生变量 {U1,U2,U3}\{U_1, U_2, U_3\}{U1,U2,U3}。表示这些因果关系的对应有向无环图如下所示。

通过采样外生变量并沿结构方程传播,SCM 可以高效地模拟观测数据和干预数据。这一能力对于设计 DGP 上的合成先验至关重要,因为真实世界的观测数据缺乏训练所需的真实干预或反事实标签。

CFMs 利用基于 transformer 的架构实现上下文学习。模型接收表格形式的观测数据集 Dobs={(xn,tn,yn)}n=1N\mathcal{D}_{\mathrm{obs}} = \{(x_n, t_n, y_n)\}_{n=1}^NDobs={(xn,tn,yn)}n=1N 作为上下文,并将因果任务作为查询。在传入 transformer 之前,输入数据会经过分词和嵌入处理。上下文数据集的嵌入包含观测协变量、处理变量和事实结果,而查询嵌入不包含任何结果信息。transformer 采用掩码注意力机制,使每个 token 只能关注上下文,而不能关注查询。这确保了查询预测仅依赖于上下文数据。架构还考虑了处理变量的特殊角色,要么将其强制放在第一列,要么在处理变量和协变量分别经过独立编码器后再进行拼接。

CFMs 使用一种修改版的先验数据损失,即因果先验数据损失进行训练。对于 CEPO-PPD,损失形式如下:

Lt(θ)=Eψπ,Dobs{x}Pobsψ[log(qθ(μt(x;Pψ)x,t,Dobs))]\mathcal{L}_t(\theta) = \mathbb{E}_{\psi \sim \pi, \mathcal{D}_{\mathrm{obs}} \cup \{x\} \sim P_{\mathrm{obs}}^\psi} \big[ - \log (q_\theta(\mu_t(x; P^\psi) \mid x, t, \mathcal{D}_{\mathrm{obs}})) \big]Lt(θ)=Eψπ,Dobs{x}Pobsψ[log(qθ(μt(x;Pψ)x,t,Dobs))]

该损失评估模型对从采样 DGP ψ\psiψ 计算出的真实因果量所赋予的似然。关键在于,它不需要以封闭形式知道真实的底层 PPD,只要能够模拟干预数据即可计算。

CFMs 的训练涉及从合成因果先验中采样,并生成观测数据和干预数据。如下图所示,一个高层级的训练流程从先验中采样 SCM、生成观测数据集、模拟干预目标,并使用因果先验数据损失进行监督学习。

为生成所需的干预真实标签,模型对 SCM 进行干预模拟。例如,为了从条件干预分布 Pψ(do(T=t),X=x)P^\psi(\cdot \mid \mathrm{do}(T=t), X=x)Pψ(do(T=t),X=x) 生成数据,将 TTT 的结构方程替换为固定处理值 T=tT=t^*T=t。下图说明了这一过程,展示了完整的原始结构方程如何将外生噪声映射为观测数据,而干预则替换处理方程,从而生成具有真实潜在结果的干预目标。

通过这一过程,CFM 学会了近似因果 PPD,从而在推理时仅通过单次前向传播即可从纯观测数据中预测因果量。

实验

评估使用半合成的 RealCause-Lalonde 基准(Lalonde-CPS 和 Lalonde-PSID 队列)来比较三种因果基础模型(CFMs),即 Do-PFN、CausalPFN 和 CausalFM,与高度调优的经典估计器(如元学习器、IPW 和 DML)进行对比,指标涵盖 CATE 准确性(PEHE)、ATE 相对误差、运行时间和平均排名。尽管未在基准数据上训练,CFMs 仍能与经典基线竞争,其中 CausalPFN 在 CFMs 中取得了最低平均排名,并与调优后的 T-Learner 非常接近,同时由于摊销推断,其 CPU 运行速度快 1 到 2 个数量级。然而,Do-PFN 和 CausalFM 会系统性地将处理效应估计向零收缩,仅恢复真实 ATE 差异的一小部分,而 CausalPFN 则能紧密恢复总体效应。

早期的因果基础模型均在合成数据上预训练并使用上下文学习,但在预测目标、先验可识别性和架构上有所不同。使用可识别后门先验并预测条件期望潜在结果的模型在基准评估中表现更好,而预测完整干预分布或处理效应且使用不可识别先验的模型则表现出向零的系统性收缩。CausalPFN 使用可识别的后门先验并预测条件期望潜在结果,在因果基础模型中取得了最佳平均排名。Do-PFN 和 CausalFM 预测干预分布或处理效应且先验可识别性较弱,仅恢复真实处理效应的一小部分,在两个队列上的误差均接近 1。所有因果基础模型均显著快于传统估计器,CPU 运行时间低 1 到 2 个数量级,且在 GPU 上还有进一步加速。

因果基础模型(CFMs)基于 transformer,在规模和深度上有所不同,其中 CausalPFN 在基准任务中是 CFMs 中最具竞争力的。CFMs 相比传统估计器具有显著的运行时间优势,但其恢复平均处理效应量级的能力因模型而异。CausalPFN 拥有 2000 万参数和 20 层 transformer 层,而 Do-PFN 较小,拥有 730 万参数和 12 层。CausalPFN 在 CFMs 中取得了最低平均排名,紧随调优后的 T-Learner 基线。CFMs 在 CPU 上的运行速度比训练和调优传统模型快 1 到 2 个数量级,其中 CausalPFN 在 CPU 上最快。CausalPFN 紧密恢复了总体效应(在 Lalonde-CPS 上 ATE 相对误差为 0.17),而 Do-PFN 和 CausalFM 则表现出向零的系统性收缩。

因果基础模型(CFMs)在 RealCause-Lalonde 基准上与经典估计器具有竞争力,其中 CausalPFN 在 CFMs 中取得了最佳平均排名,并与调优后的 T-Learner 非常接近。CFMs 的运行时间比传统估计器快 1 到 2 个数量级,但只有 CausalPFN 能恢复真实的平均处理效应量级,其他 CFMs 则系统性地将估计向零收缩。CausalPFN 在 CFMs 中取得了最低平均排名,紧随调优后的 T-Learner,并在 Lalonde-CPS 上优于后者。CFMs 在 CPU 上的运行速度比传统估计器快 1 到 2 个数量级,其中 CausalPFN 最快。Do-PFN 和 CausalFM 在两个队列上的 ATE 相对误差均接近 1,表明预测效应向零的系统性收缩。IPW 虽不提供个体层面的估计,但在所有方法中取得了最强的 ATE 相对误差。输出替代后验表示的 CFMs 性能不如 CausalPFN,但仍与调优后的 X-Learner 和 S-Learner 具有竞争力。

评估在基准数据集上将因果基础模型(CFMs)与经典估计器进行对比,重点关注预测目标、先验可识别性和运行时间。CausalPFN 使用可识别的后门先验并预测条件期望潜在结果,在 CFMs 中取得了最佳平均排名,与调优后的 T-Learner 非常接近,同时恢复了真实的处理效应量级。相比之下,预测完整干预分布或处理效应且先验可识别性较弱的模型(如 Do-PFN 和 CausalFM)会系统性地将估计向零收缩,导致相对误差接近 1。所有 CFMs 在 CPU 上的运行速度比传统估计器快 1 到 2 个数量级,其中 CausalPFN 最快,但只有它能可靠地恢复总体效应。


用 AI 构建 AI

从创意到上线——通过免费 AI 协同编码、开箱即用的环境和最优惠的 GPU 价格,加速您的 AI 开发。

AI 协同编码
开箱即用的 GPU
最优定价

HyperAI Newsletters

订阅我们的最新资讯
我们会在北京时间 每周一的上午九点 向您的邮箱投递本周内的最新更新
邮件发送服务由 MailChimp 提供