首页
/ Ultralytics SAM 模块工具函数详解:utils.py 中 12 个核心算子与位置编码实现剖析

Ultralytics SAM 模块工具函数详解:utils.py 中 12 个核心算子与位置编码实现剖析

2026-09-07 21:19:54作者:伍霜盼Ellen

Ultralytics 仓库中 ultralytics/models/sam/modules/utils.py 汇集了 SAM / SAM 2 系列模型中高度复用的 12 个底层工具函数,覆盖时间维度条件帧选取、一维正弦时间位置编码、二维轴向 RoPE(旋转位置编码)、窗口注意力切分/还原、相对与绝对位置编码等关键算子。本文以该文件为主体,逐一解读每个函数的输入输出、算法思路与调用点,并结合仓库中 blocks.pysam.pyvitdet.py 的源码揭示它们如何支撑视频分割与高分辨率注意力计算。读完本文,读者将能从算子层面理解 SAM 系列视觉 Transformer 的位置编码机制,并掌握每个函数在独立场景下的直接调用方法。

文件定位:SAM 家族共享的"无状态"工具层

utils.py 中的函数全部为纯函数式算子,不维护任何可学习参数,仅依赖 torchtorch.nn.functional(见 utils.py),因此它们可以被图像分割模型、视频分割模型以及 SAM3 的 ViT 编码器在不同位置独立复用。从源码看,该文件内的函数可分成四类能力:

能力域 函数 典型使用方
视频时序处理 select_closest_cond_framesget_1d_sine_pe sam.py 视频分割主干
二维轴向 RoPE init_t_xycompute_axial_cisreshape_for_broadcastapply_rotary_enc blocks.pysam3/vitdet.py 注意力层
窗口注意力 window_partitionwindow_unpartition 各 Transformer Block 前向
相对/绝对位置编码 get_rel_posadd_decomposed_rel_posget_abs_posconcat_rel_pos blocks.pysam3/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 位置编码一致;
  • 对归一化后的相位分别取 sincos 并在最后一维拼接。

在视频分割前向中,该函数用于为 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):

  1. theta ** (arange(0, dim, 4) / dim) 的倒数生成 x、y 两个方向各自的频率向量,维度均为 dim // 4
  2. torch.outer(t_x, freqs_x)torch.outer(t_y, freqs_y) 得到每个空间位置上的相位;
  3. 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_realflatten 还原实数张量;
  • 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_partitionwindow_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.pyblocks.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_wk_h == k_w),这符合 SAM3 ViT 的方形 patch 网格假设;
  • 可传入预计算的 relative_coords 索引表以跳过重复构造;
  • torch.eye 构造单位矩阵作为 K 的拼接块,使得点积结果中位置偏置项按索引正确对齐;
  • rescale=True(为与 PyTorch F.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_posget_abs_pos,对应其混合了"绝对位置 + 相对位置 + 窗口注意力 + 2D RoPE"的 ViT 设计;同时 set_imgsz 允许在推理时动态重算相对位置表与 RoPE 频率(vitdet.py)。

小结与工程启示

utils.py 的 12 个函数虽小,却对应着 SAM 家族在视频时序建模、高分辨率窗口注意力与位置编码上的三个核心工程决策:

  1. 位置的两种建模get_abs_pos 处理"全局绝对位置",compute_axial_cis / add_decomposed_rel_pos / concat_rel_pos 分别提供基于复数旋转与基于查表相加/拼接的相对位置方案,三者可并存于同一网络中并各司其职;
  2. 分辨率自适应的通用套路:相对位置表按需线性插值(get_rel_pos)、绝对位置按需平铺或双三次插值(get_abs_pos)、RoPE 频率按需缩放坐标(scale_pos),使模型可以摆脱训练分辨率的束缚;
  3. F.scaled_dot_product_attention 的兼容处理concat_rel_posrescaleapply_rotary_enc 对 MPS 复数 repeat 的规避,均是面向现代 PyTorch 算子与多后端部署的精细适配。

对想深入二次开发的研究者,直接阅读 utils.py 及其在 blocks.pysam.pysam3/vitdet.py 中的调用点,即可完整串起从底层算子到视频/图像分割模型前向的调用链。

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

项目优选

收起
kernelkernel
deepin linux kernel
C
33
18
ops-transformerops-transformer
本项目是CANN提供的transformer类大模型算子库,实现网络在NPU上加速计算。
C++
1.13 K
2.75 K
pytorchpytorch
作为 Ascend for PyTorch 社区的核心组件,TorchNPU 是昇腾专为 PyTorch 打造的深度学习适配插件,使 PyTorch 框架能够直接调用昇腾 NPU,为开发者提供昇腾 AI 处理器的超强算力。
Python
857
1.35 K
docsdocs
暂无描述
Markdown
897
5.8 K
kernelkernel
openEuler内核是openEuler操作系统的核心,既是系统性能与稳定性的基石,也是连接处理器、设备与服务的桥梁。
C
529
593
ops-nnops-nn
本项目是CANN提供的神经网络类计算算子库,实现网络在NPU上加速计算。
C++
916
1.83 K
jiuwenswarmjiuwenswarm
JiuwenSwarm 是一款基于openJiuwen开发的智能AI Agent,它能够将大语言模型的强大能力,通过你日常使用的各类通讯应用,直接延伸至你的指尖。
Python
3.58 K
1.01 K
ops-mathops-math
本项目是CANN提供的数学类基础计算算子库,实现网络在NPU上加速计算。
C++
1.35 K
1.46 K
cann-learning-hubcann-learning-hub
CANN 学习中心仓,支持在线互动运行、边学边练,提供教程、示例与优化方案,一站式助力昇腾开发者快速上手。
Jupyter Notebook
1.01 K
515
AscendNPU-IRAscendNPU-IR
AscendNPU-IR是基于MLIR(Multi-Level Intermediate Representation)构建的,面向昇腾亲和算子编译时使用的中间表示,提供昇腾完备表达能力,通过编译优化提升昇腾AI处理器计算效率,支持通过生态框架使能昇腾AI处理器与深度调优
C++
547
388