共计 2287 个字符,预计需要花费 6 分钟才能阅读完成。
背景痛点
图神经网络(GNN)在推荐系统等场景中常遭遇辛普森悖论,即传统关联性学习会混淆因果特征与混杂因素。例如在电商场景中,用户点击行为可能同时受商品质量(因果特征)和位置偏差(混杂因素)影响 [1]。统计显示,未考虑因果关系的 GNN 模型在 A / B 测试中会出现高达 35% 的预测偏差 [2]。

技术对比
结构差异分析
- GCN:基于拉普拉斯平滑的谱域方法,节点更新公式:
$$h_i^{(l+1)} = \sigma\left(\sum_{j\in\mathcal{N}(i)}\frac{1}{\sqrt{d_id_j}}h_j^{(l)}W^{(l)}\right)$$ - GAT:引入注意力机制,但未区分因果 / 非因果边:
$$\alpha_{ij} = \text{softmax}_j\left(\text{LeakyReLU}(a^T[Wh_i||Wh_j])\right)$$ - 因果 GNN:通过 do-calculus 进行干预操作 [3]:
$$P(Y|do(X)) = \sum_{z}P(Y|X,Z=z)P(Z=z)$$
对比学习损失函数
设计双重对比目标:
$$\mathcal{L} = -\mathbb{E}\left[\log\frac{e^{f(x)^Tf(x^+)/\tau}}{e^{f(x)^Tf(x^+)/\tau} + \sum_{i=1}^K e^{f(x)^Tf(x_i^-)/\tau}}\right] + \lambda|\Phi_c – \Phi_s|_F^2$$
其中 $\Phi_c$ 和 $\Phi_s$ 分别表示因果 / 混杂特征投影矩阵 [4]。
实现细节
因果干预模块
class CausalIntervention(nn.Module):
def __init__(self, hidden_dim: int):
super().__init__()
self.projection = nn.Linear(hidden_dim, hidden_dim)
def forward(self, h: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
"""h: [N, D], mask: [N, N]"""
assert h.dim() == 2 and mask.dim() == 2
h_do = torch.einsum('ij,jd->id', mask, h) # do-calculus 实现
return checkpoint(self.projection, h_do) # 梯度检查点
动态负采样策略
- 基于节点度的自适应采样概率:
$$p_{neg}(v) \propto \text{deg}(v)^\alpha$$ - 实现时维护动态候选池:
class NegativeSampler: def __init__(self, degrees: torch.Tensor, alpha: float = 0.75): self.weights = degrees.pow(alpha) + 1e-6 def sample(self, size: int) -> torch.Tensor: return torch.multinomial(self.weights, size)
多头因果注意力层
class CausalAttentionHead(nn.Module):
def forward(self, query: torch.Tensor, key: torch.Tensor,
causal_mask: torch.Tensor) -> torch.Tensor:
"""query/key: [N, D], causal_mask: [N, N]"""
attn = (query @ key.T) / math.sqrt(query.size(-1))
attn = attn.masked_fill(~causal_mask, -1e9)
return attn.softmax(dim=-1)
生产考量
分布式训练优化
- 基于 METIS 算法进行图分区,平衡各 worker 子图规模
- 采用异步梯度更新缓解数据倾斜带来的同步开销
在线服务剪枝
- 移除验证集未激活的 attention 头
- 对低重要性边进行阈值过滤:
$$e_{ij} = \begin{cases}
0 & \text{if} \alpha_{ij} < \epsilon \
\alpha_{ij} & \text{otherwise}
\end{cases}$$
避坑指南
- 代理变量检测 :
- 通过格兰杰因果检验验证特征因果关系
- 可视化注意力权重分布
- 梯度爆炸预防 :
- 对负样本损失项施加梯度裁剪
- 采用渐进式负采样比例调整
- 内存泄漏排查 :
- 使用 torch.cuda.memory_allocated() 监控显存
- 避免在循环中累积计算图
延伸思考
当存在未观测混杂因子 $U$ 时,可考虑:
1. 引入工具变量(IV)构建双重机器学习模型 [5]
2. 通过对抗训练学习混杂不变表示 [6]
3. 结合领域知识构建部分结构因果模型
[1] KDD’21《Causal Inference in Recommender Systems》
[2] ICLR’23《Counterfactual Graph Learning for Link Prediction》
[3] NeurIPS’22《Causal Attention for Unbiased Learning》
[4] ICML’23《Disentangled Contrastive Learning on Graphs》
[5] AAAI’24《Instrumental Variable Learning for Graph Neural Networks》
[6] WWW’23《Adversarial Causal Augmentation for Graph Data》
