SAM 自定义训练:4 步在自己的领域数据上微调分割模型
2026/8/30 21:12:34 网站建设 项目流程

SAM 自定义训练:4 步在自己的领域数据上微调分割模型

【免费下载链接】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 自定义训练为主题,用 4 步走完从数据准备、分层微调到评估部署的全过程。

先弄清训练什么:模型边界

这一节解决「SAM 到底该改哪部分」的问题。写训练代码之前先看分工,才知道微调的力气往哪使。SAM 由三部分组成,入口在 segment_anything/modeling/sam.py:

  • 图像编码器(ViT 主干):把整张图压成一张嵌入图,参数占大头,承载通用视觉特征。
  • 提示编码器:把点、框这类提示变成解码器能消费的形式。
  • 掩码解码器:参数最轻,真正负责生成掩码。

上图来自官方 predictor 示例:绿框加红星点提示车轮,掩码刚好覆盖。当你发现「提示位置对了、掩码形状却不对」时,需要修的大多是解码器,而不是编码器。

所以微调原则是:先冻结图像编码器,只训提示编码器和掩码解码器。编码器像毛坯房,解码器像精装修——精装修重做成本低,毛坯拆了再盖代价大。验证集 mIoU 不再涨时,再解冻编码器、用小学习率整体训练。

环境与数据,一次备齐

环境依赖和数据标注可以并行准备,两者就绪再开训。

环境依赖

依赖版本说明
Python≥ 3.8仓库的最低要求,见 README
PyTorch + torchvision≥ 1.7,建议带 CUDA编码器是 12 层以上 ViT,CPU 上训不动
opencv-python、pycocotools最新稳定版读图、COCO 标注解析与 RLE 解码
onnx、onnxruntime最新稳定版后面导出 ONNX 需要
git clone https://gitcode.com/GitHub_Trending/se/segment-anything cd segment-anything && pip install -e . pip install opencv-python pycocotools onnx onnxruntime

数据标注格式

推荐用 COCO 格式,train / val 各一个 JSON。关键字段:

字段含义
images.file_name图像文件名
annotations.bbox物体框 [x, y, w, h],可直接当框提示用
annotations.segmentationRLE 编码掩码,训练损失与评估的 ground truth
annotations.category_id类别 id,方便分类别统计指标

有一个坑要提前避开:bbox 是提示,segmentation 是答案,两者必须来自同一条标注,且 RLE 解码后要和图像尺寸对齐,否则训练数字看着对、实际在学错位。

数据增强

微调集通常只有几千张图,增强是用来防过拟合的。按需取用,不用全上:

增强手段参数范围什么时候用
随机翻转p=0.5默认开启,成本几乎为零
随机旋转±15°~30°目标方向多变时,如卫星图、工业零件
颜色抖动亮度/对比度 ±20%光照差异大时,如工业现场
随机缩放裁剪保留 0.8~1.0目标尺度差异大时
高斯噪声σ≈0.01图像本身带传感器噪声,如医学影像

注意:旋转和缩放必须同步变换掩码和框,否则 ground truth 就错位了。

微调一轮怎么跑通

这一节解决「训练到底怎么落地」。先说清楚:这个仓库只提供推理代码,没有官方训练脚本,下面的训练循环需要自己写。模型构建入口是 build_sam.py,预处理复用 transforms.py 里的ResizeLongestSide,保证训练和推理的输入分布一致。

整体流程是一条时间线:启动 → 训练 → 收敛:

数据集类

Dataset 的核心动作:读图、做和推理相同的预处理、从标注里取框提示、解出 ground truth 掩码。

class SamFinetuneDataset(Dataset): def __init__(self, ann_file, img_dir): self.coco, self.img_dir = COCO(ann_file), img_dir self.img_ids = list(self.coco.imgs.keys()) self.transform = ResizeLongestSide(1024) # 与推理一致 def __getitem__(self, idx): info = self.coco.imgs[self.img_ids[idx]] img = cv2.imread(os.path.join(self.img_dir, info["file_name"]))[..., ::-1] anns = self.coco.loadAnns(self.coco.getAnnIds(info["id"])) return { "image": to_tensor(self.transform.apply_image(img)), "boxes": torch.tensor([a["bbox"] for a in anns]), "gt_masks": decode_rle_masks(anns), # pycocotools 解码 RLE }

训练循环(分层冻结)

掩码头输出 logits,损失一般取BCE + Dice组合:BCE 管像素级的背景前景平衡,Dice 直接优化重叠程度。

def train(model, loader, val_loader, epochs=30, lr=1e-4): model.image_encoder.eval() # 第一阶段: 冻结编码器 [p.requires_grad_(False) for p in model.image_encoder.parameters()] opt = torch.optim.AdamW( [p for p in model.parameters() if p.requires_grad], lr=lr) for epoch in range(epochs): for batch in loader: loss = mask_loss(model, batch) # BCE + Dice opt.zero_grad(); loss.backward(); opt.step() if evaluate(model, val_loader) >= target_iou: unfreeze_encoder(model) # 第二阶段: 小 lr 整体训

超参数推荐值

参数推荐值选择理由
学习率解码器 1e-4,编码器 1e-5解码器是轻微调可以快;编码器有预训练基础,步子必须小
批大小4~8(A100)1024 输入的编码器很吃显存,不够就减半并用梯度累积补
Warmup总轮数的 5%起步阶段学习率缓慢爬升,防止前几步把解码器打乱
总轮数30~50微调收敛快,每轮都验证,早停省卡时

怎么判断训练是否有效

这一节解决「训练曲线看着对不对」。主要靠两个数字加一次肉眼抽查:

  • mIoU(平均交并比):预测掩码与真值掩码逐物体求交集并集比,再取平均。它直接回答「形状画得准不准」。
  • Dice 系数:同一件事的另一种算法,对「掩码差一点」的情况更敏感,常和 mIoU 一起看,两者走势不一致时要警惕标注有问题。

视觉上,抽几张验证图看掩码边缘。最常见的假象是「数字涨了、边缘还是毛糙」,多数源于数据量不足或标注误差,先查标注再怀疑模型。

上图是官方 自动掩码示例 的效果。微调成功的标志是:同类图像上,掩码边缘更贴合、漏检更少、假块更少。

模型版本微调前 mIoU*微调后 mIoU*推理耗时*
ViT-B0.710.86约 45 ms/图
ViT-L0.740.90约 80 ms/图
ViT-H0.780.92约 130 ms/图

* 表中数字为示例数据,实际提升取决于你的数据规模和标注质量。

把模型用起来

这一节解决「微调完怎么部署得更快」。SAM 的推理开销大头在图像编码器,解码器很轻——官方支持把解码器单独导出 ONNX,微调后的权重可以直接用这条通道:运行python scripts/export_onnx_model.py --checkpoint <你的权重> --model-type vit_b --output sam_decoder.onnx即可,详见 scripts/export_onnx_model.py。

另一个思路是缓存图像嵌入:同一张图编码一次,之后的所有提示都跳过编码器,适合「一张图、多轮提示」的交互场景:

# 同一张图编码一次, 后续提示直接复用 cache = {} def predict(sam, image, box): key = hash(image) if key not in cache: cache[key] = sam.image_encoder(preprocess(image)) return sam.mask_decoder(cache[key], encode_box(box))

部署优化清单,按投入产出排:

  • ✅ 混合精度(AMP)训练与推理,显存和耗时都省一半左右
  • 批处理推理,摊薄编码器开销
  • 高分辨率图像开单掩码输出(导出时加--return-single-mask),省上采样时间
  • ONNX 解码器做 int8 量化,掩码质量影响有限
  • NVIDIA 卡上把解码器换成 TensorRT 引擎

踩坑速查

训练微调 SAM 翻车大多集中在这几类,对着排查能省很多时间:

现象可能原因处理办法
损失不下降或剧烈震荡学习率偏高;或编码器解冻太早学习率减半;退回第一阶段只训解码器
训几轮后验证 mIoU 掉头向下过拟合,微调集太小加数据增强(见上一节表格)、早停、清洗验证集
训练 OOM1024 输入 + 批大小过大批大小减半 + 梯度累积;开混合精度
推理比预训练还慢整模型在跑,或上采样太贵导出解码器 ONNX;加嵌入缓存
⚠️ 指标好看但掩码整体错位框与掩码标注不同步抽 20 条标注人工核对,重点查 RLE 解码与坐标系

指标正常、掩码却错位是最难定位的一类,基本都出在标注管线上,别在模型上死磕。


SAM 自定义训练的路径到这里就闭环了:弄清模型边界、备齐环境与标注、分层冻结跑通训练、用 mIoU 验证效果。实践中最大的变量往往不是训练代码,而是标注质量,动手前花半小时核对 50 条标注,比调任何超参数都值。跑通一遍之后,可以接着试这三个方向:

  • 模型压缩:把 ViT-H 的嵌入蒸馏给 ViT-B,推理成本降一半
  • 多模态提示:框、点、预掩码组合输入,专治单提示搞不定的难样本
  • 跟进新一代分割模型:官方已推出支持图像的 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),仅供参考

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

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

立即咨询