[NIPS'22] Learning Causally Invariant Representations for Out-of-Distribution Generalization on Graphs

Paper: NeurIPS
Code: GitHub

1. 问题定义

给定完整图 $G$ 及其标签 $Y$,从 $G$ 中分离出跨环境稳定的因果子图 $G_c$,丢开随环境变化的伪相关部分 $G_s$,并仅根据 $G_c$ 预测 $Y$。

理想情况下,希望学习到的子图 $\hat G_c$ 同时满足:

$$ I(\hat G_c;Y) $$

尽可能大,即子图保留足够的标签信息;同时:

$$ \hat G_c\perp E $$

即子图不依赖环境 $E$。

训练时通常没有环境标签,因此 CIGA 不直接约束 $\hat G_c\perp E$,而是利用同标签样本之间的稳定共性来识别 $\hat G_c$。

2. 符号

  • $E$:环境或域
  • $C$:潜在的稳定因果因素
  • $S$:潜在的环境相关或伪相关因素
  • $Y$:图标签
  • $G$:实际观察到的完整图
  • $G_c$:由 $C$ 生成的真实稳定子图
  • $G_s$:由 $S$ 生成的真实伪相关部分
  • $g$:从完整图中选择子图的 selector
  • $\hat G_c=g(G)$:模型估计出的稳定子图
  • $\hat G_s=G-g(G)$:模型估计出的补图
  • $H\in\mathbb R^{|V|\times d_h}$:GNN 输出的节点表示
  • $a_{uv}$:边 $(u,v)$ 的可学习分数
  • $z_c\in\mathbb R^{d_z}$:稳定子图的图级表示
  • $z_s\in\mathbb R^{d_z}$:补图的图级表示
  • $f_c$:稳定子图分类器
  • $f_s$:补图分类器
  • $I(A;B\mid Y)$:给定标签 $Y$ 后 $A,B$ 的条件互信息
  • $\theta$:编码器、边选择器与分类器的参数

三组对象:

$$ C,S \quad\rightarrow\quad G_c,G_s \quad\rightarrow\quad \hat G_c,\hat G_s $$

其中:

  • $C,S$ 是生成图之前的潜变量;
  • $G_c,G_s$ 是潜变量生成的真实图结构;
  • $\hat G_c,\hat G_s$ 是神经网络实际选出的近似结果。

3. 主要目标

3.1 图生成假设

论文假设:

$$ G_c \leftarrow f_{\mathrm{gen}}^{G_c}(C) $$

$$ G_s \leftarrow f_{\mathrm{gen}}^{G_s}(S) $$

$$ G \leftarrow f_{\mathrm{gen}}^G(G_c,G_s) $$

稳定因素 $C$ 产生稳定结构 $G_c$,伪相关因素 $S$ 产生环境相关结构 $G_s$,二者共同组成观测图 $G$。环境 $E$ 主要通过改变 $S$ 引起分布偏移。

论文讨论两种生成情形。

FIIF(Fully Informative Invariant Features):

$$ Y\leftarrow f_{\mathrm{inv}}(C), \qquad S\leftarrow f_{\mathrm{spu}}(C,E) $$

对应:

$$ C\to Y, \qquad C\to S\leftarrow E $$

给定 $C$ 后,$S,E$ 不再为标签增加信息:

$$ (S,E)\perp Y\mid C $$

PIIF(Partially Informative Invariant Features):

$$ Y\leftarrow f_{\mathrm{inv}}(C), \qquad S\leftarrow f_{\mathrm{spu}}(Y,E) $$

对应:

$$ C\to Y\to S, \qquad E\to S $$

此时 $S$ 直接受到 $Y$ 影响,所以它可能非常擅长预测标签,但这种预测关系仍然会随环境变化。因此,“能够预测 $Y$”不等于“跨环境稳定”。

3.2 Better-clustered assumption

论文进一步假设,真正稳定的因素在同类样本中聚集得更紧:

$$ H(C\mid Y) \le H(S\mid Y) $$

直觉上:

  • 给定类别后,决定该类别的稳定模式变化较小;
  • 伪相关模式更容易随环境变化,类内不确定性更大。

该假设是 CIGA 使用“同标签样本相互靠近”寻找稳定子图的关键前提,并非对所有数据都无条件成立。

3.3 理想目标

如果环境标签已知,目标可以写为:

$$ \max_{f_c,g} I(\hat G_c;Y) $$

满足:

$$ \hat G_c\perp E, \qquad \hat G_c\leftarrow g(G) $$

但 $E$ 在训练时通常不可用,所以 CIGA 使用同标签样本的类内互信息作为替代。

取两个同标签样本:

$$ (G,Y), \qquad (\tilde G,Y) $$

分别选择:

$$ \hat G_c \leftarrow g(G) $$

$$ \tilde G_c \leftarrow g(\tilde G) $$

然后最大化:

$$ I(\hat G_c;\tilde G_c\mid Y) $$

它鼓励 selector 保留同类样本共享的稳定结构,并删除随未知环境变化的部分。

3.4 CIGAv1

CIGAv1 的理论目标为:

$$ \max_{f_c,g} I(\hat G_c;Y) $$

满足:

$$ \hat G_c \in \arg\max_{\substack{\hat G_c=g(G)\\|\hat G_c|\le s_c}} I(\hat G_c;\tilde G_c\mid Y) $$

其中 $s_c$ 是稳定子图的大小上限。

大小约束用于排除平凡解:

$$ \hat G_c=G $$

因为完整图通常能同时保留最多标签信息和类内信息。CIGAv1 的主要限制是需要知道或假设真实稳定子图的大小。

3.5 CIGAv2

CIGAv2 引入补图:

$$ \hat G_s \leftarrow G-g(G) $$

其理论目标为:

$$ \max_{f_c,g} I(\hat G_c;Y) + I(\hat G_s;Y) $$

满足:

$$ \hat G_c \in \arg\max_{\hat G_c=g(G)} I(\hat G_c;\tilde G_c\mid Y) $$

$$ I(\hat G_s;Y) \le I(\hat G_c;Y) $$

类内互信息负责把最稳定的结构分给 $\hat G_c$;补图目标要求剩余的标签相关信息留在 $\hat G_s$;不等式则防止真正稳定的结构被反向分给补图。

CIGAv2 去掉的是“必须准确知道真实 $G_c$ 固定大小”的理论条件。官方实现仍使用 ratio 做 Top-k 边划分,因此工程上仍有选择比例超参数。

4. 主网络:子图选择与预测

4.1 节点编码

采用赋值式表示:

$$ H \leftarrow \operatorname{GNN}_{\mathrm{enc},\theta}(G) $$

其中:

$$ H\in\mathbb R^{|V|\times d_h} $$

每一行 $h_v$ 是节点 $v$ 的上下文表示。

4.2 产生边分数

对于原图中的边 $(u,v)$,先拼接两端节点表示:

$$ r_{uv} \leftarrow [h_u\Vert h_v] $$

再计算可学习边分数:

$$ a_{uv} \leftarrow \operatorname{MLP}_{\mathrm{edge},\theta}(r_{uv}) $$

合并写为:

$$ a_{uv} \leftarrow \operatorname{MLP}_{\mathrm{edge},\theta} \left([h_u\Vert h_v]\right) $$

CIGA 直接对边评分,而不是像 GIB 一样产生节点到子图/补图的软分配矩阵。

4.3 划分稳定子图与补图

对每张图内部的边分数排序:

$$ E_c,E_s \leftarrow \operatorname{TopKSplit}(E,a,\text{ratio}) $$

其中:

  • 高分边组成 $E_c$;
  • 剩余边组成 $E_s$。

随后构造:

$$ \hat G_c \leftarrow (V_c,E_c) $$

$$ \hat G_s \leftarrow (V_s,E_s) $$

二者主要是边集合的划分,不一定是两个互不相交的节点集合。同一个节点可能同时连接稳定边和补图边。

Top-k 的索引选择是离散操作,但被选中边的连续分数会作为 edge mask 或 edge weight 输入后续 GNN。因此损失仍能通过连续权重更新边打分网络;它不能直接对“把哪一条边换进 Top-k”求导。

4.4 稳定子图表示与标签预测

对稳定子图进行编码、池化和预测:

$$ z_c \leftarrow \operatorname{Readout} \left( \operatorname{GNN}_{c,\theta}(\hat G_c,a_c) \right) $$

$$ \hat Y_c \leftarrow f_{c,\theta}(z_c) $$

分类损失为:

$$ \mathcal L_c \leftarrow \frac1B \sum_{i=1}^{B} \operatorname{CE}(\hat Y_{c,i},Y_i) $$

它近似实现:

$$ \max I(\hat G_c;Y) $$

因为:

$$ I(\hat G_c;Y) = H(Y)-H(Y\mid\hat G_c) $$

而 $H(Y)$ 对训练数据而言是常数,所以最小化分类损失近似于最小化 $H(Y\mid\hat G_c)$。

4.5 补图表示与标签预测

CIGAv2 还计算:

$$ z_s \leftarrow \operatorname{Readout} \left( \operatorname{GNN}_{s,\theta}(\hat G_s,a_s) \right) $$

$$ \hat Y_s \leftarrow f_{s,\theta}(z_s) $$

补图预测头只用于训练时约束信息分工。测试时最终预测来自 $\hat Y_c$,不与 $\hat Y_s$ 融合。

因此主网络的关键赋值关系为:

$$ H \leftarrow \operatorname{GNN}_{\mathrm{enc}}(G) $$

$$ a_{uv} \leftarrow \operatorname{MLP}_{\mathrm{edge}}([h_u\Vert h_v]) $$

$$ \hat G_c,\hat G_s \leftarrow \operatorname{TopKSplit}(G,a,\text{ratio}) $$

$$ \hat Y_c,z_c \leftarrow \operatorname{Predict}_c(\hat G_c,a_c) $$

$$ \hat Y_s,z_s \leftarrow \operatorname{Predict}_s(\hat G_s,a_s) $$

5. 类内互信息:监督式对比学习

理论上希望最大化:

$$ I(\hat G_c;\tilde G_c\mid Y) $$

实际代码不显式估计子图的概率分布,而是对稳定子图表示 $z_c$ 使用监督式对比损失。

对于 batch 中第 $i$ 个样本,定义同标签正样本集合:

$$ P(i) \leftarrow \{p\ne i\mid Y_p=Y_i\} $$

不同标签样本作为负样本。表示相似度为:

$$ s_{ij} \leftarrow \frac{z_{c,i}^{\top}z_{c,j}}{\tau} $$

若预先归一化 $z_c$,该点积等价于余弦相似度;$\tau$ 是 temperature。

监督式对比损失为:

$$ \mathcal L_{\mathrm{con}} \leftarrow -\frac1{|\mathcal I|} \sum_{i\in\mathcal I} \frac1{|P(i)|} \sum_{p\in P(i)} \log \frac{\exp(s_{ip})} {\sum_{a\ne i}\exp(s_{ia})} $$

其中 $\mathcal I$ 只包含当前 batch 中至少存在一个同标签伙伴的 anchor。

最小化该损失会:

  • 拉近同标签稳定子图的表示;
  • 相对推远不同标签子图的表示;
  • 近似最大化类内条件互信息;
  • 鼓励 selector 删除同类样本间随环境变化的结构。

这里没有 GIB 中独立的 statistics network,也没有互信息估计器的内外层优化。对比损失直接由当前 batch 的 $z_c$ 计算,并联合更新:

$$ \operatorname{GNN}_{c}, \quad \operatorname{MLP}_{\mathrm{edge}}, \quad \operatorname{GNN}_{\mathrm{enc}} $$

最大化类内互信息并不是让同类原始图完全相同,而是使选择出的稳定子图表示更加接近。如果一个类别包含多种不同的稳定机制,better-clustered assumption 可能过强。

6. 补图信息约束

CIGAv2 希望最大化补图标签信息,同时满足:

$$ I(\hat G_s;Y) \le I(\hat G_c;Y) $$

互信息越大,通常意味着可达到的预测风险越小,所以风险层面的对应关系为:

$$ R_{\hat G_c} \le R_{\hat G_s} $$

代码先计算逐样本损失:

$$ \ell_{c,i} \leftarrow \operatorname{CE}(\hat Y_{c,i},Y_i) $$

$$ \ell_{s,i} \leftarrow \operatorname{CE}(\hat Y_{s,i},Y_i) $$

再构造门控:

$$ w_i \leftarrow \mathbb I[\ell_{c,i}\le\ell_{s,i}] $$

补图损失为:

$$ \mathcal L_s \leftarrow \frac{\sum_i w_i\ell_{s,i}} {\sum_i w_i+\varepsilon} $$

  • 当补图比稳定子图更难预测,即 $\ell_s\ge\ell_c$ 时,约束满足,继续降低补图损失;
  • 当补图已经比稳定子图更容易预测时,门控关闭,不再增强它;
  • 因而在不违反信息排序的前提下,让补图尽量保留剩余的标签信息。

7. 完整训练过程

对于每个 batch,先执行节点编码与边评分:

$$ H \leftarrow \operatorname{GNN}_{\mathrm{enc},\theta}(G) $$

$$ a_{uv} \leftarrow \operatorname{MLP}_{\mathrm{edge},\theta} \left([h_u\Vert h_v]\right) $$

划分图:

$$ \hat G_c,\hat G_s \leftarrow \operatorname{TopKSplit}(G,a,\text{ratio}) $$

分别进行预测:

$$ \hat Y_c,z_c \leftarrow \operatorname{Predict}_{c,\theta}(\hat G_c,a_c) $$

$$ \hat Y_s,z_s \leftarrow \operatorname{Predict}_{s,\theta}(\hat G_s,a_s) $$

计算稳定子图分类损失:

$$ \mathcal L_c \leftarrow \operatorname{CE}(\hat Y_c,Y) $$

计算类内对比损失:

$$ \mathcal L_{\mathrm{con}} \leftarrow \operatorname{SupCon}(z_c,Y) $$

CIGAv2 再计算补图门控损失:

$$ \mathcal L_s \leftarrow \operatorname{GatedCE} \left(\hat Y_s,Y;\ell_c\right) $$

总损失为:

$$ \mathcal L_{\mathrm{total}} \leftarrow \mathcal L_c + \alpha\mathcal L_{\mathrm{con}} + \beta\mathcal L_s $$

其中:

  • $\mathcal L_c$:保证稳定子图能够预测 $Y$;
  • $\mathcal L_{\mathrm{con}}$:寻找同标签样本共享的稳定结构;
  • $\mathcal L_s$:约束稳定子图与补图之间的标签信息分配;
  • $\alpha$:代码中的 contrast
  • $\beta$:代码中的 spu_coe

CIGAv1 主要使用前两项,并通过子图大小限制排除完整图解;CIGAv2 加入第三项,减少对真实固定子图大小的理论依赖。

所有损失联合反向传播:

$$ \theta \leftarrow \theta - \eta\nabla_{\theta} \mathcal L_{\mathrm{total}} $$

梯度的主要路径为:

$$ \mathcal L_c, \mathcal L_{\mathrm{con}}, \mathcal L_s \Rightarrow z,\hat Y \Rightarrow \text{子图分类器} \Rightarrow \text{edge mask} \Rightarrow \text{边选择器与编码器} $$

最终得到:

$$ \hat G_c^* \leftarrow g(G;\theta^*) $$

它是模型在生成假设与 better-clustered assumption 下识别出的稳定子图:

  • 保留预测标签所需的信息;
  • 在同标签样本之间具有更稳定的表示;
  • 尽量把环境相关的预测信息留给补图;
  • 测试时仅使用稳定子图完成预测。
富婆饿饿饭饭