拓冰建站拓冰建站
首页 / 资讯中心 / 正文

对比学习实战:用PyTorch实现SimCLR自监督表示学习

老读者应该知道我一直在这个PyTorch实战系列里折腾各种模型从分类、检测到GAN绕了一圈这期终于写到对比学习Contrastive Learning和自监督表示学习了。之所以在这个节点专门开一期写它是因为我去年在实际项目里被标注数据卡了脖子甲方手里有几万张产品图片但标注好的只有几百张分类器怎么训都差口气。后来我试着用对比学习做无监督预训练再在下游任务上微调效果直接提升了一大截这才意识到这类方法是真能解决业务痛点的不是学术圈自嗨。这期内容打算从对比学习解决的问题讲起把InfoNCE损失、数据增强、PyTorch实现、训练避坑一条线串下来。不管你是刚入门深度学习的同学还是在业务里遇到无标注数据利用率低问题的工程师看完应该都能自己搭一套对比学习流程出来。1. 为什么自监督表示学习成了深度学习绕不开的话题1.1 标注数据的成本与瓶颈先聊一个现实问题深度学习真的很能吃数据但更吃标注好的数据。我见过太多团队卡在这个环节——数据采集容易摄像头一挂、爬虫一跑几万张甚至几十万张图都有了可让标注公司或者实习生去一张张打标签既烧钱又慢更何况有些场景压根就不是人眼能准确标注的。你用北京交通大学 深度学习 期末试题、深度学习知识点这类关键词去搜会发现教材和课程里讲的大多是监督学习的成熟套路给一堆图片标签对然后把模型训出来。可一旦落到真实业务标签这个前提往往不成立或者严重不足这时候模型就难产。自监督表示学习的思路和这个正好相反它不依赖人工标签而是想办法从数据本身构造监督信号。典型做法有预测图像的旋转角度、补全被遮挡的区域以及这期要重点讲的对比学习。它通过设计一个代理任务让模型先在大规模无标注数据上学出一套通用的特征表示之后再把这套表示迁移到具体任务上。相当于先给模型上了一堂通识课再进专业课。1.2 对比学习的基本逻辑与发展脉络对比学习的核心逻辑可以总结成一句话让模型学会分辨相似与不相似。具体操作是对同一张图做两次不同的随机增强得到两个视图它们应该被映射到特征空间中相近的位置也就是正样本对而不同图片的视图则应该被拉远成为负样本对。这个思路最早的雏形之一来自自编码器和Word2Vec的负采样思想——让模型学会通过上下文预测中心词或者区分真实样本与噪声样本。到了2020年SimCLR、MoCo、BYOL这几篇工作相继出现把对比学习推上了风口浪尖。SimCLR证明了简单的数据增强组合加上InfoNCE损失就能学到非常强的视觉特征MoCo引入了动量编码器和队列机制解决了大batch size依赖的问题BYOL更是直接抛掉了负样本纯靠预测一致性也能训练。从实际工程角度看这个方法最吸引人的地方在于预训练阶段完全不需要标注下游任务哪怕标签很少也能拿到不错的性能。我去年那个项目最后就是用SimCLR在几万张无标签产品图上预训练ResNet18再用几百张标注数据微调分类头准确率比直接训一个随机初始化的模型高了将近17个百分点。所以这绝对不是一个只能写在论文里的方法。2. InfoNCE损失函数对比学习的发动机2.1 从拉近正样本、推远负样本到形式化定义既然要说对比学习就绕不开它最核心的训练目标——InfoNCE损失。很多朋友刚开始看对比学习的代码时会陷在相似度矩阵、mask、label构造这些东西里觉得绕。实际上拆开来看它做的事情非常简单。假设一个batch里有N张图每张图经过两次数据增强变成2N个样本。模型需要判断的是每个样本应该和谁配对对第i个样本来说它的正样本是同一张图生成的另一个视图也就是第iN个样本剩下的2N-2个样本全都是负样本。InfoNCE的核心要求就是让正样本对的相似度在所有这些样本中脱颖而出越大越好。数学表达上它看起来像这样[ L_i -\log \frac{\exp(sim(z_i, z_{iN}) / \tau)}{\sum_{k1}^{2N} [k \neq i] \exp(sim(z_i, z_k) / \tau)} ]其中sim通常用余弦相似度也就是把特征向量归一化后求内积τ是温度系数控制分布的尖锐程度。从形式上看这跟一个2N类别的softmax分类器输出完全一样——只不过我们要分类的标签是样本在哪个位置。2.2 温度系数τ的直觉理解温度系数τ在训练里扮演的角色很容易被忽视但它对效果影响极大。我通常这么跟同事解释相似度除以τ之后相当于把原始相似度值放大了。τ越小logits之间的差异越显著softmax输出的概率分布越尖锐模型会变得苛刻只重点关注那个最相似的正样本对困难负样本更敏感τ越大分布越平缓所有负样本被一视同仁地拉远模型变得佛系学习信号会弱很多。实际操作里τ设置太大会导致loss降不下去特征趋向于糊在一起τ太小则模型容易过拟合到批次里的个别困难样本训练震荡。SimCLR原文在ImageNet上用的0.07在CIFAR上用的0.1我自己的经验是CIFAR级别的小数据集在0.1到0.2之间波动问题不大如果batch size比较大可以用稍小的τ比如0.07到0.1。2.3 为什么需要足够多的负样本很多初学者会有一个疑问每组正样本只对应一张图里的两个视图那负样本又有什么用答案是负样本的数量和多样性直接决定了特征空间的分布形状。如果负样本太少模型只需要识别出这两张图很像就能完成目标学到的特征是退化的、低区分度的。负样本足够多且多样时模型被迫把整个特征空间里的不同类别都推开这样学到的表示才具备泛化能力。这也是为什么SimCLR原文使用4096到8192的超大batch size——本质上不是因为显存多到没用处而是因为batch里能提供的负样本数量直接决定了表示质量。在实际代码里InfoNCE通常用全batch的相似度矩阵一次性计算这也是它比传统Siamese网络训练效率更高的原因之一。PyTorch里几行就能写出来import torch import torch.nn.functional as F def info_nce_loss(z_i, z_j, temperature0.1): batch_size z_i.size(0) # 拼接成 2N 个样本的特征 [2N, D] z torch.cat([z_i, z_j], dim0) # L2归一化保证内积等于余弦相似度 z F.normalize(z, dim1) # 相似度矩阵 [2N, 2N] logits torch.mm(z, z.t()) / temperature # 注意对角线是自己和自己的相似度要排除掉 mask torch.eye(2 * batch_size, dtypetorch.bool, devicez.device) logits.masked_fill_(mask, float(-inf)) # 每个样本的正样本索引 # 前 N 个样本(来自 z_i)的正样本是后 N 个(来自 z_j) # 后 N 个样本的正样本是前 N 个。 labels torch.cat([ torch.arange(batch_size, 2 * batch_size), torch.arange(0, batch_size) ], dim0).to(z.device) return F.cross_entropy(logits, labels)这个代码里最容易出错的点有两个。一是必须要mask掉对角线否则模型会走捷径只要学一个恒等映射就能让loss非常低但实际上啥都没学会二是labels的构造要小心别把正样本索引搞反。我刚开始复现的时候在这两个地方各栽了一次跟头特征可视化出来全是乱的排查了半天才发现是标签构造问题。3. 数据增强对比学习里最容易被低估的组件3.1 效果差异的根源不一定在模型结构而在增强SimCLR论文里有一个很反直觉的结论数据增强组合对整个表示学习效果的影响比模型结构大得多。作者甚至直接说增强的多样性是对比学习比传统方法更有效的关键原因。当时看到这个结论我还挺意外的直到自己复现实验对比之后才真正信服。为什么增强这么重要因为对比学习学习的是不变性——我们希望模型忽略掉增强带来的变化只保留图像中与语义相关的信息。如果增强方式太单一模型能轻松地找到捷径来区分正负样本但又学不到真正的语义特征。多样化的增强组合逼迫模型提取更本质的内容。3.2 常用增强组合与实战配置目前最常用的增强组合是从SimCLR总结出来的那套包括随机裁剪缩放、颜色抖动、随机灰度化和高斯模糊。但注意这些增强并非在每个数据集上都要全用。以下是我在CIFAR-10和ImageNet级数据上的配置参考增强操作CIFAR-1032×32ImageNet级别224×224随机裁剪缩放是scale(0.2, 1.0)是scale(0.08, 1.0)随机水平翻转是是颜色抖动是强度0.4~0.5概率0.8是强度0.4~0.5概率0.8随机灰度化是概率0.2是概率0.2高斯模糊否图太小模糊容易丢信息是核大小不定概率0.1~0.5CIFAR这类小分辨率图上高斯模糊作用有限有时候还会因为模糊掉边缘细节导致预训练效果下降我一般直接省略。真正核心的两个操作是随机裁剪和颜色抖动。单独用随机裁剪时模型学到的主要是物体形状和不完整结构信息加上颜色抖动后特征对颜色变化的鲁棒性会显著提升在后续下游任务中表现更稳。PyTorch的torchvision.transforms里可以直接组装import torchvision.transforms as T def get_simclr_augment(input_size32, blurFalse): augmentations [ T.RandomResizedCrop(sizeinput_size, scale(0.2, 1.0)), T.RandomHorizontalFlip(), T.RandomApply([T.ColorJitter(0.4, 0.4, 0.4, 0.1)], p0.8), T.RandomGrayscale(p0.2), ] if blur: augmentations.append(T.GaussianBlur(kernel_size3, sigma(0.1, 2.0))) augmentations.extend([ T.ToTensor(), T.Normalize([0.4914, 0.4822, 0.4465], [0.2023, 0.1994, 0.2010]), ]) return T.Compose(augmentations)注意一个容易忽略的细节同一个batch里做增强时要用两份独立的transform对象或者是每次forward独立调用transform这样同一次迭代内的两次视图才是不同的增强结果。你要是把两次视图都用同一个预先生成的Tensor那正样本对完全一致对比学习就毫无意义了。3.3 增强强度需要根据数据集调整数据增强不是越强越好这一点我在实际训练里吃过亏。有一段时间在CIFAR-10上用比较强的颜色抖动把ColorJitter的强度直接拉到0.8结果预训练出来的特征在下游分类任务上反而掉了三四个点。原因也很简单CIFAR-10物体本身颜色信息就很重要比如青蛙、鸟这类类别颜色被过度扰动之后模型学到的语义特征就会被弱化。所以实操中增强强度的选择建议是先保持SimCLR原文的默认值训练一小段时间后观察InfoNCE的loss曲线和下游任务的表现再根据数据集的特点微调。如果你的任务对颜色非常敏感就适当降低ColorJitter的强度如果对形状、纹理更敏感可以稍微加强。这里没有一劳永逸的配置要做几次消融实验才能找到适合自己数据集的组合。4. PyTorch实现SimCLR从数据管道到训练循环4.1 项目结构与核心组件到这一步我准备把一套完整的SimCLR训练流程拆开来讲。代码会以CIFAR-10为例因为它在单卡GPU上就能跑非常适合做实验验证。整个实现分四部分数据增强管道、编码器与投影头、InfoNCE损失、训练循环与日志。项目文件结构大概长这样simclr_cifar/ ├── data/ # 数据集存放 ├── augment.py # 数据增强 ├── model.py # 编码器 投影头 ├── loss.py # InfoNCE损失 ├── train.py # 训练循环 └── eval.py # 线性评测编码器用ResNet18就足够了在CIFAR-10上没必要上更大的模型反而是投影头对效果的影响值得注意。SimCLR原论文的结论是在特征输出之后接一个两层的MLP投影头效果好于直接用编码器输出的特征做对比损失。直觉上这个MLP起到一个过滤信息的作用它把编码器提取的、可能包含过多任务无关细节的特征压缩成一个128维的低维表示在这个空间里计算相似度会让模型更专注于区分正负样本。4.2 编码器与投影头的实现模型定义比较简单关键在ResNet输出的global pooling之后要接一个projection headimport torch import torch.nn as nn from torchvision.models import resnet18 class EncoderWithProjection(nn.Module): def __init__(self, base_encoderresnet18, feat_dim512, proj_dim128): super().__init__() self.encoder base_encoder(num_classesfeat_dim) self.projection nn.Sequential( nn.Linear(feat_dim, feat_dim), nn.ReLU(inplaceTrue), nn.Linear(feat_dim, proj_dim), ) def forward(self, x): feat self.encoder(x) # [N, 512] z self.projection(feat) # [N, 128] return z这里有一个小细节SimCLR原论文在计算损失时用的是投影头的输出z但在预训练结束后把投影头丢掉只用编码器输出的特征feat去做下游任务。原因在于z所在的空间是被对比损失扭曲过的投影头学到了大量与增强相关的信息这些信息对下游任务不一定有利。我第一次做下游任务时直接拿了z做分类效果比用feat差了不少后来看了代码才发现问题。4.3 训练循环的完整代码训练循环本身不复杂但要特别注意每个batch的图片要分别用增强管道处理两次得到两组视图。这两组视图在同一个模型里过前向然后送入InfoNCE损失。import torch from torch.utils.data import DataLoader from torchvision.datasets import CIFAR10 from augment import get_simclr_augment from model import EncoderWithProjection from loss import info_nce_loss device torch.device(cuda if torch.cuda.is_available() else cpu) batch_size 256 # 注意用两个独立的增强对象确保同一个batch的两组视图不同 transform_train get_simclr_augment(input_size32, blurFalse) dataset CIFAR10(root./data, trainTrue, downloadTrue, transformNone) loader DataLoader(dataset, batch_sizebatch_size, shuffleTrue, num_workers4, drop_lastTrue) model EncoderWithProjection().to(device) optimizer torch.optim.Adam(model.parameters(), lr1e-3) def train_one_epoch(epoch): model.train() total_loss, total_num 0.0, 0 for images, _ in loader: # 同一个原始batch独立做两次增强 xi torch.stack([transform_train(img) for img in images]).to(device) xj torch.stack([transform_train(img) for img in images]).to(device) zi model(xi) zj model(xj) loss info_nce_loss(zi, zj, temperature0.1) optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() * images.size(0) total_num images.size(0) avg_loss total_loss / total_num print(fEpoch [{epoch}] Loss: {avg_loss:.4f}) return avg_loss for epoch in range(1, 301): train_one_epoch(epoch) if epoch % 50 0: torch.save(model.state_dict(), fcheckpoints/simclr_epoch{epoch}.pth)这里用torch.stack([transform_train(img) for img in images]其实不太高效更好的方式是把transform放到Dataset的__getitem__里去每个样本被取出来时就被增强然后在DataLoader里加载。但为了结构简单直观我在这里展示的是循环内逐张增强的版本实际大规模训练时建议定义成transformtransform_train的Dataset并在__getitem__里调用两次增强返回(xi, xj)这样DataLoader的num_workers多进程机制能有效加速。4.4 超参数配置与优化器选择SimCLR原文用的是LARS优化器这种优化器针对超大batch训练的稳定性做了专门设计。但我们在单机单卡上跑CIFARbatch size不会太大用Adam甚至SGD都能稳定收敛。我自己试下来Adam配一个3e-4到1e-3的学习率就很稳SGD则需要配合余弦退火和momentum效果差不多但对学习率更敏感。关于训练时长对比学习比监督学习要花更多epoch才能看到效果的收敛。CIFAR-10上300个epoch是起步500到1000个epoch特征质量会更上一个台阶。如果机器资源紧张可以先用200个epoch跑通全流程确认没bug之后再加大训练规模。还有一个常用的优化技巧在训练过程中不直接使用info_nce_loss(zi, zj)而是把zi和zj拼接后统一计算损失再通过torch.utils.checkpoint来节省显存。不过这在后文踩坑部分再细说。5. 训练过程中的坑退化、batch size与温度系数5.1 模型退化学出来的特征全是常数我第三次自己手写对比学习训练时遇到了一个非常让人抓狂的现象InfoNCE的loss下降得很快很快就收敛到一个很低的值但拿编码器输出特征去做下游分类准确率和随机猜测差不多。翻开log一看特征可视化结果里所有样本都挤在了同一个点附近完全没有区分度。这个现象在对比学习里叫完全退化collapse也就是说模型找到了一个无用解——把所有输入都映射成同一个输出向量这样正样本对的相似度很高因为都是同一个点负样本对相似度也很高但softmax经过温度系数放大后依然能够区分出自己和其他人尽管区别可能很微弱loss表面上会很低但没有学到有效的语义表示。发生退化最常见的原因是几个数据增强太弱模型轻松找到捷径温度系数太小模型过度自信容易陷入局部最优投影头太简单特征表达能力不足负样本不足尤其是batch size太小。解决退化的方案主要靠上面几点逐一排查。我实验中最有效的路线是加大随机裁剪的scale范围比如从(0.2, 1.0)改成(0.08, 1.0)、把温度系数从0.1调到0.07、把投影头的hidden维度从256提升到512三个方向组合之后基本就能避免退化。另外训练初期可以打印torch.cdist特征分布的标准差如果发现特征的标准差趋近于零说明模型正在退化赶紧调整超参。5.2 batch size不够怎么办内存银行与动量编码器对于没有8卡甚至更多GPU的普通玩家来说SimCLR动辄几千的batch size实在奢侈。但如果你只是想在CIFAR或小型数据集上做实验batch size 256到512其实已经够用了。更大的问题在于如果训练数据来自不同类别且类别间相似度很高小batch提供的负样本可能不够多样导致表示质量下降。解决这个问题有两条路。第一条是MoCo的动量编码器加队列用一个队列缓存历史batch的特征作为额外负样本负样本池的大小就不受当前batch size限制了代价是缓存的特征来自旧版本的编码器所以编码器参数要缓慢更新保持一致性。第二条是memory bank思路类似但更简单把每个样本的历史特征存起来缺点是历史特征可能与当前编码器输出差异较大导致优化不稳定。我后来在项目里用的是简化版的MoCo思路实现起来也不复杂给编码器维护一份参数的moving average副本每轮参数更新之后把最旧的batch的特征从队列里替换出去。效果上差不多能用配合batch size 256在下游任务上拿到和local batch size 2048的SimCLR基本相当的结果。5.3 温度系数的敏感性一个需要实验验证的超参温度系数τ是InfoNCE中最需要认真调的超参它直接决定了模型对困难负样本的关注程度。我做一个直观的实验对比τ取值训练稳定性特征质量线性评测备注0.5很稳loss平滑下降低约50%~60%负样本被一视同仁区分度不足0.1稳定高约75%~85%多数情况下的好默认值0.05训练初期易震荡最高约80%~88%对模型与数据增强要求更高0.01极易发散或退化极低模型过度苛刻学不动当然以上数字都是CIFAR-10上单卡训练的近似参考换数据集或模型结构结论可能会变。但趋势是有参考价值的τ太大特征太平庸τ太小训练不稳。我建议新手直接先跑τ0.1稳定跑通后做一组0.07和0.15的对比用下游任务来判断优劣。5.4 SyncBN与多卡训练的小提醒如果你的数据集比较大、决定上多卡训练还有一个很容易踩的坑BN层的统计量。PyTorch默认的DataParallel或DistributedDataParallel在每个GPU上独立计算BN统计量但同一张图的两种视图不一定在同一张卡上跨卡的正负样本对信息就丢了。SimCLR原文很明确地表示使用SyncBN同步BN对训练效果提升明显。PyTorch里使用SyncBN不算复杂如果你的代码用了DistributedDataParallel把模型里的nn.BatchNorm2d换成nn.SyncBatchNorm.convert_sync_batchnorm包一层即可。我见过很多人在单卡上调试得好好的一上多卡效果反而变差查了一圈发现就是BN没有同步导致的。6. 特征质量评估线性评测与可视化验证6.1 线性评测Linear Probe最常用的定量评估方式对比学习预训练做完你不能直接说学好了得用一个靠谱的评估方法来量化表示的质量。最常用也最省事的方法叫线性评测Linear Probe把编码器参数冻结只训练一个线性分类头用分类准确率来衡量特征的可分性。为什么要用线性分类头而不是微调整个网络因为如果用全量微调编码器自身还会继续适应任务无法反映预训练特征本身的质量。线性分类头只允许在特征平面上画超平面这些特征必须具备足够的线性可分性才说明预训练学到了不错的表示。import torch from torchvision.models import resnet18 from sklearn.linear_model import LogisticRegression # 加载预训练模型 model resnet18(num_classes512) state torch.load(checkpoints/simclr_epoch300.pth, map_locationcpu) model.load_state_dict({k.replace(encoder., ): v for k, v in state.items() if encoder. in k}) model.eval() # 提取所有训练集与测试集的特征 def extract_features(dataloader): features, labels [], [] with torch.no_grad(): for x, y in dataloader: f model.encoder(x) features.append(f) labels.append(y) return torch.cat(features).numpy(), torch.cat(labels).numpy() train_feats, train_labels extract_features(train_loader) test_feats, test_labels extract_features(test_loader) # Logistic回归线性评测 clf LogisticRegression(max_iter1000, solverlbfgs) clf.fit(train_feats, train_labels) acc clf.score(test_feats, test_labels) print(fLinear Probe Accuracy: {acc:.4f})在CIFAR-10上用这个流程无监督预训练300个epoch线性评测准确率大概能到75%到85%区间具体取决于增强配置和随机种子。作为参照从随机初始化直接训练同一个线性分类头准确率通常只有20%到30%差距还是很明显的。如果你用同样的数据训一个监督ResNet18准确率可能到95%无监督预训练跟这个差距是客观存在的但它完全不需要标签这个trade-off在业务里往往非常划算。6.2 KNN评测与t-SNE可视化线性评测之外KNN分类也是一种非常直接的评估方式用测试样本的特征在训练集特征库里找K个最近邻然后投票决定类别。KNN的好处是完全不需要训练任何附加参数结果更诚实适合快速验证特征的质量。t-SNE可视化则属于定性分析手段。把测试集的特征降维到二维平面按标签着色如果学得好同一类别的点会聚成明显的簇不同类别的簇分得比较开。但t-SNE对参数的敏感性很高perplexity设置得不好容易产生误导比如看起来分散的点其实在原始空间里非常近。所以我建议把t-SNE当作调试工具而不是最终结论。6.3 消融实验怎么知道哪些模块在起作用在你调模型、调增强、调温度系数的过程中做一组简单的消融实验会非常有帮助。比如固定所有超参只去掉颜色抖动对比linear probe结果看掉了几个点或者只把投影头删掉直接用编码器特征做对比损失看看MLP到底贡献了多少。这类实验的意义不光在于发论文时能画一张好看的ablation表更在于能帮你建立直觉——知道你的模型到底在靠什么起作用。比如我就在自己的数据集上发现没有颜色抖动的话特征对光照变化极度敏感下游换一批不同光照条件下拍的照片性能就崩这是不做消融实验永远发现不了的问题。总结这部分评估环节虽然不参与训练但它是你迭代实验的眼睛。没有一套稳定的评估流程你会在超参调优里越陷越深甚至不知道自己调的方向对不对。回看这期内容从对比学习的原理、InfoNCE的实现、数据增强的策略到训练和评估的完整流程基本把SimCLR这条路线串清楚了。我在实际项目里最深的感触是对比学习的门槛不在于理解概念也不在于能不能复现代码而在于你愿不愿意在数据增强、温度系数、batch size这些细节上反复实验。这套方法不像监督学习那样有明确的标准答案很多配置需要针对数据特点去调整。如果你准备在自己业务里尝试建议先找一个小数据集把整个流程跑通再逐步放大到无标注数据上做预训练过程中时刻盯住linear probe的准确率那是检验一切超参调整的试金石。
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门