CLIP对比学习入门指南:从理论到PyTorch实战

1次阅读
没有评论

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

image.webp

技术背景

跨模态检索(比如用文字搜图片或用图片搜文字)一直是 AI 领域的难题。传统方法通常分两步走:先分别提取图像和文本特征,再训练一个模型来对齐这两种特征。这种方法的缺陷很明显——图像和文本特征来自不同的模型,天生不在同一个 ” 空间 ” 里,对齐起来就像让两个说不同语言的人强行聊天。

CLIP 对比学习入门指南:从理论到 PyTorch 实战

对比学习的突破在于,它让图像和文本在特征提取阶段就开始 ” 对话 ”。CLIP(Contrastive Language-Image Pretraining)作为代表作,通过海量的图文配对数据,让模型学会把配对的图文特征拉近,不配对的推远。这就好比教孩子认图时,同时告诉他图中物体的名称,他自然就建立了图像和文字的联系。

核心原理

CLIP 的核心是对比损失函数,数学表达式如下:

loss_i = -log[exp(sim(I_i, T_i)/τ) / Σ_j exp(sim(I_i, T_j)/τ)]
loss_t = -log[exp(sim(T_i, I_i)/τ) / Σ_j exp(sim(T_j, I_i)/τ)]
总损失 = (loss_i + loss_t)/2

其中:
sim()计算余弦相似度
– τ 是温度参数,控制分布尖锐程度
– 分母的求和是在 batch 内所有样本上进行的(这就是批次内负样本)

这个函数的直观理解是:对于每张图片,让它与正确文本的相似度远高于与错误文本的相似度(文本对图片同理)。双塔结构保证了两种模态的特征能映射到同一空间。

代码实现

1. 双塔结构搭建

import torch
import torch.nn as nn
from transformers import AutoModel, AutoTokenizer

class ImageEncoder(nn.Module):
    """ResNet50 作为图像编码器"""
    def __init__(self):
        super().__init__()
        self.model = torch.hub.load('pytorch/vision', 'resnet50', pretrained=True)
        self.mlp = nn.Sequential(nn.Linear(1000, 512),  # 将 ResNet 输出投影到共同空间
            nn.GELU(),
            nn.Linear(512, 256)   # 最终特征维度
        )

    def forward(self, x):
        return self.mlp(self.model(x))

class TextEncoder(nn.Module):
    """BERT 作为文本编码器"""
    def __init__(self):
        super().__init__()
        self.tokenizer = AutoTokenizer.from_pretrained('bert-base-uncased')
        self.model = AutoModel.from_pretrained('bert-base-uncased')
        self.mlp = nn.Sequential(nn.Linear(768, 512),  # BERT 基础版输出维度是 768
            nn.GELU(),
            nn.Linear(512, 256)
        )

    def forward(self, text):
        inputs = self.tokenizer(text, return_tensors='pt', padding=True)
        outputs = self.model(**inputs)
        # 取 [CLS] 标记对应的向量作为句子表示
        return self.mlp(outputs.last_hidden_state[:, 0, :])

2. 对比损失实现

def contrastive_loss(image_emb, text_emb, temperature=0.07):
    """对称对比损失"""
    # 归一化特征向量
    image_emb = F.normalize(image_emb, dim=1)
    text_emb = F.normalize(text_emb, dim=1)

    # 计算相似度矩阵(batch_size x batch_size)logits = torch.matmul(image_emb, text_emb.T) / temperature

    # 生成标签:对角线位置是正样本对
    labels = torch.arange(logits.size(0)).to(logits.device)

    # 计算交叉熵损失(图像 -> 文本和文本 -> 图像方向)loss_i = F.cross_entropy(logits, labels)
    loss_t = F.cross_entropy(logits.T, labels)
    return (loss_i + loss_t) / 2

3. 训练循环示例

# 初始化
image_encoder = ImageEncoder().cuda()
text_encoder = TextEncoder().cuda()
optimizer = torch.optim.AdamW(list(image_encoder.parameters()) + list(text_encoder.parameters()),
    lr=5e-5,
    weight_decay=0.01
)

# 模拟一个 batch 的数据
images = torch.randn(32, 3, 224, 224).cuda()  # 假设 batch_size=32
texts = ["a photo of a cat"] * 32  # 简化示例

# 前向传播
image_features = image_encoder(images)
text_features = text_encoder(texts)

# 计算损失
loss = contrastive_loss(image_features, text_features)

# 反向传播
loss.backward()
optimizer.step()
optimizer.zero_grad()

实战建议

  1. 数据预处理
  2. 图像:使用 RandomResizedCrop、ColorJitter 等增强
  3. 文本:保留大小写(除非特别需要),最大长度建议 64-128

  4. 超参数设置

  5. batch size 越大越好(至少 256),小 batch 时考虑梯度累积
  6. 初始学习率建议 3e- 5 到 5e-5,配合线性 warmup
  7. 温度参数 τ 通常取 0.01 到 0.1,需要验证集调优

  8. 特征可视化

    # 使用 UMAP 降维后画图
    import umap
    import matplotlib.pyplot as plt
    
    reducer = umap.UMAP()
    emb_2d = reducer.fit_transform(features.cpu().numpy())
    plt.scatter(emb_2d[:,0], emb_2d[:,1], c=labels)

避坑指南

  1. 模态不平衡
  2. 如果发现文本特征范数明显大于图像特征,可以在 MLP 后加 LayerNorm

  3. 小 batch size 问题

  4. 当 GPU 内存不足时,可以用梯度累积模拟大 batch:

    for i, (images, texts) in enumerate(dataloader):
        loss = model(images, texts)
        loss = loss / 4  # 假设累积 4 步
        loss.backward()
        if (i+1) % 4 == 0:
            optimizer.step()
            optimizer.zero_grad()

  5. 过拟合检测

  6. 监控训练 / 验证损失曲线,如果两者差距持续增大:
  7. 早停(early stopping)
  8. 增加 dropout 比率
  9. 加强数据增强

延伸阅读

  1. 原始论文: Learning Transferable Visual Models From Natural Language Supervision
  2. 对比学习综述: A Survey on Contrastive Self-supervised Learning
  3. 温度参数分析: Understanding the Behaviour of Contrastive Loss

通过这个实战指南,希望能帮助初学者快速理解 CLIP 的核心思想并实现自己的第一个跨模态检索模型。在实际项目中,可以从小的数据集(如 Flickr8k)开始实验,再逐步扩展到更大规模的应用。

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