Barycentric Fused Gromov-Wasserstein Balancing for Causal Inference under Multiple Treatments¶
作者: Yuki Murakami, Takumi Hattori, Kohsuke Kubota
主题: 因果推断
相关性: 7/10
链接: https://arxiv.org/abs/2608.22024
一、领域脉络与小综述¶
这个方向是什么¶
本文所处的子方向是基于表示平衡的因果推断,具体针对多处理(multiple treatments)设定下的异质性处理效应估计。根本问题:从观测数据中估计多个同时施加的二元处理(K个)的单个效应(CASE)和交互效应(CAIE),同时缓解因处理分配非随机导致的选择偏差。当前成熟度:单处理设定下表示平衡方法已较成熟(如CFR、TARNet),但多处理设定下仍面临全局对齐困难、局部结构保持不足和计算复杂度高三大瓶颈。
发展脉络(history)¶
- 奠基工作:Shalit, Johansson, and Sontag (2017) 提出CFR,用Wasserstein距离平衡处理组与对照组的表示分布,给出泛化误差界。这是单处理表示平衡的经典框架。留下的口子:仅适用于二元处理,无法直接处理多处理。
- 主要进展(多处理扩展):
- Saini et al. (2019) 提出TECE-VAE,用任务嵌入(task embedding)和变分自编码器处理多处理,但依赖生成模型假设,鲁棒性有限(Rissanen and Marttinen 2021指出其一致性存疑)。
- Parbhoo, Bauer, and Schwab (2021) 提出NCoRE,用单独的结果网络建模处理组合,但参数不共享,对稀有处理模式估计不稳定(Chu et al. 2022)。
- Murakami, Hattori, and Kubota (2025) 提出CISI-Net,引入任务嵌入网络和成对平衡(pairwise balancing):对每对处理模式的表示分布计算Wasserstein距离并最小化。这是本文的直接前身。留下的口子:成对平衡的计算复杂度为O(L²)(L=2^K),且一对的平衡可能恶化另一对,导致残留不平衡(Gong, Nie, and Xu 2022; Lian et al. 2021);同时成对目标不保证跨所有处理模式的局部几何一致性(Liu, Shao, and Fu 2016; Lu, Xing, and Chen 2025)。
- 当前frontier:如何实现可扩展的全局对齐同时保持局部邻近结构。本文提出BFG-WB,用Wasserstein重心作为共享锚点,将复杂度降至O(L),并用Fused Gromov-Wasserstein散度同时对齐特征和几何结构。
- 本文的位置:作者声称这是第一个将Wasserstein重心与FGW散度结合用于多处理因果推断的工作,直接解决成对平衡的三大局限。
子线索聚类¶
- 表示平衡方法(核心线索):通过最小化处理组间分布差异来缓解选择偏差。包括CFR(Wasserstein)、ACE(局部相似性保持,Yao et al. 2019)、PITE(多原型对齐,Cao, Zhang, and Li 2026)、Proximity Matters(Wang et al. 2025)。本文属于此线索,但将平衡从成对升级为重心对齐。
- 多处理深度学习方法:架构驱动的方法,包括TECE-VAE(生成模型)、NCoRE(单独网络)、CISI-Net(任务嵌入+成对平衡)、Dosage Combination Network(连续剂量,Schweisthal et al. 2023)。本文属于此线索,但用BFG-WB替代成对平衡。
- 最优传输与重心:Wasserstein重心(Agueh and Carlier 2011; Cuturi and Doucet 2014)用于多分布对齐,FGW距离(Titouan et al. 2019; Vayer et al. 2020)用于结构化数据。本文将这些工具引入因果推断。
这个方向在追问的核心问题¶
- 如何实现多处理下的全局分布对齐? 成对平衡存在优化冲突,重心对齐能否提供更一致的参考?
- 如何在平衡的同时保持局部邻近结构? Wasserstein距离只关心特征值,不关心几何关系,导致结构扭曲(如图1所示)。FGW能否弥补?
- 如何降低计算复杂度? 成对平衡O(L²)不可扩展,重心对齐O(L)是否可行且有效?
- 理论保证: 表示平衡的误差界能否推广到多处理下的CASE和CAIE?现有理论(Shalit et al. 2017)只针对单处理。
⚠️ 作者的framing¶
作者将缺口frame为:成对平衡的三大局限(全局对齐冲突、局部结构不一致、二次复杂度)→ 因此需要重心对齐+FGW。竞争路线(TECE-VAE、NCoRE)被淡化,理由是生成模型假设脆弱、单独网络参数不共享。明显该被引/该存在、却没出现在intro里:没有引用任何关于Wasserstein重心统计性质的工作(如重心估计的收敛率、重心与原始分布的距离界),也没有引用多处理因果推断的识别理论(如Egami and Imai 2019的交互效应定义虽被引用,但未讨论更一般的识别条件)。此外,没有引用任何关于深度网络泛化误差与表示平衡结合的理论工作(如Johansson et al. 2018的泛化界),尽管本文定理依赖于表示空间的Lipschitz假设。这些缺失可能意味着理论部分仍有gap。
张力¶
未见明显对立引用。所有被引工作基本一致认为成对平衡有局限,重心对齐是合理改进方向。
二、最核心、最简单的例子 / 数学问题¶
第一步:符号、模型、可观测数据交代清楚¶
符号: - \(K\):二元处理的数量(本文实验中 \(K=3\))。 - \(\mathcal{T} = \{0,1\}^K\):所有可能的处理模式(treatment pattern),共 \(L = 2^K\) 种。 - \(X \in \mathbb{R}^d\):协变量向量(可观测)。 - \(T \in \mathcal{T}\):处理分配向量(可观测)。 - \(Y(t)\):潜在结果(potential outcome),对每个 \(t \in \mathcal{T}\) 存在,但只有事实结果 \(Y = Y(T)\) 可观测。 - \(\mu(x,t) = \mathbb{E}[Y(t) \mid X=x]\):条件期望潜在结果(待估的nuisance函数)。 - \(\tau_{\text{CASE}}(k,x) = \mu(x, t^{+k}) - \mu(x, \mathbf{0})\):处理 \(k\) 的单个效应(CASE),其中 \(t^{+k}\) 是第 \(k\) 个分量为1其余为0的向量,\(\mathbf{0}\) 是全零向量(对照组)。 - \(\tau_{\text{CAIE}}(S,x) = \sum_{Q \subseteq S} (-1)^{|S|-|Q|} \mu(x, t^{(+Q)})\):处理子集 \(S\) 的交互效应(CAIE),其中 \(t^{(+Q)}\) 是 \(Q\) 中分量为1其余为0的向量。 - \(\phi: \mathbb{R}^d \to \mathbb{R}^{d_r}\):表示学习网络,将协变量映射到表示空间。 - \(R_t\):处理模式 \(t\) 下的表示分布(即 \(\phi(X) \mid T=t\) 的分布)。 - \(R_b^*\):Wasserstein重心,定义为 \(\arg\min_{R_b} \sum_{t} \lambda_t W_2^2(R_t, R_b)\)。 - \(F(R_t, R_b^*)\):Fused Gromov-Wasserstein散度,包含特征项(Wasserstein)和结构项(Gromov-Wasserstein)。 - \(L_\phi = \sum_t w_t F(R_t, R_b^*)\):BFG-WB正则项。 - \(L_y\):逆频率加权的事实预测损失。 - \(\epsilon_{\text{CASE}}(k), \epsilon_{\text{CAIE}}(S)\):估计误差的积分平方误差(类似PEHE)。
模型: - 数据生成机制:\((X_i, T_i, Y_i)\) i.i.d. 来自联合分布,满足标准因果假设(SUTVA、Ignorability、Overlap)。 - 目标:从观测数据估计 \(\tau_{\text{CASE}}\) 和 \(\tau_{\text{CAIE}}\)。 - 方法:学习表示 \(\phi(X)\),使得不同处理模式下的表示分布对齐,同时保持局部几何结构,然后用共享的结果网络 \(h([\phi(x), t_w(t)])\) 预测 \(\mu(x,t)\)。
可观测数据: - 可观测:\((x_i, t_i, y_i)\),其中 \(t_i \in \{0,1\}^K\),\(y_i\) 是事实结果。 - 不可观测:所有反事实结果 \(Y_i(t)\) for \(t \neq t_i\)。 - 关键:识别依赖于假设(Ignorability + Overlap),使得 \(\mu(x,t) = \mathbb{E}[Y \mid X=x, T=t]\)。
第二步:最小内核¶
最简特例:\(K=2\)(两个二元处理),则 \(L=4\) 种处理模式:\((0,0), (1,0), (0,1), (1,1)\)。成对平衡需要计算 \(\binom{4}{2}=6\) 个Wasserstein距离。BFG-WB的做法: 1. 计算一个Wasserstein重心 \(R_b^*\),它是四个表示分布 \(R_{00}, R_{10}, R_{01}, R_{11}\) 的“平均”分布(在Wasserstein度量下)。 2. 对每个 \(t\),计算 \(F(R_t, R_b^*)\),即FGW散度。FGW包含两项: - 特征项:\(\eta \int \|r - z\|^2 d\pi(r,z)\),对齐特征值。 - 结构项:\((1-\eta) \iint \left( \|r - r'\| - \|z - z'\| \right)^2 d\pi(r,z) d\pi(r',z')\),保持局部距离结构。 3. 总正则项 \(L_\phi = \sum_{t} w_t F(R_t, R_b^*)\),只需计算4个FGW(加上重心更新中的若干Wasserstein计算,但重心更新次数固定,总复杂度O(L))。
核心思路:用重心作为“锚点”,所有分布向它对齐,避免成对对齐的冲突。FGW的结构项惩罚那些将邻近点映射到远距离的传输计划,从而保持每个处理模式内部的局部几何。
为什么成立:定理1和2表明,CASE和CAIE的估计误差被事实预测误差加上 \(L_\phi\) 项控制。而 \(L_\phi\) 通过重心对齐和FGW同时控制了全局分布差异和局部结构扭曲。在 \(K=2\) 特例下,定理中的 \(2^K\) 因子为4,bound具体可写。
三、这篇论文做了什么¶
三句话¶
- 研究了什么问题:在多处理(multiple treatments)观测数据下,估计异质性单个处理效应(CASE)和交互处理效应(CAIE),并解决成对平衡方法的全局对齐冲突、局部结构扭曲和二次复杂度问题。
- 核心工具/方法:提出CIHSI-Net框架,核心是Barycentric Fused Gromov-Wasserstein Balancing (BFG-WB) 正则项,将每个处理模式的表示分布对齐到一个共享的Wasserstein重心,并用Fused Gromov-Wasserstein散度同时保持特征对齐和局部几何结构。
- 主要结论:在模拟实验中,CIHSI-Net在CASE和CAIE估计误差上持续优于TECE-VAE、NCoRE和CISI-Net;消融实验表明重心对齐和FGW缺一不可;真实营销数据展示了异质性效应的实际意义。
关键设定与假设¶
- 标准因果假设:SUTVA(无干扰+一致性)、Ignorability(给定协变量,处理分配独立于潜在结果)、Overlap(每个处理模式的正概率)。这些与单处理设定相同,但应用于多处理时要求对所有 \(2^K\) 种模式成立,这在 \(K\) 较大时可能不现实(作者在模拟中 \(K=3\),真实数据 \(K=3\))。
- 额外假设(用于理论):
- Assumption 4(可逆表示映射):\(\phi\) 是一一映射,存在逆映射 \(\Psi\)。这保证了表示空间与原协变量空间的信息等价,是表示平衡理论的标准假设(Shalit et al. 2017)。
- Assumption 5(Lipschitz损失):损失函数 \(l(x,t)\) 在表示空间上是 \(B_\phi\)-Lipschitz。这意味着表示空间中的邻近点应有相似的条件期望损失,是局部平滑性的形式化。
- 相比已有文献:与CISI-Net相比,本文用重心对齐替代成对对齐,用FGW替代Wasserstein。与CFR相比,本文扩展到多处理并引入结构保持项。
主要结果¶
- 定理1(CASE上界):对任意 \(k\),
\[\epsilon_{\text{CASE}}(k) \le 2\left( \frac{1}{p(t^{+k})} \epsilon_F^{(t^{+k})} + \frac{1}{p(\mathbf{0})} \epsilon_F^{(\mathbf{0})} + \frac{2^{2K}}{\eta} B_\phi L_\phi \right).\]直觉:误差由事实预测误差(被逆概率加权放大)和BFG-WB正则项控制。\(2^{2K}\) 反映了多处理的组合复杂性。
- 定理2(CAIE上界):对交互集 \(S\),
\[\epsilon_{\text{CAIE}}(S) \le \left( \sum_t a_t^2 \right) \left( \sum_t \frac{1}{p(t)} \epsilon_F^{(t)} + \frac{2^{K+1}}{\eta} B_\phi (2^K - 1) L_\phi \right).\]其中 \(a_t \in \{-1,0,1\}\) 由交互结构决定。同样,误差由事实误差和 \(L_\phi\) 控制。
- 命题2(计算效率):成对平衡每步需 \(O(L^2)\) 次OT评估,BFG-WB需 \(O(L)\) 次(重心更新次数固定)。模拟验证:\(K=8\) 时,CISI-Net每epoch 1529.83秒,CIHSI-Net仅48.32秒。
- 模拟实验:在 \(K=3, N=50000\) 的模拟数据上,CIHSI-Net在所有CASE和CAIE指标上均取得最低误差(表1)。消融实验(表2)显示:重心对齐+FGW组合最优;单独FGW在成对设定下无改善,说明重心对齐是FGW生效的前提。
- 真实数据:移动支付平台三个促销(两个线下、一个线上)。CIHSI-Net估计的CASE显示线上促销对低使用用户效果最强;CAIE显示三向交互从低使用用户的负效应变为高使用用户的正效应,提示协同与蚕食并存。
证明路线与技术技巧¶
整体路线(以定理1为例): 1. 分解误差:将 \(\epsilon_{\text{CASE}}(k)\) 分解为两个处理模式(\(t^{+k}\) 和 \(\mathbf{0}\))的预测误差之和(通过Cauchy-Schwarz或三角不等式)。 2. 将反事实误差 bound 到事实误差 + Wasserstein距离:引理2和3表明,对任意 \(t\),用其他处理模式的协变量分布预测 \(t\) 的损失,不超过事实损失加上 \(B_\phi W_1(R_t, R_{t'})\)。这是通过Kantorovich-Rubinstein对偶实现的。 3. 将成对Wasserstein距离 bound 到重心距离:引理4利用三角不等式,将所有成对Wasserstein距离之和 bound 到各分布与重心距离之和乘以 \((2^K - 1)\)。 4. 将重心距离 bound 到FGW散度:引理5利用FGW定义中特征项的下界:\(F(R_t, R_b^*) \ge \eta W_1(R_t, R_b^*)\),因此 \(W_1 \le F/\eta\)。 5. 代入得到最终bound:将步骤2-4代入步骤1,得到 \(\epsilon_{\text{CASE}}\) 由 \(\epsilon_F\) 和 \(L_\phi\) 控制。
关键跳跃点: - 从成对Wasserstein到重心Wasserstein的三角不等式应用(引理4):这是将复杂度从 \(O(L^2)\) 降到 \(O(L)\) 的理论基础。但注意,bound中出现了 \((2^K - 1)\) 因子,当 \(K\) 大时可能很松。作者在模拟中 \(K=3\),但未讨论 \(K\) 增大时bound的紧性。 - FGW下界Wasserstein(引理5):由于FGW包含非负的结构项,因此 \(F \ge \eta W_1\)。这个下界是松的,因为结构项可能很大,但作者认为结构项的实际好处(保持局部几何)超过了理论bound的松弛。
技术技巧点名: - Kantorovich-Rubinstein对偶:用于将IPM(1-Lipschitz函数族)转化为1-Wasserstein距离(引理2证明)。 - Wasserstein距离的三角不等式:用于引理4和定理证明中的多次放缩。 - FGW散度的定义与下界:引理5。 - 逆频率加权:在损失函数 \(L_y\) 中使用 \(1/\hat{p}(t)\) 加权,对应理论bound中的 \(1/p(t)\) 项。 - 重心估计的迭代算法:使用Cuturi and Doucet (2014)的Sinkhorn-based方法,每mini-batch更新重心。
真实例子与应用¶
- 数据:日本某移动支付平台的三个促销活动(两个线下同商户组CP1、CP2,一个线上不同商户组CP3)。用户协变量71维,结果变量为促销期间总支付金额(标准化)。用户按前一个月支付金额分为11组。
- 方法应用:用CIHSI-Net(超参数来自模拟最优配置)估计每个用户组的CASE和CAIE。
- 结果:
- CASE:所有三个促销效应为正。线上促销CP3对低使用用户效应最强;线下促销对高使用用户效应更明显。
- CAIE:同组促销交互 \(\tau_{\text{CAIE}}(\{1,2\})\) 为正但随使用量下降;跨组交互 \(\tau_{\text{CAIE}}(\{1,3\})\) 和 \(\tau_{\text{CAIE}}(\{2,3\})\) 模式混合;三向交互 \(\tau_{\text{CAIE}}(\{1,2,3\})\) 从低使用用户的负效应变为高使用用户的正效应,提示选择超载与协同效应并存。
- 说明什么:验证了方法能揭示有意义的异质性模式,且与营销文献一致(如线上激励对低活跃用户有效,多促销可能产生蚕食或协同)。敏感性分析(附录E.2)显示模式对超参数 \(\alpha, \eta\) 稳定。
🔎 结论是否比证明窄¶
- 定理的bound依赖于 \(B_\phi\)(Lipschitz常数)和 \(\eta\)(FGW权重),但实际中 \(B_\phi\) 未知,\(\eta\) 需调参。作者未提供如何估计或控制这些量的指导。
- 定理假设表示映射可逆(Assumption 4),但实际中深度网络 \(\phi\) 通常不是一一映射(降维)。作者承认这是理想化条件(附录A.3.3),但未讨论违反时的后果。
- bound中的 \(2^{2K}\) 因子在 \(K\) 较大时指数增长,但模拟中 \(K=8\) 时方法仍有效。作者未解释为何实际表现优于bound预示的悲观情况——可能因为bound是松的,或数据分布使某些项很小。
- 结论声称“BFG-WB提供理论上有原则的误差控制”,但bound仅涉及上界,未给出下界或minimax最优性。因此,不能声称方法是最优的,只能说误差被控制。
四、开放问题(点到为止,扎根具体语句)¶
-
用FGW重心替代Wasserstein重心:作者在结论中提出“a promising extension is to replace the Wasserstein barycenter with an FGW barycenter to unify global alignment and structural preservation”。当前方法先计算Wasserstein重心,再用FGW对齐,两步分离。FGW重心能同时优化特征和结构,但计算更复杂。扎根:Section 7 “First, a promising extension is to replace the Wasserstein barycenter with an FGW barycenter”。
-
连接下游决策:作者提到“connect CIHSI-Net to downstream decision-making, such as uplift-based allocation for combinatorial treatment”。当前只估计效应,未优化分配策略。扎根:Section 7 “Second, it will be valuable to connect CIHSI-Net to downstream decision-making”。
-
理论bound的紧化:定理中的 \(2^{2K}\) 因子可能非常松,尤其当处理模式间有结构相似性时。能否得到更紧的bound(如依赖于有效处理模式数或表示空间的固有维度)?扎根:定理1和2中的指数因子,以及作者未讨论bound的紧性。
-
弱化假设:Assumption 4(可逆表示)和Assumption 5(Lipschitz损失)在实际中难以验证。能否在更弱的条件下(如表示映射是Lipschitz但不可逆,或损失是Hölder连续)得到类似bound?扎根:附录A.3.1中Assumption 4和5,以及作者承认它们是理想化条件。
提醒:要确认这些是否真gap,建议阅读近期多处理因果推断的综述(如Schweisthal et al. 2023的讨论)以及Wasserstein重心统计性质的工作(如Bigot et al. 2013)。如果多个工作都指向同一方向,则是共识性gap;如果互相打架,则是机会。
Maintained by 陈星宇 · Homepage · Source on GitHub