对比学习需三个核心组件:数据增强策略、编码器(如resnet)、对比损失函数(如ntxentloss),缺一不可;例如漏掉两次独立随机增强,simclr将无法运行。

对比学习需要哪些核心组件?
PyTorch 本身不内置对比学习框架,必须手动组合 torch.nn.Module、torch.optim 和自定义损失。关键组件就三个:数据增强策略、编码器(如 ResNet)、对比损失函数(如 NTXentLoss)。少一个都会训练失败——比如漏掉两次随机增强,SimCLR 就根本跑不起来。
常见错误现象:RuntimeError: Expected 4D input, but got 3D,通常是因为增强后没加 batch 维度,或 torchvision.transforms 返回了 PIL 图像而非 tensor;也可能是 forward() 里忘了调用 self.encoder(x),直接返回了原始输入。
- 编码器推荐用
torchvision.models.resnet18(pretrained=False),去掉最后的fc层,输出 512 维 embedding - 两次增强必须独立采样:用
torchvision.transforms.Compose分别定义两个 pipeline,不能复用同一个 transform 对象 - batch size 至少设为 256,小 batch 下
NT-Xent的负样本不足,loss 会震荡甚至发散
NT-Xent 损失怎么手写才不出错?
官方没提供现成的 NTXentLoss,自己实现时最容易错在温度系数 tau 的位置和相似度矩阵的对角 masking。温度值一般取 0.1,不是 1.0;mask 必须排除自身与自身的相似度(即主对角线),但要保留同一图像两个增强视图之间的正样本对(即 (i, i+N) 和 (i+N, i))。
示例关键片段:
Python 3.14.2是Python编程语言在2025年12月5日发布的稳定版本,属于3.14系列的第二个维护更新。该版本包含了18项修复,重点解决了多进程、数据类及正则表达式等模块的回归问题,并修复了CVE-2025-12084等安全漏洞。此版本标志着自由线程模式(移除GIL)正式获得官方支持,是Python发展的重要里程碑。
def nt_xent_loss(z_i, z_j, tau=0.1):
N = len(z_i)
z = torch.cat([z_i, z_j], dim=0) # [2N, D]
sim_matrix = torch.cosine_similarity(z.unsqueeze(1), z.unsqueeze(0), dim=2) / tau
sim_matrix = torch.exp(sim_matrix)
# mask out self-similarity
mask = torch.eye(2*N, dtype=torch.bool)
sim_matrix = sim_matrix * ~mask
# positive pairs: (i, i+N) and (i+N, i)
pos_sim = torch.cat([sim_matrix[:N, N:], sim_matrix[N:, :N]], dim=0)
neg_sum = sim_matrix.sum(dim=1) - torch.diag(sim_matrix)
loss = -torch.log(pos_sim / neg_sum).mean()
return loss
如何避免特征坍缩(collapse)?
训练中途 embedding 全变成几乎相同的向量,loss 接近 ln(2N-1),这就是坍缩。它不是 bug,而是对比学习的固有风险,根源在于编码器没被充分约束。
- 必须加 projection head:在 encoder 后接
nn.Sequential(nn.Linear(512, 512), nn.ReLU(), nn.Linear(512, 128)),输出 128 维再算 loss;只用 encoder 原始输出会坍缩 - BN 层不能冻结:如果加载 ImageNet 预训练权重,务必设
model.encoder.layer1[0].bn1.track_running_stats = True,否则 BN 统计失效导致梯度异常 - 学习率要比监督训练高 3–5 倍,常用
3e-4到5e-4;用torch.optim.AdamW+ 余弦退火更稳
无监督提取特征后怎么用?
训练完别急着删掉 projection head。下游任务要用的是 encoder 输出(去掉 projection),但必须用原始训练时的增强方式做 inference:即单次 resize+crop+normalize,不是两次。否则特征分布偏移,KNN 分类准确率掉 10% 以上。
典型用法:
- KNN 评估:冻结 encoder,用训练集 embedding 训练
sklearn.neighbors.NearestNeighbors,在测试集上查 20-NN 投票 - 线性 probe:在 encoder 后接一个
nn.Linear(512, num_classes),只训练该层,其余 freeze,用cross_entropy训练 100 epoch - 微调(fine-tune):解冻全部参数,用小学习率(如
1e-5)在标注数据上继续训,注意要关掉 projection head
真正麻烦的是数据增强一致性——训练对比学习用的 augmentations(比如 GaussianBlur、Solarize)不能直接用于线性 probe 的训练阶段,后者应改用标准 Resize+CenterCrop,否则评估结果不可比。
Python免费学习笔记(深入):立即使用
在学习笔记中,你将探索 Python 的核心概念和高级技巧!










