概述

Siglip2 用 Sigmoid Loss 来评估图像和文本的对齐程度——这是它和 CLIP 最核心的算法差异。

CLIP 用的是 Softmax Cross-Entropy Loss:在一个 batch 内,每张图在 T 条文本中选出正确的那条(反过来也一样)。问题是:batch 里没有的正确文本怎么办? Softmax 的分母只在 batch 内归一化,batch 越大负样本越多训练越有效——batch 小则退化严重。

Siglip2 的 Sigmoid Loss 换个思路:不比较,直接判定每个 (图像, 文本) pair 是匹配还是不匹配。每对独立做二分类,不依赖 batch 内其他样本。这解耦了 batch size 和负样本数量的关系。


源码定位

核心代码在 modeling_siglip2.pySiglip2Model.forward 方法中,仅 7 行:

1
2
3
4
5
6
7
8
# 位置: Siglip2Model.forward, 行 ~870-876
# Adapted from https://github.com/google-research/big_vision/blob/.../siglip2.py#L287

eye = torch.eye(logits_per_text.size(0), device=logits_per_text.device)
m1_diag1 = -torch.ones_like(logits_per_text) + 2 * eye
loglik = torch.nn.functional.logsigmoid(m1_diag1 * logits_per_text)
nll = -torch.sum(loglik, dim=-1)
loss = nll.mean()

第一步:从图像和文本到 logits

1.1 编码

1
2
3
4
5
6
7
# 图像编码: Vision Transformer + MultiheadAttentionPoolingHead
vision_outputs = self.vision_model(pixel_values, pixel_attention_mask, spatial_shapes)
image_embeds = vision_outputs.pooler_output   # shape: (I, 768)

# 文本编码: Text Transformer + EOS pooling + Linear projection
text_outputs = self.text_model(input_ids, attention_mask, position_ids)
text_embeds = text_outputs.pooler_output      # shape: (T, 768)

1.2 L2 归一化

1
2
image_embeds = image_embeds / image_embeds.norm(p=2, dim=-1, keepdim=True)
text_embeds = text_embeds / text_embeds.norm(p=2, dim=-1, keepdim=True)

归一化后,每个嵌入向量的 L2 范数 = 1。此时:

image_embeds[i] @ text_embeds[j]  =  cos(image_i, text_j)

两个归一化向量的点积就是余弦相似度,范围 [-1, 1]。

为什么必须 L2 归一化? 如果不归一化,模型可能通过拉长向量来"作弊"——让正样本点积变大但不是真的语义对齐。归一化把所有向量约束在同一球面上,相似度只反映方向而非长度。

1.3 计算相似度矩阵

1
2
3
4
logits_per_text = torch.matmul(text_embeds, image_embeds.t())
# shape: (T, I)
#
# logits_per_text[i][j] = cos(text_i, image_j)

如果 batch 里有 3 张图和 3 条文本(对角线是匹配对):

              img_0    img_1    img_2
text_0  [    0.85   , -0.12   ,  0.03   ]   ← 正样本 (匹配)
text_1  [    0.10   ,  0.72   , -0.08   ]   ← 正样本 (匹配)
text_2  [   -0.15   ,  0.05   ,  0.91   ]   ← 正样本 (匹配)
                               ↑
                对角线 = 匹配对 = 应该判为 positive
                非对角线 = 不匹配对 = 应该判为 negative

1.4 应用可学习温度和偏置

1
2
3
4
logit_scale = self.logit_scale   # 初始化为 0 (正态分布随机 但被 _init_weights 覆盖为 0)
logit_bias = self.logit_bias     # 初始化为 0

logits_per_text = logits_per_text * logit_scale.exp() + logit_bias

logit_scalelogit_bias可学习的标量参数nn.Parameter),形状都是 (1,)

参数初始值作用训练后典型值
logit_scale0exp(0)=1,不做缩放~3 (exp(3)≈20),大幅放大余弦相似度
logit_bias0不做偏移~ -5,将整体分布左移,提高判定门槛

为什么要可学习而不是固定?

CLIP 使用固定的温度 exp(t)(t ≈ log(1/0.07) ≈ 2.66),是所有 CLIP 模型的超参数,需要人工调。Siglip2 让模型自己学——训练开始时不确定(exp(0)=1),训练结束时有信心(exp(3)≈20)

logit_bias 的负偏移(~-5)意味着:即使余弦相似度 = 0,模型也会偏向判为"不匹配"——因为数据集中负样本远比正样本多。


第二步:构造 Target 矩阵

1
2
3
4
5
6
7
8
9
eye = torch.eye(logits_per_text.size(0), device=logits_per_text.device)
# eye = [[1, 0, 0],
#        [0, 1, 0],
#        [0, 0, 1]]

m1_diag1 = -torch.ones_like(logits_per_text) + 2 * eye
# m1_diag1 = [[+1, -1, -1],
#             [-1, +1, -1],
#             [-1, -1, +1]]

m1_diag1 是一个矩阵,每个元素的值:

m1_diag1[i][j] = +1  if i == j  (对角线, 匹配对, 正样本)
m1_diag1[i][j] = -1  if i != j  (非对角线, 不匹配对, 负样本)

这就是 sigmoid loss 的关键:target 不是 one-hot 向量,而是每个 (i,j) pair 都有自己的二分类标签。


第三步:计算 Sigmoid Loss

3.1 logsigmoid 函数

1
loglik = torch.nn.functional.logsigmoid(m1_diag1 * logits_per_text)

logsigmoid(x) = log(1 / (1 + exp(-x))) = -softplus(-x)

函数图像:

logsigmoid(x)
   0 ┤························━━━━━━━━━━━━━━
     |                     ··
     |                   ··
     |                 ··
     |               ··
  -5 ┤━━━━━━━━━━━━━━·
     +-----------+-----------+-----------+---→ x
    -5          0          5         10

特征:
- x → +∞: logsigmoid(x) → 0(接近饱和,loss 很低)
- x → -∞: logsigmoid(x) → -∞(loss 很高)
- x = 0:  logsigmoid(0) = log(0.5) ≈ -0.693

3.2 逐元素计算

m1_diag1 * logits_per_text 的效果:

对于正样本 (i == j):
  target = +1
  sigmoid_input = +1 * logits[i][j]
  → logits[i][j] 越大 → sigmoid_input 越大 → logsigmoid 越接近 0 → loss 越低
  → logits[i][j] 越小 → sigmoid_input 越小 → logsigmoid 越负 → loss 越高

对于负样本 (i != j):
  target = -1
  sigmoid_input = -1 * logits[i][j]
  → logits[i][j] 越小(越负)→ sigmoid_input 越大 → logsigmoid 越接近 0 → loss 越低
  → logits[i][j] 越大(越正)→ sigmoid_input 越小 → logsigmoid 越负 → loss 越高

用一个实际的 logits 矩阵来演示:

假设 logits_per_text:
           img_0  img_1  img_2
text_0  [   3.0 ,  1.5 ,  0.2  ]
text_1  [  -0.5 ,  4.0 , -0.3  ]
text_2  [  -1.0 ,  0.8 ,  2.5  ]

m1_diag1:
         [+1, -1, -1]
         [-1, +1, -1]
         [-1, -1, +1]

sigmoid_input = m1_diag1 * logits_per_text:
           img_0  img_1  img_2
text_0  [   3.0 , -1.5 , -0.2  ]    ← text_0 vs img_0: +3.0 (正, 高分)
text_1  [   0.5 ,  4.0 ,  0.3  ]    ← text_1 vs img_1: +4.0 (正, 高分)
text_2  [   1.0 , -0.8 ,  2.5  ]    ← text_2 vs img_2: +2.5 (正, 高分)

# 对角线全部为正(高分 input → 低 loss ✓)
# 非对角线全部为负或小正(低/负 input → 适度 loss)

logsigmoid(sigmoid_input):
           img_0   img_1   img_2
text_0  [ -0.049 , -1.701 , -0.598 ]
text_1  [ -0.474 , -0.018 , -0.554 ]
text_2  [ -0.313 , -1.371 , -0.079 ]

# 对角线 (正样本) 的 -logsigmoid 都很小 → loss 低
# 非对角线 (负样本) 的 -logsigmoid 有大有小:
#   text_0 vs img_1: logits=1.5 → 余弦相似度偏高 → sigmoid_input=-1.5 → loss=1.701 (惩罚!)
#   text_1 vs img_2: logits=-0.3 → 余弦相似度低 → sigmoid_input=0.3 → loss=0.554 (OK)

3.3 聚合

1
2
nll = -torch.sum(loglik, dim=-1)   # 对图像维度求和 → (T,)
loss = nll.mean()                   # 对文本维度求平均 → 标量
nll = [-log(0.049) + -log(1.701) + -log(0.598)] = [0.049 + 1.701 + 0.598] ≈ [2.35]
     └─ text_0 对所有 3 张图像的负 log-likelihood 之和

nll = [
    2.35,   # text_0: 对角匹配好 (0.05), img_1 误识高 (1.70), img_2 OK (0.60)
    1.05,   # text_1: 对角极好 (0.02), 两个误识都低 (0.47+0.55) → 最优
    1.76,   # text_2: 对角好 (0.08), img_1 误识偏高 (1.37), img_0 OK (0.31)
]

loss = mean([2.35, 1.05, 1.76]) ≈ 1.72

第四步:与 CLIP Softmax CE Loss 的对比

4.1 CLIP 的做法

1
2
3
4
5
6
7
8
9
# CLIP 的 loss (伪代码)
logits = image_embeds @ text_embeds.T * exp(temperature)  # (I, T)

labels = torch.arange(I)  # [0, 1, 2, ..., I-1]

loss_img = CrossEntropyLoss(logits, labels)      # 每张图在 T 条文本中选对的
loss_txt = CrossEntropyLoss(logits.T, labels)     # 每条文本在 I 张图中选对的

loss = (loss_img + loss_txt) / 2

4.2 核心差异

维度CLIP (Softmax CE)Siglip2 (Sigmoid)
判定逻辑“在一堆文本中,哪个最匹配这张图?”“这个文本和这张图,是匹配还是不匹配?”
计算方式对每一行做 softmax → -log(对角)对每个 (i,j) 做 sigmoid → 逐元素 log
负样本来源仅限当前 batch。batch_size=4 → 每样本只有 3 个负样本每个 (i,j) 都是独立判定。可利用额外预计算的负样本
batch size 依赖强依赖。小 batch 效果显著下降(CLIP 需要 batch=32768)弱依赖。小 batch 也有效
温度参数固定值(超参数,如 exp(2.66))可学习(logit_scale + logit_bias
对角线 vs 非对角线对角线通过 softmax 的归一化分母间接影响对角线和非对角线各算各的,互不干扰
梯度梯度集中在对角线和与对角线竞争的非对角线上每个 (i,j) pair 都有独立的梯度信号
推理时的概率softmax → 概率加和为 1(相对)sigmoid → 每个 pair 独立概率(绝对)

4.3 数值对比示例

假设一个 batch 有 3 张猫图 + 3 条"猫"文本:

clip_logits (固定温度 exp(2.66) ≈ 14.3):
           img_0  img_1  img_2
text_0  [  12.2 ,  2.1 ,  1.5  ]    softmax: [0.990, 0.006, 0.004]
text_1  [   1.8 , 11.5 ,  2.3  ]    softmax: [0.004, 0.985, 0.011]
text_2  [   1.2 ,  2.0 , 13.0  ]    softmax: [0.001, 0.003, 0.996]

对角概率都 > 0.98,效果很好。


siglip2_logits (可学习温度 exp(3) ≈ 20):
m1_diag1 * logits:
           img_0  img_1  img_2
text_0  [ +12.2 , -2.1 , -1.5  ]
text_1  [  -1.8 , +11.5 , -2.3 ]
text_2  [  -1.2 , -2.0 , +13.0 ]

sigmoid:
           img_0  img_1  img_2
text_0  [ 1.000 , 0.109 , 0.182 ]
text_1  [ 0.142 , 1.000 , 0.091 ]
text_2  [ 0.231 , 0.119 , 1.000 ]

对角 ≈ 1.0 (正),非对角 ≈ 0.0~0.2 (负),效果很好。

两种 loss 在这个例子中表现相似——因为 batch 内的负样本足够好。

差异在极端场景

场景: batch=2, 图片=[猫, 狗], 文本=["一只猫", "一只猫"]
     (两条文本都是"一只猫",没有"狗"文本!)

CLIP Softmax CE:
  logits:    猫图    狗图
  "一只猫"   [ 8.0 ,  5.0 ]    softmax: [0.953, 0.047]
  "一只猫"   [ 7.5 ,  4.8 ]    softmax: [0.937, 0.063]

  loss_txt = -log(0.953) - log(0.063) = 0.048 + 2.764 = 2.812
  ↑ 狗图被强制在这两条"猫"文本中选一个 → softmax 给狗图分走 5% → 混乱的信号


Siglip2 Sigmoid Loss:
  m1_diag1 * logits:
              猫图    狗图
  "一只猫"   [ +8.0 , -5.0 ]
  "一只猫"   [ -7.5 , +4.8 ]

  sigmoid:
              猫图    狗图
  "一只猫"   [ 1.000 , 0.007 ]    ← 狗 vs "一只猫": sigmoid(-5) ≈ 0.007 → 低, OK!
  "一只猫"   [ 0.001 , 0.992 ]    ← 猫 vs "一只猫": sigmoid(7.5) 被判负...但这不合理!

  还是有问题:text_1 和猫图的匹配被误判为负。但至少狗图没有被强制分配概率。
  而且,如果训练数据里有"狗"的文本,sigmoid 会正确学到 text_1("一只猫") vs img_0(猫) = 1。

核心结论:CLIP 的 softmax 强迫 batch 内竞争——即使 batch 里没有正确的负样本也必须"选一个"。Siglip2 的 sigmoid 独立判定每个 pair,batch 内负样本不够也不影响判断质量。


第五步:loss 函数的完整计算图

                    pixel_values              input_ids
                    (B, N, 768)              (B, L)
                         │                       │
                    ┌────┴────┐             ┌────┴────┐
                    │ Vision  │             │  Text   │
                    │Encoder  │             │Encoder  │
                    └────┬────┘             └────┬────┘
                         │                       │
                    pooler_output            pooler_output
                    (B, 768)                 (B, 768)
                         │                       │
                    L2 normalize             L2 normalize
                    (||·||₂ = 1)             (||·||₂ = 1)
                         │                       │
                         └───────┬───────────────┘
                                 │
                    logits_per_text = text @ image.T
                    logits_per_image = logits_per_text.T
                                 │
                    ┌────────────┴────────────┐
                    │    × exp(logit_scale)   │  可学习温度
                    │    + logit_bias         │  可学习偏置
                    └────────────┬────────────┘
                                 │
                    ┌────────────┴────────────┐
                    │  m1_diag1:              │
                    │    +1  if i==j          │  正样本 target
                    │    -1  if i!=j          │  负样本 target
                    └────────────┬────────────┘
                                 │
                    ┌────────────┴────────────┐
                    │  logsigmoid(            │
                    │    m1_diag1 * logits    │  逐元素计算
                    │  )                      │
                    └────────────┬────────────┘
                                 │
                    ┌────────────┴────────────┐
                    │  nll = -sum(            │  对图像维度求和
                    │    loglik, dim=-1       │  → 每个文本的负对数似然
                    │  )                      │
                    └────────────┬────────────┘
                                 │
                    ┌────────────┴────────────┐
                    │  loss = nll.mean()      │  对 batch 求平均
                    └────────────┬────────────┘
                                 │
                               loss
                              (标量)

第六步:为什么初始化 logit_scale=0, logit_bias=0

1
2
3
4
# _init_weights 中
elif isinstance(module, Siglip2Model):
    init.zeros_(module.logit_scale)   # → exp(0) = 1
    init.zeros_(module.logit_bias)    # → 不加偏置
训练阶段exp(logit_scale)logit_bias效果
初始1.00logits = 原始余弦相似度,sigmoid(cos) ≈ 0.27~0.73
早期2~5-1~0logits 开始被放大和偏移,模型学会区分
中期5~15-3~-2正负样本差距拉大
收敛10~30-5~-3正样本 sigmoid(+20)≈1,负样本 sigmoid(-20)≈0

初始时所有 pair 的概率都在 0.5 附近(不确定),随着训练逐步变确定。从 0 初始化让模型自己"发现"合适的温度和偏置,而非人为预设。

这个设计的优势:如果 CLIP 的固定温度选错了,整个训练都受影响。Siglip2 不需要人工调温度超参数。


总结

Siglip2 的 Sigmoid Loss 用 7 行代码实现了一个简单但强大的想法:

  1. L2 归一化 保证相似度只反映方向
  2. 可学习温度和偏置 让模型自己决定置信度刻度
  3. 每对 (i,j) 独立二分类 解耦 batch size 和负样本数量的关系
  4. 逐元素 logsigmoid 替代 softmax 的全局归一化

一句话:CLIP 问的是 “这些文本里谁是正确答案”(相对排序),Siglip2 问的是 “这个 pair 匹配吗”(绝对判定)。后者更强、更灵活、更省 batch。