概述

modeling_siglip2.py 是 Siglip2 模型在 HuggingFace Transformers 库中的完整 PyTorch 实现

项目说明
路径G:\Project\blog\content\posts\VLM\Siglip2\modeling_siglip2.py
行数~600 行
类数14 个(4 个输出 dataclass + 10 个 nn.Module)
来源自动生成自 src/transformers/models/siglip2/modular_siglip2.py不可手动编辑
依赖torch, torch.nn, PIL, numpy
核心模型Siglip2VisionModel, Siglip2TextModel, Siglip2Model, Siglip2ForImageClassification

文件内的类全景图

输出层 (Output Dataclasses)
├── Siglip2VisionOutput        ← Vision 模型的输出容器
├── Siglip2TextOutput          ← Text 模型的输出容器
└── Siglip2Output              ← 联合模型的输出容器(含 loss + logits)

嵌入层 (Embedding Layers)
├── Siglip2VisionEmbeddings    ← 图像 patch 嵌入 + 可缩放位置编码
└── Siglip2TextEmbeddings      ← 文本 token 嵌入 + 位置编码

Transformer 构件 (Building Blocks)
├── Siglip2Attention           ← 多头注意力(支持 eager/sdpa/flash)
├── Siglip2MLP                 ← 前馈网络(FFN)
├── Siglip2EncoderLayer        ← 单层 Transformer(Pre-LN 架构)
└── Siglip2Encoder             ← 多层 Transformer 堆叠

模型主体 (Model Bodies)
├── Siglip2PreTrainedModel     ← 基类:权重初始化 + 属性声明
├── Siglip2VisionModel         ← 视觉编码器 + 多头注意力池化头
├── Siglip2TextModel           ← 文本编码器 + EOS 池化 + 投影头
├── Siglip2Model               ← 双塔联合模型 + Sigmoid Loss
└── Siglip2MultiheadAttentionPoolingHead ← 可学习的多头注意力池化

下游任务 (Task Head)
└── Siglip2ForImageClassification ← 视觉编码器 + 平均池化 + 分类头

数据流全貌

[Image: H x W x 3]                   [Text: "a photo of a cat"]
        |                                       |
        v                                       v
image_processor                     tokenizer
  (resize, rescale/255,               (lowercase, tokenize,
   normalize to [-1,1],                pad to max_length)
   patchify to [N, 768])
        |                                       |
        v                                       v
pixel_values (B, N, 768)            input_ids (B, L)
pixel_attention_mask (B, N)         attention_mask (B, L)
spatial_shapes (B, 2)
        |                                       |
        v                                       v
Siglip2VisionEmbeddings             Siglip2TextEmbeddings
  - patch_embedding (Linear)          - token_embedding (Embedding)
  - position_embedding (resize)       - position_embedding (Embedding)
        |                                       |
        v                                       v
Siglip2VisionEncoder                Siglip2TextEncoder
  (Pre-LN x num_layers)               (Pre-LN x num_layers)
        |                                       |
        v                                       v
post_layernorm                      final_layer_norm
        |                                       |
        v                                       v
MultiheadAttentionPoolingHead       head (Linear projection)
  (learnable probe attention)         + EOS token pooling
        |                                       |
        v                                       v
image_embeds (B, dim)               text_embeds (B, dim)
        |                                       |
        +-------- L2 normalize -----------+
                      |
                      v
              cosine_similarity
                      |
                      v
         logits * logit_scale.exp() + logit_bias
                      |
                      v
              Sigmoid Loss (training)

分析框架

对每个类按以下维度解构:

+------------------------------------------------------------------+
| 1. 类签名与角色                                                     |
| 2. __init__ 参数全量表(每个变量的含义、默认值、来源、选择原因)         |
| 3. forward 逐行解析(每一行在做什么,为什么这么做而不是另一种方式)       |
| 4. 与测试文件的对应关系(哪些测试验证了这个类的行为)                    |
| 5. 设计意图与关键抉择                                                |
+------------------------------------------------------------------+

第一部分:输出数据类 (Output Dataclasses)

1. Siglip2VisionOutput

1
2
3
4
5
6
@dataclass
class Siglip2VisionOutput(ModelOutput):
    image_embeds: torch.FloatTensor | None = None     # (B, output_dim)
    last_hidden_state: torch.FloatTensor | None = None # (B, N, hidden_size)
    hidden_states: tuple[torch.FloatTensor, ...] | None = None
    attentions: tuple[torch.FloatTensor, ...] | None = None
变量shape含义
image_embeds(B, output_dim)投影后的图像嵌入向量(当 with_projection=True 时)。经过 projection layer 的 pooler_output
last_hidden_state(B, N, hidden_size)最后一层 Transformer 的隐藏状态。N = 实际 patch 数或 max_num_patches
hidden_statestuple of (B, N, hidden_size)所有层的隐藏状态(当 output_hidden_states=True 时)
attentionstuple of (B, heads, N, N)所有层的注意力权重(当 output_attentions=True 时)

设计目的

  • 继承 ModelOutput 获得 to_tuple().values() 等序列化方法
  • image_embeds 是额外的——不是所有 ViT 都输出这个字段,但 Siglip2 的联合训练需要它。last_hidden_state 用于调试/分析,image_embeds 才是下游真正用的特征向量
  • 所有字段默认 None:推理时可以只取需要的字段,不计算的就保持 None

2. Siglip2TextOutput

1
2
3
4
5
6
@dataclass
class Siglip2TextOutput(ModelOutput):
    text_embeds: torch.FloatTensor | None = None      # (B, output_dim)
    last_hidden_state: torch.FloatTensor | None = None # (B, L, hidden_size)
    hidden_states: tuple[torch.FloatTensor, ...] | None = None
    attentions: tuple[torch.FloatTensor, ...] | None = None

Siglip2VisionOutput 对称,字段含义相同,只是 image_embedstext_embeds

设计目的:Vision 和 Text 保持相同的输出结构,联合模型可以统一处理。


3. Siglip2Output

1
2
3
4
5
6
7
8
9
@dataclass
class Siglip2Output(ModelOutput):
    loss: torch.FloatTensor | None = None              # (1,)
    logits_per_image: torch.FloatTensor | None = None   # (I, T)
    logits_per_text: torch.FloatTensor | None = None    # (T, I)
    text_embeds: torch.FloatTensor | None = None        # (B, dim)
    image_embeds: torch.FloatTensor | None = None       # (B, dim)
    text_model_output: BaseModelOutputWithPooling = None
    vision_model_output: BaseModelOutputWithPooling = None
变量shape含义
loss(1,)Sigmoid 对比损失,仅 return_loss=True 时有值
logits_per_image(I, T)每张图像对每条文本的相似度得分。I = 图像 batch 大小,T = 文本 batch 大小
logits_per_text(T, I)logits_per_image 的转置。提供给文本侧的视角
text_embeds / image_embeds(B, dim)L2 归一化后的特征向量。可直接用于检索
text_model_output / vision_model_outputBaseModelOutputWithPooling完整子模型输出。含 last_hidden_state + pooler_output,用于调试或级联任务

to_tuple() 方法

1
2
def to_tuple(self) -> tuple[Any]:
    return tuple(v.to_tuple() if isinstance(v, ModelOutput) else v for v in self.values())

递归地将嵌套的 ModelOutput(如 text_model_output)也转换为 tuple,保证 torch.jit 兼容。


第二部分:嵌入层 (Embedding Layers)

4. Siglip2VisionEmbeddings — 图像嵌入层

类签名与角色

已经 patchify 的像素值(形状 (B, N, patch_dim))映射到 embedding 空间,加上可缩放的位置编码

Siglip2 的核心创新在这里——位置编码不是固定的 (1, N, dim) 矩阵,而是根据每张图像的实际尺寸 spatial_shapes 动态缩放

__init__ 参数全量表

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
def __init__(self, config: Siglip2VisionConfig):
    self.config = config
    self.embed_dim = config.hidden_size       # 768 (base)
    self.patch_size = config.patch_size       # 16

    self.patch_embedding = nn.Linear(
        in_features=config.num_channels * self.patch_size * self.patch_size,  # 3*16*16 = 768
        out_features=self.embed_dim,                                          # 768
    )
    self.num_patches = config.num_patches     # 256
    self.position_embedding_size = int(self.num_patches**0.5)  # sqrt(256) = 16
    self.position_embedding = nn.Embedding(self.num_patches, self.embed_dim)
变量默认值 (base)含义与设计原因
embed_dim768Transformer 的隐藏维度。Vision 和 Text 必须一致
patch_size16每个 patch 的边长(像素)。16x16 是 ViT 的经典选择:224/16=14 或 384/16=24
patch_embeddingLinear(768, 768)将每个 flatten patch 向量(16x16x3=768)线性投影到 embed_dim。输入=输出维度相同,这是不做降维的设计——信息保持,靠后续 Transformer 学习
num_patches256max_num_patches:每张图像最多处理的 patch 数。数值来自 naflex 论文(256 = 16x16 个 patch 可覆盖 256x256 图像)
position_embedding_size16 (= sqrt(256))位置编码表的二维网格边长。将 (256, 768) reshape 为 (16, 16, 768) 用于后续双线性插值缩放
position_embeddingEmbedding(256, 768)可学习的 2D 位置编码表。为什么用 Embedding 而非 Parameter? Embedding 的初始化逻辑与 Transformer token embedding 一致,方便复用 HuggingFace 的 default_flax_embed_init_

resize_positional_embeddings — 核心算法

1
2
3
4
5
6
@staticmethod
def resize_positional_embeddings(
    positional_embeddings,  # (grid_h, grid_w, dim) = (16, 16, 768)
    spatial_shapes,         # (B, 2) — 每张图的 (height, width)
    max_length,             # 256 — 最大 patch 数
) -> torch.Tensor:         # (B, max_length, dim)
局部变量shape含义
batch_sizeintspatial_shapes.shape[0]
embed_dimint位置编码的维度 = 768
source_dtypedtype保存原始精度。插值可能改变 dtype
resulted_positional_embeddings(B, max_length, dim)预分配输出张量。先分配再逐样本填充
positional_embeddings (permuted)(1, dim, grid_h, grid_w)(H, W, D) 转为 (1, D, H, W),是 F.interpolate 要求的 NCHW 格式

逐行关键逻辑

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
# 第 1 步:permute → (1, dim, H, W)
positional_embeddings = positional_embeddings.permute(2, 0, 1).unsqueeze(0)

# 第 2 步:如果设备是 CPU,上转型到 float32
# 原因:CPU 上的 bilinear + antialias 不支持 bfloat16/float16
if positional_embeddings.device.type == "cpu":
    positional_embeddings = positional_embeddings.to(torch.float32)

# 第 3 步:逐样本缩放位置编码
for i in range(batch_size):
    height, width = spatial_shapes[i].tolist()
    # 双线性插值:(16, 16) → (height, width)
    resized_embeddings = F.interpolate(
        positional_embeddings,
        size=(height, width),   # 目标网格尺寸,例如 (4, 6) 表示 24 个 patch
        mode="bilinear",
        align_corners=False,    # 与 TensorFlow 兼容的关键参数
        antialias=True,         # 抗锯齿。对高频位置信息尤其重要
    )
    # reshape: (1, dim, H, W) → (dim, H*W) → (H*W, dim)
    resized_embeddings = resized_embeddings.reshape(embed_dim, height * width).transpose(0, 1)

    # 第 4 步:赋值到预分配张量的对应位置
    resulted_positional_embeddings[i, :height*width] = resized_embeddings
    # 第 5 步:padding 位置用第一个有效位置的编码填充(而非零填充)
    # 为什么是 resized_embeddings[0] 而非 0?
    # 零填充会在后续 attention 计算中产生不连续的位置编码,第一个位置编码提供平滑过渡
    resulted_positional_embeddings[i, height*width:] = resized_embeddings[0]

为什么 align_corners=False

  • align_corners=True:角点像素的中心与角点对齐,像素之间是线性插值
  • align_corners=False:像素被视为均匀网格,角点之间的插值与 TensorFlow 行为一致
  • 选择 False 是为了与 Google 的原始 JAX/TF 实现保持一致

为什么 antialias=True

  • 当目标尺寸小于源尺寸时(如 16x16 → 4x6),双线性插值会产生混叠伪影
  • antialias=True 在缩小前先做低通滤波,避免位置编码中出现高频噪声

forward 方法

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
def forward(self, pixel_values, spatial_shapes):
    # 1. 线性投影:每个 patch 向量 → embed_dim
    target_dtype = self.patch_embedding.weight.dtype
    patch_embeds = self.patch_embedding(pixel_values.to(dtype=target_dtype))

    # 2. 取出位置编码表 + reshape 回 2D → 缩放到样本尺寸 → flatten
    positional_embeddings = self.position_embedding.weight.reshape(
        self.position_embedding_size, self.position_embedding_size, -1
    )
    resized_positional_embeddings = self.resize_positional_embeddings(
        positional_embeddings, spatial_shapes, max_length=pixel_values.shape[1]
    )

    # 3. 相加
    embeddings = patch_embeds + resized_positional_embeddings
    return embeddings

关键设计选择

  • 输入 pixel_values 已经是 patchify 的形式——image processor 在预处理阶段就完成了 patch 切割。这样做的好处是:image processor 和 model 解耦,processor 可以针对不同分辨率做不同的 patchify 策略
  • target_dtype 确保计算的 dtype 与权重一致。如果权重是 fp16 但输入是 fp32,自动转换

5. Siglip2TextEmbeddings — 文本嵌入层

类签名与角色

标准的 Transformer 文本嵌入:token_embedding + position_embedding

__init__ 参数全量表

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
def __init__(self, config: Siglip2TextConfig):
    self.token_embedding = nn.Embedding(config.vocab_size, config.hidden_size)
    self.position_embedding = nn.Embedding(config.max_position_embeddings, config.hidden_size)

    # 预计算的 position_ids buffer
    self.register_buffer(
        "position_ids",
        torch.arange(config.max_position_embeddings).expand((1, -1)),
        persistent=False,
    )
变量默认值 (base)含义与设计原因
token_embeddingEmbedding(32000, 768)将 token ID 映射到嵌入向量。标准 Transformer 设计
position_embeddingEmbedding(64, 768)绝对位置编码。64 是 Siglip2 文本塔的特殊值——比 CLIP 的 77 略少,说明 Google 发现文本塔不需要太长的位置编码
position_ids(1, 64)预计算的 [0, 1, 2, ..., 63]register_buffer 确保保存/加载时跟随模型,但 persistent=False 表示不需要梯度

为什么 persistent=False

  • position_ids 是固定的整数序列,不需要梯度,不需要被 optimizer 跟踪
  • persistent=Falsestate_dict() 中排除它,减小 checkpoint 体积

forward 方法

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
17
18
19
def forward(self, input_ids, position_ids, inputs_embeds):
    seq_length = input_ids.shape[-1]
    max_position_embedding = self.position_embedding.weight.shape[0]  # 64

    # 长度检查
    if seq_length > max_position_embedding:
        raise ValueError(...)

    # 如果没有显式传入 position_ids,使用预计算的 buf
    if position_ids is None:
        position_ids = self.position_ids[:, :seq_length]

    # 如果没有显式传入 inputs_embeds(如 prompt tuning),从 token id 嵌入
    if inputs_embeds is None:
        inputs_embeds = self.token_embedding(input_ids)

    position_embeddings = self.position_embedding(position_ids)
    embeddings = inputs_embeds + position_embeddings
    return embeddings

为什么 max_position_embeddings=64 而非 CLIP 的 77?

CLIP 论文中文本 token 最大长度是 77(包括 BOS/EOS)。Siglip2 缩减到 64,因为:

  1. 图像-文本对比学习中的文本通常是短句/标签,极少超过 64 tokens
  2. 更短的位置编码 = 更少的参数,更快的训练
  3. Google 发现 64 足以覆盖所有正样本语义

第三部分:Transformer 构件 (Building Blocks)

6. eager_attention_forward — 朴素注意力实现

1
def eager_attention_forward(module, query, key, value, attention_mask, scaling, dropout=0.0):

这是 回退实现——当无法使用 Flash Attention 或 SDPA 时(如设备不支持、dtype 不兼容),使用标准 PyTorch 操作手动计算注意力。

参数shape含义
modulenn.Module调用此函数的注意力层,用于获取 module.training 状态
query(B, heads, N, head_dim)查询张量
key(B, heads, N, head_dim)键张量
value(B, heads, N, head_dim)值张量
attention_mask(B, 1, N, N) or None加性 mask(0=允许, -inf=屏蔽)
scalingfloat1/sqrt(head_dim)。缩放因子,防止点积过大导致 softmax 饱和
dropoutfloat注意力 dropout 比率。训练时生效

逐行解析

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
# 第 1 步:QK^T × scale
attn_weights = torch.matmul(query, key.transpose(-1, -2)) * scaling

# 第 2 步:加性 mask。-inf 位置的 softmax 输出为 0
if attention_mask is not None:
    attn_weights = attn_weights + attention_mask

# 第 3 步:softmax(关键 —— dtype=torch.float32)
# 为什么 softmax 要上转 float32?
# fp16/bf16 的数值范围有限,softmax 的 exp 容易溢出。
# 上转 float32 计算 softmax,再转回原始 dtype,保证数值稳定
attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query.dtype)

# 第 4 步:dropout(仅在 training=True 时生效)
attn_weights = nn.functional.dropout(attn_weights, p=dropout, training=module.training)

# 第 5 步:× V
attn_output = torch.matmul(attn_weights, value)

# 第 6 步:transpose back → contiguous
# 输入: (B, heads, N, dim) → 输出需要: (B, N, heads, dim) → contiguous
# .contiguous() 确保后续 reshape 不会失败
attn_output = attn_output.transpose(1, 2).contiguous()

return attn_output, attn_weights

7. Siglip2Attention — 多头注意力层

类签名与角色

标准的 Multi-Head Attention(论文 “Attention Is All You Need”),但不固定使用 eager 实现——通过 ALL_ATTENTION_FUNCTIONS.get_interface() 动态选择 SDPA/Flash/eager。

1
2
3
4
5
6
class Siglip2Attention(nn.Module):
    def __init__(self, config):
        self.embed_dim = config.hidden_size          # 768
        self.num_heads = config.num_attention_heads   # 12
        self.head_dim = self.embed_dim // self.num_heads  # 64
        # 除法检查:768 / 12 = 64,必须整除
变量含义默认值来源
embed_dim输入/输出维度768config.hidden_size
num_heads注意力头数12config.num_attention_heads
head_dim每个头的维度64embed_dim / num_heads,必须整除
scale1/sqrt(head_dim)1/8标准 scaled dot-product attention 的缩放因子
dropout注意力 dropout0.0config.attention_dropout。Siglip2 默认 0
is_causal是否是因果注意力False文本塔也设为 False——这是 Siglip2 与 GPT/CLIP 的关键区别之一
q_proj/k_proj/v_projQ/K/V 投影矩阵Linear(768, 768)输入 768 维,输出 768 维
out_proj输出投影矩阵Linear(768, 768)多头拼接后的投影

为什么 is_causal = False 对文本塔也成立?

这是 Siglip2 文本塔与 GPT 的根本区别:

  • GPT 是因果(自回归)模型:第 i 个 token 只能 attend 到前 i 个 token
  • Siglip2 文本塔是双向编码器:文本内部完全双向 attention
  • 批注代码也说明了这一点:"Siglip2's text model does not use a causal mask, unlike the original CLIP model."

实际上早期的 CLIP 文本塔也用了因果 mask(因为是 GPT-2 backbone),Siglip2 改用纯双向,让文本表示更充分。

forward 方法逐行解析

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
def forward(self, hidden_states, attention_mask=None):
    # Step 1: 记住输入 shape
    input_shape = hidden_states.shape[:-1]  # (B, N)

    # Step 2: Q/K/V 投影 + reshape 为多头
    # hidden_states: (B, N, 768)
    # → 投影: (B, N, 768)
    # → view(*, -1, head_dim): (B, N, 12, 64)
    # → transpose(1, 2): (B, 12, N, 64)
    hidden_shape = (*input_shape, -1, self.head_dim)
    queries = self.q_proj(hidden_states).view(hidden_shape).transpose(1, 2)
    keys    = self.k_proj(hidden_states).view(hidden_shape).transpose(1, 2)
    values  = self.v_proj(hidden_states).view(hidden_shape).transpose(1, 2)

    # Step 3: 动态选择注意力实现
    # ALL_ATTENTION_FUNCTIONS 是 HuggingFace 的注意力后端注册表
    # 根据 config._attn_implementation 选择:
    #   "sdpa" → torch.nn.functional.scaled_dot_product_attention
    #   "flash_attention_2" → flash_attn_varlen_func
    #   "eager" → eager_attention_forward (上述手动实现)
    attention_interface = ALL_ATTENTION_FUNCTIONS.get_interface(
        self.config._attn_implementation, eager_attention_forward
    )

    # Step 4: 调用注意力
    attn_output, attn_weights = attention_interface(
        self, queries, keys, values, attention_mask,
        is_causal=self.is_causal,     # False → 双向注意力
        scaling=self.scale,           # 1/sqrt(64) = 0.125
        dropout=0.0 if not self.training else self.dropout,
    )

    # Step 5: reshape 回 → out_proj
    # (B, 12, N, 64) → (B, N, 768) → Linear → (B, N, 768)
    attn_output = attn_output.reshape(*input_shape, -1).contiguous()
    attn_output = self.out_proj(attn_output)

    return attn_output, attn_weights

ALL_ATTENTION_FUNCTIONS.get_interface 机制

这是 HuggingFace Transformers 4.x 的统一注意力调度器。不同后端提供相同的函数签名,使得模型代码无需写 if/else 分支。只需在 config 中设置 _attn_implementation,调度器自动分发。


8. Siglip2MLP — 前馈网络

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
class Siglip2MLP(nn.Module):
    def __init__(self, config):
        self.activation_fn = ACT2FN[config.hidden_act]     # "gelu_pytorch_tanh" → GELU(tanh approx)
        self.fc1 = nn.Linear(config.hidden_size, config.intermediate_size)  # 768 → 3072
        self.fc2 = nn.Linear(config.intermediate_size, config.hidden_size)  # 3072 → 768

    def forward(self, hidden_states):
        hidden_states = self.fc1(hidden_states)       # 扩展 4x
        hidden_states = self.activation_fn(hidden_states)  # GELU
        hidden_states = self.fc2(hidden_states)       # 压缩回原尺寸
        return hidden_states
变量含义默认值来源
fc1扩展层Linear(768, 3072)hidden_size → intermediate_size。扩展比 4x
fc2压缩层Linear(3072, 768)intermediate_size → hidden_size
activation_fn激活函数GELU(tanh approx)ACT2FN["gelu_pytorch_tanh"]。更快但精度略低的 GELU 变体

为什么 gelu_pytorch_tanh 而非标准 GELU?

  • 标准 GELU 需要 erf() 计算,较慢
  • tanh 近似版本速度快约 15%,精度损失可忽略
  • HuggingFace 几乎所有模型都默认用此实现

9. Siglip2EncoderLayer — 单层 Transformer

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
class Siglip2EncoderLayer(GradientCheckpointingLayer):
    def __init__(self, config):
        self.layer_norm1 = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)
        self.self_attn = Siglip2Attention(config)
        self.layer_norm2 = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)
        self.mlp = Siglip2MLP(config)

    def forward(self, hidden_states, attention_mask):
        # Pre-LN 架构:先 Norm,后 Attention/MLP
        # 子层 1: Self-Attention
        residual = hidden_states
        hidden_states = self.layer_norm1(hidden_states)
        hidden_states, _ = self.self_attn(hidden_states, attention_mask)
        hidden_states = residual + hidden_states

        # 子层 2: MLP
        residual = hidden_states
        hidden_states = self.layer_norm2(hidden_states)
        hidden_states = self.mlp(hidden_states)
        hidden_states = residual + hidden_states

        return hidden_states

Pre-LN vs Post-LN 架构

Post-LN (原始 Transformer):          Pre-LN (Siglip2):
  x → Attention → Norm → +           x → Norm → Attention → +
  x → MLP → Norm → +                 x → Norm → MLP → +

Pre-LN 的优势:
- 训练更稳定(梯度不经过 Norm 层传播)
- 不需要 warmup(学习率可以直接从较高值开始)
- 是 ViT 和现代 Transformer 的默认选择

继承 GradientCheckpointingLayer:允许在反向传播时释放中间激活值,用时间换空间。对 12 层 + batch=512 的大 batch 训练至关重要。


10. Siglip2Encoder — 多层 Transformer 堆叠

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
class Siglip2Encoder(nn.Module):
    def __init__(self, config):
        self.layers = nn.ModuleList([
            Siglip2EncoderLayer(config) for _ in range(config.num_hidden_layers)
        ])
        self.gradient_checkpointing = False

    def forward(self, inputs_embeds, attention_mask):
        hidden_states = inputs_embeds
        for encoder_layer in self.layers:
            hidden_states = encoder_layer(hidden_states, attention_mask)
        return BaseModelOutput(last_hidden_state=hidden_states)

为什么 Vision 和 Text 共享同一个 Siglip2Encoder

它们的 Encoder Layer 结构完全相同(Pre-LN + Self-Attn + MLP),区别只在:

  • Vision:输入是 patch embeddings,num_hidden_layers=12
  • Text:输入是 token embeddings,num_hidden_layers=12

共享 Encoder 类避免了代码重复。


第四部分:模型主体 (Model Bodies)

11. Siglip2PreTrainedModel — 基类

1
2
3
4
5
class Siglip2PreTrainedModel(PreTrainedModel):
    config: Siglip2Config
    base_model_prefix = "siglip2"
    input_modalities = ("image", "text")
    supports_gradient_checkpointing = True

类级别属性全量表

属性含义
base_model_prefix"siglip2"state_dict 中,所有参数名前缀。如 siglip2.vision_model.embeddings...
input_modalities("image", "text")声明模型是多模态的(图像 + 文本)。HuggingFace pipeline 据此选择合适的 processor
supports_gradient_checkpointingTrue允许在 Trainer 中通过 --gradient_checkpointing 开启
_no_split_modules["Siglip2TextEmbeddings", "Siglip2VisionEmbeddings", "Siglip2EncoderLayer", "Siglip2MultiheadAttentionPoolingHead"]设备分配粒度device_map="auto" 时这些模块不会被拆分到不同设备(保持完整在一个 GPU 上)
_supports_flash_attnFalse刻意设为 False——虽然模型支持 flash attention,但 Siglip2 的实现通过 ALL_ATTENTION_FUNCTIONS 动态调度,不依赖框架级的 _supports_flash_attn 标记
_supports_sdpaTrue声明支持 PyTorch 内置 SDPA
_supports_flex_attnFalse不支持灵活的注意力 mask(nn.MultiHeadAttention 的 mask 限制为非 4D)
_supports_attention_backendTrue声明支持 config._attn_implementation 切换
_can_record_outputs{"hidden_states": Siglip2EncoderLayer, "attentions": Siglip2Attention}capture_outputs 装饰器的路由表:哪些层可以输出 hidden_states 和 attentions

_init_weights — 权重初始化策略

这是最关键的手动初始化逻辑,按模块类型不同采用不同策略:

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
def _init_weights(self, module):
    super()._init_weights(module)  # 先调用基类的默认初始化

    if isinstance(module, Siglip2VisionEmbeddings):
        # 位置编码: N(0, 1/sqrt(width))
        init.normal_(module.position_embedding.weight, std=1 / np.sqrt(width))

    elif isinstance(module, nn.Embedding):
        # 文本 token/position embedding: Flax 默认截断正态初始化
        init.default_flax_embed_init_(module.weight)

    elif isinstance(module, Siglip2Attention):
        # Q/K/V/O 投影: Xavier Uniform
        init.xavier_uniform_(module.q_proj.weight)
        init.xavier_uniform_(module.k_proj.weight)
        init.xavier_uniform_(module.v_proj.weight)
        init.xavier_uniform_(module.out_proj.weight)
        # bias 初始化为 0
        init.zeros_(module.q_proj.bias)
        init.zeros_(module.k_proj.bias)
        init.zeros_(module.v_proj.bias)
        init.zeros_(module.out_proj.bias)

    elif isinstance(module, Siglip2MLP):
        # FC1/FC2: Xavier Uniform
        init.xavier_uniform_(module.fc1.weight)
        init.xavier_uniform_(module.fc2.weight)
        # bias: N(0, 1e-6) — 极小值,几乎为 0 但有微小扰动防止对称性
        init.normal_(module.fc1.bias, std=1e-6)
        init.normal_(module.fc2.bias, std=1e-6)

    elif isinstance(module, Siglip2MultiheadAttentionPoolingHead):
        # probe (可学习查询): Xavier Uniform
        init.xavier_uniform_(module.probe)
        # MultiheadAttention 内置权重: Xavier Uniform
        init.xavier_uniform_(module.attention.in_proj_weight)
        init.zeros_(module.attention.in_proj_bias)

    elif isinstance(module, Siglip2Model):
        # logit_scale 和 logit_bias 初始化为 0
        # 训练开始时 logits = 0 * exp(0) + 0 = 0 → sigmoid(0) = 0.5
        # 即初始时模型对任何 image-text pair 都输出 0.5 概率
        init.zeros_(module.logit_scale)
        init.zeros_(module.logit_bias)

    elif isinstance(module, Siglip2ForImageClassification):
        # 分类头: N(0, 1/sqrt(hidden_size) * initializer_factor)
        init.normal_(module.classifier.weight,
                     std=config.hidden_size**-0.5 * config.initializer_factor)

    elif isinstance(module, (nn.Linear, nn.Conv2d)):
        # 所有其他线性层/卷积: LeCun 正态初始化
        init.lecun_normal_(module.weight)
        if module.bias is not None:
            init.zeros_(module.bias)
初始化策略适用模块原因
normal_(std=1/sqrt(width))Vision 位置编码与 ViT 论文一致。确保位置编码的方差与 embedding 在同一量级
default_flax_embed_init_Text token/position embedding与 Google JAX 原始实现一致(Flax 的默认 embedding 初始化)
xavier_uniform_Attention Q/K/V/O, MLP FC1/FC2, Pooling HeadTransformer 的标准选择。保持前向/反向传播的方差不变
zeros_Attention bias标准做法。bias 不需要非零初始值
normal_(std=1e-6)MLP bias几乎是零,但微小的随机扰动打破对称性(帮助不同神经元学习不同特征)
zeros_logit_scale, logit_bias从 0 开始学习,让模型自行决定最优的温度和偏置
lecun_normal_其他 Linear/Conv2d保守的默认选择。特别适合 SELU/tanh 激活函数,但对 ReLU/GELU 也有效

关键设计:logit_scale 和 logit_bias 初始化为 0

logits = text_embeds @ image_embeds.T * exp(logit_scale) + logit_bias

初始状态: logit_scale=0, logit_bias=0
→ logits = cosine_sim * exp(0) + 0 = cosine_sim * 1
→ sigmoid(cosine_sim) 范围在 0.27~0.73 之间
→ 这是一个合理的"不确定"状态
→ 随着训练,logit_scale 会增大(提高置信度),logit_bias 会偏移

这是 Siglip2 sigmoid loss 的关键设计——可学习的温度和偏置替代 CLIP 的固定温度参数。


12. Siglip2VisionModel — 视觉编码器

__init__

1
2
3
4
5
6
7
8
def __init__(self, config: Siglip2VisionConfig):
    self.embeddings = Siglip2VisionEmbeddings(config)
    self.encoder = Siglip2Encoder(config)
    self.post_layernorm = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)
    self.use_head = True if not hasattr(config, "vision_use_head") else config.vision_use_head
    if self.use_head:
        self.head = Siglip2MultiheadAttentionPoolingHead(config)
    self.post_init()
变量含义
embeddings图像 patch → embedding + 可缩放位置编码
encoder12 层 Pre-LN Transformer
post_layernorm编码器输出的后 LayerNorm
use_head是否附加池化头。默认 True
head多头注意力池化头(Siglip2MultiheadAttentionPoolingHead
_input_embed_layer"patch_embedding" — 告诉基类哪个属性是输入嵌入层,用于 get/set

forward 逐行解析

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
def forward(self, pixel_values, pixel_attention_mask, spatial_shapes):
    # 第 1 步:Embedding(线性投影 + 位置编码)
    hidden_states = self.embeddings(pixel_values, spatial_shapes)

    # 第 2 步:创建双向 attention mask
    # create_bidirectional_mask 将 1D mask (B, N) 展开为 (B, 1, N, N) 的加性 mask
    # 有效位置 = 0, padding 位置 = -inf
    encoder_attention_mask = create_bidirectional_mask(
        config=self.config,
        inputs_embeds=hidden_states,
        attention_mask=pixel_attention_mask,
    )

    # 第 3 步:Transformer 编码
    encoder_outputs = self.encoder(
        inputs_embeds=hidden_states,
        attention_mask=encoder_attention_mask,
    )

    # 第 4 步:后 LayerNorm
    last_hidden_state = self.post_layernorm(encoder_outputs.last_hidden_state)

    # 第 5 步:多头注意力池化
    pooler_output = self.head(last_hidden_state, pixel_attention_mask) if self.use_head else None

    return BaseModelOutputWithPooling(
        last_hidden_state=last_hidden_state,
        pooler_output=pooler_output,
    )

create_bidirectional_mask 做了什么?

1
2
3
4
5
6
7
8
9
# 输入: pixel_attention_mask = (B, N),值为 0 或 1
# 输出: (B, 1, N, N),格式:
#   mask[i, j] = 0     if pixel_attention_mask[i] == 1 and pixel_attention_mask[j] == 1
#   mask[i, j] = -inf  otherwise
#
# 加性 mask 在 softmax 之前与 QK^T 相加:
#   attn = softmax(QK^T/√d + mask)
# 有效位置: 加上 0,不影响
# padding 位置: 加上 -inf,softmax 后 → 0

13. Siglip2TextModel — 文本编码器

__init__

1
2
3
4
5
6
def __init__(self, config: Siglip2TextConfig):
    self.embeddings = Siglip2TextEmbeddings(config)
    self.encoder = Siglip2Encoder(config)
    self.final_layer_norm = nn.LayerNorm(embed_dim, eps=config.layer_norm_eps)
    self.head = nn.Linear(embed_dim, config.projection_size)
    self.post_init()
变量含义
embeddingstoken embedding + position embedding
encoder12 层 Pre-LN Transformer(与 Vision 共享类)
final_layer_norm编码器输出的后 LayerNorm
headLinear(768, 768) — 投影头,将 pooler 输出映射到联合嵌入空间

projection_size vs hidden_size

  • projection_size联合嵌入空间的维度,默认等于 hidden_size=768
  • 之所以分成两个参数,是为未来可能的多模态对齐(不同模态用不同 hidden_size,但投影到相同 projection_size)预留接口

forward 逐行解析

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
def forward(self, input_ids, attention_mask, position_ids):
    # 第 1 步:文本 Embedding
    hidden_states = self.embeddings(input_ids, position_ids)

    # 第 2 步:创建双向 attention mask(非因果!)
    attention_mask = create_bidirectional_mask(
        config=self.config,
        inputs_embeds=hidden_states,
        attention_mask=attention_mask,
    )

    # 第 3 步:Transformer 编码
    encoder_outputs = self.encoder(hidden_states, attention_mask)

    # 第 4 步:后 LayerNorm
    last_hidden_state = self.final_layer_norm(encoder_outputs.last_hidden_state)

    # 第 5 步:EOS token 池化 + 投影
    pooled_output = last_hidden_state[:, -1, :]      # 取最后一个 token
    pooled_output = self.head(pooled_output)          # 线性投影

    return BaseModelOutputWithPooling(
        last_hidden_state=last_hidden_state,
        pooler_output=pooled_output,
    )

为什么用 last_hidden_state[:, -1, :](最后一个 token)做池化?

这是 Siglip2 文本池化的关键设计:

  • CLIP 的做法:取 EOS token 的 hidden state,因为 CLIP 文本塔是因果的(GPT-2),最后一个 token 已经 attend 了所有前面的 token
  • Siglip2 的做法:虽然文本塔是双向的(不是因果),但仍然取最后一个 token。因为:
    1. Tokenizer 固定使用 padding="max_length",最后一个非 padding 位置恰好是 EOS 或句末 token
    2. 最后一个 token 在双向注意力中看到了全部上下文
    3. 与 CLIP 保持 API 兼容(下游用户无感知切换)

注意代码中的注释"The model uses the last token's hidden state, which may be padding."——这其实是已知问题。当文本较短时,最后一个位置可能是 padding token。但 Google 训练时使用 padding="max_length",确保 EOS 总是在序列末尾,padding 在 EOS 之前(左侧 padding)。


14. Siglip2MultiheadAttentionPoolingHead — 多头注意力池化

类签名与角色

这是 Siglip2 Vision 塔的核心池化机制——替代 CLIP 的简单 CLS token 或 mean pooling。

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
class Siglip2MultiheadAttentionPoolingHead(nn.Module):
    def __init__(self, config):
        # 可学习的 probe 向量(类比 CLS token)
        self.probe = nn.Parameter(torch.randn(1, 1, config.hidden_size))

        # 多头注意力:probe 作为 query,patch tokens 作为 key/value
        self.attention = nn.MultiheadAttention(
            config.hidden_size, config.num_attention_heads, batch_first=True
        )

        # 后处理:LayerNorm + MLP (残差连接)
        self.layernorm = nn.LayerNorm(config.hidden_size, eps=config.layer_norm_eps)
        self.mlp = Siglip2MLP(config)

工作原理(逐步)

输入:
  hidden_state: (B, N, 768)  ← 所有 patch tokens
  attention_mask: (B, N)      ← 哪些是有效 patch

步骤 1: 复制 probe
  probe = self.probe.repeat(B, 1, 1)  → (B, 1, 768)

步骤 2: 创建 cross-attention mask
  probe(1) → key(N): 只允许对有效 patch 的 attention
  target_len=1, source_len=N

步骤 3: 如果有 attention_mask
  - 构造 (B*heads, 1, N) 的 mask
  - 布尔 mask 转为加性 mask (0 / -inf)
  原因: nn.MultiheadAttention 不能直接处理布尔 mask

步骤 4: MultiheadAttention
  query  = probe       (B, 1, 768)
  key    = hidden_state (B, N, 768)
  value  = hidden_state (B, N, 768)
  output = attention(probe, hidden_state, hidden_state, attn_mask)
         → (B, 1, 768)  ← 这是池化后的图像表示

步骤 5: 后处理 (残差 + LayerNorm + MLP)
  residual = hidden_state
  hidden_state = self.layernorm(hidden_state)
  hidden_state = residual + self.mlp(hidden_state)
  → (B, 1, 768) → squeeze → (B, 768)

变量全量表

变量shape含义
self.probe(1, 1, 768)可学习的池化查询向量。类比 BERT 的 [CLS] token,但不是 token——是 learnable query
self.attentionMultiheadAttention(768, 12, batch_first=True)标准 PyTorch 多头注意力。batch_first=True 使输入=输出格式 (B, N, D) 而非 (N, B, D)
self.layernormLayerNorm(768)attention 后的 LayerNorm
self.mlpSiglip2MLP(768, 3072)attention 后的 MLP
self.num_heads12注意力头数
batch_size (local)inthidden_state.shape[0]
probe (local)(B, 1, 768)扩展后的 probe
target_len / source_len1 / Nattention mask 的维度:query 只有 1 个 token

Mask 转换的关键代码

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
if attention_mask is not None:
    # 步骤 3a: 创建 (B*heads, target_len, source_len) 的 mask
    attention_mask = attention_mask.repeat(1, self.num_heads, target_len, 1)
    attention_mask = attention_mask.reshape(-1, target_len, source_len)

    # 步骤 3b: 布尔 mask → 加性 mask
    if attention_mask.dtype == torch.bool:
        attention_mask = torch.where(
            attention_mask,
            torch.tensor(0.0, device=attention_mask.device, dtype=probe.dtype),
            torch.finfo(probe.dtype).min,  # ≈ -65504 for fp16
        )

为什么需要 reshape(-1, target_len, source_len)

nn.MultiheadAttention 期望的 attn_mask shape 是 (batch*heads, target_len, source_len)(batch*heads, target_len, source_len)。由于 batch_first=True,先 repeat 扩展 head 维度,再 reshape 合并 batch 和 head 维度。

为什么 nn.MultiheadAttention 不能处理布尔 mask?

这是 PyTorch 的限制。PyTorch 的 MultiheadAttention 内部调用 F.scaled_dot_product_attention,而 SDPA 在某些后端(Flash Attention)不支持布尔 mask,必须显式转为 0 / -inf 的浮点 mask。HuggingFace 通过这层转换兼容了所有后端。

与 CLS Token 的对比

维度CLS Token (CLIP/BERT)Multihead Attention Pooling (Siglip2)
池化方式CLS token 参与 Transformer 编码,取最后一层对应位置编码器输出后,用额外注意力层跨 patch 查询
可学习参数CLS embedding 在 input 层probe 在 pooling 层
计算开销每层都参与 attention只在最后做一次 cross-attention
灵活性CLS 看到所有层的信息probe 只看最后一层的输出
表达能力中等更强——probe 可以学会关注特定 patch
为什么 Siglip2 选后者视觉池化需要更强的聚合能力。probe 学会了关注"信息量大的 patch"(如物体中心)而非均匀聚合

15. Siglip2Model — 双塔联合模型

__init__

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
def __init__(self, config: Siglip2Config):
    text_config = config.text_config
    vision_config = config.vision_config

    # 通过 _from_config 创建子模型(确保注意力配置传递)
    self.text_model = Siglip2TextModel._from_config(text_config)
    self.vision_model = Siglip2VisionModel._from_config(vision_config)

    # 可学习的 logit_scale 和 logit_bias
    self.logit_scale = nn.Parameter(torch.randn(1))
    self.logit_bias = nn.Parameter(torch.randn(1))
变量shape含义
text_modelSiglip2TextModel文本编码器(12 层,768 维,64 个位置编码)
vision_modelSiglip2VisionModel图像编码器(12 层,768 维,256 个 patch)
logit_scale(1,)可学习的温度倒数。exp(logit_scale) 起放大/缩小 cosine similarity 的作用
logit_bias(1,)可学习的偏置。使模型可以学到某些 image-text pair 天然比另一些更相似

为什么用 _from_config 而非直接构造?

_from_config 是 HuggingFace 的内部工厂方法,确保子模型的注意力配置(_attn_implementation)从父配置正确继承。如果直接用 Siglip2VisionModel(config),可能不会正确继承 attn_implementation="flash_attention_2"

forward 完整解析(核心算法实现)

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
17
18
19
20
def forward(self, input_ids, pixel_values, pixel_attention_mask, spatial_shapes,
            attention_mask, position_ids, return_loss):
    # ===== 第 1 步:编码图像和文本 =====
    vision_outputs = self.vision_model(
        pixel_values=pixel_values,
        pixel_attention_mask=pixel_attention_mask,
        spatial_shapes=spatial_shapes,
    )
    text_outputs = self.text_model(
        input_ids=input_ids,
        attention_mask=attention_mask,
        position_ids=position_ids,
    )

    image_embeds = vision_outputs.pooler_output   # (I, dim)
    text_embeds = text_outputs.pooler_output      # (T, dim)

    # ===== 第 2 步:L2 归一化 =====
    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)
局部变量shape含义
vision_outputs.pooler_output(I, 768)多头注意力池化后的图像特征
text_outputs.pooler_output(T, 768)EOS token 投影后的文本特征
image_embeds (归一化后)(I, 768)L2 范数为 1 的图像向量
text_embeds (归一化后)(T, 768)L2 范数为 1 的文本向量

为什么 L2 归一化?

  • 余弦相似度 = image_embeds @ text_embeds.T(因为两个向量已经 L2=1)
  • 归一化确保所有嵌入在同一球面上,比较时不会被向量长度干扰
  • 这是所有 contrastive learning(CLIP、SimCLR、Siglip)的标准做法
1
2
3
    # ===== 第 3 步:计算余弦相似度矩阵 =====
    logits_per_text = torch.matmul(text_embeds, image_embeds.t().to(text_embeds.device))
    # shape: (T, I) = (文本数, 图像数)
变量shape含义
logits_per_text(T, I)文本-图像余弦相似度矩阵。第 i 行第 j 列 = text_i 与 image_j 的相似度
1
2
3
4
    # ===== 第 4 步:应用可学习温度和偏置 =====
    logit_scale = self.logit_scale.to(text_embeds.device)
    logit_bias = self.logit_bias.to(text_embeds.device)
    logits_per_text = logits_per_text * logit_scale.exp() + logit_bias

logit_scalelogit_bias 的作用:

原始余弦相似度:    [-1, 1]
乘以 exp(logit_scale):  [-exp(s), exp(s)]
加上 logit_bias:       调整决策边界

训练开始时:
  logit_scale = 0 → exp(0) = 1
  logit_bias = 0
  → logits 不动(保持原来的余弦值范围)

训练进行中:
  logit_scale 增大 → logits 被放大 → sigmoid 更陡 → 模型更确信
  logit_bias 偏移 → 适应正负样本的不平衡

这是Siglip2 相对于 CLIP 的核心创新:CLIP 使用固定的 exp(t) 温度参数(t 是常量),Siglip2 让模型自己学习温度和偏置。

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
    # ===== 第 5 步:计算 Sigmoid Loss =====
    loss = None
    if return_loss:
        # 构造标签矩阵: 对角线=1 (正样本), 非对角线=-1 (负样本)
        eye = torch.eye(logits_per_text.size(0), device=logits_per_text.device)  # (T, I) 对角线
        m1_diag1 = -torch.ones_like(logits_per_text) + 2 * eye
        # 对角线: -1 + 2*1 = +1
        # 非对角线: -1 + 2*0 = -1

        # log-sigmoid: log(1/(1+exp(-x))) = -softplus(-x)
        loglik = torch.nn.functional.logsigmoid(m1_diag1 * logits_per_text)
        # 正样本: logsigmoid(+1 * logit) → loglik 高 = loss 低
        # 负样本: logsigmoid(-1 * logit) → loglik 低 = loss 高

        nll = -torch.sum(loglik, dim=-1)  # 对图像维度求和
        loss = nll.mean()                  # 对 batch 求平均

Sigmoid Loss 的数学推导

CLIP (Softmax CE):
  对每个图像 i:
    loss_img = -log( exp(logit[i,i]) / sum_j(exp(logit[i,j])) )
  问题: 依赖 batch 中的所有样本作为负样本。batch_size 小 → 负样本少 → 效果差

Siglip2 (Sigmoid Loss):
  对每对 (i,j):
    target = +1 if i == j else -1
    loss_ij = -log(sigmoid(target * logit[i,j]))
  优点:
    1. 每个 (i,j) 对独立计算,不依赖 batch 其他样本
    2. 可以利用大量额外负样本(pre-computed cache)
    3. batch_size 可以很小而不影响效果
    4. 训练更稳定(没有 softmax 的分母竞争)

数学:
  logsigmoid(x) = log(1/(1+exp(-x))) = -softplus(-x)
  正样本 (target=+1): logsigmoid(logit) → logit越高, loss越低
  负样本 (target=-1): logsigmoid(-logit) → logit越低, loss越低

m1_diag1 矩阵的构造

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
eye = torch.eye(T, device=...)       # [[1,0,0], [0,1,0], [0,0,1]]
m1_diag1 = -torch.ones_like(T) + 2*eye

# 结果:
# [[+1, -1, -1],
#  [-1, +1, -1],
#  [-1, -1, +1]]
#
# 第 (i,j) 个元素:
#   i==j → +1 (正样本 target)
#   i!=j → -1 (负样本 target)

get_text_featuresget_image_features

1
2
3
4
5
def get_text_features(self, input_ids, attention_mask, position_ids):
    return self.text_model(input_ids, attention_mask, position_ids)

def get_image_features(self, pixel_values, pixel_attention_mask, spatial_shapes):
    return self.vision_model(pixel_values, pixel_attention_mask, spatial_shapes)

这两个方法是便捷 API,直接返回子模型的 BaseModelOutputWithPooling。用户可以用它们:

  • 离线预计算图像/文本嵌入(用于大规模检索)
  • 构建自定义的多模态 pipeline
  • 避免每次都要传 image+text 双输入

16. Siglip2ForImageClassification — 图像分类头

 1
 2
 3
 4
 5
 6
 7
 8
 9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
class Siglip2ForImageClassification(Siglip2PreTrainedModel):
    def __init__(self, config):
        self.vision_model = Siglip2VisionModel._from_config(config.vision_config)
        self.classifier = (
            nn.Linear(config.vision_config.hidden_size, config.num_labels)
            if config.num_labels > 0 else nn.Identity()
        )

    def forward(self, pixel_values, pixel_attention_mask, spatial_shapes, labels):
        outputs = self.vision_model(pixel_values, pixel_attention_mask, spatial_shapes)
        sequence_output = outputs.last_hidden_state  # (B, N, 768)

        # 平均池化(而非 MultiheadAttentionPoolingHead)
        if pixel_attention_mask is not None:
            pool_mask = pixel_attention_mask[..., None].to(sequence_output.device)  # (B, N, 1)
            sequence_output = torch.sum(sequence_output * pool_mask, dim=1) / torch.sum(pool_mask, dim=1)
        else:
            sequence_output = torch.mean(sequence_output, dim=1)

        logits = self.classifier(sequence_output)
        loss = self.loss_function(labels, logits, self.config) if labels is not None else None

        return ImageClassifierOutput(loss=loss, logits=logits)
变量shape含义
self.classifierLinear(768, num_labels)分类头。当 num_labels=0 时为 Identity(无标签模式)
sequence_output(B, N, 768)Vision encoder 的最后一层输出(未池化)
pool_mask(B, N, 1)平均池化的 mask:有效 patch=1,padding=0
sequence_output (池化后)(B, 768)加权平均池化后的图像特征
logits(B, num_labels)分类 logits

为什么分类头用平均池化而非 MultiheadAttentionPoolingHead?

MultiheadAttentionPoolingHead 是为对比学习设计的——它学会了突出与文本描述最相关的 patch。但分类任务需要均匀关注所有 patch(物体可能出现在任何位置),平均池化更鲁棒。


与测试文件的对应关系

实现类对应的测试类(在 test_modeling_siglip2.py验证什么
Siglip2VisionModelSiglip2VisionModelTestforward shape:last_hidden_state (B,N,768) + pooler_output (B,768);SDPA dispatch;Flash Attn inference
Siglip2TextModelSiglip2TextModelTestforward shape;SDPA dispatch
Siglip2ModelSiglip2ModelTestlogits shape (I,T)(T,I);config 分解为 vision/text;预训练模型加载
Siglip2ModelSiglip2ModelIntegrationTest真实图片+真实权重推理,logits 值与预期精确匹配
Siglip2ForImageClassificationSiglip2ForImageClassificationModelTest分类 forward shape;gradient checkpointing 兼容性(xfail)
Siglip2AttentionSiglip2ModelTesterMixineager vs SDPA vs Flash Attn 的数值等价性
Siglip2ImageProcessortest_image_processing_siglip2.pypatch_size=16, max_patches=256 等配置参数

关键设计决策汇总

决策位置理由
位置编码可缩放Siglip2VisionEmbeddings.resize_positional_embeddings支持灵活分辨率(naflex)。同一张位置编码表通过双线性插值适配任何 patch 网格
文本塔非因果create_bidirectional_mask in Siglip2TextModel.forward双向注意力比因果注意力提供更充分的文本表示
池化:vision 用 attention,text 用 EOSMultiheadAttentionPoolingHead vs last_hidden_state[:, -1, :]Vision 需要跨 patch 学习性聚合;Text 的 EOS token 在双向注意力中已汇总全句信息
Sigmoid Loss 替代 Softmax CESiglip2Model.forward loss 计算解耦 batch size 与负样本数量。每对 (i,j) 独立计算,允许使用预计算的负样本缓存
可学习温度和偏置logit_scale + logit_bias模型自己决定最优的温度缩放,而非使用 CLIP 的固定温度
初始化 logit_scale=0, logit_bias=0_init_weights训练从"不确定"状态开始(所有 pair 概率约 0.5),逐步学习区分正负
共享 Encoder 类Vision 和 Text 都用 Siglip2EncoderPre-LN Transformer 结构完全相同,避免重复代码
ALL_ATTENTION_FUNCTIONS 调度Siglip2Attention.forward一套代码支持 eager/SDPA/Flash Attn 三种后端,用户通过 config 切换
_from_config 构造子模型Siglip2Model.__init__确保子模型的注意力配置从父配置正确继承
CPU 上转 float32resize_positional_embeddingsCPU 上 bilinear + antialias 不支持 bf16/fp16
padding 位置用第一个有效编码填充resulted_positional_embeddings[i, height*width:] = resized_embeddings[0]避免零填充在 attention 中产生不连续的位置编码
分类头用平均池化Siglip2ForImageClassification.forward分类需要均匀关注所有 patch,Attention Pooling 过度关注特定区域