- AI 技能
- 人工智能
- 大模型
- 深度学习
【免费下载链接】AI-Research-SKILLs
Comprehensive open-source library of AI research and engineering skills for any AI model. Package the skills and your claude code/codex/gemini agent will be an AI research agent with full horsepower. Maintained by Orchestra Research.
本指南以 AI-Research-SKILLs 仓库中的 Segment Anything 技能文档(SKILL.md)为核心,系统讲解 Meta AI Segment Anything Model(SAM)的零样本图像分割能力:无需任务定制训练,即可用点、边界框或掩码提示分割任意图像中的任意物体。读完本文,你将掌握 SAM 的安装与模型加载、SamPredictor交互式分割、SamAutomaticMaskGenerator全自动掩码生成、ONNX 部署以及常见实战工作流,并能结合仓库中的 高级用法 与 故障排查 两篇参考文档完成端到端落地。
在 AI-Research-SKILLs 的技能体系中,本技能归属于18-multimodal/segment-anything/目录,是 技能路由地图 中多模态板块的重要成员。当自动研究编排器(autoresearch)在图像理解、数据标注、医学影像处理等研究任务中需要"分割任意物体"时,就会路由到本技能执行;仓库根 README.md 也将它列为 "Meta's SAM for zero-shot image segmentation with points/boxes" 的标准能力入口。
何时使用 SAM
适用场景
使用 SAM 的场景:
- 需要在无需任务定制训练的情况下分割图像中的任意物体
- 构建基于点 / 框提示的交互式标注工具
- 为其他视觉模型生成训练数据(数据飞轮)
- 需要向新的图像域进行零样本迁移
- 构建物体检测 / 分割流水线
- 处理医学、卫星或领域特定图像
关键特性
- 零样本分割:无需微调即可在任何图像域上工作
- 灵活提示:支持点(Points)、边界框(Bounding Boxes)或先前掩码(Previous Masks)
- 自动分割:自动生成图像中所有物体的掩码
- 高质量:官方在 1100 万张图像上使用 11 亿掩码进行训练(论文公开数据)
- 多种模型尺寸:ViT-B(最快)、ViT-L、ViT-H(最准确)
- ONNX 导出:可部署到浏览器与边缘设备
替代方案对比
考虑使用替代方案:
- YOLO / Detectron2:用于带类别的实时物体检测
- Mask2Former:用于带类别的语义 / 全景分割
- GroundingDINO + SAM:用于文本提示分割(组合管线见 高级用法)
- SAM 2:用于视频分割任务
快速开始
安装
SAM 官方提供 pip 安装入口,推荐安装segment-anything主包,并视需要安装opencv-python、pycocotools、matplotlib等可选依赖:
# 从 GitHub 安装 pip install git+https://github.com/facebookresearch/segment-anything.git # 可选依赖 pip install opencv-python pycocotools matplotlib # 或使用 HuggingFace transformers pip install transformers安装完成后可执行快速验证:
python -c "from segment_anything import sam_model_registry; print('OK')"下载 Checkpoint
官方发布的三个模型 checkpoint 分别对应三种主干规模,体积与精度依次递增:
# ViT-H(最大、最准确) - 2.4GB wget https://dl.fbaipublicfiles.com/segment_anything/sam_vit_h_4b8939.pth # ViT-L(中等) - 1.2GB wget https://dl.fbaipublicfiles.com/segment_anything/sam_vit_l_0b3195.pth # ViT-B(最小、最快) - 375MB wget https://dl.fbaipublicfiles.com/segment_anything/sam_vit_b_01ec64.pth注意:checkpoint 文件名中的
_h_4b8939、_l_0b3195、_b_01ec64等后缀是官方校验标识,下载后建议用md5sum校验文件完整性(如sam_vit_h_4b8939.pth的预期校验值为a7bf3b02f3ebf1267aba913ff637d9a2),并使用与模型类型完全匹配的 checkpoint,否则加载时会报unexpected key in state_dict。
使用 SamPredictor 的基础用法
SamPredictor是官方交互式推理的核心 API。其工作流是:加载模型 → 设置图像(一次性计算图像嵌入)→ 提供提示预测掩码:
import numpy as np from segment_anything import sam_model_registry, SamPredictor # 加载模型 sam = sam_model_registry"vit_h" sam.to(device="cuda") # 创建 predictor predictor = SamPredictor(sam) # 设置图像(一次性计算嵌入) image = cv2.imread("image.jpg") image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) predictor.set_image(image) # 用点提示进行预测 input_point = np.array([[500, 375]]) # (x, y) 坐标 input_label = np.array([1]) # 1 = 前景, 0 = 背景 masks, scores, logits = predictor.predict( point_coords=input_point, point_labels=input_label, multimask_output=True # 返回 3 个候选掩码 ) # 选择最佳掩码 best_mask = masks[np.argmax(scores)]predict返回三元组:masks(掩码数组,形状与multimask_output相关)、scores(每个掩码的预测 IoU 置信度)、logits(低分辨率掩码 logits,可回传用于迭代细化)。multimask_output=True时返回 3 个候选掩码供选择;当提示足够明确(如单个清晰目标)时可用False直接返回单一掩码。
HuggingFace Transformers 集成
除官方实现外,也可通过 HuggingFace 生态加载模型与处理器,适合与现有 Transformers 管线整合:
import torch from PIL import Image from transformers import SamModel, SamProcessor # 加载模型与处理器 model = SamModel.from_pretrained("facebook/sam-vit-huge") processor = SamProcessor.from_pretrained("facebook/sam-vit-huge") model.to("cuda") # 用点提示处理图像 image = Image.open("image.jpg") input_points = [[[450, 600]]] # 点需嵌套列表以包含 batch 维度 inputs = processor(image, input_points=input_points, return_tensors="pt") inputs = {k: v.to("cuda") for k, v in inputs.items()} # 生成掩码 with torch.no_grad(): outputs = model(**inputs) # 后处理:将掩码恢复到原始尺寸 masks = processor.image_processor.post_process_masks( outputs.pred_masks.cpu(), inputs["original_sizes"].cpu(), inputs["reshaped_input_sizes"].cpu() )需要注意:SamProcessor返回的pred_masks是在模型输入分辨率下的低分辨率掩码,必须通过post_process_masks结合original_sizes(原始尺寸)与reshaped_input_sizes(重采样尺寸)恢复到原图大小,否则无法直接叠加到原图上。
核心概念
模型架构
SAM 由三个可组合的模块构成:图像编码器(Image Encoder)与提示编码器(Prompt Encoder)将输入映射为嵌入,掩码解码器(Mask Decoder)基于两者预测掩码与 IoU 分数:
SAM Architecture: ┌─────────────────┐ ┌─────────────────┐ ┌─────────────────┐ │ Image Encoder │────▶│ Prompt Encoder │────▶│ Mask Decoder │ │ (ViT) │ │ (Points/Boxes) │ │ (Transformer) │ └─────────────────┘ └─────────────────┘ └─────────────────┘ │ │ │ Image Embeddings Prompt Embeddings Masks + IoU (computed once) (per prompt) predictions理解这一架构对性能优化至关重要:图像嵌入只计算一次,可被多个提示反复复用——这正是"一张图、多个提示"批量推理的高效基础;提示编码与掩码解码相对轻量,单次预测开销远小于图像编码。
模型变体
| 模型 | Checkpoint 注册名 | 大小 | 速度 | 精度 |
|---|---|---|---|---|
| ViT-H | vit_h | 2.4 GB | 最慢 | 最佳 |
| ViT-L | vit_l | 1.2 GB | 中等 | 良好 |
| ViT-B | vit_b | 375 MB | 最快 | 良好 |
提示类型
| 提示 | 描述 | 使用场景 |
|---|---|---|
| 点(前景) | 点击物体内部 | 单个物体选择 |
| 点(背景) | 点击物体外部 | 排除区域 |
| 边界框 | 物体外接矩形 | 较大物体 |
| 先前掩码 | 低分辨率掩码输入 | 迭代细化 |
交互式分割
点提示
单个前景点即可触发分割;多个点(含背景点)可进一步提高精度:
# 单个前景点 input_point = np.array([[500, 375]]) input_label = np.array([1]) masks, scores, logits = predictor.predict( point_coords=input_point, point_labels=input_label, multimask_output=True ) # 多个点(前景 + 背景) input_points = np.array([[500, 375], [600, 400], [450, 300]]) input_labels = np.array([1, 1, 0]) # 2 个前景, 1 个背景 masks, scores, logits = predictor.predict( point_coords=input_points, point_labels=input_labels, multimask_output=False # 提示清晰时返回单一掩码 )坐标约定:点坐标必须是(x, y)格式——x是列索引、y是行索引,且需落在图像边界内(0 <= x < w且0 <= y < h),否则会出现index out of bounds或掩码位置错误。
框提示
边界框格式为[x1, y1, x2, y2](左上 + 右下),适合框选较大物体:
# 边界框 [x1, y1, x2, y2] input_box = np.array([425, 600, 700, 875]) masks, scores, logits = predictor.predict( box=input_box, multimask_output=False )注意x1 < x2且y1 < y2是合法框的前提,否则会报invalid box coordinates。
组合提示
框与点可同时使用,兼顾整体范围与局部精确控制:
# 框 + 点实现精确控制 masks, scores, logits = predictor.predict( point_coords=np.array([[500, 375]]), point_labels=np.array([1]), box=np.array([400, 300, 700, 600]), multimask_output=False )迭代细化
SAM 支持把上一次预测的 logits 作为mask_input回传,在交互标注中实现"点一点、改一处"的连续细化:
# 初始预测 masks, scores, logits = predictor.predict( point_coords=np.array([[500, 375]]), point_labels=np.array([1]), multimask_output=True ) # 用先前掩码 + 额外点细化 masks, scores, logits = predictor.predict( point_coords=np.array([[500, 375], [550, 400]]), point_labels=np.array([1, 0]), # 添加背景点 mask_input=logits[np.argmax(scores)][None, :, :], # 使用最佳掩码 multimask_output=False )自动掩码生成
基础自动分割
SamAutomaticMaskGenerator在图像上铺设规则点网格,自动分割出全部物体掩码,无需任何人工提示:
from segment_anything import SamAutomaticMaskGenerator # 创建生成器 mask_generator = SamAutomaticMaskGenerator(sam) # 生成全部掩码 masks = mask_generator.generate(image) # 每个掩码包含: # - segmentation: 二值掩码 # - bbox: [x, y, w, h] # - area: 像素数量 # - predicted_iou: 质量分数 # - stability_score: 稳定性分数 # - point_coords: 生成点自定义生成参数
核心可调参数包括网格密度(points_per_side)、质量阈值(pred_iou_thresh)、稳定性阈值(stability_score_thresh)、多尺度裁剪(crop_n_layers)以及最小掩码面积(min_mask_region_area):
mask_generator = SamAutomaticMaskGenerator( model=sam, points_per_side=32, # 网格密度(越大掩码越多) pred_iou_thresh=0.88, # 质量阈值 stability_score_thresh=0.95, # 稳定性阈值 crop_n_layers=1, # 多尺度裁剪层数 crop_n_points_downscale_factor=2, min_mask_region_area=100, # 移除过小掩码 ) masks = mask_generator.generate(image)参数调节的经验方向:掩码过多时降低points_per_side、提高pred_iou_thresh/stability_score_thresh、增大min_mask_region_area,并可引入box_nms_thresh做更激进的 NMS;掩码过少或漏检小物体时,提高points_per_side、降低阈值、增加crop_n_layers多尺度裁剪(详见 故障排查)。
过滤掩码
生成结果可按面积、IoU 分数与稳定性分数做二次筛选:
# 按面积排序(大的在前) masks = sorted(masks, key=lambda x: x['area'], reverse=True) # 按预测 IoU 过滤 high_quality = [m for m in masks if m['predicted_iou'] > 0.9] # 按稳定性分数过滤 stable_masks = [m for m in masks if m['stability_score'] > 0.95]批量推理
多图像处理
对多张图像循环调用,注意每张图像需重新set_image(重新计算嵌入):
# 高效处理多张图像 images = [cv2.imread(f"image_{i}.jpg") for i in range(10)] all_masks = [] for image in images: predictor.set_image(image) masks, _, _ = predictor.predict( point_coords=np.array([[500, 375]]), point_labels=np.array([1]), multimask_output=True ) all_masks.append(masks)单图像多提示
得益于"嵌入只算一次"的架构,同一图像上的多个提示可以共享一次图像编码,批量预测非常高效:
# 高效处理多个提示(一次图像编码) predictor.set_image(image) # 批量点提示 points = [ np.array([[100, 100]]), np.array([[200, 200]]), np.array([[300, 300]]) ] all_masks = [] for point in points: masks, scores, _ = predictor.predict( point_coords=point, point_labels=np.array([1]), multimask_output=True ) all_masks.append(masks[np.argmax(scores)])ONNX 部署
导出模型
官方提供scripts/export_onnx_model.py导出脚本,可导出用于浏览器与边缘设备部署的 ONNX 模型:
python scripts/export_onnx_model.py \ --checkpoint sam_vit_h_4b8939.pth \ --model-type vit_h \ --output sam_onnx.onnx \ --return-single-mask--return-single-mask选项让解码器只输出单个掩码,推理更快(对多候选交互场景可去掉该参数);导出失败时可显式指定--opset 17,并固定onnx==1.14.0、onnxruntime==1.15.0等版本组合(详见 故障排查)。
使用 ONNX 模型推理
ONNX 模型仅包含提示编码器 + 掩码解码器,图像嵌入需在 PyTorch 侧预先计算后作为输入传入:
import onnxruntime # 加载 ONNX 模型 ort_session = onnxruntime.InferenceSession("sam_onnx.onnx") # 运行推理(图像嵌入需单独计算) masks = ort_session.run( None, { "image_embeddings": image_embeddings, "point_coords": point_coords, "point_labels": point_labels, "mask_input": np.zeros((1, 1, 256, 256), dtype=np.float32), "has_mask_input": np.array([0], dtype=np.float32), "orig_im_size": np.array([h, w], dtype=np.float32) } )若 GPU 上 ONNX Runtime 报错,可用onnxruntime.get_available_providers()查看可用 Provider,并显式指定providers=['CPUExecutionProvider']回退到 CPU。
常见工作流
工作流 1:交互式标注工具
用 OpenCV 鼠标回调实现"点击即分割"的标注体验:
import cv2 # 加载模型 predictor = SamPredictor(sam) predictor.set_image(image) def on_click(event, x, y, flags, param): if event == cv2.EVENT_LBUTTONDOWN: # 前景点 masks, scores, _ = predictor.predict( point_coords=np.array([[x, y]]), point_labels=np.array([1]), multimask_output=True ) # 显示最佳掩码 display_mask(masks[np.argmax(scores)])工作流 2:物体提取
点击物体内部即可输出带透明背景的 RGBA 抠图:
def extract_object(image, point): """提取点击点处的物体,输出透明背景 RGBA。""" predictor.set_image(image) masks, scores, _ = predictor.predict( point_coords=np.array([point]), point_labels=np.array([1]), multimask_output=True ) best_mask = masks[np.argmax(scores)] # 创建 RGBA 输出 rgba = np.zeros((image.shape[0], image.shape[1], 4), dtype=np.uint8) rgba[:, :, :3] = image rgba[:, :, 3] = best_mask * 255 return rgba工作流 3:医学图像分割
医学影像通常为灰度图,需先转成 RGB 三通道再送入模型,用 ROI 框提示分割感兴趣区域:
# 处理医学图像(灰度转 RGB) medical_image = cv2.imread("scan.png", cv2.IMREAD_GRAYSCALE) rgb_image = cv2.cvtColor(medical_image, cv2.COLOR_GRAY2RGB) predictor.set_image(rgb_image) # 分割感兴趣区域 masks, scores, _ = predictor.predict( box=np.array([x1, y1, x2, y2]), # ROI 边界框 multimask_output=True )输出格式
掩码数据结构
SamAutomaticMaskGenerator生成的每个掩码字典包含:
{ "segmentation": np.ndarray, # H×W 二值掩码 "bbox": [x, y, w, h], # 边界框 "area": int, # 像素数量 "predicted_iou": float, # 0-1 质量分数 "stability_score": float, # 0-1 稳定性分数 "crop_box": [x, y, w, h], # 生成裁剪区域 "point_coords": [[x, y]], # 输入点 }其中predicted_iou是模型对自身掩码质量的估计,stability_score衡量掩码对点扰动的鲁棒性,两者是自动化数据标注中的关键筛选依据。
COCO RLE 格式
与 COCO 数据集生态互通时,可将掩码编码为 RLE 格式存储:
from pycocotools import mask as mask_utils # 掩码编码为 RLE rle = mask_utils.encode(np.asfortranarray(mask.astype(np.uint8))) rle["counts"] = rle["counts"].decode("utf-8") # RLE 解码为掩码 decoded_mask = mask_utils.decode(rle)性能优化
GPU 内存
- 显存受限时改用 ViT-B 小模型
- 大批量处理之间调用
torch.cuda.empty_cache()清理缓存 - 超大图像先等比缩放到最长边 1024 以内再送入模型
# 显存有限时使用小模型 sam = sam_model_registry"vit_b" # 分批量处理图像,清理 CUDA 缓存 torch.cuda.empty_cache()速度优化
# 使用半精度 sam = sam.half() # 减少自动生成的网格点数 mask_generator = SamAutomaticMaskGenerator( model=sam, points_per_side=16, # 默认是 32 ) # 部署用 ONNX # 导出时加 --return-single-mask 加快推理常见问题速查
| 问题 | 解决方案 |
|---|---|
| 内存不足 | 使用 ViT-B 模型,缩小图像尺寸 |
| 推理缓慢 | 使用 ViT-B,减少 points_per_side |
| 掩码质量差 | 尝试不同提示,使用框 + 点组合 |
| 边缘伪影 | 使用 stability_score 过滤 |
| 小物体漏检 | 增大 points_per_side |
进阶主题:生产级集成与微调
以下内容来自 高级用法指南,适合把 SAM 接入真实系统。
SAM 2 视频分割
SAM 2 通过流式记忆(streaming memory)架构把 SAM 扩展到视频域,使用sam2包提供视频预测器:
from sam2.build_sam import build_sam2_video_predictor predictor = build_sam2_video_predictor("sam2_hiera_l.yaml", "sam2_hiera_large.pt") # 用视频初始化 predictor.init_state(video_path="video.mp4") # 在首帧添加提示 predictor.add_new_points( frame_idx=0, obj_id=1, points=[[100, 200]], labels=[1] ) # 在视频中传播 for frame_idx, masks in predictor.propagate_in_video(): # masks 包含所有跟踪物体的分割 process_frame(frame_idx, masks)SAM 与 SAM 2 的关键差异:SAM 仅支持图像输入(ViT + Decoder),无跨帧记忆与跟踪能力;SAM 2 支持图像 + 视频(Hiera + Memory 架构),通过流式记忆库实现跨帧物体跟踪,模型系列为 Hiera-T/S/B+/L。
Grounded SAM:文本提示分割
组合 GroundingDINO(文本 → 框)与 SAM(框 → 掩码)即可实现"输入一句文本、输出对应掩码":
from groundingdino.util.inference import load_model, predict from segment_anything import sam_model_registry, SamPredictor import cv2 # 加载 Grounding DINO grounding_model = load_model("groundingdino_swint_ogc.pth", "GroundingDINO_SwinT_OGC.py") # 加载 SAM sam = sam_model_registry"vit_h" predictor = SamPredictor(sam) def text_to_mask(image, text_prompt, box_threshold=0.3, text_threshold=0.25): """从文本描述生成掩码。""" # 从文本获取边界框 boxes, logits, phrases = predict( model=grounding_model, image=image, caption=text_prompt, box_threshold=box_threshold, text_threshold=text_threshold ) # 用 SAM 生成掩码 predictor.set_image(image) masks = [] for box in boxes: # 归一化框转像素坐标 h, w = image.shape[:2] box_pixels = box * np.array([w, h, w, h]) mask, score, _ = predictor.predict( box=box_pixels, multimask_output=False ) masks.append(mask[0]) return masks, boxes, phrases # 使用示例 image = cv2.imread("image.jpg") image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) masks, boxes, phrases = text_to_mask(image, "person . dog . car")批量处理封装
可将预测器封装为批处理类,统一处理"多图像 + 多提示"的常见任务:
class BatchedSAM: def __init__(self, checkpoint, model_type="vit_h", device="cuda"): self.sam = sam_model_registrymodel_type self.sam.to(device) self.predictor = SamPredictor(self.sam) self.device = device def process_batch(self, images, prompts): """用对应提示处理多张图像。""" results = [] for image, prompt in zip(images, prompts): self.predictor.set_image(image) if "point" in prompt: masks, scores, _ = self.predictor.predict( point_coords=prompt["point"], point_labels=prompt["label"], multimask_output=True ) elif "box" in prompt: masks, scores, _ = self.predictor.predict( box=prompt["box"], multimask_output=False ) results.append({ "masks": masks, "scores": scores, "best_mask": masks[np.argmax(scores)] }) return results并行处理多张图像时,每个线程需持有独立的模型实例(SAM 模型非线程安全共享),可用ThreadPoolExecutor配合多实例实现:
from concurrent.futures import ThreadPoolExecutor from segment_anything import SamAutomaticMaskGenerator def generate_masks_parallel(images, num_workers=4): """并行为多张图像生成掩码。""" # 注意:每个 worker 需要自己的模型实例 def worker_init(): sam = sam_model_registry"vit_b" return SamAutomaticMaskGenerator(sam) generators = [worker_init() for _ in range(num_workers)] def process_image(args): idx, image = args generator = generators[idx % num_workers] return generator.generate(image) with ThreadPoolExecutor(max_workers=num_workers) as executor: results = list(executor.map(process_image, enumerate(images))) return results服务化部署:FastAPI 与 Gradio
模型常驻内存(仅加载一次),通过 FastAPI 暴露点提示 / 自动分割接口,即可构建图像分割微服务:
from fastapi import FastAPI, File, UploadFile from pydantic import BaseModel import numpy as np import cv2 import io app = FastAPI() # 模型只加载一次 sam = sam_model_registry"vit_h" sam.to("cuda") predictor = SamPredictor(sam) class PointPrompt(BaseModel): x: int y: int label: int = 1 @app.post("/segment/point") async def segment_with_point( file: UploadFile = File(...), points: list[PointPrompt] = [] ): # 读取图像 contents = await file.read() nparr = np.frombuffer(contents, np.uint8) image = cv2.imdecode(nparr, cv2.IMREAD_COLOR) image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB) # 设置图像 predictor.set_image(image) # 准备提示 point_coords = np.array([[p.x, p.y] for p in points]) point_labels = np.array([p.label for p in points]) # 生成掩码 masks, scores, _ = predictor.predict( point_coords=point_coords, point_labels=point_labels, multimask_output=True ) best_idx = np.argmax(scores) return { "mask": masks[best_idx].tolist(), "score": float(scores[best_idx]), "all_scores": scores.tolist() }交互式标注界面的轻量方案是 Gradio:通过gr.SelectData捕获点击坐标,返回叠加了掩码的预览图:
import gradio as gr def segment_image(image, evt: gr.SelectData): """分割点击处的物体。""" predictor.set_image(image) point = np.array([[evt.index[0], evt.index[1]]]) label = np.array([1]) masks, scores, _ = predictor.predict( point_coords=point, point_labels=label, multimask_output=True ) best_mask = masks[np.argmax(scores)] # 在图像上叠加掩码 overlay = image.copy() overlay[best_mask] = overlay[best_mask] * 0.5 + np.array([255, 0, 0]) * 0.5 return overlay with gr.Blocks() as demo: gr.Markdown("# SAM Interactive Segmentation") gr.Markdown("Click on an object to segment it") with gr.Row(): input_image = gr.Image(label="Input Image", interactive=True) output_image = gr.Image(label="Segmented Image") input_image.select(segment_image, inputs=[input_image], outputs=[output_image]) demo.launch()微调 SAM
借助peft库可对 SAM 做参数高效的 LoRA 微调(实验性方案),把注意力层的qkv作为目标模块:
from peft import LoraConfig, get_peft_model from transformers import SamModel # 加载模型 model = SamModel.from_pretrained("facebook/sam-vit-base") # 配置 LoRA lora_config = LoraConfig( r=16, lora_alpha=32, target_modules=["qkv"], # 注意力层 lora_dropout=0.1, bias="none", ) # 应用 LoRA model = get_peft_model(model, lora_config) # 简化训练循环 optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4) for batch in dataloader: outputs = model( pixel_values=batch["pixel_values"], input_points=batch["input_points"], input_labels=batch["input_labels"] ) # 自定义损失(例如与真值掩码的 IoU 损失) loss = compute_loss(outputs.pred_masks, batch["gt_masks"]) loss.backward() optimizer.step() optimizer.zero_grad()领域化微调的典型代表是 MedSAM(医学影像微调版),将通用 checkpoint 替换为医学专用权重后,同样通过sam_model_registry加载,配合框提示完成 CT / 超声等影像的 ROI 分割。
掩码后处理
对模型输出可做形态学后处理,如闭运算填补孔洞、开运算去除噪点、填充内部空洞、移除过小连通域:
import cv2 from scipy import ndimage def refine_mask(mask, kernel_size=5, iterations=2): """用形态学运算细化掩码。""" kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (kernel_size, kernel_size)) # 闭运算填补小孔 closed = cv2.morphologyEx(mask.astype(np.uint8), cv2.MORPH_CLOSE, kernel, iterations=iterations) # 开运算去除小噪点 opened = cv2.morphologyEx(closed, cv2.MORPH_OPEN, kernel, iterations=iterations) return opened.astype(bool) def fill_holes(mask): """填充掩码孔洞。""" filled = ndimage.binary_fill_holes(mask) return filled def remove_small_regions(mask, min_area=100): """移除过小的不连通区域。""" labeled, num_features = ndimage.label(mask) sizes = ndimage.sum(mask, labeled, range(1, num_features + 1)) mask_clean = np.zeros_like(mask) for i, size in enumerate(sizes, 1): if size >= min_area: mask_clean[labeled == i] = True return mask_cleanTensorRT 加速
对 ONNX 模型可进一步转换为 TensorRT engine(支持 FP16)以获得 GPU 端极致推理性能:
import tensorrt as trt def export_to_tensorrt(onnx_path, engine_path, fp16=True): """将 ONNX 模型转换为 TensorRT engine。""" logger = trt.Logger(trt.Logger.WARNING) builder = trt.Builder(logger) network = builder.create_network(1 << int(trt.NetworkDefinitionCreationFlag.EXPLICIT_BATCH)) parser = trt.OnnxParser(network, logger) with open(onnx_path, 'rb') as f: if not parser.parse(f.read()): for error in range(parser.num_errors): print(parser.get_error(error)) return None config = builder.create_builder_config() config.max_workspace_size = 1 << 30 # 1GB if fp16: config.set_flag(trt.BuilderFlag.FP16) engine = builder.build_engine(network, config) with open(engine_path, 'wb') as f: f.write(engine.serialize()) return engine故障排查要点
完整排障手册见 故障排查指南,这里提炼最常遇到的几类问题。
环境与安装
RuntimeError: CUDA not available:先print(torch.cuda.is_available())与print(torch.version.cuda)检查 CUDA;按需用pip install torch torchvision --index-url https://download.pytorch.org/whl/cu121安装带 CUDA 的 PyTorch;模型需显式sam.to("cuda")。ModuleNotFoundError: No module named 'segment_anything':从 GitHub 安装,或git clone后pip install -e .。- 缺少
cv2/pycocotools等依赖:pip install opencv-python pycocotools matplotlib onnxruntime onnx;Windows 上pycocotools可换pycocotools-windows。
模型加载
- checkpoint 找不到:使用绝对路径,并用
md5sum校验文件完整性。 KeyError: 'unexpected key in state_dict':模型类型与 checkpoint 必须一一对应(vit_h↔sam_vit_h_4b8939.pth,vit_l↔sam_vit_l_0b3195.pth,vit_b↔sam_vit_b_01ec64.pth)。- 加载时 CUDA 内存溢出:改用 ViT-B;先
sam.to("cpu")再torch.cuda.empty_cache()后转 GPU;或sam.half()半精度。
推理
expected input to have 3 channels:统一转 RGB(cv2.COLOR_BGR2RGB),灰度图用COLOR_GRAY2RGB,RGBA 丢弃 alpha 通道(image[:, :, :3])。- 坐标越界 / 掩码位置错误:确认点是
(x, y)而非(row, col),并断言0 <= x < w and 0 <= y < h;框需满足x1 < x2 and y1 < y2。 - 掩码不匹配目标:加多前景点、加背景点、换框提示、框 + 点组合,并打印
scores取np.argmax。 - 推理慢:用 ViT-B、复用图像嵌入(
set_image一次多次predict)、降低points_per_side、ONNX 部署。
自动掩码生成
- 掩码过多:
points_per_side=16、pred_iou_thresh=0.92、stability_score_thresh=0.98、box_nms_thresh=0.5、min_mask_region_area=500。 - 掩码过少 / 漏检小物体:
points_per_side=64、降低阈值、crop_n_layers=2多尺度、min_mask_region_area=0;或将大图切块处理(patch_size=512, overlap=64)并偏移回原坐标。
内存
- CUDA 内存不足:小模型、逐图
torch.cuda.empty_cache()、顺序处理、图像等比缩放到最长边 1024 以内。 - RAM 内存不足:逐图处理并
del+gc.collect(),或改用生成器惰性产出结果。
常见错误速查
| 错误 | 原因 | 解决方案 |
|---|---|---|
CUDA out of memory | GPU 内存占满 | 使用小模型、清理缓存 |
expected 3 channels | 图像格式错误 | 转换为 RGB |
index out of bounds | 坐标非法 | 检查点 / 框边界 |
checkpoint not found | 路径错误 | 使用绝对路径 |
unexpected key | 模型与 checkpoint 不匹配 | 匹配模型类型 |
invalid box coordinates | x1 > x2 或 y1 > y2 | 修正框格式 |
在 AI-Research-SKILLs 中的定位
本技能是仓库多模态技能族(18-multimodal/)的组成部分。根据 技能路由文档 的路由原则:"当你遇到领域特定任务时,在技能库中搜索合适的工具,并在开始前阅读对应 SKILL.md——它包含工作流、常见问题与生产级代码示例"。
在实际研究流程中,autoresearch 编排器会在以下典型环节路由到本技能:
- 需要为下游检测 / 分割模型构建训练数据时,用
SamAutomaticMaskGenerator自动生成标注(可结合 数据标注生成示例 中的数据集生成代码); - 处理医学、卫星等特定域图像需要零样本分割时;
- 构建交互式标注或物体提取工具时。
配合仓库内其他技能可形成完整流水线:用 CLIP 做图文检索、用本技能做像素级分割、用 学术绘图 将掩码结果可视化到论文图表中。
参考资源
- 本技能主文档:SKILL.md
- 进阶用法(视频分割、Grounded SAM、服务化、微调、掩码后处理、TensorRT):advanced-usage.md
- 故障排查手册(安装、加载、推理、内存、ONNX、质量优化):troubleshooting.md
- 技能路由地图(多模态板块):skill-routing.md
- 官方公开资料(供自行检索):SAM 论文(arXiv 2304.02643)、Segment Anything 官方演示站、SAM 2 视频分割仓库、HuggingFace 上的
facebook/sam-vit-huge模型卡
- AI 技能
- 人工智能
- 大模型
- 深度学习
【免费下载链接】AI-Research-SKILLs
Comprehensive open-source library of AI research and engineering skills for any AI model. Package the skills and your claude code/codex/gemini agent will be an AI research agent with full horsepower. Maintained by Orchestra Research.
相关推荐
Kornia 中的 Segment Anything(SAM)实战:VisualPrompter 点/框提示分割全指南
Kornia 中的 Segment Anything(SAM)实战:VisualPrompter 点/框提示分割全指南 导读 :本文基于 Segment Any
计算机视觉人工智能深度学习图像处理Kornia 中 Segment Anything (SAM) 的提示式分割:VisualPrompter 与 Sam 模型实战指南
Kornia 中 Segment Anything SAM 的提示式分割:VisualPrompter 与 Sam 模型实战指南 Segment Anythin
计算机视觉深度学习人工智能图像处理ComfyUI Segment Anything 终极指南:用文本提示实现智能图像分割
ComfyUI Segment Anything 终极指南:用文本提示实现智能图像分割 想要通过简单的文本描述就能精确分割图像中的任何元素吗?🤔 ComfyU
人工智能计算机视觉AI 应用
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考