对比学习(Clap)入门指南:从零构建高效特征提取模型

1次阅读
没有评论

共计 2722 个字符,预计需要花费 7 分钟才能阅读完成。

image.webp

为什么我们需要对比学习?

在传统监督学习中,我们严重依赖标注数据来训练模型。但在实际项目中,标注数据往往面临两大难题:

对比学习 (Clap) 入门指南:从零构建高效特征提取模型

  • 标注成本高:专业领域的标注需要专家参与,时间和经济成本都很高
  • 样本稀缺:某些场景(如医疗图像)获取大量样本本身就困难

这些问题导致监督学习方法在真实场景中常常 ” 巧妇难为无米之炊 ”。

三种学习方式的对比

  1. 监督学习:需要完整标注,直接学习输入到标签的映射
  2. 优点:目标明确,评估直观
  3. 缺点:依赖标注数据量

  4. 自监督学习:利用数据自身结构创建伪标签

  5. 优点:减少标注依赖
  6. 缺点:预训练任务可能偏离下游任务

  7. 对比学习(Clap):通过比较样本间的相似度学习特征表示

  8. 核心思想:相似样本在表征空间中靠近,不相似样本远离
  9. 优势:特别适合特征提取任务,对标注数据依赖小

动手构建第一个 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 之间(避免坍缩)

进一步探索方向

  1. 不同的数据增强组合如何影响最终特征质量?
  2. 投影头的深度与特征维度是否存在最优配置?
  3. 如何将对比学习与其他自监督方法结合?

希望这篇指南能帮助你快速上手对比学习。记住,实践出真知 – 现在就用你的数据集试试这些方法吧!

正文完
 0
评论(没有评论)