首页
/ Ultralytics SAM3 视觉主干解析:ViTDet Backbone(vitdet.py)架构与参数全解读

Ultralytics SAM3 视觉主干解析:ViTDet Backbone(vitdet.py)架构与参数全解读

2026-09-07 11:42:57作者:滕妙奇

本文是 docs/en/reference/models/sam/sam3/vitdet.md 这一 API 参考页面的深度展开。该页面通过 mkdocstrings 索引了 ultralytics/models/sam/sam3/vitdet.py 模块中的 AttentionBlockViT 三个类。本文以这些类的源码 docstring 与实际实现为主体,结合 SAM3 在仓库中的真实装配方式,完整解读这套 ViTDet 视觉主干的模块职责、每个构造参数的语义,以及它在 SAM3 多模态分割链路中扮演的角色。

一、模块定位:SAM3 的纯 Transformer 视觉特征提取器

SAM3 是 Ultralytics 仓库中负责视觉语言分割与视频跟踪的一族模型(源码位于 ultralytics/models/sam/sam3)。其图像侧编码器并不使用 CNN 金字塔,而是直接采用一套"纯 ViT"主干——也就是 vitdet.py 中实现的 ViTDet Backbone

源码文件头说明了它的来源与定位:

  • 该主干改编自 Detectron2 的 ViTDet 实现(原始论文题为 Exploring Plain Vision Transformer Backbones for Object Detection),即"把普通 ViT 当作检测 / 分割骨干直接使用"的路线;
  • 模块内的 RoPE(旋转位置编码) 实现参考了多份公开代码(文件注释列出的 3 个参考来源,均为公开的 rope/rotary embedding 开源实现);
  • 文件同时携带 Copyright (c) Meta Platforms, Inc. 的版权声明,表明其上游血统。

在 SAM3 中,这个主干被包在 Sam3DualViTDetNeck 颈部结构内,再与文本编码器组成 SAM3VLBackbone,最后接入掩码 Transformer 解码链路。因此,理解 vitdet.py 就等于理解 SAM3 图像特征从 patch 化到多尺度特征输出的全过程。

模块从内到外包含三个层次,与 reference 页面索引一一对应:

  1. Attention —— 带相对位置编码 + 2D RoPE的多头注意力单元;
  2. Block —— 支持**窗口注意力(window attention)**的标准 Transformer Block;
  3. ViT —— 由上述 Block 堆叠而成的完整主干,负责输出特征图。

三者均依赖同目录工具函数(modules/utils.py 中的 compute_axial_cisapply_rotary_encwindow_partitionwindow_unpartitionconcat_rel_posget_abs_pos),并在构造时引用 sam3/model_misc.py 中的 LayerScale

二、Attention:融合相对位置编码与 2D RoPE 的多头注意力

Attention 类(vitdet.py 中 L40 起)是一个"即插即用"的多头注意力模块。它最大的特点是在标准 QKV 注意力之外,提供两条正交的位置信息注入通道:可学习的分解式相对位置编码(rel_pos)与旋转位置编码(RoPE)。二者互不冲突,源码注释明确写到 use_rope"independent of use_rel_pos, as it can be used together"。

2.1 构造参数

参数 默认值 含义
dim 必填 输入通道数
num_heads 8 注意力头数,head_dim = dim // num_heads
qkv_bias True 是否在 QKV 线性投影中加入可学习偏置
use_rel_pos False 是否向注意力图加入相对位置编码
rel_pos_zero_init True 是否将相对位置参数零初始化(否则用 trunc_normal_,std=0.02)
input_size None 输入分辨率,用于计算相对位置参数尺寸或 RoPE 尺寸
cls_token False 是否存在 cls token(存在时 attention map 需相应处理)
use_rope False 是否使用 2D RoPE
rope_theta 10000.0 控制 RoPE 频率的基频
rope_pt_size None 前一训练阶段(pretrain)的 RoPE 尺寸,供插值/平铺使用
rope_interp False 是否把 RoPE 插值(而非外推)到目标输入尺寸

2.2 相对位置编码的初始化(_setup_rel_pos

开启 use_rel_pos 后:

  • 要求提供 input_size,且不允许与 cls_token 共存(源码断言 cls_token is False);
  • 注册两个可学习参数 rel_pos_hrel_pos_w,形状均为 (2 * input_size - 1, head_dim),即对 H、W 两个轴向分别建模相对距离;
  • rel_pos_zero_init=False,用 trunc_normal_ 以 std=0.02 初始化;
  • 预计算 relative_coords 索引张量,供前向时直接从 embedding 查表,避免逐次计算。

2.3 2D RoPE 频域的初始化(_setup_rope_freqs

开启 use_rope 后:

  • rope_pt_sizeNone 时以 input_size 兜底;
  • 通过 functools.partial 固定 compute_axial_cis(dim=head_dim, theta=rope_theta) 作为频域计算函数;
  • rope_interp=True 时,以 scale_pos = rope_pt_size[0] / input_size[0] 缩放坐标,实现"把预训练分辨率下学习到的频域插值到当前分辨率"的目的;
  • 若携带 cls token,会额外拼接一个恒等旋转的极坐标项 cls_freqs_cis,保证 cls 位置不参与旋转。

2.4 前向过程

forward 同时支持两种布局的输入:

  • 4D 张量 (B, H, W, C):特征图形式(此时不允许 cls token);
  • 3D 张量 (B, L, C):序列形式,并自动由 L 推导 H = W = sqrt(L - s)s 为 cls token 个数)。

计算流程(见 vitdet.py):

  1. 一次 qkv 线性投影得到 (B, L, 3, num_heads, head_dim),拆出 q、k、v;
  2. 对 q、k 施加 RoPE(_apply_ropeapply_rotary_enc);
  3. use_rel_pos,调用 concat_rel_pos(..., rescale=True) 把相对位置编码"拼接"进 q、k——这是一种用拼接维度代替加法偏置的等效实现,并通过 rescale 修正缩放因子以配合 F.scaled_dot_product_attention(详见后文工具函数部分);
  4. 用 PyTorch 的 F.scaled_dot_product_attention 完成高效注意力计算;
  5. 按原布局(4D/3D)重组并经过 proj 输出。

三、Block:带窗口注意力与 LayerScale 的 Transformer 单元

Block 类(vitdet.py)把上面的 Attention 包装成标准的前置归一化(pre-norm)Transformer Block,并叠加三项 ViTDet/SAM 体系常用的工程化组件:窗口注意力Stochastic Depth(DropPath)LayerScale

3.1 构造参数

Block 在透传 Attention 全部位置编码相关参数的基础上,新增了:

参数 默认值 含义
mlp_ratio 4.0 MLP 隐藏维与嵌入维之比
drop_path 0.0 Stochastic Depth(随机深度)概率
norm_layer nn.LayerNorm 归一化层构造器
act_layer nn.GELU MLP 激活函数
window_size 0 窗口注意力窗口尺寸;0 表示不启用窗口注意力
dropout 0.0 残差路径与 MLP 内部 dropout
init_values None LayerScale 初始值;为 None 时不启用 LayerScale

注意依赖关系:Block 内部 check_requirements("timm") 并从 timm.layers 导入 DropPathMlp,即该模块的运行环境需要安装 timm

3.2 窗口注意力机制

window_size > 0 时,Block 在前向中先调用 window_partition(x, window_size)(B, H, W, C) 特征切成互不重叠的若干 window_size × window_size 窗口(必要时自动 padding),再送入注意力;注意力计算完成后用 window_unpartition 还原并裁掉 padding(见 vitdet.py)。窗口注意力的价值在于把 O((H·W)²) 的全局注意力复杂度降为窗口内的局部计算,让深层高分辨率特征成为可行。

同时 Block 构造 Attention 时,若处于窗口模式,会把 Attentioninput_size 固定为 (window_size, window_size),使位置编码参数大小只与窗口相关。

3.3 残差结构

shortcut = x
x = norm1(x)                 # 前置归一化
x = (window_partition) → attn → (window_unpartition)
x = ls1(x)                   # LayerScale(若启用)
x = shortcut + dropout(drop_path(x))
x = x + dropout(drop_path(ls2(mlp(norm2(x)))))   # 第二条残差

其中 LayerScale 来自 sam3/model_misc.py,以逐通道可学习缩放系数(init_values 为初值)作用在注意力/MLP 输出上,是稳定深层 Transformer 训练的重要技巧。每个 Block 的 drop_path 概率由上层 ViT 按 stochastic depth 递减规则逐一分配。

四、ViT:由 32 层 Block 组成的完整主干

ViT 类(vitdet.py)是模块的顶层封装,docstring 明确引用了 ViTDet 论文。它负责把原始图像送入 Patch Embedding,逐层经过 Block 后返回(多层)特征图。

4.1 构造参数全表

参数 默认值 含义
img_size 1024 输入图像尺寸(仅影响相对位置/RoPE 计算)
patch_size 16 Patch 尺寸
in_chans 3 输入通道数
embed_dim 768 Patch 嵌入维度
depth 12 Transformer 深度(Block 数)
num_heads 12 每 Block 注意力头数
mlp_ratio 4.0 MLP 隐藏维与嵌入维之比
qkv_bias True QKV 投影是否带偏置
drop_path_rate 0.0 Stochastic Depth 总速率(按深度线性衰减分配)
norm_layer "LayerNorm" 归一化层,可为构造器或名称字符串
act_layer nn.GELU MLP 激活函数
use_abs_pos True 是否使用绝对位置编码
tile_abs_pos True 绝对位置编码尺寸不匹配时用平铺而非插值
rel_pos_blocks (2, 5, 8, 11) 启用相对位置编码的 Block 索引;bool=True 时全部启用
rel_pos_zero_init True 相对位置参数零初始化
window_size 14 窗口尺寸;配合 global_att_blocks 区分全局/窗口块
global_att_blocks (2, 5, 8, 11) 使用全局注意力的 Block 索引,其余块使用窗口注意力
use_rope False 是否使用 2D RoPE(可与 rel_pos 并存)
rope_pt_size None 预训练阶段的 RoPE 尺寸,供插值/平铺
use_interp_rope False 是否插值 RoPE 到目标尺寸
pretrain_img_size 224 预训练输入尺寸,决定绝对位置编码 patch 数
pretrain_use_cls_token True 预训练模型是否带 cls token
retain_cls_token True 当前模型是否保留 cls token
dropout 0.0 残差及 MLP 内 dropout
return_interm_layers False 是否返回所有全局注意力块的中间特征
init_values None LayerScale 初始值
ln_pre False 是否在 Block 前施加 LayerNorm
ln_post False 是否在末层输出前施加 LayerNorm
bias_patch_embed True Patch Embedding 卷积是否带偏置
compile_mode None torch.compile 编译模式;None 表示不编译
use_act_checkpoint True 训练时是否启用激活检查点(activation checkpointing)

4.2 构造逻辑要点

窗口与全局块划分。 window_block_indexes = [i for i in range(depth) if i not in global_att_blocks],即只有落在 global_att_blocks 索引上的块做全局注意力,其余块都使用窗口注意力。相对位置编码块的标记则按 rel_pos_blocks(tuple 或 bool)展开到 depth 长度的布尔列表。

cls token 约束。retain_cls_token=True 时:

  • 要求预训练本就使用 cls token;
  • 由于窗口化特征图会被展平,源码断言"windowing 不支持与 cls token 共存"且"rel pos 不支持与 cls token 共存";
  • 新增 class_embedding 可学习参数,初始缩放为 embed_dim ** -0.5

Patch Embedding。 复用 modules/blocks.py 中的 PatchEmbed(卷积实现,kernel/stride 均为 patch_size),可通过 bias_patch_embed 控制偏置。

绝对位置编码。use_abs_pos=True,按 pretrain_img_size // patch_size 计算 num_patches,再依 pretrain_use_cls_token 决定是否多留一个 cls 位置,注册为 pos_embed 参数并 trunc_normal_(std=0.02)初始化。当前向遇到尺寸与预训练不一致的特征时,get_abs_pos 会依据 tile_abs_pos 选择平铺或双三次插值。

Stochastic Depth 衰减。 dpr = torch.linspace(0, drop_path_rate, depth),逐 Block 分配递增的丢弃概率。

LayerScale 与可选编译。 每块按 init_values 装配 LayerScale;当 compile_mode 非空时,会把整个 forwardtorch.compile(..., mode=compile_mode, fullgraph=True) 编译,并在训练 + 激活检查点场景下关闭 DDP 图优化(torch._dynamo.config.optimize_ddp = False)。

4.3 前向输出与多尺度特性

前向流程(见 vitdet.py):

  1. patch_embed(x) 得到 (B, H, W, C) 的 patch 序列(SAM3 中通常为 14×14 卷积下采样);
  2. 若保留 cls token,则把 class_embedding 拼接在序列头部;
  3. 叠加绝对位置编码(经 get_abs_pos 按需平铺/裁剪);
  4. 逐 Block 前向;训练模式下若 use_act_checkpoint=True,用 torch.utils.checkpoint.checkpoint(blk, x, use_reentrant=False) 包住每个 Block 以换取显存;
  5. 在每个(或最后一个)全局注意力块后收集输出,剥离 cls token 并还原为 NCHW 特征图。

返回值为 list[torch.Tensor]。若 return_interm_layers=False,只在最后一个全局块处返回一层特征;若为 True,则返回全部全局注意力块对应的中间层特征——这正是为多尺度下游(neck + 多层级 Transformer encoder)预留的接口。

4.4 动态分辨率适配:set_imgsz

set_imgsz(imgsz=None)vitdet.py)允许在推理前把主干切换到新的输入分辨率(默认回退到 [1008, 1008])。它遍历所有 Block,对非窗口块重新执行 _setup_rel_pos_setup_rope_freqs,按 imgsz // patch_size 重建相对位置参数与 RoPE 频域。这是 SAM3 在不重新训练的前提下适配不同分辨率图像、并让 RoPE 通过插值保持平移等变性的关键机制。

五、底层位置编码与窗口工具:vitdet 的"基础设施"

ViT/Attention/Block 频繁调用的六个工具函数全部集中在 modules/utils.py,理解它们才能真正读懂 vitdet 的数值流:

  • compute_axial_cis(dim, end_x, end_y, theta, scale_pos)(L119):对 H、W 两轴分别构造轴向(axial)旋转频率,用极坐标形式 torch.polar 生成复数旋转因子,输出 (end_x*end_y, dim//2)。SAM2 系列的 RoPEAttentionmodules/blocks.py 中)也复用它,theta=10000scale_pos 为插值缩放。
  • apply_rotary_enc(xq, xk, freqs_cis, repeat_freqs_k)(L175):把 q/k 视为复数,按 freqs_cis 做旋转乘法。含对 MPS 设备不支持复数 repeat 的特殊处理。
  • window_partition / window_unpartition(L225/L255):特征图 ⇆ 不重叠窗口的互逆变换,自动 padding 对齐并裁剪还原。
  • concat_rel_pos(q, k, q_hw, k_hw, rel_pos_h, rel_pos_w, rescale, relative_coords)(L454):把相对位置 bias 以"拼接 q、k 的扩展维度"方式注入注意力——q 拼上 rel_hrel_w,k 拼上单位阵。这样 qkᵀ 的乘积中自然出现位置偏置项;rescale=True 时按新增维度重新校准缩放因子,确保与 F.scaled_dot_product_attention 内部缩放配合正确。
  • get_abs_pos(abs_pos, has_cls_token, hw, retain_cls_token, tiling)(L389):把预训练分辨率下学习的绝对位置编码适配到当前分辨率——tiling 模式下用平铺复制,否则用双三次插值,并正确处理 cls 位置。

六、生产装配:SAM3 里真实使用的 ViTDet 配置

源码文件 build_sam3.py_create_vision_backbone 给出了这套主干在 SAM3 中的真实实例化参数(L37-L62),与上一节的默认值有显著差异,充分体现了 ViTDet 的设计意图:

vit_backbone = ViT(
    img_size=1008,          # SAM3 工作分辨率
    pretrain_img_size=336,  # 预训练分辨率(patch 化后为 24×24 patch 网格)
    patch_size=14,          # 特征步长 14(tracker 侧 backbone_stride=14 与此对应)
    embed_dim=1024,
    depth=32,               # 32 层
    num_heads=16,           # head_dim = 1024 / 16 = 64
    mlp_ratio=4.625,
    norm_layer="LayerNorm",
    drop_path_rate=0.1,
    qkv_bias=True,
    use_abs_pos=True,
    tile_abs_pos=True,      # 预训练→运行分辨率用平铺而非插值
    global_att_blocks=(7, 15, 23, 31),  # 每 8 层做一次全局注意力
    rel_pos_blocks=(),      # 全部块关闭 rel_pos,改为纯 RoPE
    use_rope=True,
    use_interp_rope=True,   # 336→1008 的 3× 分辨率迁移靠 RoPE 插值
    window_size=24,         # 1008/14=72,72 可被 24 整除
    pretrain_use_cls_token=True,
    retain_cls_token=False, # 部署时丢弃 cls token
    ln_pre=True,
    ln_post=False,
    return_interm_layers=False,
    bias_patch_embed=False,
    compile_mode=compile_mode,
)

从配置可以看出几条设计链:输入 1008×1008patch_size=14 卷积后得到 72×72 的 patch 网格;除第 7、15、23、31 层做全局注意力外,其余 28 层使用 24×24 窗口注意力以控制计算量;预训练于 336 分辨率(24×24 网格)后,通过 tile_abs_pos=True 平铺绝对位置编码、use_interp_rope=True 插值 RoPE 频域,完成到 3× 更高分辨率的迁移。

该主干随后被封装进 necks.pySam3DualViTDetNeck

Sam3DualViTDetNeck(
    position_encoding=PositionEmbeddingSine(num_pos_feats=256, normalize=True, scale=None, temperature=10000),
    d_model=256,          # 下游统一特征维
    scale_factors=[4.0, 2.0, 1.0, 0.5],  # 上采样/下采样尺度因子,构建多尺度 level
    trunk=vit_backbone,
    add_sam2_neck=enable_inst_interactivity,
)

从构造参数可以推断:neck 负责把 ViT 输出投影/缩放到 256 维,并按 scale_factors 生成多个分辨率的特征 level,供 encoder.py 中基于 Deformable-DETR 风格的多层 Transformer encoder(TransformerEncoderFusionnum_feature_levels=1)使用。需要交互式实例/视频分割时(add_sam2_neck=True),会额外附加 SAM2 风格的 neck 输出(高层特征 stride 16/stride 8 等)。

七、结语与延伸阅读

AttentionBlockViT 三层递进构成了 SAM3 视觉侧的主干答案:一个把 rel_pos、RoPE、窗口注意力、Stochastic Depth、LayerScale、绝对位置平铺、激活检查点、torch.compile 等多种工程手段集成于一体、且支持动态分辨率切换的纯 Transformer 检测骨干。它既承载了 ViTDet"plain ViT as backbone"的简洁理念,也体现了 SAM 系列针对高分辨率分割 / 跟踪场景的精细化调优。

想继续深入,可在仓库内按如下路径追踪:

说明:本文涉及的一切参数默认值、断言约束与调用关系均以当前仓库 vitdet.pybuild_sam3.py 的实际代码为准;文中凡属推断性质的表述(如 neck 的缩放行为)均已使用"可以推断/从构造参数看"等限定语。

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