对比学习实战:SimCLR框架的表征学习与下游任务迁移

对比学习的基本原理与SimCLR框架定位

对比学习(Contrastive Learning)是机器学习领域中一种自监督表征学习方法,核心思想是通过拉近正样本对、推远负样本对来学习数据的深层特征表示。SimCLR(Simple Contrastive Learning of Visual Representations)由Google Research团队提出,是对比学习领域最具代表性的框架之一,其简洁的架构设计和优异的性能表现使其成为视觉表征学习的主流方案。人工智能领域的自监督学习近年来受到广泛关注,SimCLR通过数据增强构造正样本对,利用对比损失函数优化编码器,无需人工标注即可学到高质量的视觉特征。

SimCLR框架架构详解

SimCLR的整体流程分为四个关键步骤:数据增强、编码器提取特征、投影头映射到对比空间、对比损失计算。每个Batch的每张图像经过两种不同的随机增强操作,生成两个增强视图,同一图像的两个视图构成正样本对,同一Batch内其他图像的增强视图构成负样本对。

数据增强策略对SimCLR的性能影响显著。常用增强包括随机裁剪、颜色抖动、高斯模糊、水平翻转等。其中颜色抖动对性能贡献最大,单一裁剪加颜色抖动的组合已经能取得不错的效果。多种增强组合的叠加使用可以进一步提升表征质量。

import torch
import torch.nn as nn
import torchvision.transforms as transforms
from torchvision.models import resnet50

class SimCLRAugment:
    '''SimCLR数据增强模块'''
    def __init__(self, image_size=224):
        self.transform = transforms.Compose([
            transforms.RandomResizedCrop(image_size, scale=(0.2, 1.0)),
            transforms.RandomHorizontalFlip(p=0.5),
            transforms.RandomApply(
                [transforms.ColorJitter(0.4, 0.4, 0.4, 0.1)], p=0.8
            ),
            transforms.RandomGrayscale(p=0.2),
            transforms.RandomApply([transforms.GaussianBlur(kernel_size=3)], p=0.5),
            transforms.ToTensor(),
            transforms.Normalize(mean=[0.485, 0.456, 0.406],
                               std=[0.229, 0.224, 0.225])
        ])
    
    def __call__(self, x):
        return self.transform(x), self.transform(x)

编码器与投影头设计

SimCLR的编码器通常采用ResNet系列网络,输入图像经过编码器后得到2048维特征向量。投影头(Projection Head)是一个两层的MLP结构,将编码器输出映射到128维对比空间。投影头在训练阶段使用,在下游任务迁移时丢弃,直接使用编码器输出特征。

class SimCLRModel(nn.Module):
    '''SimCLR完整模型:编码器 + 投影头'''
    def __init__(self, proj_dim=128, hidden_dim=512):
        super().__init__()
        self.encoder = resnet50(pretrained=False)
        self.feature_dim = self.encoder.fc.in_features
        self.encoder.fc = nn.Identity()
        
        self.projector = nn.Sequential(
            nn.Linear(self.feature_dim, hidden_dim),
            nn.BatchNorm1d(hidden_dim),
            nn.ReLU(inplace=True),
            nn.Linear(hidden_dim, proj_dim)
        )
    
    def forward(self, x1, x2):
        h1 = self.encoder(x1)
        h2 = self.encoder(x2)
        z1 = self.projector(h1)
        z2 = self.projector(h2)
        return h1, h2, z1, z2

NT-Xent对比损失函数实现

NT-Xent(Normalized Temperature-scaled Cross Entropy)损失是SimCLR的核心,通过温度系数控制正负样本对的区分程度。温度参数越小,模型对困难负样本的关注度越高。对于Batch Size为N的情况,每个样本有2N-2个负样本对。

class NTXentLoss(nn.Module):
    '''NT-Xent对比损失'''
    def __init__(self, temperature=0.5):
        super().__init__()
        self.temperature = temperature
        self.cos_sim = nn.CosineSimilarity(dim=-1)
    
    def forward(self, z1, z2):
        batch_size = z1.size(0)
        z1 = nn.functional.normalize(z1, dim=-1)
        z2 = nn.functional.normalize(z2, dim=-1)
        z = torch.cat([z1, z2], dim=0)
        sim = torch.matmul(z, z.T) / self.temperature
        mask = torch.eye(2 * batch_size, device=z.device).bool()
        sim.masked_fill_(mask, -1e9)
        labels = torch.cat([
            torch.arange(batch_size, 2 * batch_size),
            torch.arange(0, batch_size)
        ], dim=0).to(z.device)
        loss = nn.functional.cross_entropy(sim, labels)
        return loss

训练流程与超参数配置

SimCLR的训练高度依赖Batch Size,论文实验表明Batch Size越大,负样本越多,性能越好。Batch Size从256到8192不等,学习率与Batch Size线性缩放。训练轮数通常为500-1000 epoch,采用LARS优化器和余弦退火学习率调度。

def train_simclr(model, dataloader, epochs=100, lr=0.001, temperature=0.5):
    optimizer = torch.optim.LARS(
        model.parameters(), lr=lr, weight_decay=1e-6
    )
    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
        optimizer, T_max=epochs
    )
    criterion = NTXentLoss(temperature=temperature)
    model = model.cuda()
    
    for epoch in range(epochs):
        total_loss = 0
        for batch_idx, (x1, x2) in enumerate(dataloader):
            x1, x2 = x1.cuda(), x2.cuda()
            h1, h2, z1, z2 = model(x1, x2)
            loss = criterion(z1, z2)
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()
            total_loss += loss.item()
        scheduler.step()
        avg_loss = total_loss / len(dataloader)
        if (epoch + 1) % 10 == 0:
            print(f"Epoch {epoch+1}/{epochs}, Loss: {avg_loss:.4f}")

下游任务迁移与线性评估

SimCLR训练完成后,冻结编码器参数,在其输出特征上训练一个线性分类器,评估表征质量。线性评估协议是对比学习的标准评估方式,能直接反映预训练特征的可分性。在ImageNet数据集上,SimCLR线性评估的Top-1准确率可达76.5%,接近全监督训练水平。

def linear_evaluation(encoder, train_loader, val_loader, lr=0.01, epochs=100):
    encoder.eval()
    for param in encoder.parameters():
        param.requires_grad = False
    
    classifier = nn.Linear(2048, 1000).cuda()
    optimizer = torch.optim.SGD(classifier.parameters(), lr=lr, momentum=0.9)
    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs)
    criterion = nn.CrossEntropyLoss()
    
    for epoch in range(epochs):
        classifier.train()
        for images, labels in train_loader:
            images, labels = images.cuda(), labels.cuda()
            with torch.no_grad():
                features = encoder(images)
            outputs = classifier(features)
            loss = criterion(outputs, labels)
            optimizer.zero_grad()
            loss.backward()
            optimizer.step()
        scheduler.step()
        if (epoch + 1) % 10 == 0:
            validate(classifier, encoder, val_loader)

SimCLR的工程优化建议

大Batch训练对GPU显存消耗很大,8192的Batch Size需要128张TPU V3或大量A100 GPU。对于资源有限的场景,可以采用Memcached-based负样本缓存策略,将历史Batch的特征向量缓存为额外负样本,用较小的Batch Size模拟大Batch效果。另一种方案是使用MoCo的动量编码器思路,维护一个特征队列作为负样本池。

数据增强的选择需要根据任务特点调整。自然图像推荐使用裁剪加颜色抖动加模糊的组合;医学影像等灰度图像则应移除颜色相关增强,增加旋转、弹性形变等几何变换。温度系数的推荐范围为0.1到0.5,值越小对困难负样本越敏感,过大则损失对正负样本的区分能力下降。

原创文章,作者:小编,如若转载,请注明出处:https://www.yunthe.com/dui-bi-xue-xi-shi-zhan-simclr-kuang-jia-de-biao-zheng-xue/

(0)
小编小编
上一篇 5小时前
下一篇 3小时前

相关推荐