Ultralytics SAM 模块工具函数详解:utils.py 中 12 个核心算子与位置编码实现剖析
Ultralytics 仓库中 ultralytics/models/sam/modules/utils.py 汇集了 SAM / SAM 2 系列模型中高度复用的 12 个底层工具函数,覆盖时间维度条件帧选取、一维正弦时间位置编码、二维轴向 RoPE(旋转位置编码)、窗口注意力切分/还原、相对与绝对位置编码等关键算子。本文以该文件为主体,逐一解读每个函数的输入输出、算法思路与调用点,并结合仓库中 blocks.py、sam.py 与 vitdet.py 的源码揭示它们如何支撑视频分割与高分辨率注意力计算。读完本文,读者将能从算子层面理解 SAM 系列视觉 Transformer 的位置编码机制,并掌握每个函数在独立场景下的直接调用方法。
文件定位:SAM 家族共享的"无状态"工具层
utils.py 中的函数全部为纯函数式算子,不维护任何可学习参数,仅依赖 torch 与 torch.nn.functional(见 utils.py),因此它们可以被图像分割模型、视频分割模型以及 SAM3 的 ViT 编码器在不同位置独立复用。从源码看,该文件内的函数可分成四类能力:
| 能力域 | 函数 | 典型使用方 |
|---|---|---|
| 视频时序处理 | select_closest_cond_frames、get_1d_sine_pe |
sam.py 视频分割主干 |
| 二维轴向 RoPE | init_t_xy、compute_axial_cis、reshape_for_broadcast、apply_rotary_enc |
blocks.py、sam3/vitdet.py 注意力层 |
| 窗口注意力 | window_partition、window_unpartition |
各 Transformer Block 前向 |
| 相对/绝对位置编码 | get_rel_pos、add_decomposed_rel_pos、get_abs_pos、concat_rel_pos |
blocks.py、sam3/vitdet.py 注意力层 |
下文按这四个能力域分别展开。
视频时序处理:条件帧选取与时间正弦编码
这两组函数服务于 SAM 2 / SAM3 的视频(Memory Attention)路径,解决"当前帧应该参考哪些历史记忆帧"以及"如何给每帧/每个 object pointer 注入时间位置信息"两个问题。
select_closest_cond_frames:挑选时间上最邻近的条件帧
条件帧(conditioning frame)是视频分割中承载显式 prompt 的帧。该函数根据当前帧索引从 cond_frame_outputs 中挑出时间距离最近的至多 max_cond_frame_num 帧:
def select_closest_cond_frames(frame_idx: int, cond_frame_outputs: dict[int, Any], max_cond_frame_num: int):
关键逻辑(见 utils.py):
- 当
max_cond_frame_num == -1或候选数不超过上限时,全部选中,无未选帧; - 否则断言上限至少为 2(保证可同时取前后两个方向的条件帧);
- 先取两侧锚点:优先选中
frame_idx之前最近的一帧与frame_idx起最近的一帧(含当前帧本身),这是"边界稳定"策略; - 再按时间距离补足:对剩余帧按
abs(t - frame_idx)升序取足名额; - 返回
(selected_outputs, unselected_outputs)两个字典。
源码内给出了可复现的示例(utils.py):
frame_idx = 5
cond_frame_outputs = {1: "a", 3: "b", 7: "c", 9: "d"}
max_cond_frame_num = 2
selected, unselected = select_closest_cond_frames(frame_idx, cond_frame_outputs, max_cond_frame_num)
# selected -> {3: 'b', 7: 'c'}(5 前后最近的各一帧)
# unselected -> {1: 'a', 9: 'd'}
在 sam.py 的 _prepare_memory_conditioned_features 中,该函数按 self.max_cond_frames_in_attn 挑选条件帧用于跨帧注意力(cross attention);值得注意的是,未被选中的条件帧如果恰好落在最近 num_maskmem - 1 帧内,仍会作为普通记忆帧参与融合(sam.py),可见"是否被选中"只影响其是否参与强约束的显式条件注意力,不影响其作为通用记忆的可用性。
get_1d_sine_pe:一维正弦位置编码
def get_1d_sine_pe(pos_inds: torch.Tensor, dim: int, temperature: float = 10000):
该函数为任意位置索引生成维度为 dim 的一维正弦位置编码,输出形状为 (pos_inds.shape[0], dim)。其内部实现(utils.py):
pe_dim = dim // 2,仅支持偶数维度;- 频率底数
dim_t = temperature ** (2 * (dim_t // 2) / pe_dim),与经典 Transformer 位置编码一致; - 对归一化后的相位分别取
sin与cos并在最后一维拼接。
在视频分割前向中,该函数用于为 object pointer(对象指针) 构造时间位置编码。sam.py 中先把每个历史指针相对当前帧的时间差除以最大指针跨度做归一化(obj_pos / t_diff_max),再调用 get_1d_sine_pe(...) 生成编码并经 MLP 投影后广播到 batch 维度(见 sam.py),从而让记忆编码器知道每个历史对象指针距离当前帧有多远。
二维轴向 RoPE:从坐标网格到复数旋转
SAM3 的 ViT 主干(vitdet.py 中的 Attention)支持 2D RoPE 作为位置信息注入方式。这组函数完成"坐标初始化 → 轴向复数频率 → 广播对齐 → 旋转 Q/K"的完整链路。
init_t_xy:生成网格的 x/y 坐标
def init_t_xy(end_x: int, end_y: int, scale: float = 1.0, offset: int = 0):
返回长度为 end_x * end_y 的展平坐标张量对:t_x 表示列坐标,t_y 表示行坐标(utils.py):
t_x, t_y = init_t_xy(3, 2)
# t_x -> tensor([0., 1., 2., 0., 1., 2.])
# t_y -> tensor([0., 0., 0., 1., 1., 1.])
坐标可统一乘 scale 并加 offset,这一设计为后续做基于分辨率的坐标缩放(例如 RoPE 插值)预留了接口。
compute_axial_cis:轴向复数指数编码
def compute_axial_cis(dim: int, end_x: int, end_y: int, theta: float = 10000.0, scale_pos: float = 1.0):
在 2D 网格上构造复数指数位置编码(utils.py):
- 按
theta ** (arange(0, dim, 4) / dim)的倒数生成 x、y 两个方向各自的频率向量,维度均为dim // 4; - 用
torch.outer(t_x, freqs_x)、torch.outer(t_y, freqs_y)得到每个空间位置上的相位; - 用
torch.polar将相位转为单位模长的复数,最终cat成形状(end_x * end_y, dim // 2)的复数张量。
从文档字符串给出的示例看,dim=128, end_x=8, end_y=8 时输出形状为 (64, 64),即 64 个空间位置对应 64 个复数频率分量(utils.py)。
reshape_for_broadcast:为广播重塑频率张量
def reshape_for_broadcast(freqs_cis: torch.Tensor, x: torch.Tensor):
将复数频率张量从 (H, W) 重塑为可与 x(通常是 [batch, heads, seq, dim] 结构)广播的维度,即把前 ndim - 2 个维度全部置 1,并断言 freqs_cis.shape == (x.shape[-2], x.shape[-1])(utils.py)。
apply_rotary_enc:对 Q/K 应用旋转编码
def apply_rotary_enc(xq, xk, freqs_cis, repeat_freqs_k: bool = False):
这是 RoPE 的核心执行点(utils.py):
- 用
torch.view_as_complex把 Q/K 最后一维拆成两两一组并视为复数; - 与广播后的
freqs_cis做复数乘法,实现向量旋转,再经view_as_real与flatten还原实数张量; - 当
xk序列为空(例如 key dropout)时直接返回原始xk; repeat_freqs_k=True且 key 序列长于 query(如跨分辨率注意力、或 key 含 memory 拼接)时,沿序列维度重复频率分量;针对 MPS 平台不支持复数repeat的限制,代码先view_as_real拆分后再 repeat 并还原(utils.py),这是相当实用的硬件兼容细节;- 返回值通过
type_as(...)保持与输入相同的浮点类型与设备。
调用示例(来自源码文档字符串 utils.py):
xq = torch.randn(2, 8, 16, 64) # [batch, heads, seq_len, dim]
xk = torch.randn(2, 8, 16, 64)
freqs_cis = compute_axial_cis(64, 4, 4) # 4x4 空间网格、dim=64
q_encoded, k_encoded = apply_rotary_enc(xq, xk, freqs_cis)
在 SAM2 记忆注意力中,blocks.py 会根据 num_k_exclude_rope 将 key 中"不应做旋转的部分"(如显式条件帧)切出,只对前 num_k_rope 个 key 施加旋转,并在 key 序列更长时开启 repeat_freqs_k;SAM3 的 ViT 则在 _apply_rope 中按需把预计算的 freqs_cis 搬运到 Q 所在设备后调用(vitdet.py)。
窗口注意力:切分与还原的非重叠窗口
高层级特征图的全局注意力计算成本随分辨率平方级增长,Swin/ViTDet 式窗口注意力通过把特征切分成固定大小的非重叠窗口来约束计算范围。window_partition 与 window_unpartition 是一对互逆算子。
window_partition
def window_partition(x: torch.Tensor, window_size: int):
输入形状为 (B, H, W, C)(注意是 HWC 布局),输出 (B * num_windows, window_size, window_size, C) 的窗口序列,并返回补齐后的 (Hp, Wp)(utils.py):
- 当 H、W 不能被
window_size整除时,先在右侧/下方用零补齐; - 通过一次
view把特征张量拆成 6 维(B, Hp//ws, ws, Wp//ws, ws, C),再permute + contiguous + view聚合成窗口张量,整个过程无需unfold循环,且完全无参。
window_unpartition
def window_unpartition(windows: torch.Tensor, window_size: int, pad_hw: tuple[int, int], hw: tuple[int, int]):
反向还原(utils.py):由窗口序列依据 (Hp, Wp) 反推出 batch 数,重排回 (B, Hp, Wp, C),再按原始 (H, W) 裁掉补齐的边缘:
windows = torch.rand(32, 8, 8, 64) # 32 个 8x8 窗口、64 通道
pad_hw = (16, 16); hw = (15, 14)
x = window_unpartition(windows, window_size=8, pad_hw=pad_hw, hw=hw)
# x.shape -> torch.Size([8, 15, 14, 64])
这一对函数在 SAM3 的 Block.forward 中按 window_size > 0 包裹注意力前后(vitdet.py),在 SAM2 主干的多尺度 Block 与标准 Block 中同样出现(blocks.py、blocks.py)。
相对与绝对位置编码:给注意力注入空间关系
get_rel_pos:按 q/k 尺寸抽取相对位置嵌入
def get_rel_pos(q_size: int, k_size: int, rel_pos: torch.Tensor):
输入相对位置嵌入表 rel_pos,形状 (L, C),其中 L = 2 * max(q_size, k_size) - 1 是最大相对距离覆盖范围,输出形状 (q_size, k_size, C) 的相对位置编码张量(utils.py)。算法要点:
- 若预训练嵌入长度与实际最大相对距离不符,用
F.interpolate(..., mode="linear")沿长度维线性插值重采样; - 对 q、k 尺寸不等的情形,坐标先按
max(k/q, 1.0)等比例缩放,构造(q_coords - k_coords) + (k_size - 1) * max(q_size/k_size, 1.0)的索引矩阵,最后用索引查表得到每个 (q 位置, k 位置) 对的编码向量。
该查表技巧是 SAM 系列支持"训练分辨率与推理分辨率不一致"的关键——尺寸变化时无需重新学习,只需插值表并重算索引。
add_decomposed_rel_pos:MViTv2 式分解相对位置注意力
def add_decomposed_rel_pos(attn, q, rel_pos_h, rel_pos_w, q_size, k_size):
按 MViTv2 论文思想,把二维相对位置分解为高度、宽度两个轴分别处理(utils.py):
- 先由
get_rel_pos分别得到高度轴Rh(q_h,k_h,C)与宽度轴Rw(q_w,k_w,C); - 用
torch.einsum("bhwc,hkc->bhwk", r_q, Rh)与einsum("bhwc,wkc->bhwk", r_q, Rw)把查询与相对编码内积,得到逐位置的空间偏置; - 最后将
attn展开为(B, q_h, q_w, k_h, k_w)后与两个偏置相加再还原。
文档字符串中的示例:B=1, C=64, q_h=q_w=k_h=k_w=8 时输出形状保持 (1, 64, 64)(utils.py),即空间注意力图在网络前向中带上了相对位置偏置。该方法在 SAM2 的全局注意力分支中用于提升没有使用窗口时的高分辨率注意力质量(blocks.py)。
get_abs_pos:绝对位置嵌入的尺寸自适应
def get_abs_pos(abs_pos, has_cls_token: bool, hw: tuple[int, int], retain_cls_token: bool = False, tiling: bool = False):
输入预训练的绝对位置嵌入(可为 (1, num_position, C) 含或不含 cls token),输出适配当前 hw 分辨率的嵌入(utils.py)。关键分支:
- 含 cls token 时先剥离
cls_pos,处理完再按retain_cls_token决定是否拼接回首位; - 尺寸不匹配时提供两种适配策略:
tiling=True用平铺裁剪(abs_win风格,适应高分辨率推理),否则用F.interpolate(mode="bicubic")双三次插值; - 尺寸匹配时直接 reshape 为
(1, H, W, C)。
在 vitdet.py 的前向中,patch embedding 的输出会在进入 Transformer Block 前加上由 get_abs_pos 生成的绝对位置嵌入(是否保留 cls token 与平铺策略均可配置)。
concat_rel_pos:把相对位置偏置"拼进" Q/K
def concat_rel_pos(q, k, q_hw, k_hw, rel_pos_h, rel_pos_w, rescale=False, relative_coords=None):
与 add_decomposed_rel_pos 的"加偏置"思路不同,此函数把相对位置系数拼接到 Q 和 K 的特征维度上,使 q @ k^T 在点积时天然把相对位置信息包含进去(utils.py):
- 仅支持方形输入(
q_h == q_w且k_h == k_w),这符合 SAM3 ViT 的方形 patch 网格假设; - 可传入预计算的
relative_coords索引表以跳过重复构造; - 用
torch.eye构造单位矩阵作为 K 的拼接块,使得点积结果中位置偏置项按索引正确对齐; - 若
rescale=True(为与 PyTorchF.scaled_dot_product_attention的缩放因子对齐),会按新增维度重新计算缩放比例(dim + k_h + k_w) ** 0.5并补偿到 Q 上——否则 SDPA 内部的1/sqrt(d)缩放会因拼接后的维度膨胀而失配,这是与F.scaled_dot_product_attention组合使用时的关键细节。
该函数被 SAM3 的 Attention.forward 用于 use_rel_pos=True 的路径,先拼好 Q/K,再送入 F.scaled_dot_product_attention(q, k, v) 完成高效注意力(vitdet.py)。
在仓库中观察到的整体调用关系
- SAM2 图像/视频主干(blocks.py):集中导入
add_decomposed_rel_pos, apply_rotary_enc, compute_axial_cis, window_partition, window_unpartition,它们服务于记忆注意力(RoPE + key 部分旋转)与多尺度/高分辨率窗口注意力。 - SAM2 视频分割模型(sam.py):仅导入
get_1d_sine_pe, select_closest_cond_frames,分别承担对象指针的时间编码与条件帧筛选。 - SAM3 图像编码器(vitdet.py):组合使用 RoPE 全套、窗口切分、
concat_rel_pos与get_abs_pos,对应其混合了"绝对位置 + 相对位置 + 窗口注意力 + 2D RoPE"的 ViT 设计;同时set_imgsz允许在推理时动态重算相对位置表与 RoPE 频率(vitdet.py)。
小结与工程启示
utils.py 的 12 个函数虽小,却对应着 SAM 家族在视频时序建模、高分辨率窗口注意力与位置编码上的三个核心工程决策:
- 位置的两种建模:
get_abs_pos处理"全局绝对位置",compute_axial_cis/add_decomposed_rel_pos/concat_rel_pos分别提供基于复数旋转与基于查表相加/拼接的相对位置方案,三者可并存于同一网络中并各司其职; - 分辨率自适应的通用套路:相对位置表按需线性插值(
get_rel_pos)、绝对位置按需平铺或双三次插值(get_abs_pos)、RoPE 频率按需缩放坐标(scale_pos),使模型可以摆脱训练分辨率的束缚; - 与
F.scaled_dot_product_attention的兼容处理:concat_rel_pos的rescale与apply_rotary_enc对 MPS 复数 repeat 的规避,均是面向现代 PyTorch 算子与多后端部署的精细适配。
对想深入二次开发的研究者,直接阅读 utils.py 及其在 blocks.py、sam.py、sam3/vitdet.py 中的调用点,即可完整串起从底层算子到视频/图像分割模型前向的调用链。
atomcodeClaude Code 的开源替代方案。连接任意大模型,编辑代码,运行命令,自动验证 — 全自动执行。用 Rust 构建,极致性能。 | An open-source alternative to Claude Code. Connect any LLM, edit code, run commands, and verify changes — autonomously. Built in Rust for speed. Get StartedRust0629
Hy4-previewHy4 preview 是由腾讯混元团队研发的新一代混合专家(MoE)旗舰模型。模型总参数量 770B,每个 token 激活 49B,主干共包含78层,第一层采用标准 FFN,其余 77 层均为 MoE 结构,每层包含 256 个路由专家与 1 个共享专家,每个 token 激活 top-8 路由专家及共享专家。主干之外原生内置 1 层 MTP(总参数量 10B,激活 0.7B)以支持投机解码。Python00
GLM-5.3GLM-5.3 与 GLM-5.2 使用相同的基座模型——所有提升均来自后训练。与 GLM-5.2 相比,它在复杂编程和长程任务上的表现显著提升。Jinja00
GLM-5.3-FlashGLM-5.3-Flash (320B-A18B),是GLM-5系列的首个原生多模态模型。320B总参数,能力超过GLM-5.2Jinja00
Spark-X2.5-4BSpark-X2.5-4B 旨在让强大的 AI 更实用、更高效、更易获得。在广泛日常任务中表现强劲,涵盖对话、写作、翻译、推理、编码、工具调用以及智能体工作流,并在同等规模的开源模型中取得领先成绩。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00
Spark-X2.5-1.7BSpark-X2.5-1.7B 旨在让强大的 AI 更加实用、高效且易于获取。这些模型在广泛的日常任务中表现出色,涵盖对话、写作、翻译、推理、编程、工具调用和智能体工作流,并在同等规模的开源模型中取得领先结果。Spark-X2.5 将面向效率的架构与最高 1M tokens 的原生上下文窗口相结合,并支持 200 多种语言。Python00