Segment Anything SAM 微调完整指南:3 个阶段跑通只懂你领域的分割模型
【免费下载链接】segment-anythingThe repository provides code for running inference with the SegmentAnything Model (SAM), links for downloading the trained model checkpoints, and example notebooks that show how to use the model.项目地址: https://gitcode.com/GitHub_Trending/se/segment-anything
通用版 Segment Anything(SAM)什么都能分,但放到工业零件、医疗影像这类领域图上,掩码边缘经常打飘。这份指南基于 segment-anything 仓库,带你走一遍在自己数据集上微调的完整链路:先建立可对比的基线,再准备标注数据,最后分层微调。全程不追求理论推导,只关心每步做完之后"长什么样算对了"。
如果本地没有代码,先执行git clone https://gitcode.com/GitHub_Trending/se/segment-anything拿到仓库。核心代码都在 segment_anything/ 目录,演示样例在 notebooks/。
原理速览:三个模块各管一件事
理解 SAM 的机制,微调时才知道该动哪里。它的推理链路是三段式:
- 图像编码器:像一个先把整张图"看熟"的摄影师,把图像压缩成一份 64×64 的图像 embedding。同一张图只算一次,之后所有提示都基于它。
- 提示编码器:把你给的点、框转换成和 embedding 同一"语言"的向量。
- 掩码解码器:把图像 embedding 和提示向量拼起来,吐出候选掩码和每个掩码的可信度分数。
这个结构决定了微调的核心逻辑:图像编码器负责"看懂像素",提示编码器+掩码解码器负责"听懂指令"。领域适配通常只需要后者学你的标注风格,所以标准打法是冻结编码器、只训轻量的两个小模块,显存省一大截。三个模块的源码分别在 modeling/image_encoder.py、modeling/prompt_encoder.py、modeling/mask_decoder.py。
🧭 上手三阶段
阶段一:微调前,先跑通基线
目标:拿到"未微调 SAM"在你领域上的表现存档,后面所有效果对比都依赖它。
动作:用仓库自带的 predictor_example.ipynb 思路,最小调用只有几行:
import numpy as np from segment_anything import SamPredictor, sam_model_registry sam = sam_model_registry"vit_b" predictor = SamPredictor(sam) predictor.set_image(image_rgb) masks, scores, logits = predictor.predict( point_coords=np.array([[400, 300]]), point_labels=np.array([1]), multimask_output=False, )从你的测试集里挑 30~50 张有代表性的图,逐张记录:点提示能否框住目标、scores分布如何、掩码和人工标注的重合度(IoU)。notebooks/ 里的三个 notebook 分别演示了点框提示、稠密预测和 ONNX 推理,照着改就行。
做对了长什么样:一张基线表——每张图一个 IoU 均值和一个典型 badcase 截图。此时你应该能直观说出 SAM 在你领域"错在哪":是边缘毛糙、小目标漏检、还是把相邻零件连成一片。
阶段二:准备领域数据集
目标:把 badcase 变成训练样本,标注格式对齐 COCO。
动作:
- 按 train/val 切分,比例 8:2 起步。val 集里的图必须和基线评估用同一批,否则数字没法比。
- 标注用 COCO 格式:每张图给
file_name、width、height;每个实例给bbox(x, y, w, h)、segmentation(多边形或 RLE)和iscrowd。用 pycocotools 读写,不要自己造格式。 - 数量门槛:单一目标类别 500 张以上才有稳定提升的迹象;只有几十张时,微调大概率过拟合,优先靠阶段一的提示技巧补救。
- 数据增强只做几何类(随机裁剪、翻转、小角度旋转 ≤30°),别上重度色彩变换——SAM 的图像归一化参数是写死的(
pixel_mean=[123.675, 116.28, 103.53],见 build_sam.py),色彩分布漂移会伤到预训练特征。
做对了长什么样:用 pycocotools 能无损读回全部标注;随机抽 20 张把掩码叠回原图肉眼检查,没有错位、漏标。
阶段三:分层微调
目标:先训轻量模块,再决定是否动编码器。
动作,分两轮走:
第一轮冻结图像编码器,只训提示编码器 + 掩码解码器。损失用BCEWithLogitsLoss对掩码 logits 计算,配合 IoU 分支的回归损失(掩码解码器本身会输出 IoU 预测,见 mask_decoder.py)。超参从默认值起步:AdamW,学习率 1e-4,权重衰减 1e-4,batch 4,跑 20~30 个 epoch。
第二轮视验证集情况决定:若还有明显提升空间,解冻编码器,学习率降到 1e-5 量级再跑 5~10 个 epoch。
做对了长什么样:验证集 IoU 曲线平稳上行、没有出现训完 val 反超 train 的过拟合形态;最后一轮保存的 checkpoint 在阶段一的 30 张测试图上跑分高于基线。
⚖️ 选型与取舍
三个模型版本怎么选,看 build_sam.py 里的注册表就能确认:
| 版本 | 参数量 | 什么时候选它 |
|---|---|---|
| vit_b | 91M | 微调首选,单卡 24G 显存可跑 batch 4 |
| vit_l | 308M | vit_b 调不动上限后的升级项 |
| vit_h | 636M | 对精度极端敏感且推理资源充足 |
关键超参数给推荐值,附一句理由:
- 学习率 1e-4(编码器 1e-5):解码器是轻模块可以大一点,编码器预训练特征贵,大了就毁掉。
- batch 4:图像输入 1024 分辨率,显存是第一约束,不够就降到 2 + 梯度累积。
- 权重衰减 1e-4:默认值即可,微调场景别折腾。
- 训练轮数 20~30:轻量模块收敛快,更多轮数主要贡献过拟合。
- pred_iou_thresh 0.88:官方默认,低于该分数的掩码直接丢弃,调低它只影响召回不影响精度。
✅ 效果验证:怎么判断微调真的成功了
别只看训练曲线,用三组数字交叉验证:
- 同图对比:拿阶段一那 30~50 张测试图,微调前后各跑一遍,算平均 IoU。有效微调的门槛:均值提升 ≥5 个百分点,且 badcase 类别(阶段一里说错的那类)改善最明显。如果均值没动、只有个别图变好,多半是数据量不够。
- 分数校准:微调后
scores分布应该和掩码真实质量更对齐——高分掩码更准、低分掩码更差。抽查 20 个高分预测,若边缘仍然明显错位,说明模型在"自信地错",回查标注一致性。 - 稠密预测兜底:用 SamAutomaticMaskGenerator 对一张 val 图跑全图掩码,看掩码是否碎片化、是否把目标切碎。这一步抓的是点提示测试发现不了的回归问题。
下图是稠密掩码的典型效果,你可以用它校准自己"分得对"的视觉标准:
⚠️ 避坑清单
损失不降原因:学习率过高,或第一轮就解冻了编码器。 处理:先把学习率降到 1e-5 观察 100 step,确认编码器确实处于冻结状态再逐步加回。
显存溢出(OOM)原因:vit_h 的图像编码器 + 1024 输入 + batch 偏大。 处理:换 vit_b、batch 降到 1 再上梯度累积,图像编码器支持只算一次后复用,别每个 prompt 都重编码。
微调后还不如预训练原因:数据量不足或标注噪声大,过拟合抹掉了通用能力。 处理:砍到第一轮配置(冻结编码器)、加早停、用 500 张门槛重新评估数据规模。
掩码边缘锯齿、抖动原因:解码后的掩码分辨率低于原图,上采样后边缘失真。 处理:检查transform的缩放链路,确认坐标映射回原图;对精度敏感的场景用 ONNX 导出时加--return-single-mask走高分辨率分支,见 export_onnx_model.py。
🔭 延伸方向
- 推理加速:把微调好的模型用 scripts/export_onnx_model.py 导出 ONNX,配合 onnx_model_example.ipynb 的加载方式接入线上服务,图像 embedding 可以按图缓存,二次提示几乎零成本。
- 交互式前端:demo/ 里自带一个基于 ONNX 的网页 Demo,把你的 checkpoint 塞进去,就能给同事做个"点一下出掩码"的内网演示。
下一步动作:今天就从阶段一开始——挑 30 张你领域里最刁钻的图,跑基线、录 badcase。基线表有了,后面所有调参决策都有了标尺。
【免费下载链接】segment-anythingThe repository provides code for running inference with the SegmentAnything Model (SAM), links for downloading the trained model checkpoints, and example notebooks that show how to use the model.项目地址: https://gitcode.com/GitHub_Trending/se/segment-anything
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考