CLIP 深度剖析:对比语言-图像预训练的完整技术栈
1. CLIP 是什么:打破数据标注的瓶颈
1.1 传统视觉训练的困境
在 CLIP 出现之前,视觉模型的训练严重依赖人工标注数据:
ImageNet 分类: 1.2M 图像, 1,000 个类别, 人工标注COCO 检测: 330K 图像, 80 个物体类别, 每张图多人标注
问题: 1. 标注成本极高——ImageNet 用了 2.5 万名众包工人 2. 类别固定——无法泛化到 ImageNet 之外的任何类别 3. 任务单一——分类模型无法直接用于检测、分割、OCR 4. 分布偏置——互联网图片 ≠ ImageNet 风格1.2 CLIP 的核心洞察
互联网上天然存在大量的图像-文本配对数据**(网页 alt 标签、标题描述、用户上传说明等)。与其让人工给图像打标签,不如直接用自然语言 supervision——让模型学会”图像和什么文字描述最匹配”。**
这带来了一个根本性的范式转变:
传统: 图像 → (类别标签) → 预测类别 # 需要人工标注CLIP: 图像 + 文本 → 匹配程度预测 # 用自然语言 supervision
文本本身就是标签——不需要预先定义类别。新的类别 = 新的文本描述 = 零成本。1.3 CLIP 的能力
CLIP 训练完成后,可以在零样本(zero-shot)条件下完成从未见过的分类任务:
任务: 分类一张猫的照片
传统做法: - 需要在 ImageNet-1K 的 1000 个类上微调 - 如果类别不是"猫",模型完全无法识别
CLIP 零样本: - 类别描述: ["a photo of a cat", "a photo of a dog", "a photo of a bird"] - 把这些文本编码成向量 - 把图像编码成向量 - 计算余弦相似度 → 最高者胜出 - 无需任何微调,即插即用1.4 论文信息
论文: "Learning Transferable Visual Models From Natural Language Supervision"作者: Alec Radford, Jong Wook Kim, Chris Hallacy, Aditya Ramesh, Gabriel Goh, Sandhini Agarwal, Girish Sastry, Amanda Askell, Pamela Mishkin, Jack Clark, Gretchen Krueger, Ilya Sutskever单位: OpenAI发表于: ICML 2021引用: > 20,000 次(截至 2024)★★★★★开源: https://github.com/openai/CLIP1.5 一句话概括 CLIP
CLIP 通过一个双塔对比学习框架,用自然语言作为监督信号,让视觉编码器和文本编码器在同一个特征空间中对齐——从而实现零样本图像分类、开放词汇检测、多模态理解等下游任务,并成为几乎所有 VLM 和 VLA 的视觉编码器基础。
2. 数学框架:对比学习目标
2.1 双塔架构
CLIP 由两个编码器组成:
┌──────────────────────────────────────────────────────────────┐│ ││ [图像] ──→ 视觉编码器 ──→ I(x) ∈ ℝ^D (视觉特征向量) ││ ViT 或 ResNet ││ ││ [文本] ──→ 文本编码器 ──→ T(t) ∈ ℝ^D (文本特征向量) ││ Transformer ││ ││ 对比: 最大化匹配的 I·T,余弦相似度矩阵 ││ │└──────────────────────────────────────────────────────────────┘2.2 InfoNCE 对比损失
给定一个 batch 的图文对 ,CLIP 最大化正样本对的相似度,最小化负样本对的相似度。
余弦相似度矩阵:
其中 是第 张图像的归一化特征, 是第 个文本的归一化特征。
对称的对比损失(两方向都要优化):
图像→文本损失(Image-to-Text):
文本→图像损失(Text-to-Image):
其中 是温度系数(可学习的参数),控制 softmax 的锐度。
2.3 损失的直觉解释
def info_nce_loss(image_embeds, text_embeds, temperature): """ InfoNCE 对比损失(简化版)。
参数: image_embeds: (B, D) 图像特征 text_embeds: (B, D) 文本特征 temperature: 温度系数 τ """ # 归一化 image_embeds = F.normalize(image_embeds, dim=-1) text_embeds = F.normalize(text_embeds, dim=-1)
# 计算余弦相似度矩阵 (B, B) logits = torch.matmul(image_embeds, text_embeds.T) / temperature
# labels: 对角线上是正样本 batch_size = image_embeds.shape[0] labels = torch.arange(batch_size, device=image_embeds.device)
# 两方向的对称损失 loss_i2t = F.cross_entropy(logits, labels) # 图像→文本 loss_t2i = F.cross_entropy(logits.T, labels) # 文本→图像
loss = (loss_i2t + loss_t2i) / 2 return loss直观理解:
对角线 (正样本): S_ii → 最大 (希望)非对角线 (负样本): S_ij (i≠j) → 最小 (希望)
每个图像 i 应该与对应的文本 j 最匹配每个文本 j 应该与对应的图像 i 最匹配
这就是"对比"——把正确配对从 N 个配对中对比出来。2.4 温度系数的意义
def temperature_effect(): """ 温度系数 τ 对对比损失的影响。 τ 是可学习的参数,初始化通常在 0.01~0.1 之间。 """ configs = { "τ 极小 (0.01)": "softmax 极度尖锐 → 只关注最难的负样本", "τ 适中 (0.07)": "平衡——正确配对突出,负样本被适度压制", "τ 过大 (1.0)": "softmax 趋于均匀 → 无法区分正负样本", } return configs
# CLIP 原始实现中 τ 被建模为 log(scaled_logits) 的可学习参数 # logit_scale = nn.Parameter(torch.ones([]) * np.log(1/0.07)) # temperature = 1 / torch.exp(logit_scale) ≈ 0.073. 视觉编码器:Vision Transformer
3.1 图像分块
CLIP 的视觉编码器使用 ViT(Vision Transformer) 而非 ResNet。图像被切分成固定大小的 patch:
输入图像: 224 × 224 × 3
patch_size = 16 × 16
图像分块: - 224 / 16 = 14 - 总 patch 数 = 14 × 14 = 196 - 每个 patch = 16 × 16 × 3 = 768 维
加上 [CLS] token → 197 个视觉 tokenclass ImagePatcher(nn.Module): """ 图像分块:把图像切成 patch 并线性投影。 """ def __init__(self, image_size=224, patch_size=16, in_channels=3, embed_dim=768): super().__init__() self.image_size = image_size self.patch_size = patch_size self.num_patches = (image_size // patch_size) ** 2 # 196
# 线性投影层(等价于 2D conv with kernel_size=patch_size) self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=patch_size)
def forward(self, x): """ 参数: x: (B, 3, 224, 224) 图像 batch 返回: patches: (B, num_patches, embed_dim) 每个 patch 的嵌入 """ # Conv2d: (B, 3, 224, 224) → (B, embed_dim, 14, 14) x = self.proj(x) # 展平空间维度: (B, embed_dim, 14, 14) → (B, embed_dim, 196) x = x.flatten(2) # 调换维度: (B, embed_dim, 196) → (B, 196, embed_dim) x = x.transpose(1, 2) return x3.2 ViT Transformer 编码器
class ViTEncoder(nn.Module): """ CLIP 的 ViT 视觉编码器。 """ def __init__(self, image_size=224, patch_size=16, embed_dim=768, num_layers=12, num_heads=12, mlp_ratio=4.0): super().__init__() self.patch_embed = ImagePatcher(image_size, patch_size, 3, embed_dim) self.num_patches = self.patch_embed.num_patches
# [CLS] token self.cls_token = nn.Parameter(torch.zeros(1, 1, embed_dim))
# 位置编码(可学习) self.pos_embed = nn.Parameter(torch.zeros(1, self.num_patches + 1, embed_dim))
# Transformer 编码器 encoder_layer = nn.TransformerEncoderLayer( d_model=embed_dim, nhead=num_heads, dim_feedforward=int(embed_dim * mlp_ratio), activation="gelu", batch_first=True, norm_first=True, # Pre-LN (CLIP 使用) ) self.transformer = nn.TransformerEncoder(encoder_layer, num_layers)
# LayerNorm self.ln = nn.LayerNorm(embed_dim)
def forward(self, x): """ 参数: x: (B, 3, 224, 224) 返回: pooled: (B, embed_dim) [CLS] token 的最终表示 sequence: (B, num_patches+1, embed_dim) 完整序列 """ B = x.shape[0]
# 1) 分块 + 投影 patches = self.patch_embed(x) # (B, 196, embed_dim)
# 2) 添加 [CLS] token cls_tokens = self.cls_token.expand(B, -1, -1) # (B, 1, embed_dim) x = torch.cat([cls_tokens, patches], dim=1) # (B, 197, embed_dim)
# 3) 添加位置编码 x = x + self.pos_embed # (B, 197, embed_dim)
# 4) Transformer 编码 x = self.transformer(x) # (B, 197, embed_dim)
# 5) LayerNorm x = self.ln(x)
# [CLS] token pooled = x[:, 0, :] # (B, embed_dim) sequence = x # (B, 197, embed_dim)
return pooled, sequence3.3 CLIP 使用的两种视觉编码器
| 编码器 | 规模 | patch size | 层数 | 参数量 | 特点 |
|---|---|---|---|---|---|
| ViT-B/32 | Base | 32×32 | 12 | 86M | CLIP 默认, 速度最快 |
| ViT-B/16 | Base | 16×16 | 12 | 86M | 更高分辨率 |
| ViT-L/14 | Large | 14×14 | 24 | 304M | CLIP 最强, 最常用 |
| ViT-L/14@336 | Large | 14×14 | 24 | 304M | 微调版, 336px |
| RN50 | ResNet-50 | - | - | 38M | 论文消融实验 |
4. 文本编码器:Text Transformer
4.1 文本预处理
class TextTokenizer: """ CLIP 使用 BPE (Byte Pair Encoding) 分词器。 与 GPT-2 / RoBERTa 相同。 """ def __init__(self, vocab_size=50257): self.tokenizer = GPT2Tokenizer.from_pretrained("gpt2") self.tokenizer.add_special_tokens(["<|startoftext|>", "<|endoftext|>"])
def encode(self, texts): """ 参数: texts: str 或 list[str] 返回: input_ids: (B, L) token IDs attention_mask: (B, L) 注意力掩码 """ if isinstance(texts, str): texts = [texts]
encoded = self.tokenizer( texts, padding=True, truncation=True, max_length=77, # CLIP 固定文本长度 return_tensors="pt", ) return encoded["input_ids"], encoded["attention_mask"]4.2 Transformer 文本编码器
class TextTransformerEncoder(nn.Module): """ CLIP 的文本编码器。 Transformer Encoder 处理文本 token 序列。 """ def __init__(self, vocab_size=50257, embed_dim=512, context_length=77, num_layers=12, num_heads=8): super().__init__() self.context_length = context_length
# Token 嵌入 self.token_embedding = nn.Embedding(vocab_size, embed_dim)
# 位置编码(可学习,固定长度) self.pos_embedding = nn.Parameter( torch.zeros(1, context_length, embed_dim) # (1, 77, 512) )
# Transformer 编码器 encoder_layer = nn.TransformerEncoderLayer( d_model=embed_dim, nhead=num_heads, dim_feedforward=embed_dim * 4, activation="gelu", batch_first=True, norm_first=True, ) self.transformer = nn.TransformerEncoder(encoder_layer, num_layers)
# LayerNorm self.ln = nn.LayerNorm(embed_dim)
def forward(self, input_ids, attention_mask=None): """ 参数: input_ids: (B, L) token IDs attention_mask: (B, L) 可选 返回: pooled: (B, embed_dim) [EOS] token 的表示 sequence: (B, L, embed_dim) 完整序列 """ B, L = input_ids.shape
# 1) Token 嵌入 x = self.token_embedding(input_ids) # (B, L, embed_dim)
# 2) 添加位置编码 x = x + self.pos_embedding[:, :L, :] # (B, L, embed_dim)
# 3) Transformer 编码 x = self.transformer(x) # (B, L, embed_dim)
# 4) LayerNorm x = self.ln(x)
# 5) 取 [EOS] token(通常是文本序列最后一个 token)作为句子表示 # 或者用 attention pooling pooled = x[:, -1, :] # (B, embed_dim)
return pooled, x4.3 文本表示的选择
CLIP 原文使用 [EOS] token 的输出作为整个文本的表示,而非 attention pooling:
方法 1: [EOS] token → 简单, 最常用 "hello world [EOS]" ↑ 取这个位置的特征
方法 2: Attention Pooling → 更灵活 attention_weights = softmax(w · tanh(W · h_i^T)) text_repr = Σ attention_weights_i · h_i
方法 3: Mean Pooling → 无参数 text_repr = mean(h_1, h_2, ..., h_L)CLIP 实验发现 [EOS] token 效果最好,原因可能是:Transformer 的 [EOS] 位置已经”看过”了完整序列。
5. 训练细节:从零开始的工程实践
5.1 训练数据
CLIP 使用 WIT-400M 数据集(WebImageText):
WIT-400M: - 400,000,000 图像-文本对 - 来源: 从互联网上抓取的图像+alt 文本/title/caption - 文本平均长度: ~12 个词 - 覆盖类别: 自然场景、物体、人物、动作、场景描述等
与 ImageNet 的对比: ImageNet: 1.2M 图像, 1,000 类, 人工标注 WIT-400M: 400M 图像, 无固定类别, 自然语言描述数据预处理:
class WITDataset(Dataset): """ CLIP 的训练数据加载器。 """ def __init__(self, image_paths, texts, image_size=224): self.image_paths = image_paths self.texts = texts self.image_size = image_size self.img_transform = transforms.Compose([ transforms.Resize(image_size, interpolation=transforms.InterpolationMode.BICUBIC), transforms.CenterCrop(image_size), transforms.ToTensor(), transforms.Normalize( mean=[0.48145466, 0.4578275, 0.40821073], # ImageNet RGB 统计 std=[0.26862954, 0.26130258, 0.27577711], ), ])
def __getitem__(self, idx): # 加载图像 image = Image.open(self.image_paths[idx]).convert("RGB") image = self.img_transform(image)
# 获取文本 text = self.texts[idx][:77] # 截断到 77 tokens
return {"image": image, "text": text}5.2 训练超参数
@dataclassclass CLIPTrainingConfig: """CLIP 训练配置。""" # 模型 vision_model: str = "ViT-L/14" # CLIP 默认使用 ViT-L/14 embed_dim: int = 768 # 视觉/文本特征的维度
# 训练 batch_size: int = 32768 # CLIP 原始实现用 32768(很大!) num_epochs: int = 32 learning_rate: float = 5e-4 # AdamW 学习率 weight_decay: float = 0.1 warmup_steps: int = 5000
# 优化器 optimizer: str = "AdamW" beta1: float = 0.9 beta2: float = 0.98 eps: float = 1e-6
# 损失 temperature_init: float = 0.07 # 温度系数初始化 label_smoothing: float = 0.0
# 正则化 gradient_clip: float = 1.0 mixed_precision: bool = True # FP16 混合精度
# 数据 image_size: int = 224 context_length: int = 77 max_text_length: int = 77注意:CLIP 原始实现的 batch_size 高达 32768!这需要分布式训练(多机多卡)。社区复现通常用 256-1024。
5.3 训练循环
def train_clip(model, dataloader, config): """ CLIP 训练循环(简化版)。 """ optimizer = torch.optim.AdamW( model.parameters(), lr=config.learning_rate, weight_decay=config.weight_decay, betas=(config.beta1, config.beta2), eps=config.eps, )
# 学习率调度器 def lr_lambda(step): if step < config.warmup_steps: return step / config.warmup_steps else: progress = (step - config.warmup_steps) / (config.total_steps - config.warmup_steps) return 0.5 * (1 + np.cos(np.pi * progress))
scheduler = torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambda)
# 可学习的温度系数 logit_scale = nn.Parameter(torch.ones([]) * np.log(1 / config.temperature_init)) logit_scale = logit_scale.to(device)
model.train() for epoch in range(config.num_epochs): epoch_loss = 0.0 for batch in tqdm(dataloader): images = batch["image"].to(device) texts = batch["text"]
# 编码 image_features = model.visual(images) # (B, D) text_features = model.text(texts) # (B, D)
# L2 归一化 image_features = F.normalize(image_features, dim=-1) text_features = F.normalize(text_features, dim=-1)
# 对比 logits temperature = (1 / torch.exp(logit_scale)) logits_per_image = (image_features @ text_features.T) / temperature # (B, B) logits_per_text = logits_per_image.T # (B, B)
# Labels: 对角线是正样本 batch_size = images.shape[0] labels = torch.arange(batch_size, device=device)
# 对称交叉熵损失 loss_i2t = F.cross_entropy(logits_per_image, labels) loss_t2i = F.cross_entropy(logits_per_text, labels) loss = (loss_i2t + loss_t2i) / 2
# 反向传播 optimizer.zero_grad() loss.backward()
# 梯度裁剪 torch.nn.utils.clip_grad_norm_(model.parameters(), config.gradient_clip)
optimizer.step() scheduler.step()
epoch_loss += loss.item()
# 打印 print(f"Epoch {epoch}: loss={epoch_loss / len(dataloader):.4f}, " f"temp={temperature.item():.4f}")6. 零样本分类:CLIP 的杀手锏
6.1 零样本分类流程
@torch.no_grad()def zero_shot_classify(model, image, class_names, device): """ CLIP 零样本分类。
参数: model: CLIP 模型 image: (C, H, W) 图像 tensor class_names: list[str] 类别名称,如 ["cat", "dog", "bird"] device: 计算设备 返回: 预测的类别索引和概率 """ model.eval()
# 1) 构建文本提示(class name → sentence) # CLIP 原论文发现"a photo of a {class}"效果最好 prompts = [f"a photo of a {c}." for c in class_names]
# 2) 编码文本 text_tokens = model.tokenizer(prompts).to(device) text_features = model.encode_text(text_tokens) text_features = F.normalize(text_features, dim=-1) # (num_classes, D)
# 3) 编码图像 if len(image.shape) == 3: image = image.unsqueeze(0) # (1, C, H, W) image = image.to(device) image_features = model.encode_image(image) image_features = F.normalize(image_features, dim=-1) # (1, D)
# 4) 计算余弦相似度 similarity = (image_features @ text_features.T) # (1, num_classes) probs = F.softmax(similarity * model.logit_scale, dim=-1) # (1, num_classes)
# 5) 返回结果 probs = probs[0].cpu().numpy() pred_idx = probs.argmax() return pred_idx, class_names[pred_idx], probs[pred_idx], probs6.2 提示工程(Prompt Engineering)
CLIP 对文本提示的措辞非常敏感:
def prompt_ensemble(model, class_names, image, device): """ 提示集成:使用多个提示模板提升零样本分类精度。 CLIP 原论文发现这一招效果显著。 """ # 多种提示模板 prompt_templates = [ "a photo of a {}.", "a blurry photo of a {}.", "a low resolution photo of a {}.", "a bright photo of a {}.", "a dark photo of a {}.", "a photo of a small {}.", "a photo of a large {}.", "a photo of the {}.", "a photo of a {} in the scene.", "a photo of the {} in the scene.", ]
all_text_features = []
for template in prompt_templates: prompts = [template.format(c) for c in class_names] text_tokens = model.tokenizer(prompts).to(device) text_features = model.encode_text(text_tokens) text_features = F.normalize(text_features, dim=-1) all_text_features.append(text_features)
# 平均所有提示的特征 text_features_avg = torch.stack(all_text_features, dim=0).mean(dim=0) # (num_classes, D) text_features_avg = F.normalize(text_features_avg, dim=-1)
# 图像编码 image_features = model.encode_image(image.to(device).unsqueeze(0)) image_features = F.normalize(image_features, dim=-1)
# 相似度 similarity = image_features @ text_features_avg.T probs = F.softmax(similarity * model.logit_scale, dim=-1)
return probs[0].cpu().numpy()提示工程的效果(来自 CLIP 论文):
ImageNet 零样本分类准确率对比:
单一提示 "a photo of a {}": 56.3%集成 80 种提示: 59.1% (+2.8%)+ 描述性提示 (如 "a blurry photo..."): +2.0-4.0%+ 类别同义词 ("cat, kitty"): +1.0-2.0%6.3 零样本迁移的理论解释
def why_zero_shot_works(): """ CLIP 零样本分类有效的原因。 """ explanations = { "语义丰富": "文本编码器能理解语义——'golden retriever' 和 'labrador' 都指向同一种狗", "分布匹配": "训练数据是互联网图文对,测试分布与训练分布相似", "语言先验": "语言模型带来了世界知识——知道什么是'斑马'的描述", "特征泛化": "ViT 学到了通用的视觉表示,适用于任何视觉概念", "开放词汇": "不需要固定类别,任意文本都可以作为查询", } return explanations7. 扩展:ALIGN / EVA-CLIP / SigLIP
7.1 ALIGN:更大数据 + 更强骨干
ALIGN (Jia et al., 2021): - 数据: 1.8B 图像-文本对(比 CLIP 多 4.5 倍) - 骨干: EfficientNet-L2 (更大的视觉编码器) - 结果: 在 ImageNet 上零样本 85.5% (CLIP: 76.2%)
核心发现: 数据规模 > 模型架构 400M 图像 + ViT-L = 76.2% 1.8B 图像 + ViT-L = 85.5% (+9.3%)
→ 更多的图文数据比更大的模型更重要7.2 EVA-CLIP:打补丁的 CLIP
EVA-CLIP (Sun et al., 2023): - 数据: LAION-5B (50 亿图文对) - 骨干: EVA-02 (视觉编码器改进版) - 关键改进: 1. CLIP 预训练 → EVA 预训练 → CLIP 对比微调(两阶段) 2. 更大的 ViT-G (2B 参数) 3. 渐进式训练策略
结果: - EVA-CLIP-18B: ImageNet 零样本 82.0% - EVA-CLIP-18B + 融合: 86.5%7.3 SigLIP:sigmoid 损失替代 softmax
SigLIP(Google, 2024)用成对 sigmoid 损失替代 CLIP 的 softmax InfoNCE:
其中 当且仅当 (正样本对)。
def siglip_loss(image_features, text_features, temperature): """ SigLIP 的 Sigmoid 对比损失。
特点: - 每个样本独立优化(不需要 global softmax) - 可以处理更大的 batch size - 避免正样本被其他正样本"稀释" """ image_features = F.normalize(image_features, dim=-1) text_features = F.normalize(text_features, dim=-1)
logits = image_features @ text_features.T / temperature # (B, B)
# 正样本标签 B = image_features.shape[0] labels = torch.eye(B, device=image_features.device)
# Sigmoid 交叉熵(而非 softmax) loss = F.binary_cross_entropy_with_logits(logits, labels)
return lossSigLIP vs CLIP:
| 维度 | CLIP (softmax) | SigLIP (sigmoid) |
|---|---|---|
| 损失函数 | InfoNCE (softmax) | Pairwise sigmoid |
| Batch 内交互 | 归一化(全局 softmax) | 独立(成对) |
| 最优温度 | ~0.07 | ~0.1 |
| 大 batch 依赖 | 高(需要大 batch 提供足够负样本) | 低 |
| 精度 | 76.2% | 78.4% (ImageNet zero-shot) |
8. 完整实现
import torchimport torch.nn as nnimport torch.nn.functional as Ffrom dataclasses import dataclass
@dataclassclass CLIPConfig: """CLIP 模型配置。""" # 视觉编码器 image_size: int = 224 patch_size: int = 14 vision_embed_dim: int = 768 # ViT-L/14 vision_layers: int = 24 vision_heads: int = 16
# 文本编码器 vocab_size: int = 49408 text_embed_dim: int = 512 text_layers: int = 12 text_heads: int = 8 context_length: int = 77
# 融合 embed_dim: int = 768 # 视觉和文本特征最终维度
class ImagePatcher(nn.Module): """图像分块层。""" def __init__(self, patch_size=14, in_channels=3, embed_dim=768): super().__init__() self.proj = nn.Conv2d(in_channels, embed_dim, kernel_size=patch_size, stride=patch_size)
def forward(self, x): x = self.proj(x) # (B, embed_dim, H/P, W/P) x = x.flatten(2).transpose(1, 2) # (B, num_patches, embed_dim) return x
class CLIPVisionEncoder(nn.Module): """CLIP 的 ViT 视觉编码器。""" def __init__(self, config: CLIPConfig): super().__init__() self.config = config self.num_patches = (config.image_size // config.patch_size) ** 2
self.patch_embed = ImagePatcher(config.patch_size, 3, config.vision_embed_dim)
# [CLS] token self.cls_token = nn.Parameter(torch.zeros(1, 1, config.vision_embed_dim))
# 位置编码 self.pos_embed = nn.Parameter( torch.zeros(1, self.num_patches + 1, config.vision_embed_dim) )
# Transformer 编码器 encoder_layer = nn.TransformerEncoderLayer( d_model=config.vision_embed_dim, nhead=config.vision_heads, dim_feedforward=config.vision_embed_dim * 4, activation="gelu", batch_first=True, norm_first=True, ) self.transformer = nn.TransformerEncoder(encoder_layer, config.vision_layers)
self.ln_post = nn.LayerNorm(config.vision_embed_dim)
def forward(self, x): B = x.shape[0]
# 分块 x = self.patch_embed(x) # (B, num_patches, embed_dim)
# [CLS] token cls_tokens = self.cls_token.expand(B, -1, -1) x = torch.cat([cls_tokens, x], dim=1)
# 位置编码 x = x + self.pos_embed
# Transformer x = self.transformer(x) x = self.ln_post(x)
return x
class CLIPTextEncoder(nn.Module): """CLIP 的文本编码器。""" def __init__(self, config: CLIPConfig): super().__init__() self.config = config
# Token 嵌入 self.token_embedding = nn.Embedding(config.vocab_size, config.text_embed_dim)
# 位置编码 self.pos_embedding = nn.Parameter( torch.zeros(1, config.context_length, config.text_embed_dim) )
# Transformer encoder_layer = nn.TransformerEncoderLayer( d_model=config.text_embed_dim, nhead=config.text_heads, dim_feedforward=config.text_embed_dim * 4, activation="gelu", batch_first=True, norm_first=True, ) self.transformer = nn.TransformerEncoder(encoder_layer, config.text_layers)
self.ln_final = nn.LayerNorm(config.text_embed_dim)
def forward(self, input_ids): B, L = input_ids.shape
# Token 嵌入 + 位置编码 x = self.token_embedding(input_ids) + self.pos_embedding[:, :L, :]
# Transformer x = self.transformer(x) x = self.ln_final(x)
# 取 [EOS] token 作为句子表示 return x
class CLIP(nn.Module): """ CLIP 完整模型。 论文: "Learning Transferable Visual Models From Natural Language Supervision" """
def __init__(self, config: CLIPConfig): super().__init__() self.config = config
# 编码器 self.visual = CLIPVisionEncoder(config) self.transformer = CLIPTextEncoder(config)
# 投影层:统一特征维度 self.visual_projection = nn.Linear(config.vision_embed_dim, config.embed_dim) self.text_projection = nn.Linear(config.text_embed_dim, config.embed_dim)
# 可学习的温度(logit scale) self.logit_scale = nn.Parameter(torch.ones([]) * 2.659)
# 分词器 from transformers import GPT2Tokenizer self.tokenizer = GPT2Tokenizer.from_pretrained("gpt2") self.tokenizer.add_special_tokens(["<|startoftext|>", "<|endoftext|>"]) self.tokenizer.pad_token = self.tokenizer.eos_token
def encode_image(self, image): """编码图像。""" x = self.visual(image) # (B, num_patches+1, vision_embed_dim) x = x[:, 0, :] # [CLS] token: (B, vision_embed_dim) x = self.visual_projection(x) # (B, embed_dim) return F.normalize(x, dim=-1)
def encode_text(self, text): """ 编码文本。
参数: text: str 或 list[str] 或 token IDs """ # Tokenize if isinstance(text, str): text = [text]
if isinstance(text, list) and isinstance(text[0], str): encoded = self.tokenizer( text, padding=True, truncation=True, max_length=self.config.context_length, return_tensors="pt", ) input_ids = encoded["input_ids"].to(next(self.parameters()).device) else: input_ids = text.to(next(self.parameters()).device)
x = self.transformer(input_ids) # (B, L, text_embed_dim) x = x[torch.arange(x.shape[0]), input_ids.argmax(dim=-1)] # [EOS] token x = self.text_projection(x) # (B, embed_dim) return F.normalize(x, dim=-1)
def forward(self, image, text, return_loss=False): """ 完整前向。
参数: image: (B, 3, H, W) 图像 text: list[str] 文本列表 return_loss: 是否返回 InfoNCE 损失 """ # 编码 image_features = self.encode_image(image) # (B, embed_dim) text_features = self.encode_text(text) # (B, embed_dim)
if return_loss: # 计算对比损失 logit_scale = self.logit_scale.exp() logits_per_image = image_features @ text_features.T * logit_scale logits_per_text = logits_per_image.T
B = image.shape[0] labels = torch.arange(B, device=image.device)
loss_i2t = F.cross_entropy(logits_per_image, labels) loss_t2i = F.cross_entropy(logits_per_text, labels) loss = (loss_i2t + loss_t2i) / 2 return {"loss": loss, "image_features": image_features, "text_features": text_features}
return {"image_features": image_features, "text_features": text_features}
@torch.no_grad() def zero_shot_classify(self, image, class_names, prompt_template="a photo of a {}."): """ 零样本分类。
参数: image: (3, H, W) 图像 tensor class_names: list[str] 类别名称 prompt_template: str 带占位符的模板 """ # 构建提示 prompts = [prompt_template.format(c) for c in class_names]
# 编码 image_features = self.encode_image(image.unsqueeze(0)) text_features = self.encode_text(prompts)
# 相似度 logit_scale = self.logit_scale.exp() similarity = (image_features @ text_features.T) * logit_scale probs = F.softmax(similarity, dim=-1)[0]
return probs.cpu().numpy()9. 应用场景
9.1 开放词汇检测与分割
CLIP 的特征可以泛化到任意视觉概念,因此被广泛用于:
def clip_applications(): """ CLIP 的主要应用场景。 """ apps = { "开放词汇检测 (OVOD)": { "代表": "GLIP, MDETR, OWLVIT", "方法": "用 CLIP 特征替换检测器的类别嵌入", "效果": "检测任意文本描述的物体", }, "零样本分割 (ZS3)": { "代表": "Lseg, ZeroSeg, SA-V", "方法": "CLIP 特征 + 分割模型", "效果": "用文本描述来分割图像区域", }, "图像检索": { "代表": "LAION-5B 检索", "方法": "图像-图像 / 文本-图像相似度搜索", "效果": "给定图像找相似图 / 给定文本找图像", }, "图像生成条件": { "代表": "Stable Diffusion, DALL-E 3", "方法": "CLIP text encoder 作为生成模型的条件", "效果": "CLIP 理解文本 → 指导图像生成", }, "分布外检测": { "代表": "CLIP Score, Winoground", "方法": "用 CLIP 相似度衡量图像-文本一致性", "效果": "判断模型是否正确理解图像内容", }, } return apps9.2 CLIP 作为视觉特征提取器
def extract_visual_features(model, image_paths): """ 用 CLIP 提取视觉特征,用于下游任务。 """ features = [] for path in image_paths: image = Image.open(path).convert("RGB") image_tensor = preprocess(image).unsqueeze(0)
with torch.no_grad(): feat = model.encode_image(image_tensor)
features.append(feat)
return torch.cat(features, dim=0) # (N, embed_dim)
def extract_text_features(model, texts): """ 用 CLIP 提取文本特征。 """ with torch.no_grad(): features = model.encode_text(texts) return features10. 局限性
def clip_limitations(): """ CLIP 的主要局限性。 """ limitations = { "分布偏置": { "问题": "训练数据偏向英语和西方文化,对其他语言和文化适应性差", "数据": "WIT-400M 主要是英语图文对", }, "长尾分布": { "问题": "对常见物体分类好,对稀有/抽象概念分类差", "原因": "训练数据中稀有概念出现少", }, "文字识别": { "问题": "CLIP 不擅长 OCR,文字区域的特征不佳", "改进": "需要额外的 OCR 预训练(如 CLIP-OCR)", }, "空间关系": { "问题": "CLIP 不擅长理解物体之间的空间关系", "示例": "'红在蓝左边' vs '红在蓝右边' 难以区分", "原因": "图文对数据中空间描述较少", }, "计数": { "问题": "CLIP 不擅长精确计数", "示例": "'三只猫' vs '两只猫' 容易混淆", "原因": "自然语言中精确数量描述少", }, "对抗鲁棒性": { "问题": "对对抗噪声的鲁棒性较低", "原因": "对比学习没有显式防御机制", }, "组合推理": { "问题": "不擅长理解属性的组合(颜色+形状+材质)", "示例": "'蓝色的圆形' vs '红色的方形'", "改进": "需要组合推理能力的改进(LEO, COMET)", }, } return limitations11. 总结
11.1 核心要点
| 维度 | 关键要点 |
|---|---|
| 核心创新 | 用自然语言作为监督信号,打破类别标签的限制 |
| 模型架构 | 双塔:ViT 视觉编码器 + Transformer 文本编码器 |
| 损失函数 | 对称 InfoNCE(两方向的对比损失) |
| 训练数据 | WIT-400M(4 亿图文对) |
| 零样本 | 任何文本描述都可以作为分类标签 |
| 提示工程 | 多种提示模板集成可显著提升精度 |
| 扩展方向 | ALIGN(更大数据)/ EVA-CLIP(两阶段训练)/ SigLIP(sigmoid 损失) |
| 下游应用 | 开放词汇检测、分割、图像检索、生成条件、图文匹配 |
11.2 一句话总结
CLIP 的核心思想简单而深刻:与其让人类给图像打上固定的类别标签,不如直接用自然语言作为监督——让视觉编码器和文本编码器在同一个特征空间里”对齐”。通过 4 亿图文对的对比学习,CLIP 学到了通用的视觉-语义表示,实现了真正的零样本迁移,成为几乎所有多模态模型(VLM、VLA、图像生成)的视觉编码器基础。
11.3 推荐资源
论文: - CLIP (Radford et al., 2021): "Learning Transferable Visual Models From Natural Language Supervision" - ALIGN (Jia et al., 2021): "Scaling Up Visual and Vision-Language Representation Learning" - EVA-CLIP (Sun et al., 2023): "EVA-CLIP: Improved Training of CLIP on Large-scale Web Data" - SigLIP (Zhai et al., 2024): "Sigmoid Loss for Language-Image Pre-Training" - OpenCLIP (Ilharco et al., 2021): "OpenCLIP"
代码: - openai/CLIP (官方实现) - mlfoundations/open_clip (OpenCLIP, 支持更多模型) - facebookresearch/CLIP (PyTorch 官方重实现)文章分享
如果这篇文章对你有帮助,欢迎分享给更多人!

