InfoNCE Loss
pytorch
本文字数:1.8k 字 | 阅读时长 ≈ 6 min

InfoNCE Loss

pytorch
本文字数:1.8k 字 | 阅读时长 ≈ 6 min

InfoNCE Loss 是对比学习中常用的一种损失函数。简单来说,我们希望模型把相互匹配的样本表示得更相似,同时区分不匹配的样本。下面先看它在比较什么,再通过一个具体例子理解公式和 PyTorch 写法。

1. 基本概念

假设我们有一张猫的图片,对它做两次随机裁剪,得到两个不同的视图。虽然像素不完全相同,但它们来自同一张图片,我们希望模型提取出的特征仍然比较接近。

这里会用到三个名称:

模型先将图片编码为特征向量,再计算 Query 与每个候选样本的相似度。InfoNCE 要做的事情就是:给定一个正样本和若干负样本,让模型从中选出正样本。

这里的正负关系由训练数据的构造方式决定。比如按“是否来自同一张图片”配对时,另一张猫图也可能被当成负样本,并不等同于按猫、狗这样的类别来划分。

为什么叫 InfoNCE?

这个名字可以拆成两部分:

联系上面的猫图例子,为了匹配同一张图的两次裁剪,模型需要提取它们共有、能够帮助识别配对的信息。InfoNCE 通过正负样本之间的比较来提供这样的学习信号。

2. 计算公式

InfoNCE 在 CPC 论文中被提出。下面使用带温度参数的常见相似度写法,考虑一个 Query、一个正样本和 \(K\) 个负样本:

\[ L = -\log \frac{\exp(\operatorname{sim}(q,k^+)/\tau)}{\exp(\operatorname{sim}(q,k^+)/\tau)+\sum_{j=1}^{K}\exp(\operatorname{sim}(q,k_j^-)/\tau)} \]

其中,\(q\) 是 Query 的特征,\(k^+\) 是正样本特征,\(k_j^-\) 是第 \(j\) 个负样本特征,\(\tau>0\) 是温度参数。\(\operatorname{sim}\) 表示相似度,下面使用余弦相似度:

\[ \operatorname{sim}(q,k)=\frac{q^\top k}{\|q\|_2\|k\|_2} \]

先将特征做 L2 归一化,再求点积,就能得到余弦相似度。

我们可以把损失的计算拆成三步:

  1. 计算 Query 与各个候选样本的相似度,再除以温度 \(\tau\),得到分数 logits
  2. 对所有候选的分数做 Softmax,得到选中每个候选的概率。
  3. 取正样本对应的概率 \(p_+\),计算 \(-\log p_+\)。

注意,分母包含正样本和全部负样本。正样本获得的概率越大,Loss 就越小。多个 Query 组成一个 batch 时,通常对各自的 Loss 求平均。

温度控制 Softmax 对分数差异的敏感程度。对于同一组分数,\(\tau\) 越小,概率越集中到最高分的候选上;如果最高分恰好是负样本,Loss 也会变大,所以温度并不是越小越好。

3. 举个例子

假设 Query 与三个候选样本的余弦相似度如下,温度取 \(\tau=0.5\):

候选样本 关系 相似度 除以温度后的分数
同一张猫图的另一次裁剪 正样本 0.8 1.6
狗的图片 负样本 0.2 0.4
汽车的图片 负样本 0.0 0.0

这里的相似度是为了演示计算而设定的。正样本对应的概率为:

\[ p_+=\frac{e^{1.6}}{e^{1.6}+e^{0.4}+e^0}\approx 0.6653 \]

因此:

\[ L=-\log(0.6653)\approx 0.4075 \]

如果三个候选的相似度完全相同,那么正样本的概率就是 \(1/3\),Loss 为 \(-\log(1/3)\approx 1.0986\)。相比之下,上面的分数已经让模型更倾向于选中正样本。

4. PyTorch 实现

先直接用刚才的相似度复现计算:

import torch
import torch.nn.functional as F

similarities = torch.tensor([[0.8, 0.2, 0.0]])
temperature = 0.5
logits = similarities / temperature
labels = torch.tensor([0])  # 正样本位于第 0 列

loss = F.cross_entropy(logits, labels)
print(f"loss: {loss.item():.4f}")

按上一节的手算结果,这里应得到约 0.4075

为什么可以直接用交叉熵?因为这里相当于一个三选一的分类问题,只不过每一列代表一个候选样本,标签表示正样本所在的位置。计算方式和之前的 分类损失:交叉熵与 BCE 一样,都是取正确位置对应的负对数概率。

这里需要注意,传给 F.cross_entropy 的是 logits,不要提前做 Softmax。它与 nn.CrossEntropyLoss 的计算对应,内部已经包含 LogSoftmax 和负对数似然这两步,输入要求也可以查看 PyTorch 文档

实际训练时,我们一般拿到的是模型输出的特征向量。下面使用 batch 内的其他样本作为负样本,约定 query[i]key[i] 是一对正样本:

import torch
import torch.nn.functional as F

def info_nce_loss(query, key, temperature=0.5):
    # query、key 的形状均为 [B, D],每个 query 对应一个正样本
    if temperature <= 0:
        raise ValueError("temperature must be positive")

    query = F.normalize(query, dim=-1)
    key = F.normalize(key, dim=-1)
    logits = query @ key.T / temperature  # [B, B]
    labels = torch.arange(query.size(0), device=query.device)
    return F.cross_entropy(logits, labels)

# 用三个二维向量演示,实际使用时替换为模型对两组视图的输出
query = torch.tensor([[1.0, 0.0], [0.0, 1.0], [-1.0, 0.0]])
key = torch.tensor([[0.8, 0.6], [0.6, 0.8], [-1.0, 0.0]])
loss = info_nce_loss(query, key)
print(f"loss: {loss.item():.4f}")

在这个例子中,除以温度后的相似度矩阵为:

[[ 1.6,  1.2, -2.0],
 [ 1.2,  1.6,  0.0],
 [-1.6, -1.2,  2.0]]

每一行对应一个 Query,每一列对应一个 Key。正样本在对角线上,所以标签为 [0, 1, 2],其余位置都是当前行的负样本。按公式计算,三个 Query 的平均 Loss 约为 0.4074

这里计算的是 Query 到 Key 这一个方向的损失,每个 Query 有 \(B-1\) 个负样本。对角线是两个视图之间的正确配对,需要保留;如果改成同一组特征与自身比较,才需要排除“自己和自己”的位置。实际训练还需要保证正样本顺序对齐,否则标签就会指向错误的候选。

Aug 01, 2026
Mar 13, 2026
ufw