首页
/ Ultralytics SAM3 几何编码器源码剖析:从几何 Prompt 到 Transformer 可读序列的完整实现解读

Ultralytics SAM3 几何编码器源码剖析:从几何 Prompt 到 Transformer 可读序列的完整实现解读

2026-09-07 12:02:01作者:蔡丛锟

文章导读

本文基于 Ultralytics 仓库 geometry_encoders.py 的公开接口与真实实现,深入剖析 SAM3(Segment Anything Model 3)中"几何提示(geometric prompt)"的表示与编码链路:Prompt 数据结构如何承载 box/point 几何提示,is_right_paddedconcat_padded_sequences 两个工具函数如何维护右填充(right-padded)的不定长批序列,以及 SequenceGeometryEncoder 如何将"归一化的 CxCyWH 框"投影、RoI 池化并融合位置编码后送入 Transformer。读完本文,你将能够准确理解 geometry_encoders.md 所列四个公开符号的输入输出约定、维度规则与参数语义,并能在二次开发(接入自定义几何提示、调整编码配置)时直接定位到相关源码与调用点。

说明:docs/en/reference/... 下以 ::: 模块路径.符号 形式呈现的是由 mkdocstrings 从源码 docstring 自动生成的 API 参考页,其技术细节的权威来源正是本文所引用的源码模块与上游调用点。


一、模块定位:SAM3 中"框提示"通往 Transformer 的必经之路

ultralytics/models/sam/sam3/geometry_encoders.py 位于 SAM3 子包内,是标准 SAM3 语义模型(SAM3SemanticModel)中"文本提示 + 几何提示"双通道提示编码的关键一环。从 sam3_image.py 可以看到,模型在编码阶段先将几何提示交给 self.geometry_encoder(即 SequenceGeometryEncoder 实例),得到 geo_feats, geo_masks,再与视觉提示 embedding 拼接成完整的 prompt 序列,随后被 Transformer encoder 消费:

# ultralytics/models/sam/sam3/sam3_image.py
geo_feats, geo_masks = self.geometry_encoder(
    geo_prompt=geometric_prompt,
    img_feats=img_feats,
    img_sizes=vis_feat_sizes,
    img_pos_embeds=img_pos_embeds,
)
prompt = torch.cat([geo_feats, visual_prompt_embed], dim=0)
prompt_mask = torch.cat([geo_masks, visual_prompt_mask], dim=1)

由此可见该模块在整个 SAM3 前向中的位置:图像特征由 backbone 产出,几何提示(框/角点)由该模块编码为与图像特征同维的 token 序列,供 Transformer 在 cross-attention 中作为 query/condition 使用。

在调用链另一侧,build_sam3.py 以"模块配置"的方式给出了标准 SAM3 语义模型的几何编码器组装实例(详见第四节);predict.py 在推理器(SAM3Predictor)中通过 Prompt(box_embeddings=torch.zeros(0, num_prompts, 4), ...) 构造"零框 dummy 几何提示",用于无框提示时的推理路径。这两处共同构成该模块在"模型构建"与"推理"两条链路上的真实用法证据。


二、Prompt:统一承载几何提示的数据容器

Promptgeometry_encoders.py 的类 docstring 中被定位为"操作几何提示的工具类"。它是介于上游"归一化坐标输入"与下游编码器之间的统一数据容器,其维度约定如下(源码注释原文要点):

张量 形状 含义 允许取值
box_embeddings N_boxes × B × C_box 每个框的几何特征 归一化框坐标或预计算 embedding
box_mask B × N_boxes attention mask(PyTorch 约定,1 表示被 mask/pad) None 表示无 mask 项
box_labels N_boxes × B 正/负样本标签(long 型) None 表示全部视为正样本

三个关键约定在类 docstring 中被强调:

  1. 序列维度在前:所有序列张量按 PyTorch 约定组织——序列长度(N)在前、batch 维度(B)在后;而 mask 张量则 batch-first。
  2. 盒坐标采用归一化 CxCyWH 格式:即 (center_x, center_y, width, height),坐标值归一化到 [0, 1],这一点由 SequenceGeometryEncoder 的类 docstring 声明并由构造函数断言 box_embeddings.shape[-1] == 4 来保证(见 geometry_encoders.py)。
  3. 标签默认全正、mask 默认全非 padbox_labels 缺省时填充 torch.ones(...)box_mask 缺省时填充 torch.zeros(...)geometry_encoders.py)。

2.1 构造与形状校验

Prompt.__init__ 接受 box_embeddings / box_mask / box_labels 三个可选参数。当 box_embeddings is None 时构造出一个"空 prompt"(三字段全为 None),可用于零框推理;否则依次补齐缺失的 labelsmask,并对以下条件做断言校验:

  • box_embeddings 前两维为 [N_boxes, B],末维必须为 4(四个几何量);
  • box_mask 形状恰为 [B, N_boxes]
  • box_labels 形状恰为 [N_boxes, B]
  • 三者 device 一致。

源码中这些校验一方面保证数据自洽,另一方面与 concat_padded_sequences 的断言共同构成编码前最后一道"维度防火墙"。

2.2 append_boxes:动态追加框提示

Prompt.append_boxes(boxes, labels=None, mask=None)geometry_encoders.py)支持两种场景:

  • 首框初始化:当 self.box_embeddings is None 时,直接以本次 boxes 初始化并补齐 labels/mask;
  • 追加既有:校验 batch 大小与 shapes 一致后,借助 concat_padded_sequences 分别对 box_labels(先 unsqueeze(-1) 再拼接、再 squeeze(-1))与 box_embeddings 完成右填充序列拼接。

该方法与第一节所述 _get_dummy_prompt(零框)形成互补:推理中通过逐帧"追加框"即可在保持右填充不变式的前提下累积提示。从源码结构看,这类逐次追加的能力为视频/交互式场景下提示的增量维护提供了原语。


三、两个基础工具函数:右填充判定与不等长序列拼接

3.1 is_right_padded:判定 padding 是否在右侧

def is_right_padded(mask: torch.Tensor):
    return (mask.long() == torch.sort(mask.long(), dim=-1)[0]).all()

按 PyTorch 约定,padding mask 中 1 表示被 pad 的占位。若 padding 位于序列右侧,则 mask 中先出现一段 0、后出现一段 1,恰好是**非递减(升序)**序列。该函数通过比较 mask 与其升序排序结果是否逐元素相等,来判断"整批序列是否都是右填充"(geometry_encoders.py)。它被 concat_padded_sequences 内部以 torch._assert 调用,作为拼接操作的前置不变量检查。

3.2 concat_padded_sequences:两条右填充序列的无缝拼接

concat_padded_sequences(seq1, mask1, seq2, mask2, return_index=False)geometry_encoders.py)是 Prompt.append_boxesSequenceGeometryEncoder.forward(追加 CLS token 时)共同依赖的核心拼接原语。其输入输出约定如下:

参数 形状 说明
seq1 (L1, B, H) 序列优先、特征在末维
mask1 (B, L1) 1 表示 pad
seq2 (L2, B, H) 同上
mask2 (B, L2) 同上
return_index bool 是否额外返回 seq2 在拼接序列中的索引

实现要点(算法层面):

  1. 前置断言:核对 batch、hidden、序列长度两两匹配,并断言 mask1/mask2 均为右填充。
  2. 计算真实长度actual_seqN_lengths = (~maskN).sum(dim=-1) 统计每样本非 pad 的真实 token 数;拼接后每样本真实长度相加为 final_lengths,最大可能长度为 max_length = L1 + L2
  3. 构造拼接 mask:利用广播比较 torch.arange(max_length) >= final_lengths 生成新的右填充 mask——凡超过该样本真实总长度的位置都置 1。
  4. 移位放置 seq2:先新建 (max_length, B, H) 的全零张量,把 seq1 直接放进前 L1 行;随后计算 seq2 各行应落入的目标行号 index = arange(L2)[:,None] + actual_seq1_lengths[None](即"在 seq1 实际长度基础上偏移"),用 scatter 将 seq2 写入对应位置。
  5. 可选返回 indexreturn_index=True 时额外返回 index(形状 (L2, B)),用于从拼接序列中精确取回 seq2 的元素。

正是"mask 右填充 + 每样本真实长度可推导"这一不变式,使得该函数无需逐样本循环即可高效完成变长序列的批式拼接,拼接结果天然仍是右填充序列,可直接馈给下游 Transformer 的 key_padding_mask


四、SequenceGeometryEncoder:构造参数与三种框编码路径

SequenceGeometryEncoder 的完整 docstring 与构造函数位于 geometry_encoders.py。它声明接受"归一化 CxCyWH"格式的框,框可被三种方式编码,三者互不排斥、可叠加求和

  • direct projection(线性投影):对 4 维坐标做线性投影到 d_model
  • pooling(RoI align):从 backbone 特征图做 RoI align,汇聚框内区域特征;
  • pos encoder(位置编码):对框中心做正余弦位置编码(复用 PositionEmbeddingSine)。

作为替代方案,框还可以被拆解为左上/右下两个角点来编码(encode_boxes_as_points=True)。

4.1 构造参数语义

参数 类型 语义与影响
encode_boxes_as_points bool 是否把框拆为两个角点编码。True 时使用 (左上, 右下) 两组点
boxes_direct_project bool 线性投影路径,对应 nn.Linear(4, d_model)
boxes_pool bool RoI 路径,对应 nn.Conv2d(d_model, d_model, roi_size)
boxes_pos_enc bool 位置编码路径,对应 nn.Linear(d_model + 2, d_model)
d_model int 模型宽度,所有编码输出的公共通道维度
pos_enc nn.Module 位置编码器(如 PositionEmbeddingSine),用于框中心编码
num_layers int 后续 Transformer 编码层数量;> 0 时强制建议开启 CLS
layer nn.Module 单个 Transformer 编码层(由 _get_clones 深拷贝复制)
roi_size int=7 RoI align 输出尺寸(高/宽)
add_cls bool=True 是否在序列头部加入可学习的 CLS token
add_post_encode_proj bool=True 是否追加 Linear + LayerNorm 作为编码后精化
use_act_ckpt bool=False 是否在多层编码器上启用激活检查点(省显存)

构造函数中还蕴含两个与配置一致性相关的细节:

  • 标签 embedding 数量动态化:编码为框时每 token 只有正/负 2 类标签;编码为角点时每点可能出现"普通正负、左上正负、右下正负"共 6 类,故 label_embed = nn.Embedding(num_labels, d_model)num_labels = 6 if encode_boxes_as_points else 2geometry_encoders.py)。
  • 非角点模式至少需要一种框编码方式:若 encode_boxes_as_points=False 且三种框编码开关全为 False,则直接断言报错 "Error: need at least one way to encode boxes"
  • RoI 相关模块附带输入归一化:当任一 pooling 路径启用时,img_pre_normnn.Identity() 切换为 nn.LayerNorm(d_model),在池化前对特征做逐层归一。

4.2 仓库中的真实组装示例

build_sam3.pybuild_semantic_sam3 给出的标准配置可作为理解各参数的权威样例:

input_geometry_encoder = SequenceGeometryEncoder(
    pos_enc=PositionEmbeddingSine(
        num_pos_feats=256, normalize=True, scale=None, temperature=10000,
    ),
    encode_boxes_as_points=False,
    boxes_direct_project=True,
    boxes_pool=True,
    boxes_pos_enc=True,
    d_model=256,
    num_layers=3,
    layer=TransformerEncoderLayer(
        d_model=256,
        dim_feedforward=2048,
        dropout=0.1,
        pos_enc_at_attn=False,
        pre_norm=True,
        pos_enc_at_cross_attn_queries=False,
        pos_enc_at_cross_attn_keys=True,
    ),
    use_act_ckpt=True,
    add_cls=True,
    add_post_encode_proj=True,
)

可以看到生产级语义模型默认三条框编码路径全部开启并求和d_model=256、3 层 Transformer 编码层(num_layers=3)、激活检查点开启(use_act_ckpt=True),且使用正弦位置编码 PositionEmbeddingSine(num_pos_feats=256, temperature=10000)。其中 _get_clones(layer, num_layers) 负责把同一个 layer 深拷贝出多层堆叠(该工具函数定义在 nn/modules/utils.py,是 Ultralytics 内部通用的模块克隆助手)。


五、三种框编码方式的底层实现

5.1 直接线性投影(direct projection)

_encode_boxes 中(geometry_encoders.py),若 boxes_direct_project 开启,则将归一化 4 维坐标直接送入 nn.Linear(4, d_model)

proj = self.boxes_direct_project(boxes.to(img_feats.dtype))

这里 boxes.to(img_feats.dtype) 表明坐标会先被转换成与图像特征一致的精度(如混合精度下为 fp16)。

5.2 RoI align 特征汇聚(pooling)

boxes_pool 开启,其流程(geometry_encoders.py)为:

  1. img_featsH, W
  2. xywh2xyxy 把归一化 CxCyWH 框转为 xyxy,再按 [W, H, W, H] 缩放反归一化到像素坐标;
  3. 调用 torchvision.ops.roi_align(延迟导入以加快 ultralytics 包加载)在特征图上采样,得到 (B*N, d_model, roi_size, roi_size) 的 RoI 特征;
  4. nn.Conv2d(d_model, d_model, roi_size) 将每个 RoI 汇聚成 d_model 维向量(roi_size=7 时等价于 7×7 全局卷积池化);
  5. view(bs, n_boxes, d_model).transpose(0, 1) 还原为 (N, B, d_model) 的序列优先格式。

注意 xywh2xyxy 来自 utils/ops.py,是该仓库全局复用的坐标转换工具,与 SAM/SAM2 其他模块保持一致。

5.3 框中心位置编码(pos encoder)

boxes_pos_enc 开启(geometry_encoders.py),会把框解绑为 cx, cy, w, h 四个标量组,调用位置编码器的 encode_boxes

enc = self.pos_enc.encode_boxes(cx.flatten(), cy.flatten(), w.flatten(), h.flatten())
proj = self.boxes_pos_enc_project(enc.to(img_feats.dtype))

boxes_pos_enc_projectnn.Linear(d_model + 2, d_model)——多出的 2 维来自 encode_boxes 末尾直接拼接的 (h, w) 原始宽度/高度。PositionEmbeddingSine.encode_boxes 的实现在 nn/.../blocks.py(注:该文件真实路径为 ultralytics/models/sam/modules/blocks.py),其做法是把中心点按 y, x 顺序拼接正弦编码后再接 h, w

pos_x, pos_y = self._encode_xy(x, y)
return torch.cat((pos_y, pos_x, h[:, None], w[:, None]), dim=1)

三路结果(直接投影 / RoI / 位置编码)在 _encode_boxes 中通过"先到先得、后到累加"的方式求和;最后统一加上 type_embed = self.label_embed(boxes_labels.long()) 的标签 embedding,作为几何 token 的最终表示。这印证了类 docstring 中"三种编码互不排斥、多选即求和"的表述。


六、角点编码模式:encode_boxes_as_points=True 的分支

encode_boxes_as_points=True 时,forward_encode_points 分支(geometry_encoders.py),将每个框"升级"成一对角点 token:

  1. boxes_xyxy = xywh2xyxy(boxes) 转为归一化 xyxy,再 split(split_size=2, dim=-1) 拆成 top_left(前两维)与 bottom_right(后两维);
  2. 对角点标签做偏移区分来源labels_tl = boxes_labels + 2labels_br = boxes_labels + 4,配合构造时预留的 6 类 label_embed,使 Transformer 能区分"左上正/负"与"右下正/负";
  3. 两组点按序列维 torch.cat 拼接成 (2*N, B, 2) 的点序列,mask 相应横向拼接;
  4. 交由 _encode_pointsnn.Linear(2, d_model) 直接投影 2 维坐标并叠加 6 类标签 embedding。

这种模式下序列长度翻倍(每框两 token),换来的是模型对框两角位置更细粒度的关注,适合需要更强空间定位能力的设定。


七、forward 主流程:CLS、归一化精化与 Transformer 编码层

SequenceGeometryEncoder.forward(geo_prompt, img_feats, img_sizes, img_pos_embeds=None)geometry_encoders.py)的完整流水线为:

  1. 取数:从 geo_prompt 解出 boxes / boxes_mask / boxes_labels;同时取 img_feats[-1] 作为"序列优先"(H*W, B, C)的跨模态记忆,供后续 cross-attention 使用。
  2. 池化前的特征准备:若启用了任一条 pooling 路径,则用 img_pre_norm(LayerNorm)对最后一层图像特征归一化,并由 (H*W, B, C) 重排为 (N, C, H, W) 图像格式以配合 RoI align。
  3. 按模式编码encode_boxes_as_points 为 True 走角点路径,否则走框编码路径,得到 final_embeds (L', B, d_model)final_mask (B, L')
  4. 追加 CLS:若 add_cls=True,用可学习 cls_embed 生成 1 个全 batch 共享的 CLS token(mask 位为 0,永不被 pad),并通过 concat_padded_sequences 将其拼接在序列头部(geometry_encoders.py)。这也解释了构造器中"使用 Transformer 时强烈建议开启 CLS"的断言——CLS 是编码层输出汇聚的聚合位。
  5. 后编码精化:若 final_proj 存在,执行 norm(final_proj(final_embeds))(Linear + LayerNorm)。
  6. 堆叠 Transformer 编码层:将 num_layers 个克隆的 layer 逐层作用——每层以图像特征为 memory、几何序列为 tgttgt_key_padding_mask 传入右填充 mask、pos 传入图像侧位置编码,最终经 encode_norm(LayerNorm)输出。

返回的 (final_embeds, final_mask) 即是第二节 sam3_image.pygeo_feats, geo_masks 的来历:前者为几何 token 序列,后者为其对应的 padding mask,二者一起作为 prompt 参与后续 Transformer 的文本/几何联合编码。use_act_ckpt 在构造时被保存但不在本模块内显式包装,从源码结构看它由外部的封装层结合 torch.utils.checkpoint 机制统一启用。


八、维度约定与使用要点速查

综合 Promptconcat_padded_sequencesSequenceGeometryEncoder 三者的 docstring 与断言,可提炼出以下必须遵守的约定(也是二次开发时最容易出错之处):

  1. 序列优先、批次第二:embedding/坐标类张量形状为 (seq_len, batch, feat);mask 与绝大多数标签张量批次优先。
  2. mask 的 1 表示 pad,且必须是右填充:所有交给编码器 / 拼接函数的 mask 都需满足 is_right_padded
  3. 框坐标使用归一化 CxCyWH:末维为 4;转 xyxy、反归一化等由编码器内部按 H/W 完成,外部只需保证归一化。
  4. d_model 贯穿始终:图像特征、位置编码输出、标签 embedding、投影输出与图像侧记忆共享同一维度,改动时需保证 backbone 特征通道与 d_model 匹配。
  5. 框与角点编码二选一encode_boxes_as_points=True 时内部忽略 boxes_direct_project/boxes_pool/boxes_pos_enc 三开关(此时点编码仅用线性投影),其标签字典扩大为 6 类。

推理侧的最小化示例可以参考 _get_dummy_promptpredict.py):构造零框 Prompt 时传入 box_embeddings=torch.zeros(0, B, 4)box_mask=torch.zeros(B, 0, dtype=torch.bool) 即可,编码器对空序列亦能正常前向(直接拼接 CLS 后进入 Transformer)。


九、小结

geometry_encoders.py 是 SAM3 几何提示链路上封装最完整、契约最严格的模块之一:Prompt 统一了输入数据的形状与语义,两个工具函数维护了变长序列批式运算的核心不变式,SequenceGeometryEncoder 则以"多路编码求和 + 可选 CLS + 后置精化 + Transformer 堆叠"的模块化设计,把任意数量的归一化框提示转换为与图像/文本特征同构的 token 序列。本文所引的接口签名与行为均以 reference 文档源码 及其在 build_sam3.pysam3_image.pypredict.py 中的真实调用为准;对"为何如此设计"等推断性结论,已在文中以"从源码结构看/可以推断"等措辞明确标注。读者如需深入,可直接以上述文件为入口研读完整实现。

登录后查看全文
热门项目推荐
相关项目推荐