共计 2277 个字符,预计需要花费 6 分钟才能阅读完成。
一、为什么我们需要 CLIP?
传统跨模态学习任务(比如图像描述生成)通常需要大量 成对标注数据(即每张图片都要配人工写的文本描述)。这种数据不仅收集成本高,还存在两个明显缺陷:

- 标注质量不稳定(不同人写的描述差异很大)
- 难以覆盖长尾场景(比如医疗图像的精准描述需要专业医师)
CLIP(Contrastive Language-Image Pretraining)通过对比学习解决了这个问题——它只需要 图像和文本的弱关联(比如网页中同时出现的图片和周边文字),就能让模型自动学习两者的语义对应关系。
二、监督学习 vs 对比学习
传统监督学习
假设我们要训练一个猫狗分类器:
- 输入:图片
- 输出:概率分布(如 [0.9, 0.1] 表示 90% 概率是猫)
- 需要每张图片都有明确的类别标签
CLIP 的对比学习
- 输入:图片 + 文本的 组合对(不需要严格匹配)
- 目标:让匹配的图片 - 文本对在特征空间中更接近
- 关键优势:利用大量互联网天然存在的图文数据
三、CLIP 核心实现
1. 双编码器结构
import torch
import torch.nn as nn
# 简化版图像编码器(实际 CLIP 使用 ViT 或 ResNet)class ImageEncoder(nn.Module):
def __init__(self):
super().__init__()
self.cnn = nn.Sequential(nn.Conv2d(3, 64, kernel_size=3),
nn.ReLU(),
nn.MaxPool2d(2),
nn.Flatten(),
nn.Linear(64*14*14, 512) # 输出 512 维特征
)
def forward(self, x):
return self.cnn(x)
# 简化版文本编码器(实际 CLIP 使用 Transformer)class TextEncoder(nn.Module):
def __init__(self, vocab_size=10000):
super().__init__()
self.embed = nn.Embedding(vocab_size, 512)
self.rnn = nn.LSTM(512, 512, batch_first=True)
def forward(self, x):
x = self.embed(x)
_, (h, _) = self.rnn(x)
return h.squeeze(0)
2. 对比损失函数
核心思想:让正样本对的相似度远高于随机负样本对
def contrastive_loss(image_emb, text_emb, temperature=0.07):
# 计算余弦相似度矩阵
logits = image_emb @ text_emb.T / temperature
# 对角线元素是正样本对
labels = torch.arange(len(logits)).to(logits.device)
# 对称计算两个方向的交叉熵
loss_i = nn.CrossEntropyLoss()(logits, labels) # 图像 -> 文本
loss_t = nn.CrossEntropyLoss()(logits.T, labels) # 文本 -> 图像
return (loss_i + loss_t) / 2
3. 负样本采样
- 隐式负样本:同一个 batch 内的其他样本自动作为负样本
- 关键技巧:batch size 越大,负样本越丰富(通常需要≥256)
四、完整训练示例
# 超参数设置
batch_size = 256
lr = 5e-5
temp = 0.07
epochs = 10
# 初始化
img_encoder = ImageEncoder().cuda()
txt_encoder = TextEncoder().cuda()
optimizer = torch.optim.Adam(list(img_encoder.parameters()) + list(txt_encoder.parameters()),
lr=lr
)
# 训练循环
for epoch in range(epochs):
for images, texts in dataloader: # 假设已实现数据加载
# 前向传播
img_feat = img_encoder(images.cuda())
txt_feat = txt_encoder(texts.cuda())
# 计算损失
loss = contrastive_loss(img_feat, txt_feat, temp)
# 反向传播
optimizer.zero_grad()
loss.backward()
optimizer.step()
print(f"Epoch {epoch}, Loss: {loss.item():.4f}")
五、性能优化技巧
Batch Size 选择
- 32-64:适合调试
- 256-1024:实际训练推荐
-
2048:需要多 GPU 并行
学习率设置
- 先用
lr=5e-5测试模型能否正常下降 - 稳定后可尝试
lr=3e-4(CLIP 原论文设置) - 配合线性 warmup 效果更好
梯度爆炸预防
# 在优化器之前添加梯度裁剪
torch.nn.utils.clip_grad_norm_(list(img_encoder.parameters()) + list(txt_encoder.parameters()),
max_norm=1.0
)
六、实际应用场景
- 零样本分类:输入图片和类别文本描述,直接预测类别
- 图文检索:用文本搜索相关图片(或反向)
- 内容审核:检测违规图文组合
七、下一步建议
尝试在自定义数据集上微调:
- 准备至少 1 万对相关图文数据
- 冻结部分层(例如图像编码器的浅层)
- 使用更小的学习率(例如原值的 1 /10)
小技巧:先用 CLIP 预训练权重初始化,再微调效果更佳
正文完
