共计 1368 个字符,预计需要花费 4 分钟才能阅读完成。
什么是对比学习?
对比学习就像教 AI 玩『找不同』游戏。假设你有一张猫的照片,我们通过裁剪、变色等手段生成它的两个『变体』(正样本),其他动物的照片就是负样本。模型的目标是让正样本在特征空间里紧紧挨着,同时把负样本推得远远的。这种『同类靠近,异类远离』的思想,就是对比学习的核心。

主流模型家族
- SimCLR:谷歌提出的『简单粗暴』代表
- 优点:结构清晰,单批次内所有样本互为负样本
-
缺点:显存消耗大(因为需要大批次)
-
MoCo:Facebook 的『记忆银行』方案
- 优点:用队列存储历史负样本,小批次也能获得大量负样本
- 缺点:实现复杂度较高,需维护动量编码器
PyTorch 实战 SimCLR
数据增强模块
import torchvision.transforms as transforms
train_transform = transforms.Compose([transforms.RandomResizedCrop(32), # 随机裁剪到 32x32
transforms.RandomApply([transforms.ColorJitter(0.8,0.8,0.8,0.2)], p=0.8), # 颜色抖动
transforms.RandomGrayscale(p=0.2), # 20% 概率转灰度
transforms.ToTensor()])
投影头网络(关键!)
class ProjectionHead(nn.Module):
def __init__(self, input_dim=2048, hidden_dim=512, output_dim=128):
super().__init__()
# 两层 MLP,输出 128 维对比特征
self.net = nn.Sequential(nn.Linear(input_dim, hidden_dim),
nn.ReLU(),
nn.Linear(hidden_dim, output_dim)
)
NT-Xent 损失函数
def nt_xent_loss(z1, z2, temperature=0.5):
# 拼接所有特征
z = torch.cat([z1, z2], dim=0) # [2*bsz, dim]
# 计算余弦相似度矩阵
sim = torch.matmul(z, z.T) / temperature # [2*bsz, 2*bsz]
# 对角线位置是正样本对
labels = torch.arange(z.size(0)).to(device)
return F.cross_entropy(sim, labels)
训练避坑指南
-
批次大小 :SimCLR 需要大批次(至少 512),否则负样本不足会导致性能骤降
-
学习率 :推荐使用带 warmup 的余弦退火调度
-
梯度爆炸 :对投影头输出做 L2 归一化(torch.nn.functional.normalize)
CIFAR-10 基准测试
| 模型 | 线性评估准确率 | 训练耗时(GPU 小时) |
|---|---|---|
| SimCLR | 78.2% | 3.5 |
| 有监督基线 | 85.1% | 1.2 |
留给读者的思考题
-
当显存不足时,有哪些替代方案可以增加有效负样本数量?
-
为什么颜色抖动对图像对比学习特别重要?
-
如何设计针对文本数据的对比学习增强策略?
写完这个 demo 后发现,对比学习就像 AI 界的『无师自通』高手——不需要人工标注,自己就能从数据中挖掘规律。虽然调参过程有点头疼,但看到模型自动学会区分不同物体时,那种成就感绝对值得一试!
正文完
