CLIP双塔结构图解析:从对比学习到跨模态检索实战

1次阅读
没有评论

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

image.webp

跨模态检索的挑战与机遇

在推荐系统和智能搜索场景中,我们常常需要处理图像和文本之间的关联匹配。传统方法依赖于手工设计的特征工程,比如用 SIFT 提取图像特征,用 TF-IDF 处理文本。这种方式存在明显的局限性:

CLIP 双塔结构图解析:从对比学习到跨模态检索实战

  • 手工特征难以捕捉高层次语义信息
  • 不同模态的特征空间维度不一致,无法直接比较
  • 特征提取流程与后续任务分离,无法端到端优化

CLIP 的创新设计

CLIP(Contrastive Language-Image Pretraining)通过对比学习解决了这些问题。相比 VSE++ 等早期方案,CLIP 的优势在于:

  1. 双塔结构
  2. 图像编码器(通常用 ResNet 或 ViT)
  3. 文本编码器(通常用 Transformer)
  4. 两个编码器并行处理不同模态数据

  5. 计算效率

  6. 离线计算特征向量,在线检索只需计算余弦相似度
  7. 适合大规模部署

  8. 对比学习目标

  9. 通过 InfoNCE 损失拉近正样本对距离
  10. 推开负样本对距离

PyTorch 实现详解

模型构建

import torch
import torch.nn as nn
from torchvision.models import resnet50
from transformers import AutoTokenizer, AutoModel

class ImageEncoder(nn.Module):
    def __init__(self):
        super().__init__()
        self.model = resnet50(pretrained=True)
        self.model.fc = nn.Identity()  # 移除最后的全连接层

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

class TextEncoder(nn.Module):
    def __init__(self):
        super().__init__()
        self.tokenizer = AutoTokenizer.from_pretrained('bert-base-uncased')
        self.model = AutoModel.from_pretrained('bert-base-uncased')

    def forward(self, text):
        inputs = self.tokenizer(text, return_tensors='pt', padding=True, truncation=True)
        outputs = self.model(**inputs)
        return outputs.last_hidden_state[:,0,:]  # 取 [CLS] token 作为句子表示 

对比损失实现

InfoNCE 损失的数学表达:

$$
\mathcal{L} = -\log\frac{\exp(s_{i,j}/\tau)}{\sum_{k=1}^N \exp(s_{i,k}/\tau)}
$$

其中 $\tau$ 是温度系数,代码实现:

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

训练技巧与优化

批处理负采样

  • 利用当前 batch 内的其他样本作为负样本
  • 无需额外存储负样本队列

混合精度训练

from torch.cuda.amp import autocast, GradScaler

scaler = GradScaler()

with autocast():
    image_emb = image_encoder(images)
    text_emb = text_encoder(texts)
    loss = contrastive_loss(image_emb, text_emb)

scaler.scale(loss).backward()
scaler.step(optimizer)
scaler.update()

特征归一化

# 在计算相似度前先做 L2 归一化
image_emb = nn.functional.normalize(image_emb, dim=-1)
text_emb = nn.functional.normalize(text_emb, dim=-1)

实践中的经验总结

  1. 数据预处理
  2. 图像缩放保持长宽比,用 letterbox 填充
  3. 文本统一小写处理

  4. 超参数调优

  5. 初始学习率建议 3e-5
  6. batch size 尽可能大(至少 256)

  7. 调试技巧

  8. 用 t -SNE 可视化特征空间分布
  9. 定期检查 top- k 检索准确率

扩展与展望

  1. 视频文本检索
  2. 用 3D CNN 处理视频
  3. 时间维度上做 pooling

  4. 模型蒸馏

  5. 用大模型生成伪标签
  6. 训练轻量化的学生模型

完整代码已上传 Colab: 实践链接

通过 CLIP 的双塔结构,我们实现了高效的跨模态检索。这种设计不仅适用于图文场景,经过适当调整还能扩展到音频、视频等多模态领域,为构建更智能的推荐系统提供了坚实基础。

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