Transformers Sam3TrackerVideo 实战:面向视频的可提示对象分割与追踪
Transformers Sam3TrackerVideo 实战面向视频的可提示对象分割与追踪【免费下载链接】transformers Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers本文基于 Transformers 仓库中 SAM3 Tracker Video 模型文档 展开完整讲解Sam3TrackerVideoModel的视频可提示分割Promptable Visual Segmentation工作流从初始化推理会话、注入点/框/掩码提示到整段视频传播与流式推理并结合 模型实现、配置类 与 处理器 源码剖析记忆编码器、对象指针与设备管理等底层机制。一、模型定位SAM 3 的视频追踪器Sam3TrackerVideo 面向视频场景下的可提示视觉分割PVS输入可以是交互式的视觉提示点、框、掩码或文本模型针对每个提示追踪视频帧序列中同一个特定对象实例。它是 SAM2 Video 的更新版本保持相同的 API 但提供改进的性能与能力。其上游来自 Meta 的 SAM 3Segment Anything with Concepts工作该模型由一个图像级检测器与一个基于记忆memory-based的视频追踪器组成二者共享同一个 backbone并通过 presence head 解耦识别与定位。SAM3TrackerVideo 正是其中负责视频传播追踪的那一部分。该模型于 2025-11-19 由社区贡献者加入 Transformers仓库中对应模型目录为 src/transformers/models/sam3_tracker_video从源码结构看它与sam3_vision视觉 backbone配合使用通过AutoModel.from_config(config.vision_config)加载视觉编码器。模型预训练权重以 Hub 上的facebook/sam3检查点形式提供文档与测试均使用该检查点from transformers import Sam3TrackerVideoModel, Sam3TrackerVideoProcessor model Sam3TrackerVideoModel.from_pretrained(facebook/sam3, device_mapauto) processor Sam3TrackerVideoProcessor.from_pretrained(facebook/sam3)从 PreTrainedModel 基类定义 可见Sam3TrackerVideoPreTrainedModel声明了_supports_sdpa True与_supports_flash_attn True即同时支持 SDPA 与 FlashAttention 两种注意力后端这与文档页眉的 SDPA/FlashAttention 徽章一致。二、核心 API 总览Sam3TrackerVideo 的推理围绕**处理器管理会话 模型执行前向** 展开关键入口如下均可在 模型文档 的 autodoc 章节中找到对应条目API所属作用processor.init_video_sessionSam3TrackerVideoProcessor创建视频推理会话整段视频或流式processor.add_inputs_to_inference_session处理器向指定帧注入点/框/掩码提示model(inference_session, frame_idx...)Sam3TrackerVideoModel.forward单帧推理返回该帧所有对象的掩码model.propagate_in_video_iterator传播迭代器将已确认的对象逐帧传播到整个视频processor.post_process_masks后处理去 padding 并上采样到原始分辨率inference_session.reset_inference_session会话清空追踪状态与缓存开启新追踪Sam3TrackerVideoInferenceSession是贯穿全流程的状态容器。从 其构造实现 看它持有对象映射obj_id_to_idx/obj_idx_to_id把用户自定义的任意整数对象 ID 映射到内部索引obj_id_to_idx会在首次出现时自动登记新对象提示存储point_inputs_per_obj、mask_inputs_per_obj按对象 × 帧两级组织输出历史output_dict_per_obj区分cond_frame_outputs条件帧即有用户提示的帧与non_cond_frame_outputs被追踪帧帧数据整段视频以 dict 形式存入processed_frames源码注释说明这是为避免torch.cat带来的双重内存分配。会话还内置了设备分层管理inference_device计算、inference_state_device大张量状态如掩码特征、video_storage_device视频帧可分别指定。输出存储逻辑 会把小张量object_pointer、object_score_logits留在推理设备把大张量掩码、特征异步搬运到 state 设备读取时再自动移回——这对 GPU 显存受限场景很关键。视觉特征缓存则由 Sam3TrackerVideoInferenceCache 管理max_vision_features_cache_size默认 1控制缓存帧数上限超限时淘汰最旧帧避免重复跑视觉编码器。三、配置体系一个主配置 三个子配置Sam3TrackerVideoConfig配置源码采用主配置 子配置结构vision_config视觉 backbone 配置模型类型为sam3_vision_model含 FPN 特征图尺寸backbone_feature_sizesprompt_encoder_configSam3TrackerVideoPromptEncoderConfighidden_size256、image_size1008、patch_size14、mask_input_channels16、num_point_embeddings4等mask_decoder_configSam3TrackerVideoMaskDecoderConfigmlp_dim2048、num_hidden_layers2、num_multimask_outputs3、iou_head_depth3以及动态多掩码稳定性参数dynamic_multimask_stability_delta0.05、dynamic_multimask_stability_thresh0.98。主配置中最值得关注的是视频追踪专属参数它们直接决定了传播阶段的行为参数默认值含义num_maskmem7记忆掩码槽位数量即每帧可供后续帧注意力访问的记忆条数max_cond_frame_num4记忆注意力中参与的最大条件帧数由_select_closest_cond_frames选取最近的记忆帧max_object_pointers_in_encoder16编码器中可容纳的对象指针object pointer上限sigmoid_scale_for_mem_enc/sigmoid_bias_for_mem_enc20.0 / -10.0记忆编码器中对掩码概率做 sigmoid 前后的缩放/偏移enable_occlusion_spatial_embeddingTrue为目标被遮挡/消失的帧注入遮挡空间嵌入multimask_output_in_samTrue图像SAM阶段是否输出多掩码multimask_min_pt_num/multimask_max_pt_num0 / 1触发多掩码输出的点数区间memory_attention_num_layers/_hidden_size/_num_attention_heads4 / 256 / 1记忆注意力模块规模memory_attention_rope_theta/_rope_feat_sizes10000 / [72, 72]记忆注意力的 RoPE 旋转位置编码参数memory_encoder_output_channels64记忆编码器输出的通道数mem_dimmemory_fuser_num_layers/_embed_dim/_intermediate_dim2 / 256 / 1024记忆融合模块含 CXBlock 结构的规模image_size是一个联动属性其 setter 在修改时会同步更新prompt_encoder_config、vision_config并按patch_size重新计算三级 FPN 特征图尺寸与 RoPE 特征尺寸——因此修改分辨率后无需手工维护其他字段。配置也支持由三个子配置直接组装 from transformers import ( ... Sam3TrackerVideoConfig, ... Sam3TrackerVideoPromptEncoderConfig, ... Sam3TrackerVideoMaskDecoderConfig, ... Sam3TrackerVideoModel, ... ) # 以 facebook/sam3 风格初始化配置 configuration Sam3TrackerVideoConfig() # 用随机权重初始化模型 model Sam3TrackerVideoModel(configuration) # 也可由三个子配置组装 vision_config Sam3TrackerVideoVisionConfig() prompt_encoder_config Sam3TrackerVideoPromptEncoderConfig() mask_decoder_config Sam3TrackerVideoMaskDecoderConfig() config Sam3TrackerVideoConfig(vision_config, prompt_encoder_config, mask_decoder_config)四、实战基本视频追踪完整流程分四步加载视频帧 → 初始化会话 → 注入提示 → 传播。以下代码继承自 官方文档 的 Basic Video Tracking 示例from transformers import Sam3TrackerVideoModel, Sam3TrackerVideoProcessor import torch model Sam3TrackerVideoModel.from_pretrained(facebook/sam3, device_mapauto) processor Sam3TrackerVideoProcessor.from_pretrained(facebook/sam3) # 加载视频帧这里使用仓库自带的视频加载工具 from transformers.video_utils import load_video video_url https://huggingface.co/datasets/hf-internal-testing/sam2-fixtures/resolve/main/bedroom.mp4 video_frames, _ load_video(video_url) # 1. 初始化视频推理会话 inference_session processor.init_video_session( videovideo_frames, inference_devicedevice, ) # 2. 在第 0 帧点击选择目标 ann_frame_idx 0 ann_obj_id 1 points [[[[210, 350]]]] labels [[[1]]] processor.add_inputs_to_inference_session( inference_sessioninference_session, frame_idxann_frame_idx, obj_idsann_obj_id, input_pointspoints, input_labelslabels, ) # 3. 先在该帧做一次单帧分割可选也可以直接传播 outputs model( inference_sessioninference_session, frame_idxann_frame_idx, ) video_res_masks processor.post_process_masks( [outputs.pred_masks], original_sizes[[inference_session.video_height, inference_session.video_width]], binarizeFalse )[0] print(fSegmentation shape: {video_res_masks.shape}) # Segmentation shape: torch.Size([1, 1, 480, 854]) # 4. 传播到整个视频 video_segments {} for sam3_tracker_video_output in model.propagate_in_video_iterator(inference_session): video_res_masks processor.post_process_masks( [sam3_tracker_video_output.pred_masks], original_sizes[[inference_session.video_height, inference_session.video_width]], binarizeFalse )[0] video_segments[sam3_tracker_video_output.frame_idx] video_res_masks print(fTracked object through {len(video_segments)} frames) # Tracked object through 180 frames提示坐标的层级格式是理解这套 API 的关键。处理器对输入做了严格的嵌套校验见_validate_single_inputinput_points4 层嵌套[图像级, 对象级, 点级, 坐标(2)]例如[[[[210, 350]]]]表示1 帧上 1 个对象的 1 个点input_labels3 层嵌套[图像级, 对象级, 点级]1为正点击、0为负点击框的两个角点会被转换为内部标签2、3见 process_new_points_or_boxes_for_video_frame;input_boxes3 层嵌套[图像级, 框级, 坐标(4)]且由于模型限制不同图像级的框数量必须一致不允许 padding。坐标会由处理器统一归一化到target_size默认取图像处理器的尺寸配置_normalize_coordinates按原始(H, W)做线性缩放点 padding 值默认为-10point_pad_value归一化时会通过preserve_padding避免误缩放填充坐标。add_inputs_to_inference_session在入口处做了一组约束校验源码点与标签必须同时提供点、框、掩码三者至少提供其一掩码不能与点或框混合提供clear_old_inputs默认True控制是覆盖还是追加已有提示若追加框提示到已有点上会抛出错误因为框提示必须先于点提示给出。五、多对象追踪与交互式精修一次添加多个对象对象 ID 可以是任意整数处理器按对象级维度把提示批量分配源码 中逐对象切分input_points[:, idx]# 重置会话开启新追踪 inference_session.reset_inference_session() # 在第 0 帧同时添加两个对象 ann_frame_idx 0 obj_ids [2, 3] input_points [[[[200, 300]], [[400, 150]]]] # 两个对象各 1 个点批量 input_labels [[[1], [1]]] processor.add_inputs_to_inference_session( inference_sessioninference_session, frame_idxann_frame_idx, obj_idsobj_ids, input_pointsinput_points, input_labelsinput_labels, ) # 一次性获得两个对象在第 0 帧的掩码 outputs model( inference_sessioninference_session, frame_idxann_frame_idx, ) # 传播两个对象 video_segments {} for sam3_tracker_video_output in model.propagate_in_video_iterator(inference_session): video_res_masks processor.post_process_masks( [sam3_tracker_video_output.pred_masks], original_sizes[[inference_session.video_height, inference_session.video_width]], binarizeFalse )[0] video_segments[sam3_tracker_video_output.frame_idx] { obj_id: video_res_masks[i] for i, obj_id in enumerate(inference_session.obj_ids) } print(fTracked {len(inference_session.obj_ids)} objects through {len(video_segments)} frames) # Tracked 2 objects through 180 frames多对象场景下forward的实现揭示了性能设计单帧前向 按对象逐个运行单对象推理batch_size1因为不同对象的点击/掩码输入可以不同但记忆编码是跨对象批处理的——_batch_encode_memories把所有需要记忆编码的对象的高分辨率掩码拼成一个 batch 一次性送入记忆编码器再按对象切分回写。此外若某对象没有新提示且该帧已有条件帧输出会直接复用缓存掩码而不重算。任意帧追加点击精修追踪中途可以在任意帧对某个对象追加正/负点击来纠正# 在第 50 帧追加一个正点击精修对象 2 refine_frame_idx 50 ann_obj_id 2 points [[[[220, 280]]]] # 额外点 labels [[[1]]] # 正点击 processor.add_inputs_to_inference_session( inference_sessioninference_session, frame_idxrefine_frame_idx, obj_idsann_obj_id, input_pointspoints, input_labelslabels, ) # 携带新信息重新传播 video_segments {} for sam3_tracker_video_output in model.propagate_in_video_iterator(inference_session): video_res_masks processor.post_process_masks( [sam3_tracker_video_output.pred_masks], original_sizes[[inference_session.video_height, inference_session.video_width]], binarizeFalse )[0] video_segments[sam3_tracker_video_output.frame_idx] video_res_masks注意forward中有一个细节初始条件帧首次出现提示的帧 会强制reverseFalse即新提示帧总是作为新的前向起点这保证了追加点击后传播方向的正确性。六、传播迭代器与单帧前向的参数细节propagate_in_video_iterator实现的完整参数start_frame_idx传播起点。缺省时自动取最早存在提示输入的帧若从未对任何帧调用过forward则必须手动指定否则会抛出 Cannot determine the starting frame indexmax_frame_num_to_track最多追踪的帧数缺省时追踪到视频末尾reverse是否反向向前方帧传播start_frame_idx为 0 时反向传播为空操作show_progress_bar是否显示 tqdm 进度条。它逐帧调用model(inference_session, frame_idx..., reverse...)并yieldSam3TrackerVideoSegmentationOutput输出字段为object_ids本帧追踪中的对象 ID 列表、pred_masks模型分辨率下的掩码 logits形状(num_objects, num_masks, H, W)、object_score_logits对象是否存在的 logit、frame_idx。forward本身支持两种模式整段视频模式传frame_idx视频已存入会话frameNone流式模式传frame张量未提供frame_idx由会话的add_new_frame自动分配索引。若会话中还没有任何对象就传入新帧会直接报 No objects are provided for tracking; please add inputs first。forward还有run_mem_encoder参数默认True控制是否对预测掩码运行记忆编码器源码注释说明记忆编码器会跨对象批处理以提升效率。七、流式Streaming视频推理对实时场景会话初始化时不传video帧随到随处理。文档示例省略号处为按帧循环的省略部分# 流式会话不一次性提供视频 inference_session processor.init_video_session( inference_devicedevice, ) # 逐帧处理 for frame_idx, frame in enumerate(video_frames[:10]): inputs processor(imagesframe, devicedevice, return_tensorspt).to(model.device) if frame_idx 0: # 在第 0 帧注入点提示 processor.add_inputs_to_inference_session( inference_sessioninference_session, frame_idx0, obj_ids1, input_points[[[[210, 350], [250, 220]]]], input_labels[[[1, 1]]], original_sizeinputs.original_sizes[0], # 流式推理时必须提供 ) # 处理当前帧 sam3_tracker_video_output model(inference_sessioninference_session, frameinputs.pixel_values[0]) video_res_masks processor.post_process_masks( [sam3_tracker_video_output.pred_masks], original_sizesinputs.original_sizes, binarizeFalse )[0] print(fFrame {frame_idx}: mask shape {video_res_masks.shape})与整段模式的两个关键差异必须提供original_size因为会话初始化时没有视频可供推断video_height/video_widthprocess_new_points_or_boxes_for_video_frame 会在首个流式帧缺少original_size时直接抛错前向传frame张量而非frame_idx帧会经add_new_frame写入会话实现 会自动 squeeze 掉 4D 输入多余的 batch 维并落到video_storage_device。八、内部机制记忆、对象指针与遮挡从 模型构造 可以看出传播阶段的组件构成视觉编码器AutoModel 加载的sam3_vision_model输出三级 FPN 特征get_image_features会预先对 level 0/1 特征跑conv_s0/conv_s1投影避免每次点击重复计算源码注释 明确说明此意图记忆注意力Sam3TrackerVideoMemoryAttention4 层、带 RoPE 旋转位置编码以当前帧视觉特征 对象指针为 query对过去条件帧的记忆做注意力记忆编码器Sam3TrackerVideoMemoryEncoder把当前帧顶层视觉特征与预测掩码编码成mem_dim64通道的记忆特征。_encode_new_memory中有几个可对照配置理解的细节点击产生的掩码会先二值化再经sigmoid_scale_for_mem_enc(20.0)/sigmoid_bias_for_mem_enc(-10.0)缩放目标消失object_score_logits 0的帧会被加上occlusion_spatial_embedding_parameter以标记遮挡记忆特征最终以bfloat16存储以节省显存源码注释指出这是与原实现保持一致的做法对象指针object pointerobject_pointer_proj3 层前馈把 SAM 解码器输出 token 变成固定维度的对象标识配合temporal_positional_encoding_projection_layer注入时间位置编码对应enable_temporal_pos_encoding_for_object_pointers供记忆注意力区分不同历史帧中的同一对象无记忆占位no_memory_embedding/no_object_pointer等零初始化参数为该对象尚无历史的情况提供统一表示。post_process_masks本身是对图像处理器同名方法的直接转发源码支持mask_threshold、binarize、max_hole_area、max_sprinkle_area、apply_non_overlapping_constraints等参数。文档示例统一使用binarizeFalse保留 logit 值集成测试 中同样以binarizeFalse断言了 3x3 像素的精确 logit 数值atol1e-4可作为结果回归的参考。九、测试与验证路径集成测试位于 tests/models/sam3_tracker_video/test_modeling_sam3_tracker_video.py基于facebook/sam3检查点slow标记需要联网与 GPU 环境。以单点视频测试为例其断言链路可以当作最小验证用例outputs self.video_model(inference_sessioninference_session, frame_idxann_frame_idx) low_res_masks outputs.pred_masks self.assertEqual(low_res_masks.shape, (1, 1, 288, 288)) video_res_masks self.processor.post_process_masks( [low_res_masks], [raw_video.shape[-3:-1]], binarizeFalse )[0] self.assertEqual(video_res_masks.shape, (1, 1, raw_video.shape[-3], raw_video.shape[-2]))pred_masks的低分辨率形状为(1, 1, 288, 288)对象数 × 掩码数 × 模型内部分辨率post_process_masks后再恢复到原始视频分辨率测试使用与文档示例相同的测试视频bedroom.mp4与坐标(210, 350)并用max_frame_num_to_track2做短程传播验证。十、参考文件索引内容路径模型文档本文依据docs/source/en/model_doc/sam3_tracker_video.md模型实现会话、前向、记忆、传播迭代器src/transformers/models/sam3_tracker_video/modeling_sam3_tracker_video.py配置类主配置 提示编码器 掩码解码器src/transformers/models/sam3_tracker_video/configuration_sam3_tracker_video.py处理器会话初始化、提示注入、坐标归一化、掩码后处理src/transformers/models/sam3_tracker_video/processing_sam3_tracker_video.py视频加载工具load_videosrc/transformers/video_utils.py集成测试tests/models/sam3_tracker_video/test_modeling_sam3_tracker_video.py适用前提与限制以上流程以 PyTorch 环境、facebook/sam3检查点与 1008 输入分辨率的默认配置为前提流式推理必须提供original_size框提示不能 padding 且必须先于点提示给出掩码提示与点/框提示互斥。理解这些约束均可在处理器源码的校验逻辑中直接看到能显著减少 API 调用时的试错成本。【免费下载链接】transformers Transformers: the model-definition framework for state-of-the-art machine learning models in text, vision, audio, and multimodal models, for both inference and training.项目地址: https://gitcode.com/GitHub_Trending/tra/transformers创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
上一篇/下一篇内容由系统自动关联
返回资讯列表 →