共计 3257 个字符,预计需要花费 9 分钟才能阅读完成。
技术背景
跨模态检索(比如用文字搜图片或用图片搜文字)一直是 AI 领域的难题。传统方法通常分两步走:先分别提取图像和文本特征,再训练一个模型来对齐这两种特征。这种方法的缺陷很明显——图像和文本特征来自不同的模型,天生不在同一个 ” 空间 ” 里,对齐起来就像让两个说不同语言的人强行聊天。

对比学习的突破在于,它让图像和文本在特征提取阶段就开始 ” 对话 ”。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()
实战建议
- 数据预处理
- 图像:使用 RandomResizedCrop、ColorJitter 等增强
-
文本:保留大小写(除非特别需要),最大长度建议 64-128
-
超参数设置
- batch size 越大越好(至少 256),小 batch 时考虑梯度累积
- 初始学习率建议 3e- 5 到 5e-5,配合线性 warmup
-
温度参数 τ 通常取 0.01 到 0.1,需要验证集调优
-
特征可视化
# 使用 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)
避坑指南
- 模态不平衡
-
如果发现文本特征范数明显大于图像特征,可以在 MLP 后加 LayerNorm
-
小 batch size 问题
-
当 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() -
过拟合检测
- 监控训练 / 验证损失曲线,如果两者差距持续增大:
- 早停(early stopping)
- 增加 dropout 比率
- 加强数据增强
延伸阅读
- 原始论文: Learning Transferable Visual Models From Natural Language Supervision
- 对比学习综述: A Survey on Contrastive Self-supervised Learning
- 温度参数分析: Understanding the Behaviour of Contrastive Loss
通过这个实战指南,希望能帮助初学者快速理解 CLIP 的核心思想并实现自己的第一个跨模态检索模型。在实际项目中,可以从小的数据集(如 Flickr8k)开始实验,再逐步扩展到更大规模的应用。
