共计 2722 个字符,预计需要花费 7 分钟才能阅读完成。
为什么我们需要对比学习?
在传统监督学习中,我们严重依赖标注数据来训练模型。但在实际项目中,标注数据往往面临两大难题:

- 标注成本高:专业领域的标注需要专家参与,时间和经济成本都很高
- 样本稀缺:某些场景(如医疗图像)获取大量样本本身就困难
这些问题导致监督学习方法在真实场景中常常 ” 巧妇难为无米之炊 ”。
三种学习方式的对比
- 监督学习:需要完整标注,直接学习输入到标签的映射
- 优点:目标明确,评估直观
-
缺点:依赖标注数据量
-
自监督学习:利用数据自身结构创建伪标签
- 优点:减少标注依赖
-
缺点:预训练任务可能偏离下游任务
-
对比学习(Clap):通过比较样本间的相似度学习特征表示
- 核心思想:相似样本在表征空间中靠近,不相似样本远离
- 优势:特别适合特征提取任务,对标注数据依赖小
动手构建第一个 Clap 模型
环境准备
import torch
import torch.nn as nn
import torchvision.transforms as transforms
from torchvision.models import resnet18
数据增强 Pipeline
对比学习的效果高度依赖数据增强策略。以下是经典组合:
train_transform = transforms.Compose([transforms.RandomResizedCrop(224), # 随机裁剪并缩放到 224x224
transforms.RandomHorizontalFlip(), # 50% 概率水平翻转
transforms.ColorJitter(0.4, 0.4, 0.4, 0.1), # 颜色扰动
transforms.GaussianBlur(kernel_size=21), # 高斯模糊
transforms.ToTensor(),
transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])
])
模型架构
关键组件是双分支结构和投影头(projection head):
class ClapModel(nn.Module):
def __init__(self, feature_dim=128):
super().__init__()
# 骨干网络(这里使用 ResNet18)self.encoder = resnet18(pretrained=False)
self.encoder.fc = nn.Identity() # 移除原始全连接层
# 投影头
self.projector = nn.Sequential(nn.Linear(512, 256), # ResNet18 输出 512 维
nn.ReLU(),
nn.Linear(256, feature_dim)
)
def forward(self, x):
features = self.encoder(x)
return self.projector(features)
实现 InfoNCE 损失
这是对比学习的核心损失函数,数学表达式为:
$$
\mathcal{L} = -\log\frac{\exp(sim(z_i,z_j)/\tau)}{\sum_{k=1}^{2N} \mathbb{1}_{k\neq i} \exp(sim(z_i,z_k)/\tau)}
$$
代码实现:
def info_nce_loss(features, temperature=0.1):
"""
features: 投影后的特征,形状为[2*N, D](N 是 batch size)前 N 个是原始样本,后 N 个是对应的增强样本
"""
n = features.shape[0] // 2
labels = torch.cat([torch.arange(n) for _ in range(2)], dim=0)
labels = (labels.unsqueeze(0) == labels.unsqueeze(1)).float().to(device)
# 计算相似度矩阵
features = F.normalize(features, dim=1)
similarity = torch.matmul(features, features.T) / temperature
# 排除对角线(自己与自己比较)mask = torch.eye(labels.shape[0], dtype=torch.bool).to(device)
labels = labels[~mask].view(labels.shape[0], -1)
similarity = similarity[~mask].view(similarity.shape[0], -1)
# 选择正样本对
positives = similarity[labels.bool()].view(labels.shape[0], -1)
# 选择负样本对
negatives = similarity[~labels.bool()].view(similarity.shape[0], -1)
logits = torch.cat([positives, negatives], dim=1)
labels = torch.zeros(logits.shape[0], dtype=torch.long).to(device)
return F.cross_entropy(logits, labels)
实战调优建议
数据增强策略选择
- 基本原则:增强应该改变外观但保留语义
- 图像推荐组合:裁剪 + 翻转 + 颜色抖动 + 模糊
- 文本推荐:同义词替换 + 随机掩码 + 词序调整
批量大小与负样本
- 更大的 batch size 意味着更多负样本,但受限于显存
- 实用技巧:使用梯度累积模拟大批量
温度系数 τ 的经验值
- 通常设在 0.05 到 0.2 之间
- 太高导致学习缓慢,太低导致难收敛
常见问题与解决方案
梯度爆炸
- 症状:损失突然变为 NaN
- 对策:
- 梯度裁剪:
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) - 调小学习率
- 检查数据归一化
特征坍缩
- 表现:所有样本的特征高度相似
- 诊断:计算特征间的平均相似度(接近 1 表示坍缩)
- 解决方法:
- 增加更多样的负样本
- 使用可学习的温度参数
- 尝试解耦对比损失
计算资源不足
替代方案:
1. 使用更小的骨干网络(如 ResNet18)
2. 降低图像分辨率
3. 尝试内存高效的实现(如 FAIR 的 SwAV)
训练监控技巧
理想的训练曲线应呈现:
– 训练损失稳定下降(可能有波动)
– 验证准确率持续提升
– 特征相似度在 0.5-0.8 之间(避免坍缩)
进一步探索方向
- 不同的数据增强组合如何影响最终特征质量?
- 投影头的深度与特征维度是否存在最优配置?
- 如何将对比学习与其他自监督方法结合?
希望这篇指南能帮助你快速上手对比学习。记住,实践出真知 – 现在就用你的数据集试试这些方法吧!
正文完
