Ultralytics SAM Predictor 深度解析:从提示分割到全图自动分割的完整预测流程
2026/9/15 20:12:05 网站建设 项目流程

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 Model(SAM)的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.py),SAM 预测器会强制写入三项配置:

overrides.update(dict(task="segment", mode="predict", imgsz=1024)) super().__init__(cfg, overrides, _callbacks) self.args.retina_masks = True
  • task="segment"mode="predict"表明其只服务于分割推理场景;
  • imgsz=1024是 SAM 图像编码器的固定输入分辨率(与 build.py 中image_size = 1024一一对应);
  • retina_masks = True用于在结果可视化时保留高分辨率掩码细节,保证分割边缘质量。

这种"强制覆盖"意味着无论上层传入什么配置,SAM 推理都会被锁定在正确的任务与输入尺寸上,降低误用风险。

模型构建与设备分配:setup_model

setup_model(predict.py)负责把模型放到目标设备并完成归一化参数初始化:

device = select_device(self.args.device, verbose=verbose) 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 = Falseself.model.stride = 32self.model.fp16 = Falseself.done_warmup = True

模型构建由 build.py 的sam_model_map按权重文件名分派:

权重构建函数编码器
sam_h.ptbuild_sam_vit_hViT-H(embed_dim=1280, depth=32)
sam_l.ptbuild_sam_vit_lViT-L(embed_dim=1024, depth=24)
sam_b.ptbuild_sam_vit_bViT-B(embed_dim=768, depth=12)
mobile_sam.ptbuild_mobile_samTinyViT(embed_dims=[64,128,160,320])

若传入不支持的权重名,build_sam会抛出FileNotFoundError并列出可用模型。

预处理与输入变换

preprocess:图像的标准化流水线

preprocess(predict.py)支持torch.Tensor(BCHW)与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_transform:LetterBox 与单图限制

pre_transform(predict.py)使用LetterBox(self.args.imgsz, auto=False, center=False)将图像等比缩放到 1024 分辨率并填充至正方形,auto=False意味着不自动选择步长对齐。值得注意的是它带有硬性断言:

assert len(im) == 1, "SAM model does not currently support batched inference"

即当前 SAM 预测器不支持批处理推理,一次只能处理一张图像,这是由 SAM 的提示编码机制决定的。

提示分割:prompt_inference 的完整链路

inference(predict.py)是外部调用的入口,它会先从self.prompts弹出预设的 bboxes/points/masks;若三者皆为空,则转入全图自动分割generate();否则调用prompt_inference走提示分割路径。

prompt_inference(predict.py)完整展示了 SAM 的三段式架构(图像编码器 → 提示编码器 → 掩码解码器):

features = self.model.image_encoder(im) if self.features is None else self.features ... sparse_embeddings, dense_embeddings = self.model.prompt_encoder(points=points, boxes=bboxes, masks=masks) pred_masks, pred_scores = self.model.mask_decoder( image_embeddings=features, image_pe=self.model.prompt_encoder.get_dense_pe(), sparse_prompt_embeddings=sparse_embeddings, dense_prompt_embeddings=dense_embeddings, multimask_output=multimask_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 格式并乘r
  • masks:作为低分辨率提示输入(SAM 中 H=W=256),会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 背景;缺省全为 1
masks(N, H, W)前一轮低分辨率掩码 logits(H=W=256),支持迭代精化
multimask_outputbool为 True 时返回多个掩码(默认 False)

返回三元组:输出掩码CxHxW、每个掩码的质量分数C、以及供后续推理使用的低分辨率 logitsCxHxW

掩码阈值与结果生成:postprocess

postprocess(predict.py)接收(pred_masks, pred_scores[, pred_bboxes]),完成:坐标反缩放(ops.scale_boxes)、掩码缩放回原图尺寸(ops.scale_masks)、以self.model.mask_threshold二值化,最终组装为Results对象(ultralytics/engine/results.py),携带masksboxesnames元数据。全图模式下pred_bboxes还会拼入分数与类别索引(pred_scores[:, None]cls[:, None]),供下游使用;结束后self.segment_all复位为 False,避免污染下一次推理。

全图自动分割:generate 参数全解

当不提供任何提示时,inference调用generate(predict.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 阈值

执行流程分四步:

  1. 生成裁剪区域generate_crop_boxes(ultralytics/models/sam/amg.py)生成原图 + 各层裁剪块(XYWH 格式);
  2. 生成点网格build_all_layer_point_grids(amg.py)按层生成[0,1]×[0,1]均匀点阵,配合batch_iterator分批送入prompt_inference
  3. 逐裁剪块过滤:先按conf_thres过滤低质量掩码,再用calculate_stability_score(amg.py,计算高/低阈值二值掩码的 IoU)过滤不稳定掩码,随后移除贴近裁剪边缘但不贴近图像边缘的掩码(is_box_near_crop_edge),块内执行torchvision.ops.nms
  4. 跨裁剪块合并:将各块掩码通过uncrop_masks/uncrop_boxes_xyxy还原到全图坐标,若存在多个裁剪区域,则以scores = 1 / region_areas为权重再做一次 NMS(crop_nms_thresh)去除重复掩码。

该方法最终返回(pred_masks, pred_scores, pred_bboxes)三元组,分别对应分割掩码、置信度分数与边界框。

高效交互模式:set_image / set_prompts / reset_image

SAM 推理的最大开销在图像编码器。Predictor为此提供了"一次编码、多次提示"的模式:

from ultralytics.models.sam import Predictor as SAMPredictor # 创建 SAMPredictor(conf=0.25, task='segment', mode='predict', imgsz=1024) overrides = dict(conf=0.25, task='segment', mode='predict', imgsz=1024, model="mobile_sam.pt") predictor = SAMPredictor(overrides=overrides) # 设置图像:支持文件路径或 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_image(predict.py):未初始化模型时自动build_samsetup_model;随后setup_source(image)加载数据源并断言仅单张图像(assert len(self.dataset) == 1);最后运行一次preprocessimage_encoder,把结果分别缓存到self.imself.features
  • 之后每次predictor(bboxes=...)predictor(points=..., labels=...)调用都会经inferenceself.prompts弹出提示,并在prompt_inference跳过图像编码器,直接复用self.features
  • set_prompts(predict.py)允许预先批量注册提示字典(含bboxes/points/masks键),随后正常调用推理;
  • reset_image(predict.py)将self.imself.features置空,释放缓存。

这一模式在tests/test_cuda.pytest_predict_sam(tests/test_cuda.py)中有对应验证:加载sam_b.pt后依次执行全图推理、bbox 提示、点提示,再创建SAMPredictorset_image→ 提示推理 →reset_image的完整链路。

顶层调用方式与提示传入

除了直接实例化SAMPredictor,还可以通过 SAM 模型接口 以model(bboxes=..., points=..., labels=...)方式调用。其predict方法(model.py)同样强制conf=0.25, task="segment", mode="predict", imgsz=1024,并把提示打包为prompts字典传给BasePredictortask_map(model.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')(不传提示)
多次提示复用编码SAMPredictor+set_image+ 多次调用
全图分割增强参数predictor(source=..., crop_n_layers=1, points_stride=64)

相关文件索引

  • SAM 预测器核心实现:ultralytics/models/sam/predict.py
  • SAM 模型接口与task_map:ultralytics/models/sam/model.py
  • 模型构建与权重分派:ultralytics/models/sam/build.py
  • 全图分割辅助工具(点网格、裁剪、稳定性分数、NMS):ultralytics/models/sam/amg.py
  • 基础预测器(生命周期与回调框架):ultralytics/engine/predictor.py
  • 结果对象(掩码/框封装):ultralytics/engine/results.py
  • 测试用例(提示分割与 SAMPredictor 流程):tests/test_cuda.py
  • SAM 使用指南与模型对比: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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询