PyTorch 2.1+ 内置 torch.nn.ContrastiveLoss,支持 margin 和平方项,仅接受 1D 特征对输入;input1、input2 形状为 [N, D],target 为 [N] 的 0/1 张量;适用于 Siamese 网络等 pair-wise 场景。

PyTorch 没有内置的 ContrastiveLoss 类,但官方提供了现成可用的 torch.nn.ContrastiveLoss —— 仅限 PyTorch 2.1+;旧版本必须手动实现或从 torchvision / 第三方库导入。
PyTorch 2.1+ 直接用 torch.nn.ContrastiveLoss
这是最省事的方式,但要注意它和经典 Hadsell 等人论文中定义的公式略有差异(默认带 margin 和平方项),且只支持 1D 特征对输入。
-
margin参数控制负样本距离下界,默认为 1.0;设太小会导致负样本梯度消失,太大则抑制收敛 - 输入
input1和input2必须是形状一致的[N, D]张量,target是[N]的 0/1 张量:1 表示正样本对,0 表示负样本对 - 该 loss 对 batch 内每对样本独立计算,不建模全局相似性,适合 pair-wise 场景(如 Siamese 网络)
- 示例:
loss_fn = torch.nn.ContrastiveLoss(margin=1.0)<br>output1 = model(img1) # [32, 128]<br>output2 = model(img2) # [32, 128]<br>targets = torch.randint(0, 2, (32,)) # [32]<br>loss = loss_fn(output1, output2, targets)
PyTorch
旧版本用户或想显式控制公式细节(比如去掉平方、换 L1 距离、加温度系数)就得自己写。核心是:正样本拉近,负样本推远,但推远不能无限——靠 margin 截断。
- 别直接用
torch.dist或torch.norm在 batch 上逐对算,会触发 Python 循环;应向量化计算欧氏距离矩阵:torch.cdist(output1, output2, p=2)或手动广播:torch.sqrt(torch.sum((output1 - output2)**2, dim=1)) - 常见错误:把
target当作类别 ID 传入,实际它必须是 0/1 布尔语义;错传会静默导致 loss 值异常但不报错 - 数值不稳定点:当正样本对距离接近 0 时,
(d - margin)**2可能因浮点精度产生负值,建议加torch.relu包裹二次项 - 轻量实现示例:
def contrastive_loss(output1, output2, target, margin=1.0):<br> dist = torch.norm(output1 - output2, dim=1)<br> loss_pos = (1 - target) * dist**2<br> loss_neg = target * torch.relu(margin - dist)**2<br> return torch.mean(loss_pos + loss_neg)
注意这里target==1表示负对(与 PyTorch 官方相反),按需调整逻辑
对比学习中 Contrastive Loss 的典型使用陷阱
它不是 SimCLR、MoCo 那类 instance discrimination loss 的替代品,设计目标不同:Contrastive Loss 天然依赖人工构造的正/负样本对,泛化能力受限于配对质量。
立即学习“Python免费学习笔记(深入)”;
- 训练前必须确保正样本对确实语义一致(如同一图像的不同增强、同一 ID 的两张人脸),否则 loss 会强行扭曲特征空间
- 负样本不能全来自当前 batch(即 no cross-batch negative)——这会让模型退化为记忆 batch ID;如有条件,应构建 external memory bank 或用
torch.nn.CrossEntropyLoss配合 large-scale logits(如 InfoNCE) - batch size 过小(如 ≤16)时,负样本多样性不足,loss 容易震荡;建议 ≥64,并配合梯度累积模拟大 batch
- 特征维度 D > 512 时,欧氏距离易受维度灾难影响,此时改用余弦相似度更稳:
1 - F.cosine_similarity(output1, output2)
真正难的不是写出 loss 函数,而是让正负样本对的构造方式和模型 capacity 匹配——比如人脸场景用 bounding box 对齐做正样本,而细粒度分类可能得靠 attention mask 切分局部区域。这些不在 loss 里,但在数据 pipeline 里。


















