[NIPS'23] Learning Invariant Molecular Representation in Latent Discrete Space
Paper: NIPS
Code: GitHub
1. 问题定义
给定分子图 G 及其标签或性质 Y,希望学习一个预测器:
$$ f:G\rightarrow Y $$
训练分子和测试分子可能来自不同环境,因此:
$$ P_{\mathrm{train}}(G,Y) \neq P_{\mathrm{test}}(G,Y) $$
环境变化可以来自:
- scaffold:分子的核心骨架发生变化;
- size:分子的原子数量发生变化;
- assay:实验方法、靶点或测量条件发生变化。
论文希望从完整分子的表示中分离出:
$$ z^{\mathrm{Inv}} $$
即跨环境相对稳定、足以预测 Y 的不变表示;以及:
$$ z^{\mathrm{Spu}} $$
即容易随环境变化的伪相关表示。
最终只使用不变表示预测:
$$ \widehat Y \leftarrow \rho(z^{\mathrm{Inv}}) $$
理想情况下,希望:
$$ I(z^{\mathrm{Inv}};Y) $$
尽可能大,即不变表示包含足够的标签信息;同时希望:
$$ z^{\mathrm{Inv}}\perp E $$
即它尽量不依赖环境 E。
训练时没有环境标签,因此 iMoLD 不直接约束 z^{\mathrm{Inv}}\perp E,而是通过随机替换伪相关表示,要求预测出的不变表示保持一致。
2. 符号
- G=(V,E):输入分子图
- V:原子节点集合
- E\subseteq V\times V:化学键集合
- Y:分子标签或连续性质
- B:batch size
- d:节点表示维度
- \operatorname{GNN}_E:encoding GNN,负责编码完整分子
- \operatorname{GNN}_S:scoring GNN,负责产生分离分数
- H\in\mathbb R^{|V|\times d}:encoding GNN 输出的连续节点表示
- h_v\in\mathbb R^d:节点 v 的连续表示
- \mathcal C={e_1,\ldots,e_{|\mathcal C|}}:向量量化 codebook
- e_k\in\mathbb R^d:第 k 个离散 code
- Q(h_v):与 h_v 最近的 code
- H'\in\mathbb R^{|V|\times d}:经过残差向量量化后的节点表示
- S\in(0,1)^{|V|\times d}:节点与特征维度上的软分离分数
- H^{\mathrm{Inv}}:不变节点表示
- H^{\mathrm{Spu}}:伪相关节点表示
- z^{\mathrm{Inv}}\in\mathbb R^d:图级不变表示
- z^{\mathrm{Spu}}\in\mathbb R^d:图级伪相关表示
- \rho:下游标签预测器
- \omega:自监督目标中的 MLP predictor
- \operatorname{sg}[\cdot]:stop-gradient 操作
- \gamma\in(0,1):期望选为不变特征的平均比例
- \lambda_1,\lambda_2,\lambda_3:不同损失项的权重
完整表示的分工为:
$$ H' \quad\longrightarrow\quad H^{\mathrm{Inv}},H^{\mathrm{Spu}} \quad\longrightarrow\quad z^{\mathrm{Inv}},z^{\mathrm{Spu}} $$
3. 主要思想:First-Encoding-Then-Separation
已有图 OOD 方法通常采用:
$$ G \xrightarrow{\text{先分离}} G_c,G_s \xrightarrow{\text{再编码}} z_c,z_s $$
即先在原始图结构中选择节点、边或 motif,再分别编码稳定子图和伪相关子图。
iMoLD 采用相反顺序:
$$ G \xrightarrow{\text{先编码}} H \xrightarrow{\text{RVQ}} H' \xrightarrow{\text{再分离}} H^{\mathrm{Inv}},H^{\mathrm{Spu}} $$
原因是分子性质可能依赖多个原子和基团之间的整体相互作用。若在理解完整分子之前就删除部分结构,可能过早损失有用信息。
因此 iMoLD 不直接寻找一个离散的“因果子图”,而是在上下文化的潜表示中,对每个节点、每个特征维度进行软分离。
与 CIGA 的主要区别是:
$$ \text{CIGA:原始边空间中的 Top-k 分离} $$
$$ \text{iMoLD:潜在表示空间中的逐元素软分离} $$
4. 主网络:编码、量化与分离
4.1 编码完整分子
首先使用 encoding GNN 编码完整分子:
$$ H \leftarrow \operatorname{GNN}_E(G) $$
其中:
$$ H = [h_1,h_2,\ldots,h_{|V|}]^\top \in \mathbb R^{|V|\times d} $$
每个 h_v 已经聚合节点邻域信息,是上下文化的原子表示。
此时完整图还没有被划分,所以每个节点表示可以利用整个消息传递范围中的结构信息。
4.2 向量量化
引入共享且可学习的 codebook:
$$ \mathcal C \leftarrow \{e_1,e_2,\ldots,e_{|\mathcal C|}\} $$
对于每个节点表示 h_v,寻找欧氏距离最近的 code:
$$ k(v) \leftarrow \arg\min_{k\in\{1,\ldots,|\mathcal C|\}} \|h_v-e_k\|_2 $$
$$ Q(h_v) \leftarrow e_{k(v)} $$
连续空间中略有差异的节点表示可能被映射到同一个 code。因此,codebook 相当于一组可复用的原子上下文原型。
离散瓶颈的作用是:
- 减少模型对训练样本细节的记忆;
- 让新环境中的表示复用训练阶段出现过的原型;
- 使训练环境和测试环境的特征分布更容易对齐。
但纯 VQ 会把 h_v 完全替换为有限 code,可能造成欠拟合。
4.3 残差向量量化 RVQ
为了保留连续表示的表达能力,论文加入残差连接:
$$ h'_v \leftarrow h_v+Q(h_v) $$
合并所有节点:
$$ H' \leftarrow [Q(h_1)+h_1,\ldots,Q(h_{|V|})+h_{|V|}]^\top $$
其中:
- Q(h_v):提供离散原型和泛化瓶颈;
- h_v:保留具体分子的连续细节;
- h'_v:在泛化能力和表达能力之间折中。
需要注意,H' 仍包含连续的 H,所以它不是严格的全离散表示,而是连续表示与离散 code 的残差组合。
4.4 Codebook 的 EMA 更新
对于第 t 个 mini-batch,设分配给 code e_k 的节点数量为:
$$ n_k^{(t)} $$
维护该 code 的指数移动计数:
$$ N_k^{(t)} \leftarrow \eta N_k^{(t-1)} + (1-\eta)n_k^{(t)} $$
同时维护分配给该 code 的节点表示之和:
$$ m_k^{(t)} \leftarrow \eta m_k^{(t-1)} + (1-\eta) \sum_{v:k(v)=k}h_v^{(t)} $$
再更新 code:
$$ e_k^{(t)} \leftarrow \frac{m_k^{(t)}}{N_k^{(t)}} $$
其含义类似在线聚类中心更新:code 会逐渐靠近被分配到该 code 的节点表示均值。
4.5 产生节点—特征分离矩阵
另一个 scoring GNN 根据原始分子图产生分数:
$$ S \leftarrow \sigma \left( \operatorname{GNN}_S(G) \right) $$
其中:
$$ S\in(0,1)^{|V|\times d} $$
S_{vk} 表示节点 v 的第 k 个潜在特征被分配给不变表示的软权重。
这与只产生一个节点分数不同。对于同一个节点:
$$ S_{v1},S_{v2},\ldots,S_{vd} $$
可以互不相同,因此同一原子的某些特征维度可以进入不变部分,另一些维度进入伪相关部分。
4.6 分离不变表示与伪相关表示
逐元素划分:
$$ H^{\mathrm{Inv}} \leftarrow H'\odot S $$
$$ H^{\mathrm{Spu}} \leftarrow H'\odot(1-S) $$
二者满足:
$$ H^{\mathrm{Inv}}+H^{\mathrm{Spu}}=H' $$
因此它是互补的软分解,而不是删除原始图中的节点或边。
随后执行 permutation-invariant readout:
$$ z^{\mathrm{Inv}} \leftarrow \operatorname{READOUT}(H^{\mathrm{Inv}}) \in\mathbb R^d $$
$$ z^{\mathrm{Spu}} \leftarrow \operatorname{READOUT}(H^{\mathrm{Spu}}) \in\mathbb R^d $$
主网络的关键赋值关系为:
$$ H \leftarrow \operatorname{GNN}_E(G) $$
$$ H' \leftarrow H+Q(H) $$
$$ S \leftarrow \sigma(\operatorname{GNN}_S(G)) $$
$$ H^{\mathrm{Inv}},H^{\mathrm{Spu}} \leftarrow H'\odot S, H'\odot(1-S) $$
$$ z^{\mathrm{Inv}},z^{\mathrm{Spu}} \leftarrow \operatorname{READOUT}(H^{\mathrm{Inv}}), \operatorname{READOUT}(H^{\mathrm{Spu}}) $$
5. Task-agnostic Self-supervised Invariant Learning
论文没有环境标签,因此无法直接比较不同环境中同一个稳定因素,也不能直接优化:
$$ z^{\mathrm{Inv}}\perp E $$
作者采用 batch 内随机交换伪相关表示的方法,检验当前分离是否可靠。
5.1 Batch 内打乱伪相关表示
对于 batch 中第 i 个样本,保留其自身的不变表示:
$$ z_i^{\mathrm{Inv}} $$
从打乱后的 batch 中取另一个样本 j=\pi(i) 的伪相关表示:
$$ z_{\pi(i)}^{\mathrm{Spu}} $$
将二者拼接为增强视图:
$$ \widetilde z_i^{\mathrm{Inv}} \leftarrow z_i^{\mathrm{Inv}} \oplus z_{\pi(i)}^{\mathrm{Spu}} $$
其中:
$$ \widetilde z_i^{\mathrm{Inv}} \in\mathbb R^{2d} $$
这里的上标 \mathrm{Inv} 表示该增强视图的目标仍然是恢复不变表示,并不表示拼接后的整个向量已经是不变的。
5.2 构造正样本对
把下面两个视图视为正样本:
$$ z_i^{\mathrm{Inv}} $$
$$ \widetilde z_i^{\mathrm{Inv}} = z_i^{\mathrm{Inv}} \oplus z_{\pi(i)}^{\mathrm{Spu}} $$
使用 MLP predictor 将拼接表示映射回 d 维:
$$ p_i \leftarrow \omega(\widetilde z_i^{\mathrm{Inv}}) $$
其中:
$$ \omega:\mathbb R^{2d}\rightarrow\mathbb R^d $$
然后计算负余弦相似度:
$$ \mathcal L_{\mathrm{inv}} \leftarrow -\sum_{i=1}^{B} \operatorname{sim} \left( \operatorname{sg}[z_i^{\mathrm{Inv}}], \omega(\widetilde z_i^{\mathrm{Inv}}) \right) $$
其中:
$$ \operatorname{sim}(a,b) = \frac{a^\top b}{\|a\|_2\|b\|_2} $$
stop-gradient 表示目标分支不从这一项接收梯度:
$$ \frac{\partial\operatorname{sg}[z_i^{\mathrm{Inv}}]} {\partial z_i^{\mathrm{Inv}}} =0 $$
其作用类似 SimSiam,用非对称 predictor 和 stop-gradient 降低表示坍塌的风险。
5.3 该目标如何影响分离矩阵
如果 z_i^{\mathrm{Inv}} 确实已经捕获稳定信息,而 z_j^{\mathrm{Spu}} 是可替换的环境相关信息,那么无论拼入哪个 z_j^{\mathrm{Spu}},都应该恢复相同的目标:
$$ \omega \left( z_i^{\mathrm{Inv}}\oplus z_j^{\mathrm{Spu}} \right) \approx z_i^{\mathrm{Inv}} $$
最小化 \mathcal L_{\mathrm{inv}} 的梯度会经过:
$$ \mathcal L_{\mathrm{inv}} \Rightarrow \widetilde z^{\mathrm{Inv}} \Rightarrow z^{\mathrm{Inv}},z^{\mathrm{Spu}} \Rightarrow H^{\mathrm{Inv}},H^{\mathrm{Spu}} \Rightarrow S,H' $$
从而联合更新:
- encoding GNN;
- scoring GNN;
- MLP predictor;
- 与 RVQ 输入有关的表示参数。
这一目标不要求标签是单分类标签,也不需要根据类别构造正负样本,因此可以与二分类、多标签分类和回归任务共同使用。这是论文所说的 task-agnostic。
但它仍然只是间接诱导不变性,并没有严格证明:
$$ z^{\mathrm{Inv}}\perp E $$
也不能保证 z^{\mathrm{Spu}} 中只包含环境信息。
6. 预测与正则化损失
6.1 任务预测损失
只使用不变表示进行预测:
$$ \widehat Y_i \leftarrow \rho(z_i^{\mathrm{Inv}}) $$
分类任务使用交叉熵或二元交叉熵:
$$ \mathcal L_{\mathrm{pred}} \leftarrow -\sum_{i=1}^{B} \left[ y_i\log\widehat y_i + (1-y_i)\log(1-\widehat y_i) \right] $$
回归任务使用:
$$ \mathcal L_{\mathrm{pred}} \leftarrow \frac1B \sum_{i=1}^{B} \|\widehat y_i-y_i\|_2^2 $$
该损失保证 z^{\mathrm{Inv}} 具有充分的预测能力,近似实现:
$$ \max I(z^{\mathrm{Inv}};Y) $$
论文公式(13)的二元交叉熵没有写负号;若总目标按最小化训练,标准实现应包含负号,这应是论文中的符号疏漏。
6.2 Scoring GNN 正则化
如果没有额外约束,scoring GNN 可能产生平凡分配:
$$ S\approx\mathbf 1 $$
即把几乎所有信息都分给不变部分;或者:
$$ S\approx\mathbf 0 $$
即几乎不选择不变信息。
先计算 S 的平均值:
$$ \bar S \leftarrow \frac{\langle J,S\rangle_F}{|V|d} $$
其中 J\in\mathbb R^{|V|\times d} 是全 1 矩阵。
再约束平均选择比例接近预设阈值 \gamma:
$$ \mathcal L_{\mathrm{reg}} \leftarrow \left| \frac{\langle J,S\rangle_F}{|V|d} -\gamma \right| $$
它约束的是所有节点—特征权重的平均值,而不是规定每个节点必须选择相同比例。
与 CIGA 的 Top-k ratio 相比:
- CIGA 对每张图的边执行硬排序和固定比例划分;
- iMoLD 使用连续的 S,仅对平均权重施加软约束。
6.3 Codebook commitment loss
最近邻查找 Q(\cdot) 是离散操作。为了让连续编码 h_v 靠近被选中的 code,并避免它在多个 code 之间频繁跳动,加入:
$$ \mathcal L_{\mathrm{cmt}} \leftarrow \sum_{v\in V} \| \operatorname{sg}[e_{k(v)}]-h_v \|_2^2 $$
对 code 使用 stop-gradient,所以该损失主要把 encoder 输出 h_v 拉向当前选中的 code:
$$ h_v \longrightarrow e_{k(v)} $$
codebook 本身主要通过 EMA 更新,而不是由此项直接进行梯度更新。
6.4 完整目标
最终损失为:
$$ \mathcal L \leftarrow \mathcal L_{\mathrm{pred}} + \lambda_1\mathcal L_{\mathrm{inv}} + \lambda_2\mathcal L_{\mathrm{reg}} + \lambda_3\mathcal L_{\mathrm{cmt}} $$
四个损失分别负责:
- \mathcal L_{\mathrm{pred}}:保证不变表示能够预测标签;
- \mathcal L_{\mathrm{inv}}:使表示对随机替换的伪相关部分保持稳定;
- \mathcal L_{\mathrm{reg}}:防止 scoring GNN 产生全选或全不选的平凡解;
- \mathcal L_{\mathrm{cmt}}:稳定 encoder 与 codebook 之间的分配。
7. 训练过程
对于每个 batch,先用两套 GNN 分别编码和评分:
$$ H \leftarrow \operatorname{GNN}_E(G) $$
$$ S \leftarrow \sigma(\operatorname{GNN}_S(G)) $$
对每个节点执行最近邻 code 查找:
$$ Q(h_v) \leftarrow e_{\arg\min_k\|h_v-e_k\|_2} $$
构造残差量化表示:
$$ H' \leftarrow H+Q(H) $$
分离潜在表示:
$$ H^{\mathrm{Inv}} \leftarrow H'\odot S $$
$$ H^{\mathrm{Spu}} \leftarrow H'\odot(1-S) $$
池化为图级表示:
$$ z^{\mathrm{Inv}} \leftarrow \operatorname{READOUT}(H^{\mathrm{Inv}}) $$
$$ z^{\mathrm{Spu}} \leftarrow \operatorname{READOUT}(H^{\mathrm{Spu}}) $$
使用不变表示预测:
$$ \widehat Y \leftarrow \rho(z^{\mathrm{Inv}}) $$
$$ \mathcal L_{\mathrm{pred}} \leftarrow \operatorname{TaskLoss}(\widehat Y,Y) $$
在 batch 内打乱伪相关表示:
$$ \widetilde z_i^{\mathrm{Inv}} \leftarrow z_i^{\mathrm{Inv}} \oplus z_{\pi(i)}^{\mathrm{Spu}} $$
计算自监督不变损失:
$$ \mathcal L_{\mathrm{inv}} \leftarrow -\sum_i \operatorname{sim} \left( \operatorname{sg}[z_i^{\mathrm{Inv}}], \omega(\widetilde z_i^{\mathrm{Inv}}) \right) $$
计算选择比例约束:
$$ \mathcal L_{\mathrm{reg}} \leftarrow \left| \frac{\langle J,S\rangle_F}{|V|d} -\gamma \right| $$
计算 commitment loss:
$$ \mathcal L_{\mathrm{cmt}} \leftarrow \sum_v \| \operatorname{sg}[Q(h_v)]-h_v \|_2^2 $$
组合总损失:
$$ \mathcal L \leftarrow \mathcal L_{\mathrm{pred}} + \lambda_1\mathcal L_{\mathrm{inv}} + \lambda_2\mathcal L_{\mathrm{reg}} + \lambda_3\mathcal L_{\mathrm{cmt}} $$
梯度下降更新主网络:
$$ \theta \leftarrow \theta - \eta_\theta \nabla_\theta\mathcal L $$
同时,根据当前 batch 中的节点分配,通过 EMA 更新 codebook:
$$ \mathcal C^{(t-1)} \longrightarrow \mathcal C^{(t)} $$
测试时不需要打乱伪相关表示,也不需要 MLP predictor。执行:
$$ G \rightarrow H \rightarrow H' \rightarrow S \rightarrow z^{\mathrm{Inv}} \rightarrow \widehat Y $$
最终预测始终只来自:
$$ \widehat Y = \rho(z^{\mathrm{Inv}}) $$
