1. 对比学习:从“找不同”到“学特征”

如果你玩过“找不同”游戏,或者教过小朋友认识动物,那你其实已经接触过对比学习的核心思想了。想象一下,你面前有一堆猫和狗的图片,你不需要别人告诉你“这是猫,那是狗”,你只需要知道“这两张图是同一只猫的不同角度”和“这张猫图和那张狗图不是一回事”。通过不断地比较“相同”与“不同”,你的大脑自然而然就学会了区分猫和狗的特征。这就是对比学习最直观的比喻。

在AI的世界里,尤其是在自监督学习领域,对比学习已经成为一种强大的“无师自通”方法。传统的监督学习需要海量的人工标注数据,成本高昂。而对比学习的目标,是让模型在没有标签的情况下,自己从数据中“悟”出有用的特征表示。它通过构造“正样本对”(相似的、相关的)和“负样本对”(不相似的、无关的),让模型学会一个核心能力:把相似的样本在特征空间里拉近,把不相似的样本推远。

这个过程中,一个关键的“裁判”角色,就是损失函数。它负责告诉模型:“你这次拉近和推远的动作,做得对不对,好不好。” 而NCEInfoNCE,就是两位在对比学习赛场上表现极其出色的“明星裁判”。它们都源于同一个思想——通过对比来学习,但在设计理念、计算方式和适用场景上,又有着微妙的区别。很多刚入门的朋友容易把它们搞混,或者不清楚在什么情况下该请哪位“裁判”出马。今天,我就结合自己这些年踩过的坑和实战经验,带你彻底搞懂这两位“裁判”的看家本领。

2. NCE:化繁为简的“高效裁判”

2.1 它要解决什么“老大难”问题?

要理解NCE的妙处,得先看它要解决什么麻烦。在自然语言处理(NLP)里训练一个语言模型,比如预测下一个词是什么,模型最后通常会接一个Softmax层。这个层的计算有个让人头疼的地方:它需要对整个词汇表的所有词进行计算和归一化。词汇表有多大呢?动辄几万、几十万甚至上百万个词。每次预测,都要算一遍几十万次的指数运算和求和,这计算开销,想想都肉疼。这个问题被称为“归一化常数难计算”问题。

NCE的聪明之处在于,它把这个问题巧妙地“转换”了。它不再让模型直接去估计一个词在百万词汇中的精确概率,而是把它变成了一个二分类问题:给定一个上下文和一个候选词,模型只需要判断“这个候选词是来自真实数据分布的正样本,还是来自某个噪声分布的负样本?” 这样一来,我们就不用去算那个庞大的归一化分母了,计算量瞬间降了下来。

2.2 NCE的原理:一场“真假美猴王”的鉴别游戏

我们可以把NCE想象成一个鉴宝专家训练过程。我们给专家(模型)看很多真品(正样本,来自真实数据分布),同时也混入一些高仿赝品(负样本,来自我们设定的噪声分布,比如一个均匀分布)。专家的任务不是去给每件古董估一个精确的市场价(概率值),而是简单地判断:“眼前这件,是真品还是赝品?”

NCE的损失函数公式,看起来有点复杂,但拆解开来就很好懂:

L_NCE = - (1/N) Σ [ log(P_model(x_i) / (P_model(x_i) + k * P_n(x_i))) + Σ log(k * P_n(x_ij) / (P_model(x_ij) + k * P_n(x_ij))) ]

我来翻译一下:

  • P_model(x):模型认为样本x是“真品”(来自真实分布)的概率。我们希望这个值对于真正的正样本越大越好。
  • P_n(x):样本x来自我们设定的噪声分布的概率。这个分布是我们事先选好的,比如一个简单的均匀分布。
  • k:对于每一个真品(正样本),我们会配套拿出k个赝品(负样本)来一起让模型鉴别。
  • 公式的第一部分,是针对正样本的:鼓励模型提高P_model(正样本)
  • 公式的第二部分,是针对那k个负样本的:鼓励模型认识到它们来自噪声分布,即降低P_model(负样本),使其更接近k * P_n(负样本)

整个训练过程,就是让模型在这个“真假鉴别游戏”中越玩越溜,最终,模型不仅学会了区分真假,它内部的P_model函数也间接地学到了真实数据的概率分布形状。这就是“曲线救国”的智慧。

2.3 实战代码与细节剖析

光说不练假把式,我们来看一个PyTorch实现的简化版NCE损失。在实际应用中,比如Word2Vec的负采样,就是NCE思想的一个经典变体。

import torch
from torch import nn

class NCELoss(nn.Module):
    def __init__(self, noise_distribution_size):
        super().__init__()
        self.noise_distribution_size = noise_distribution_size

    def forward(self, model_scores, targets):
        """
        model_scores: 模型对正样本和负样本的打分,形状 [batch_size, 1 + k]
                      第0列是正样本的分数,第1到k列是k个负样本的分数。
        targets: 占位符,为了接口统一,实际NCE不依赖具体标签。
        """
        batch_size = model_scores.size(0)
        k = model_scores.size(1) - 1  # 负样本数量

        # 假设噪声分布是均匀的,每个噪声样本的先验概率
        P_n = 1.0 / self.noise_distribution_size

        # 提取正样本分数(第一列)
        pos_scores = model_scores[:, 0].unsqueeze(1)  # [batch_size, 1]
        # 提取负样本分数
        neg_scores = model_scores[:, 1:]  # [batch_size, k]

        # 计算正样本部分的损失:log( P_model / (P_model + k*P_n) )
        # 这里将模型分数视为未归一化的logit,通过sigmoid转换为概率
        pos_probs = torch.sigmoid(pos_scores)
        # 注意:实际实现中,P_model和k*P_n需要是概率值,这里做了简化示意
        # 更严谨的实现会涉及将分数转换为概率密度比
        pos_term = torch.log(pos_probs / (pos_probs + k * P_n + 1e-8))

        # 计算负样本部分的损失:sum( log( k*P_n / (P_model + k*P_n) ) )
        neg_probs = torch.sigmoid(neg_scores)
        neg_term = torch.sum(torch.log((k * P_n + 1e-8) / (neg_probs + k * P_n + 1e-8)), dim=1, keepdim=True)

        # 合并损失
        loss = - (pos_term + neg_term).mean()
        return loss

# 模拟使用场景
batch_size = 32
vocab_size = 50000  # 词汇表大小
k = 10  # 每个正样本对应10个负样本

# 模型输出分数(例如,经过一层线性变换后的值)
model_output = torch.randn(batch_size, 1 + k)
# 初始化损失函数,传入噪声分布大小(这里用词汇表大小近似)
nce_loss = NCELoss(noise_distribution_size=vocab_size)

loss = nce_loss(model_output, None)
print(f"NCE Loss: {loss.item():.4f}")

在实际踩坑中,有两点需要特别注意:

  1. 噪声分布的选择:噪声分布P_n(x)不能随便选。理想情况下,它应该尽可能接近真实的数据分布,这样“鉴别游戏”才更有挑战性,模型学到的特征也更鲁棒。在词向量训练中,常根据词频的3/4次方来采样,这就是为了让噪声分布更贴近真实。
  2. 负样本数量k:k是一个超参数。k太小,游戏太简单,模型学不到东西;k太大,计算量会增加,可能收敛变慢。通常需要根据任务和资源进行调优。

3. InfoNCE:专注表征学习的“互信息裁判”

3.1 从NCE到InfoNCE的进化

如果说NCE是一位致力于解决具体计算难题(归一化)的效率专家,那么InfoNCE就是一位专注于学习高质量数据表征的特征大师。InfoNCE继承了NCE“通过对比来学习”的基因,但它的目标更加明确和纯粹:最大化正样本对之间的互信息

互信息是信息论里的概念,简单理解就是“知道了一个信息,能给你带来多少关于另一个信息的信息量”。在对比学习中,我们希望锚点样本(比如一张图片)和它的正样本(同一张图片的不同增强视图)之间的互信息尽可能大,这意味着它们的表征非常相似;而锚点样本和负样本之间的互信息尽可能小,这意味着它们的表征差异很大。

3.2 InfoNCE原理:一个“择优录取”的竞赛

InfoNCE的公式比NCE看起来更清爽,也更有“对比学习”的味道:

L_InfoNCE = - E [ log( exp(sim(z_i, z_j) / τ) / ( exp(sim(z_i, z_j) / τ) + Σ exp(sim(z_i, z_k) / τ) ) ) ]

这里的sim是相似度函数,常用余弦相似度或点积。τ是一个温度超参数,我习惯叫它“宽容度调节器”。

这个过程就像一个竞赛:

  • 锚点样本z_i是评委。
  • 一个正样本z_jK个负样本z_k是参赛选手。
  • 评委根据与每位选手的“相似度”sim来打分。
  • 最后,通过一个Softmax操作,将分数转换为概率分布。InfoNCE损失的目标,就是让正样本选手z_j在这个分布中的概率尽可能接近1。

温度参数τ非常关键。τ值小,Softmax分布会更“尖锐”,模型会拼命拉大正样本与最难负样本(相似度最高的负样本)之间的差距,这对学习非常精细的区分特征有好处,但也可能让训练不稳定。τ值大,分布更“平滑”,模型对所有负样本一视同仁地推开,学习过程更温和。我一般在图像任务上从0.1开始调,文本任务上可能会试试0.05或0.2。

3.3 手把手实现InfoNCE

下面是一个功能比较完整的InfoNCE实现,它考虑了两种负样本模式:‘paired’(每个锚点有自己专属的一组负样本)和‘unpaired’(所有锚点共享同一组负样本)。

import torch
import torch.nn.functional as F

def info_nce_loss(query, positive_key, negative_keys=None, temperature=0.1, mode='unpaired'):
    """
    query: 锚点样本特征,形状 [batch_size, feature_dim]
    positive_key: 正样本特征,形状 [batch_size, feature_dim]
    negative_keys: 负样本特征。
        若 mode='unpaired', 形状为 [num_negatives, feature_dim]
        若 mode='paired', 形状为 [batch_size, num_negatives, feature_dim]
    temperature: 温度参数
    mode: 负样本模式,'paired' 或 'unpaired'
    """
    batch_size = query.shape[0]
    feature_dim = query.shape[1]

    # 第一步:归一化特征向量(非常重要!通常使用余弦相似度)
    query = F.normalize(query, dim=-1)
    positive_key = F.normalize(positive_key, dim=-1)

    # 计算锚点与正样本的相似度 [batch_size, 1]
    pos_sim = torch.sum(query * positive_key, dim=-1, keepdim=True)  # 点积即余弦相似度(因为已归一化)

    if negative_keys is not None:
        negative_keys = F.normalize(negative_keys, dim=-1)
        if mode == 'unpaired':
            # 计算锚点与所有负样本的相似度 [batch_size, num_negatives]
            # 这里用了矩阵乘法,等价于批量点积
            neg_sim = torch.matmul(query, negative_keys.T)
        elif mode == 'paired':
            # 每个锚点与自己那组负样本计算相似度
            # query: [batch_size, 1, feature_dim], negative_keys: [batch_size, num_negatives, feature_dim]
            neg_sim = torch.sum(query.unsqueeze(1) * negative_keys, dim=-1)
        else:
            raise ValueError("mode must be 'unpaired' or 'paired'")

        # 拼接相似度:第一列是正样本,后面是负样本
        logits = torch.cat([pos_sim, neg_sim], dim=1)
    else:
        # 如果没有显式提供负样本,则默认将批次内其他样本作为负样本(SimCLR做法)
        # 计算所有样本两两之间的相似度 [batch_size, batch_size]
        logits = torch.matmul(query, positive_key.T)
        # 对角线位置是正样本对,其余是负样本对
        pos_sim = torch.diag(logits).unsqueeze(1)  # 提取对角线作为正样本相似度
        # 这里logits矩阵本身就已经是我们要的形式了

    # 除以温度参数
    logits = logits / temperature

    # 标签:对于每个锚点,正样本在logits中的索引是0(当提供负样本时)
    # 如果没有提供负样本(批次内负样本),则标签是对角线的位置索引
    if negative_keys is not None:
        labels = torch.zeros(batch_size, dtype=torch.long, device=query.device)
    else:
        labels = torch.arange(batch_size, dtype=torch.long, device=query.device)

    # 使用交叉熵损失
    loss = F.cross_entropy(logits, labels)
    return loss

# 示例1:使用显式负样本(unpaired模式)
batch_size = 16
feature_dim = 128
num_negatives = 255

query = torch.randn(batch_size, feature_dim)
positive = torch.randn(batch_size, feature_dim)
negatives = torch.randn(num_negatives, feature_dim)  # 所有查询共享这255个负样本

loss_unpaired = info_nce_loss(query, positive, negatives, temperature=0.07, mode='unpaired')
print(f"InfoNCE Loss (unpaired negatives): {loss_unpaired.item():.4f}")

# 示例2:使用批次内其他样本作为负样本(SimCLR经典做法)
loss_implicit = info_nce_loss(query, positive, negative_keys=None, temperature=0.07)
print(f"InfoNCE Loss (in-batch negatives): {loss_implicit.item():.4f}")

在真实项目中,我更喜欢使用“批次内负样本”的模式,因为它实现简单,且能高效利用批量数据。但要注意,当批次内存在“假负样本”(即本质是同类但被误当作负样本)时,会影响效果。这时可能需要更复杂的采样策略。

4. NCE vs InfoNCE:如何选择你的“裁判”?

了解了两位“裁判”的看家本领后,最关键的问题来了:我该用哪一个?下面这个表格从几个核心维度进行了对比,帮你快速决策。

特性维度NCE (Noise Contrastive Estimation)InfoNCE (Information Noise Contrastive Estimation)
核心目标估计概率密度,解决归一化常数计算难题。学习特征表示,最大化正样本对间的互信息。
问题视角将密度估计转化为二分类(数据 vs 噪声)。将表示学习转化为多分类(正样本 vs 众多负样本)。
输出意义模型输出是“样本来自真实分布”的概率模型输出是特征向量间的相似度
主要战场自然语言处理:语言模型、词向量训练(如Word2Vec的负采样)。自监督学习:图像、语音、多模态的对比学习(如SimCLR, MoCo)。
依赖先验需要显式定义噪声分布 P_n(x)不依赖具体的噪声分布,通过采样构建负样本。
计算关注点侧重于高效地近似归一化常数侧重于构造有效的正负样本对和调节温度参数τ
通俗比喻鉴宝专家:学习区分“真品”(数据)和“赝品”(噪声)。选秀评委:学习在众多选手中找出“最匹配”的那一个。

4.1 选型决策指南

根据上面的对比,我们可以得出一些实用的选型原则:

什么时候用NCE?

  • 当你真正关心概率值本身时。比如你在训练一个语言模型,最终需要它输出下一个词的确切概率分布来做预测或评估困惑度。NCE通过逼近归一化常数,其输出更接近真实的概率。
  • 当你有不错的噪声分布先验时。如果你对数据背后的噪声有一定了解(比如在NLP中,词频分布是一个很好的参考),利用这个先验的NCE可以学得更高效。
  • 经典场景:训练Word2Vec的Skip-Gram模型(负采样就是NCE的特例),以及一些早期的神经语言模型

什么时候用InfoNCE?

  • 当你只关心学到的特征表示好坏时。这是目前对比学习中最主流的场景。你不在乎模型输出的具体概率值,只在乎编码器提取的特征向量是否语义相关、是否在下游任务(如图像分类、目标检测)上表现好。
  • 当你的数据模态多样,难以定义噪声分布时。对于图像、音频、视频等,定义一个好的P_n(x)很难,不如直接采样构建负样本对来得直接。
  • 经典场景SimCLRMoCo等图像自监督学习框架;SimCSE等句子表示学习模型;以及CLIP等多模态对齐模型。

4.2 一个融合理解的视角

其实,从理论上看,InfoNCE可以看作是NCE的一个特殊变体。当我们把NCE中的噪声分布看作是一个均匀分布,并且把分类问题扩展到“一个正样本 vs 多个负样本”时,经过一些推导,NCE的目标就趋近于最大化互信息的下界,而这正是InfoNCE的目标。所以,你可以认为InfoNCE是NCE思想在表示学习领域的一个成功演进和特化

在我自己的项目中,如果做NLP相关的概率模型,我依然会考虑NCE或其变种(如负采样)。但只要是做视觉、语音或多模态的对比学习,InfoNCE几乎是默认的首选,它的社区生态、代码资源和调参经验都丰富得多。

5. 超越理论:实战中的技巧与坑点

理论懂了,代码看了,但一上手还是容易翻车。这里分享几个我从实战中总结出来的经验。

技巧1:温度参数τ是“灵魂” 温度τ是InfoNCE中最需要精心调校的超参数。它控制着对困难负样本的关注程度。我的经验是:

  • 图像领域τ通常在[0.05, 0.2]之间。SimCLR用的就是0.1。如果任务中类内差异大(比如同一只狗的不同姿态),可以稍微调大τ(如0.2),让学习更平滑。
  • 文本领域,由于语义空间更离散,有时用更小的τ(如0.05)效果更好,以放大细微的语义差异。
  • 监控损失值:如果训练初期损失就非常大或非常小,很可能是τ设置不当。一个健康的InfoNCE损失在训练初期应该在一个合理的范围(例如几到十几),然后稳步下降。

技巧2:负样本的质量和数量是关键 “垃圾进,垃圾出”在对比学习中尤其明显。

  • 数量:更多的负样本通常能带来更好的性能,因为它提供了更丰富的“对比背景”。MoCo框架的核心就是维护一个巨大的负样本队列。但也要考虑计算资源,一般256到4096都是常见范围。
  • 质量:要小心“假负样本”。比如在图像批次内采样负样本时,如果同一个批次里恰好有另一张“黑猫”的图片,它对于“黑猫”锚点来说其实是正样本,却被当成了负样本,这会混淆模型。解决方法包括使用更复杂的采样策略,或者在MoCo中使用动量编码器来生成更一致的负样本特征。

技巧3:特征归一化是“标配” 在使用余弦相似度计算InfoNCE时,务必对特征向量进行L2归一化。这能确保相似度得分被限制在[-1, 1]区间,让训练更稳定。从上面的代码你也可以看到,F.normalize是必不可少的一步。

踩过的坑:梯度爆炸与消失 早期实验时,我曾忘记除以温度τ,或者τ设得太小(如0.01),导致相似度得分经过Softmax后产生极端的概率分布(正样本概率接近1,负样本概率接近0)。这会造成两个问题:一是梯度变得极小(消失),模型停止更新;二是在某些情况下可能导致数值不稳定。所以,始终要对相似度得分进行缩放(除以τ),并从经验值开始调试。

6. 展望:更广阔的对比学习天地

NCE和InfoNCE为我们打开了对比学习的大门,但门后的世界更加广阔。基于InfoNCE,研究者们发展出了许多强大的变体和改进:

  • NT-Xent损失:其实就是归一化温度标度的交叉熵损失,是InfoNCE在多视图下的具体实现,SimCLR用的就是它。
  • Margin-based对比损失:如Triplet Loss,它引入了一个边际(margin),要求正负样本对之间的相似度差距至少超过这个值,目标更明确。
  • Soft-Nearest Neighbors Loss:考虑所有样本对的相似度,而不仅仅是一个正样本对,对噪声更鲁棒。

选择哪种损失,最终还是要回到你的任务本质:你是要一个概率模型,还是要一个特征提取器?你的数据模态是什么?你的计算预算有多少?想清楚这些问题,选择就变得清晰了。在我个人看来,InfoNCE及其衍生思想已经成为自监督学习的基石,它的简洁、有效和通用性,让“让模型自己通过对比来学习”这个梦想,变得越来越触手可及。下次当你需要从无标注数据中挖掘宝藏时,不妨试试这位强大的“裁判”。

Logo

北京人形旗下天工造物具身智能开源社区,聚焦具身天工与慧思开物两大平台

更多推荐