Math-05.信息论-04.互信息与对比学习

本页讲解互信息(mutual information,MI)及在对比学习(contrastive learning)中的核心损失 InfoNCE

段末注释互信息 $I(X;Y)$ 度量知道 $Y$ 后 $X$ 不确定性减少的量;InfoNCE(Noise Contrastive Estimation)用分类式目标估计或最大化表示与正样本的 MI 下界。

系列入口00.系列规划 | 前置:01 总论02 熵


1. 互信息定义(D2–D3)

图 1 I(X;Y) 减少的不确定性

$$
I(X;Y) = H(X) - H(X \mid Y) = H(Y) - H(Y \mid X) = H(X) + H(Y) - H(X,Y)
$$

等价 KL 形式

$$
I(X;Y) = D_{\mathrm{KL}}(p(x,y) | p(x)p(y)) = \mathbb{E}_{p(x,y)}\left[\log \frac{p(x,y)}{p(x)p(y)}\right]
$$

性质 内容
非负 $I(X;Y) \ge 0$
独立 $I(X;Y)=0 \Leftrightarrow X \perp Y$
对称 $I(X;Y)=I(Y;X)$

与相关:高斯变量下 MI 与 Pearson 相关单调相关;一般关系非线性,MI 更一般(Math-02 互补)。


2. 条件互信息与链式(D3)

$$
I(X;Y,Z) = H(X) - H(X \mid Y,Z)
$$

数据处理不等式:若 $X \to Y \to Z$ 形成 Markov 链,则 $I(X;Z) \le I(X;Y)$——中间表示不能创造新信息。对比学习希望学 $Z$ 保留与标签/正样本的 MI,丢弃 nuisance。


3. InfoNCE 与对比学习(D3–D7)

图 2 InfoNCE:正样本 vs 负样本

设定:anchor 表示 $q$,正样本 $k^+$,$N-1$ 个负样本 ${k_j^-}$。SimCLR / MoCo 风格损失:

$$
\mathcal{L}{\mathrm{InfoNCE}} = -\log \frac{\exp(\mathrm{sim}(q, k^+) / \tau)}{\exp(\mathrm{sim}(q, k^+) / \tau) + \sum{j=1}^{N-1} \exp(\mathrm{sim}(q, k_j^-) / \tau)}
$$

理论:InfoNCE 是 $I(q; k^+)$ 的下界;负样本越多、估计越紧($N$ 大 → 更接近真实 MI)。


4. ML 场景(D7)

图 3 对比学习应用

场景 正/负对构造
SimCLR 同图不同增强 = 正;batch 内其他 = 负
MoCo 动量编码器 + 队列负样本
CLIP 图文配对 = 正;错配 = 负
蛋白表示 同蛋白序列增强 / 同源 = 正
酶-底物 天然底物对 = 正;随机配对 = 负
特征选择 估计 $I(X_j; Y)$ 选特征(Math-02/20

5. 局限与估计(D8)

图 4 局限

问题 说明
高维 MI 难估 直方图/KDE 失效;用 InfoNCE、MINE 等
负样本不足 下界松、表示 collapse
collapse 所有 $q$ 相同 → 需 projection head、stop-gradient
批大小依赖 大 batch 或 memory bank 改善
$I$ 大 ≠ 下游好 需 linear probe / 微调验证

6. PyTorch 示例(D12)

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
import torch
import torch.nn.functional as F

def info_nce(q, k_pos, k_neg, temperature=0.07):
"""
q: (B, d), k_pos: (B, d), k_neg: (B, N, d)
"""
q = F.normalize(q, dim=-1)
k_pos = F.normalize(k_pos, dim=-1)
k_neg = F.normalize(k_neg, dim=-1)

pos = (q * k_pos).sum(-1, keepdim=True) / temperature # (B,1)
neg = torch.bmm(k_neg, q.unsqueeze(-1)).squeeze(-1) / temperature # (B,N)
logits = torch.cat([pos, neg], dim=1)
labels = torch.zeros(q.size(0), dtype=torch.long, device=q.device)
return F.cross_entropy(logits, labels)

B, d, N = 32, 128, 64
q = torch.randn(B, d)
k_pos = q + 0.1 * torch.randn(B, d)
k_neg = torch.randn(B, N, d)
print("InfoNCE:", info_nce(q, k_pos, k_neg).item())

7. 小结

MI 量化变量间统计依赖;InfoNCE 把「拉近正样本、推远负样本」变成可训练的 CE。RLHF 中的偏好对齐见 20 RLHF

系列导航03 CE/KL | 20 RLHF

-------------本文结束感谢您的阅读-------------