Segment Anything 自定义训练:从标注到上线的 SAM 微调实操指南
【免费下载链接】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
SAM 零样本能"把图里的狗切出来",但你的业务要的是"这批零件里哪个有划痕"。通用模型不懂你的语义,只能靠 Segment Anything 自定义训练(SAM 微调)补上:用自己的数据把预训练权重推到业务域。本文按数据 → 训练 → 评估 → 上线的顺序,把每一步该做什么、坑在哪讲清楚。
先把 SAM 跑起来
环境要求不高:python ≥ 3.8,PyTorch 带 CUDA。要注意一点,这个仓库只提供推理代码和示例 notebook,没有训练脚本——训练代码得自己写,但模型构建、预处理、ONNX 导出都是现成可 import 的。
安装命令就三行,装完先跑通一个 inference 再谈微调:
git clone https://gitcode.com/GitHub_Trending/se/segment-anything cd segment-anything && pip install -e . pip install opencv-python pycocotools matplotlib onnxruntime onnxSAM 由三部分组成:image encoder(ViT,参数占比超过 95%,是"重"的部分)、prompt encoder、mask decoder(只有几 M 参数,是"轻"的部分,也是微调主战场)。模型构建入口在 sam_model_registry,有 vit_b / vit_l / vit_h 三档;预处理统一走 ResizeLongestSide,把长边缩到 1024。
数据怎么标、怎么喂
标注格式推荐 COCO,理由很实际:pycocotools 生态成熟,而且和官方 notebook 的读取习惯一致。标注内容必须是实例级 mask(RLE 或 polygon),因为 SAM 推理时不直接吃 mask,你得把标注转成 prompt。
SAM 微调数据集准备的核心就两件事:标注 prompt 化、增强时 mask 与 prompt 同步。
- 点提示:mask 内采 k 个点作正点(label=1),mask 外采 k 个作负点(label=0),k 取 4~8 通常够用。
- box 提示:直接用标注的 bbox,适合形状规则的场景。
数据集__getitem__的伪代码,看懂这四步就行:
img = cv2.imread(path) # 原始图像 x = ResizeLongestSide(1024).apply_image(img) # 模型输入 pts, labels = sample_prompts(rle_ann) # 标注 → 正/负提示点 gt_mask = decode_rle(rle_ann) # 损失监督用 # 返回 dict: image=x, prompts=(pts,labels), gt_mask, original_size增强策略别贪多,这四条性价比最高:
| 增强 | 推荐参数 | 备注 |
|---|---|---|
| 水平 / 垂直翻转 | 各 50% | mask 和 prompt 坐标必须同步翻转 |
| 颜色抖动 | 亮度 / 对比度 ±20% | 只动图像,标注零成本 |
| 随机旋转 | ±30° 以内 | 点坐标要重投影,实现略繁琐 |
| 随机裁剪 | 保留 ≥80% 目标 | 需重裁标注,收益一般,默认不开 |
上图是官方 automatic mask generator 的输出,标注时可以参考这种效果:让 SAM 先自动粗切,人工修边界,比从零画快得多。
分层微调:先解码器,再碰编码器
image encoder 又贵又通用,所以策略是:先冻结它,只训 prompt encoder + mask decoder;等验证集 mIoU 平台期了,再考虑解冻。这就是 SAM 分层微调策略的全部。
别急,先别调参。SAM 学习率怎么设?记住一句:解码器可以用常规 lr,编码器必须小一个量级。基线配置如下:
| 超参 | 起点 | 说明 |
|---|---|---|
| learning rate | 解码器 1e-4 / 编码器 1e-5 | 编码器对 lr 敏感,差一个量级就会崩 |
| weight decay | 1e-4 | 常规值即可 |
| batch size | 2~4 | 1024 输入下显存是硬约束,不够就上 AMP |
| 调度 | warmup 5% steps + cosine | 长训练建议加 |
| 优化器 | AdamW | 搭配 weight decay |
训练循环给到伪代码级:
for epoch in range(EPOCHS): for batch in loader: masks, iou = model(batch["image"], batch["prompts"]) # multimask_output=True loss = bce(masks, gt) + iou_reg(iou, gt_iou) # 损失组合参考论文 loss.backward(); optimizer.step() if val_plateau and phase == 1: unfreeze_encoder(); set_lr(1e-5) # 进入阶段二说明一下:仓库本身不带训练 loss,上面是按论文思路的参考写法,模型前向细节看 modeling/sam.py。
📊 练到什么程度算好
看四个数,不用贴代码,理解含义就行:
| 指标 | 看什么 | 参考线 |
|---|---|---|
| mIoU | 预测 mask 与标注的整体重合度,主指标 | 业务可用一般 >0.85 |
| Dice | 对边界更敏感,比 IoU 严格 | 比 mIoU 低 5 个点以内算正常 |
| Precision / Recall | 差值大说明边缘没学干净(过分割或欠分割) | 两者差 >0.1 要回头查数据 |
| iou 预测值 | 模型自己估的置信度,线上拿它做阈值 | 与真实 mIoU 相关性 >0.9 算校准好 |
看 SAM 训练 mIoU 提升的曲线时,前几个 epoch 涨幅大是正常现象,判断收敛看验证集平台期。微调前后的典型量级(示例数字,用于建立预期):
| 场景 | 零样本 mIoU | 微调后 | 备注 |
|---|---|---|---|
| 通用物体 | 0.78 | 0.80 | 通用域本身已强,提升有限 |
| 工业小目标 | 0.61 | 0.87 | 典型收益场景 |
| 纹理单一(医疗类) | 0.55 | 0.84 | 提示点质量影响很大 |
📦 上线:导出 ONNX 与缓存 embedding
SAM ONNX 部署的关键不在导出,在拆分。仓库自带 export_onnx_model.py,会把 SAM 拆成 image encoder / prompt encoder / mask decoder 三个 ONNX 子模型,推理封装在 segment_anything/utils/onnx.py,用法和 SamPredictor 对齐。
拆分的好处:image embedding 是个 64×64×256 的张量,算一次就能反复用,之后每个 prompt 只跑 mask decoder。同一张图多次提示时,推理开销差一个数量级。
predictor = SamPredictor(sam) predictor.set_image(img) # embedding 只算这一次 for p in prompts: masks, _, _ = predictor.predict(points=p) # 后续都只走 decoder⚠️ 踩坑速查
这个坑我踩过不少,直接给对照表:
| 现象 | 原因 | 解法 |
|---|---|---|
| loss 卡住不降 | lr 偏大,先训的模块在震荡 | 冻结编码器,解码器从 5e-5 起试 |
| 微调后通用能力掉点 | 数据少却全参微调 | 回分层策略,编码器 lr ≤1e-5 |
| 开翻转后 mask 整体错位 | 只翻了图,没翻 prompt 坐标 | 增强同时作用于 mask 与 points |
| 显存 OOM | batch=4 @1024 偏贪 | batch=2 + 梯度累积,或开 AMP |
| 多掩码排序乱 | 只训了 mask,没训 iou head | 损失里保留 iou 回归项 |
收尾:还能往前走哪一步
- SAM 微调的内核就是八个字:先冻重的,先训轻的;数据量上去了再逐步靠近全参。
- 数据偏少时,可以借 automatic_mask_generator.py 的思路让 SAM 自动生成提示样本,扩充训练对。
- 下一站是 SAM 2(视频分割),训练思路与本文一致,这套经验可以直接复用。
【免费下载链接】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),仅供参考