从零手写 CLIP:对比学习如何让模型"看懂"图文匹配?|多模态大模型专栏②

一句话讲透 CLIP:别去预测图像里有什么词,而是让匹配的图文在向量空间里靠得最近。 这就是对比学习(contrastive learning),也是所有现代 VLM「那双眼睛」的训练方法。


写在前面

上一篇 建立了 VLM 的三组件心智模型:图像 → 视觉编码器 → Projector → LLM。其中视觉编码器被形容为"租来的眼睛"——其权重在 VLM 训练开始前便已固定。

那这双眼睛是怎么训练出来的?答案就是 CLIP(Radford et al., 2021)。它开创的对比学习范式,至今仍是绝大多数 VLM(LLaVA、Qwen-VL、InternVL)视觉编码器的训练基础,SigLIP 更是它的直接改进版、当前 VLM 的默认编码器。

本篇从零手写一个 CLIP,阅读后可获得:

  • ✅ 画出 CLIP 的双塔架构,说清图文如何被拉到同一个向量空间
  • ✅ 逐行理解 InfoNCE 损失,知道"对角线"为什么是核心
  • ✅ 搞懂 CLIP(softmax)和 SigLIP(sigmoid)的本质区别,以及为什么 SigLIP 不需要大 batch
  • ✅ 跑通一份不到 100 行的 PyTorch demo,亲眼看见 loss 下降

本文代码对应 code/01_vision/02_clip_from_scratch.py,可直接 python 运行。


一、CLIP 的核心思想:不预测,只对比

传统视觉模型是预测式的:给定图像,要求输出"这是一只猫"之类的判断。其局限在于类别必须预先定义(猫、狗、车……),难以泛化到新类别。

CLIP 换了个思路——对比式

给模型一批 (图像, 文本) 配对,不要求它说出图像里有什么,只要求它判断哪段文本和哪张图是配对的

具体说,同一批里有 N 个配对,模型要在一个 N×N 的匹配矩阵里,把对角线(正确配对)的分数拉高,把非对角线(错误配对)的分数压低。

🖼️ 图像 batch
(N, 3, H, W)

Image Encoder
图像编码器

img_emb
(N, d)

L2 归一化

📝 文本 batch
token ids (N, L)

Text Encoder
文本编码器

text_emb
(N, d)

L2 归一化

相似度矩阵
S = img_emb · text_embᵀ
(N, N)

对称 InfoNCE 损失
对角线 = 正样本

关键设计:两个编码器把图像和文本分别压成 d 维向量,并且共享同一个向量空间。训练后,“一只猫的图片"和文本"a photo of a cat"的向量会非常接近——这就是"图文对齐”。

💡 这个"共享向量空间"是 CLIP 留给后世最重要的遗产。VLM 的 Projector 干的活,本质就是把 LLM 的向量空间"对接"到这个视觉向量空间上。


二、逐行实现

Step 1:图像编码器

真实 CLIP 的图像编码器是 ViT(不熟的话,可把它理解成"把图切成 patch、当 token 喂给 Transformer",详见 01_vit_from_scratch.py)。这里为了 demo 跑得快,用几层卷积替代,原理一样:把图像压成 d 维向量

class ImageEncoder(nn.Module):
    def __init__(self, d: int = 256):
        super().__init__()
        self.conv = nn.Sequential(
            nn.Conv2d(3, 32, 3, 2, 1), nn.GELU(),    # 32x32 -> 16x16
            nn.Conv2d(32, 64, 3, 2, 1), nn.GELU(),   # 16x16 -> 8x8
            nn.Conv2d(64, 128, 3, 2, 1), nn.GELU(),  # 8x8 -> 4x4
            nn.AdaptiveAvgPool2d(1),                 # 全局池化 -> 1x1
            nn.Flatten(),
        )
        self.proj = nn.Linear(128, d)                # 压到 d 维

    def forward(self, x):
        return self.proj(self.conv(x))               # (N, d)

Step 2:文本编码器

文本侧更简单:词嵌入 + masked 平均池化 + 一层线性。真实 CLIP 用 Transformer,这里同样简化。

class TextEncoder(nn.Module):
    def __init__(self, vocab_size=1000, d=256):
        super().__init__()
        self.embed = nn.Embedding(vocab_size, d)
        self.proj = nn.Linear(d, d)

    def forward(self, token_ids, attention_mask):
        x = self.embed(token_ids) * attention_mask.unsqueeze(-1).float()
        # masked mean pool:只在非 padding 位置上取平均
        x = x.sum(dim=1) / attention_mask.sum(dim=1, keepdim=True).clamp(min=1)
        return self.proj(x)                          # (N, d)

Step 3:归一化 + 可学习温度(CLIP 的两个灵魂细节)

把两路输出都 L2 归一化,这样点积就是余弦相似度,取值落在 [-1, 1],训练稳定。再乘一个可学习的温度 logit_scale

class CLIP(nn.Module):
    def __init__(self, d=256, vocab_size=1000, init_temperature=0.07):
        super().__init__()
        self.image_encoder = ImageEncoder(d)
        self.text_encoder = TextEncoder(vocab_size, d)
        # 可学习温度,初始 0.07(论文设置)
        self.logit_scale = nn.Parameter(torch.tensor(1.0 / init_temperature).log())

    def encode_image(self, x):
        return F.normalize(self.image_encoder(x), dim=-1)   # ← L2 归一化

    def encode_text(self, ids, mask):
        return F.normalize(self.text_encoder(ids, mask), dim=-1)

    def forward(self, image, text_ids, text_mask):
        img_emb = self.encode_image(image)            # (N, d)
        txt_emb = self.encode_text(text_ids, text_mask)
        logit_scale = self.logit_scale.exp().clamp(max=100.0)
        logits_per_image = logit_scale * img_emb @ txt_emb.t()  # (N, N)
        return logits_per_image, logits_per_image.t()

温度为什么重要? 温度越小,softmax 越尖锐,模型对"正确配对"和"错误配对"的区分越狠。CLIP 训练后学到的温度 ≈ 0.01(非常尖锐),这说明模型变得极度自信。

Step 4:对称 InfoNCE 损失

这是 CLIP 的心脏。给定 N×N 相似度矩阵,每一行的正确答案就是对角线那个位置

def clip_loss(logits_per_image, logits_per_text):
    N = logits_per_image.shape[0]
    labels = torch.arange(N)                          # [0, 1, 2, ..., N-1] ← 对角线
    loss_i = F.cross_entropy(logits_per_image, labels)  # 图→文
    loss_t = F.cross_entropy(logits_per_text, labels)   # 文→图
    return (loss_i + loss_t) / 2                       # 对称

labels = [0,1,...,N-1] 这行是精髓:它告诉交叉熵"第 0 行的正确列是第 0 列,第 1 行是第 1 列……"——正好是对角线。两个方向的 loss 取平均,保证图能找文、文也能找图。


三、一个 N×N 矩阵,看懂全部

光看代码不直观。假设 batch 里有 3 个图文对,训练后(理想状态下)的相似度矩阵长这样:

文本A「一只猫」 文本B「沙滩海浪」 文本C「城市夜景」
图A 🐱 0.85 0.10 0.12
图B 🏖️ 0.08 0.90 0.15
图C 🌃 0.11 0.13 0.88
  • 对角线 = 正样本配对(图A 配 文本A),CLIP 要把它们拉到最大
  • 其余格子 = 负样本配对(图A 配 文本B/C),要压到最小

InfoNCE 在每一行做 softmax,目标是让对角线那一列的概率最大。用公式写就是:

L img = − 1 N ∑ i = 1 N log ⁡ exp ⁡ ( s i i / τ ) ∑ j = 1 N exp ⁡ ( s i j / τ ) \mathcal{L}_{\text{img}} = -\frac{1}{N}\sum_{i=1}^{N} \log \frac{\exp(s_{ii}/\tau)}{\sum_{j=1}^{N}\exp(s_{ij}/\tau)} Limg=N1i=1Nlogj=1Nexp(sij/τ)exp(sii/τ)

其中 s i j s_{ij} sij 是第 i i i 张图和第 j j j 段文本的相似度, τ \tau τ 是温度。文本侧对称地做一遍,两者求平均。

📌 配图说明:上表用 emoji + ✅ 高亮对角线,发布即可。想更直观可改成 mermaid 热力图,但表格在移动端可读性更好,建议保留表格

直觉记忆:CLIP 的全部工作,就是把这个矩阵的对角线擦亮、非对角线擦暗


四、InfoNCE vs SigLIP:为什么后者成了 VLM 默认编码器

CLIP 虽经典,但有个硬伤:softmax 把一行的 N 个格子耦合在一起,意味着必须有足够多的负样本(大 batch)才能学到好的对比。OpenAI 训 CLIP 用了 batch size = 32768,普通实验室根本玩不起。

SigLIP(Zhai et al., 2023)的改进极简却致命:别 softmax 了,每个格子独立做二分类

def siglip_loss(img_emb, txt_emb, logit_scale, bias):
    N = img_emb.shape[0]
    logits = logit_scale * (img_emb @ txt_emb.t()) + bias   # (N, N)
    labels = torch.eye(N)            # 单位阵:对角线=1(正),其余=0(负)
    return F.binary_cross_entropy_with_logits(logits, labels)

torch.eye(N) 生成单位阵——对角线是 1(正样本),其余是 0(负样本),每个格子独立算 BCE,互不耦合。完整对比:

CLIP(InfoNCE / softmax) SigLIP(sigmoid)
每行操作 softmax over N 个候选 每个 cell 独立 BCE
正样本 对角线是 argmax 目标 对角线 label = 1
负样本 同行其余 N-1 个 非对角线 label = 0
对 batch 的要求 大(数千级) 小 batch 也行,靠大累计
随机初始化 loss ≈ ln(N)(N=8 时 ≈2.08) ≈ ln(2) ≈ 0.69
当前地位 经典开山 VLM 默认编码器

💡 这就是为什么第一篇里说 SigLIP-SO400M 是当前 VLM 的"租来的眼睛"——它把 CLIP 的对齐思想保留下来,又摆脱了对超大 batch 的依赖,让中等规模数据/算力也能训出强编码器。Qwen-VL、InternVL 的视觉侧都用它。


五、跑起来:完整 demo

把上面拼起来,跑一个 N=8 的假数据 demo:

def demo():
    torch.manual_seed(0)
    N, d, vocab = 8, 256, 100
    images = torch.randn(N, 3, 32, 32)
    text_ids = torch.randint(0, vocab, (N, 10))
    text_mask = torch.ones(N, 10, dtype=torch.long)

    model = CLIP(d=d, vocab_size=vocab)
    logits_i, logits_t = model(images, text_ids, text_mask)

    print("=== CLIP (softmax / InfoNCE) ===")
    print(f"similarity matrix shape: {tuple(logits_i.shape)}")
    print(f"CLIP loss: {clip_loss(logits_i, logits_t).item():.4f}")

    # 同一组 embedding 上对比 SigLIP
    img_emb = model.encode_image(images)
    txt_emb = model.encode_text(text_ids, text_mask)
    print("\n=== SigLIP (sigmoid) ===")
    print(f"SigLIP loss: {siglip_loss(img_emb, txt_emb, model.logit_scale.exp(), torch.zeros(1)).item():.4f}")

预期输出(随机初始化时):

=== CLIP (softmax / InfoNCE) ===
similarity matrix shape: (8, 8)
CLIP loss: ~2.20        # ≈ ln(8),因为 8 个候选随机猜

=== SigLIP (sigmoid) ===
SigLIP loss: ~0.69      # ≈ ln(2),每个 cell 像抛硬币

这两个初始 loss 值非常值得记住

  • CLIP 初始 ≈ ln(N):因为每行要在 N 个候选里 softmax,全随机时正确概率 1/N,loss = -ln(1/N) = ln(N)
  • SigLIP 初始 ≈ ln(2):每个 cell 是独立的二分类,全随机时就是抛硬币,loss = -ln(0.5) = ln(2)

随着训练,对角线被擦亮,两个 loss 都会趋近 0。

🔧 动手建议:将 N 依次改为 16、32,可观察到 CLIP 的初始 loss 随 ln(N) 增长,而 SigLIP 始终保持在 0.69 附近。这一对比能直观揭示"为何 CLIP 依赖大 batch"。


六、CLIP 何以成为所有 VLM 的"眼睛"

回到专栏主线。CLIP/SigLIP 训练完后,图像编码器的权重就冻结了,它会把任何图像编码成一个高质量的视觉向量。后续所有 VLM 工作:

  • LLaVA:直接拿 CLIP-ViT-L 当眼睛,后面接 MLP Projector → LLM
  • Qwen-VL / InternVL:拿 SigLIP-SO400M 当眼睛,配 M-RoPE + 动态分辨率

所以搞懂 CLIP 的对比学习,就搞懂了 VLM 视觉侧的根基。Projector 要做的"维度翻译",翻译的源语言正是 CLIP 学到的这个视觉向量空间。


📎 本篇涉及代码

文件 说明
code/01_vision/02_clip_from_scratch.py 本文核心:手写 CLIP + SigLIP 损失,含 InfoNCE vs sigmoid 对比
code/01_vision/01_vit_from_scratch.py 最小 ViT 实现(CLIP 图像编码器的真实形态)

小结

概念 一句话
CLIP 目标 把匹配图文拉近、不匹配拉远,不预测具体类别
双塔结构 图像编码器 + 文本编码器,输出到同一 d 维空间
InfoNCE N×N 矩阵每行 softmax,对角线是正样本
温度 τ 可学习,越小对比越尖锐
SigLIP 改进 每格独立 BCE,摆脱大 batch 依赖 → VLM 默认编码器

写在最后

下一篇是专栏的"张量级 walkthrough":我们会用一个零依赖纯 Python demo,跟着一个 batch 的张量从头走到尾,亲眼看「图像 → patch → ViT → Projector → 拼进文本流 → LLM」每一步的 shape 变化。关注专栏,更新第一时间通知。

专栏:《从 NLP 到 VLM:多模态大模型研发实战》
上一篇:① 多模态大模型全景图 | 本篇:② 手写 CLIP | 下一篇:③ 张量级端到端 walkthrough

Logo

汇聚全球AI编程工具,助力开发者即刻编程。

更多推荐