Ultralytics SAM Predictor 深度解析:从提示分割到全图自动分割的完整预测流程

发布时间:2026/9/15 20:12:15
Ultralytics SAM Predictor 深度解析:从提示分割到全图自动分割的完整预测流程
Ultralytics SAM Predictor 深度解析从提示分割到全图自动分割的完整预测流程【免费下载链接】yolov10YOLOv10: Real-Time End-to-End Object Detection [NeurIPS 2024]项目地址: https://gitcode.com/GitHub_Trending/yo/yolov10本文以 SAM 预测器参考文档 为核心结合仓库源码 ultralytics/models/sam/predict.py 与配套模块系统讲解 Ultralytics 框架中 Segment Anything ModelSAM的Predictor类包括其继承关系、预处理与推理生命周期、bbox/point/mask 三种提示分割机制、全图自动分割generate()的全部可调参数以及set_image/set_prompts/reset_image的高效调用模式。读完本文你将能理解 SAM 预测器内部每个环节的实现细节并独立编写可复用的 SAM 提示分割与全图分割代码。Predictor 类的定位与设计Predictor定义在 ultralytics/models/sam/predict.py 中继承自 BasePredictor是 SAM 模型在 Ultralytics 框架内的推理接口。从源码结构看它的职责被明确定义为生成 Segment Anything 模型的分割预测提供可提示分割promptable segmentation与全图自动分割两类能力支持框bounding box、点point、低分辨率掩码low-resolution mask三种输入提示。类的公开属性在 docstring 中有清晰界定属性含义cfg模型与任务相关的配置字典overrides覆盖默认配置的字典_callbacks用户自定义回调函数集合args命令行参数或运行变量的命名空间im预处理后的输入图像张量features图像编码器提取的特征供推理复用prompts各类提示bboxes、points、masks的集合segment_all是否分割图像中全部对象全图模式标志构造器强制覆盖的任务配置在__init__中predict.pySAM 预测器会强制写入三项配置overrides.update(dict(tasksegment, modepredict, imgsz1024)) super().__init__(cfg, overrides, _callbacks) self.args.retina_masks Truetasksegment、modepredict表明其只服务于分割推理场景imgsz1024是 SAM 图像编码器的固定输入分辨率与 build.py 中image_size 1024一一对应retina_masks True用于在结果可视化时保留高分辨率掩码细节保证分割边缘质量。这种强制覆盖意味着无论上层传入什么配置SAM 推理都会被锁定在正确的任务与输入尺寸上降低误用风险。模型构建与设备分配setup_modelsetup_modelpredict.py负责把模型放到目标设备并完成归一化参数初始化device select_device(self.args.device, verboseverbose) if model is None: model build_sam(self.args.model) model.eval() self.model model.to(device) self.mean torch.tensor([123.675, 116.28, 103.53]).view(-1, 1, 1).to(device) self.std torch.tensor([58.395, 57.12, 57.375]).view(-1, 1, 1).to(device)这里mean/std是 SAM 预训练时使用的像素归一化参数与 build.py 中Sam模块的pixel_mean/pixel_std一致保证预处理与模型训练分布对齐。方法末尾还设置了若干 Ultralytics 兼容标记self.model.pt False、self.model.stride 32、self.model.fp16 False、self.done_warmup True。模型构建由 build.py 的sam_model_map按权重文件名分派权重构建函数编码器sam_h.ptbuild_sam_vit_hViT-Hembed_dim1280, depth32sam_l.ptbuild_sam_vit_lViT-Lembed_dim1024, depth24sam_b.ptbuild_sam_vit_bViT-Bembed_dim768, depth12mobile_sam.ptbuild_mobile_samTinyViTembed_dims[64,128,160,320]若传入不支持的权重名build_sam会抛出FileNotFoundError并列出可用模型。预处理与输入变换preprocess图像的标准化流水线preprocesspredict.py支持torch.TensorBCHW与List[np.ndarray]HWC两种输入。对 numpy 输入流程为先经pre_transform变换 → BGR 转 RGB → BHWC 转 BCHW → 转 torch.Tensor → 移到self.device→ 依据self.model.fp16选择 half/float 精度 → 最后用(im - mean) / std做 SAM 风格归一化注意与通用 YOLO 预测器的im / 255不同。方法开头有if self.im is not None: return self.im的缓存短路当通过set_image预先设置图像后后续推理直接复用已预处理图像避免重复计算。pre_transformLetterBox 与单图限制pre_transformpredict.py使用LetterBox(self.args.imgsz, autoFalse, centerFalse)将图像等比缩放到 1024 分辨率并填充至正方形autoFalse意味着不自动选择步长对齐。值得注意的是它带有硬性断言assert len(im) 1, SAM model does not currently support batched inference即当前 SAM 预测器不支持批处理推理一次只能处理一张图像这是由 SAM 的提示编码机制决定的。提示分割prompt_inference 的完整链路inferencepredict.py是外部调用的入口它会先从self.prompts弹出预设的 bboxes/points/masks若三者皆为空则转入全图自动分割generate()否则调用prompt_inference走提示分割路径。prompt_inferencepredict.py完整展示了 SAM 的三段式架构图像编码器 → 提示编码器 → 掩码解码器features self.model.image_encoder(im) if self.features is None else self.features ... sparse_embeddings, dense_embeddings self.model.prompt_encoder(pointspoints, boxesbboxes, masksmasks) pred_masks, pred_scores self.model.mask_decoder( image_embeddingsfeatures, image_peself.model.prompt_encoder.get_dense_pe(), sparse_prompt_embeddingssparse_embeddings, dense_prompt_embeddingsdense_embeddings, multimask_outputmultimask_output, )其内部处理细节如下特征缓存若已通过set_image提取过self.features则跳过昂贵的图像编码器直接复用特征坐标缩放提示坐标按r图像缩放比例换算到 1024 输入坐标系。r 1.0 if self.segment_all else min(...)全图模式因裁剪已对齐无需缩放points 规范化一维数组自动补为二维(N, 2)用户未传labels时默认全为前景正例np.ones(N)坐标统一乘r后 reshape 为(N, 1, 2)bboxes 规范化同样支持一维/二维输入要求 XYXY 格式并乘rmasks作为低分辨率提示输入SAM 中 HW256会unsqueeze(1)增加通道维输出展平(N, d, H, W)→(N*d, H, W)其中d为 1 或 3取决于multimask_output——多掩码输出有助于消解歧义提示。三种提示的参数规范参数形状说明bboxes(N, 4)或(4,)边界框XYXY 像素坐标points(N, 2)或(2,)提示点像素坐标labels(N,)点标签1 前景、0 背景缺省全为 1masks(N, H, W)前一轮低分辨率掩码 logitsHW256支持迭代精化multimask_outputbool为 True 时返回多个掩码默认 False返回三元组输出掩码CxHxW、每个掩码的质量分数C、以及供后续推理使用的低分辨率 logitsCxHxW。掩码阈值与结果生成postprocesspostprocesspredict.py接收(pred_masks, pred_scores[, pred_bboxes])完成坐标反缩放ops.scale_boxes、掩码缩放回原图尺寸ops.scale_masks、以self.model.mask_threshold二值化最终组装为Results对象ultralytics/engine/results.py携带masks、boxes与names元数据。全图模式下pred_bboxes还会拼入分数与类别索引pred_scores[:, None]、cls[:, None]供下游使用结束后self.segment_all复位为 False避免污染下一次推理。全图自动分割generate 参数全解当不提供任何提示时inference调用generatepredict.py实现分割图像中所有对象。其核心策略是网格点采样 多尺度裁剪 稳定性过滤 NMS 去重。完整参数表如下参数默认值作用crop_n_layers0额外裁剪层数第i层产生2**i_layer数量的裁剪块crop_overlap_ratio512/1500裁剪块间重叠比例后续层按比例递减crop_downscale_factor1每层点数采样侧的缩放因子point_gridsNone自定义点网格归一化到 [0,1]用于指定裁剪层points_stride32图像每侧采样的点数与point_grids互斥points_batch_size64每批并行处理的点数conf_thres0.88基于掩码质量分数的置信度过滤阈值stability_score_thresh0.95基于掩码稳定性的过滤阈值stability_score_offset0.95稳定性分数计算的偏移量crop_nms_thresh0.7裁剪块之间掩码去重的 NMS IoU 阈值执行流程分四步生成裁剪区域generate_crop_boxesultralytics/models/sam/amg.py生成原图 各层裁剪块XYWH 格式生成点网格build_all_layer_point_gridsamg.py按层生成[0,1]×[0,1]均匀点阵配合batch_iterator分批送入prompt_inference逐裁剪块过滤先按conf_thres过滤低质量掩码再用calculate_stability_scoreamg.py计算高/低阈值二值掩码的 IoU过滤不稳定掩码随后移除贴近裁剪边缘但不贴近图像边缘的掩码is_box_near_crop_edge块内执行torchvision.ops.nms跨裁剪块合并将各块掩码通过uncrop_masks/uncrop_boxes_xyxy还原到全图坐标若存在多个裁剪区域则以scores 1 / region_areas为权重再做一次 NMScrop_nms_thresh去除重复掩码。该方法最终返回(pred_masks, pred_scores, pred_bboxes)三元组分别对应分割掩码、置信度分数与边界框。高效交互模式set_image / set_prompts / reset_imageSAM 推理的最大开销在图像编码器。Predictor为此提供了一次编码、多次提示的模式from ultralytics.models.sam import Predictor as SAMPredictor # 创建 SAMPredictorconf0.25, tasksegment, modepredict, imgsz1024 overrides dict(conf0.25, tasksegment, modepredict, imgsz1024, modelmobile_sam.pt) predictor SAMPredictor(overridesoverrides) # 设置图像支持文件路径或 cv2 读取的 np.ndarray predictor.set_image(ultralytics/assets/zidane.jpg) results predictor(bboxes[439, 437, 524, 709]) # 框提示 results predictor(points[900, 370], labels[1]) # 点提示 # 重置图像与特征缓存 predictor.reset_image()源码行为如下set_imagepredict.py未初始化模型时自动build_sam并setup_model随后setup_source(image)加载数据源并断言仅单张图像assert len(self.dataset) 1最后运行一次preprocess与image_encoder把结果分别缓存到self.im与self.features之后每次predictor(bboxes...)或predictor(points..., labels...)调用都会经inference从self.prompts弹出提示并在prompt_inference中跳过图像编码器直接复用self.featuresset_promptspredict.py允许预先批量注册提示字典含bboxes/points/masks键随后正常调用推理reset_imagepredict.py将self.im与self.features置空释放缓存。这一模式在tests/test_cuda.py的test_predict_samtests/test_cuda.py中有对应验证加载sam_b.pt后依次执行全图推理、bbox 提示、点提示再创建SAMPredictor走set_image→ 提示推理 →reset_image的完整链路。顶层调用方式与提示传入除了直接实例化SAMPredictor还可以通过 SAM 模型接口 以model(bboxes..., points..., labels...)方式调用。其predict方法model.py同样强制conf0.25, tasksegment, modepredict, imgsz1024并把提示打包为prompts字典传给BasePredictor。task_mapmodel.py将 segment 任务映射到本文的Predictor类构成模型 → 预测器的完整闭环。两套 API 的对应关系如下场景API单次推理 提示model(zidane.jpg, bboxes[439, 437, 524, 709])单次推理 点提示model(zidane.jpg, points[900, 370], labels[1])全图自动分割model(path/to/image.jpg)不传提示多次提示复用编码SAMPredictorset_image 多次调用全图分割增强参数predictor(source..., crop_n_layers1, points_stride64)相关文件索引SAM 预测器核心实现ultralytics/models/sam/predict.pySAM 模型接口与task_mapultralytics/models/sam/model.py模型构建与权重分派ultralytics/models/sam/build.py全图分割辅助工具点网格、裁剪、稳定性分数、NMSultralytics/models/sam/amg.py基础预测器生命周期与回调框架ultralytics/engine/predictor.py结果对象掩码/框封装ultralytics/engine/results.py测试用例提示分割与 SAMPredictor 流程tests/test_cuda.pySAM 使用指南与模型对比docs/en/models/sam.md小结Predictor类以图像编码器 提示编码器 掩码解码器为骨架通过prompt_inference支持框、点、掩码三类提示的灵活组合通过generate实现全图自动分割并以set_image/set_prompts/reset_image提供特征级复用能力。理解它的预处理归一化、坐标缩放、多裁剪层 NMS 与稳定性过滤等实现细节是在实际项目中高效调优 SAM 推理性能与分割质量的关键。【免费下载链接】yolov10YOLOv10: Real-Time End-to-End Object Detection [NeurIPS 2024]项目地址: https://gitcode.com/GitHub_Trending/yo/yolov10创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考