Segment Anything 微调完整指南:让 SAM 在自己的数据集上学会领域分割
2026/8/30 14:14:38 网站建设 项目流程

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

把 Segment Anything(简称 SAM,"万物分割模型")下载下来跑第一张图时,效果往往很惊艳——点一下,就把目标抠了出来。但换成自己手头的数据就不一样了:医疗影像里的细微病灶被忽略,工业产线上的划痕漏检一大片,遥感图里的地物边界糊成一团。SAM 是在海量通用图片上预训练的,它"见过"的东西和你的垂直领域之间存在天然差距,这时候靠提示工程已经榨不出多少油水了。

本文是一份 SAM 自定义训练(微调)实操教程:用你自己领域的数据集对 SAM 做微调,让分割模型在你的场景里表现更稳。读完你可以:

  • ✅ 用白话看懂 SAM 的三大模块,知道微调时该动哪一块
  • ✅ 把领域数据整理成 SAM 能吃的 COCO 标注格式
  • ✅ 跑通一个最小可运行的微调骨架,并掌握分层微调策略
  • ✅ 看懂 mIoU、Dice 等指标,判断模型是不是真的学会了
  • ✅ 用 ONNX 导出和推理缓存,把微调后的模型部署上线

快速看懂 Segment Anything:SAM 的三大件

SAM 的代码就在 segment_anything/ 目录里,模型由三个模块拼成,源码入口在 segment_anything/build_sam.py:

  • 图像编码器:先把整张图"读完",压缩成一块图像嵌入。同一张图只算一次,后面无论点多少次提示都不用重算。
  • 提示编码器:把你点的点、画的框这类"指哪打哪"的信号,翻译成模型能理解的语言。
  • 掩码解码器:拿着图像嵌入 + 提示,吐出分割掩码,并顺手给每个掩码打个"我觉得像不像"的质量分。

官方给出的数据流示意图(图像 → 编码器 → 提示 → 解码 → 掩码):

交互式提示的效果长这样,绿框是"提示",蓝色区域是解码出来的掩码:

SAM 有三种规格,微调时先选对体型:

规格参数量编码器规模微调建议
vit_h636M1280 维 × 32 层数据多、精度要求高时用
vit_l308M1024 维 × 24 层精度与显存的折中
vit_b91M768 维 × 12 层显存紧张、想快速迭代首选

开工前的准备:SAM 微调环境搭建与依赖安装

微调 SAM 的依赖不多,核心就是 PyTorch 加上这个仓库本身。

# 1. 建独立环境,避免污染 conda create -n sam-ft python=3.9 -y conda activate sam-ft # 2. PyTorch(按自己的 CUDA 版本选,这里给通用写法) pip install torch torchvision # 3. Segment Anything 本体 pip install git+https://gitcode.com/GitHub_Trending/se/segment-anything # 4. 数据标注与预处理 pip install opencv-python pycocotools

建议的项目目录结构,训练产物和原始数据分开放:

sam_finetune/ ├── data/ │ ├── images/ # 领域图像 │ └── labels/ │ └── train.json # COCO 标注 ├── src/ │ ├── dataset.py # 数据加载 │ └── train.py # 训练入口 ├── checkpoints/ # 预训练权重 + 训练产出 └── configs/ # 训练参数

喂给模型什么样的数据:SAM 数据集格式与增强策略

标注格式:用 COCO 就行

微调 SAM 最省心的标注格式是 COCO JSON(LabelMe、CVAT 等工具都能导出)。每条标注包含一个 RLE 编码的多边形/掩码,外加一个外接框——后者恰好可以直接转成训练用的提示点。一个最小示例:

{ "images": [ {"id": 1, "file_name": "crack_0001.png", "width": 1024, "height": 768} ], "annotations": [ { "id": 1, "image_id": 1, "category_id": 1, "bbox": [412, 305, 180, 96], "area": 12400, "segmentation": {"size": [768, 1024], "counts": "RLE编码字符串"}, "iscrowd": 0 } ], "categories": [{"id": 1, "name": "crack"}] }

一张待标注的领域图片(本仓库 notebook 里就有一张这样的示例):

数据增强:按你的领域"对症下药"

增强别一上来就全套拉满,先加低风险、高收益的,效果不行再加码:

增强手段建议幅度适合的场景风险等级
随机水平/垂直翻转100% 概率各 50%目标方向不敏感(如工业缺陷)
亮度/对比度抖动±15% ~ 20%光照不稳定的采集设备
小角度旋转±10°卫星遥感、显微图像
随机缩放裁剪0.8 ~ 1.0目标尺度变化大
高斯模糊/噪声σ ≤ 2低质量图像为主高,小目标会受伤

⚠️ 一条经验:如果标注框/掩码精度有限,别用大幅度几何变换,标错的标签比噪声更毒。

让模型先跑起来:最小可运行微调骨架

下面是一个能跑通的最小骨架,只覆盖三块:训练配置、数据加载、训练循环。它刻意简化了提示的构造(每张图取第一条标注,用 bbox 中心当正提示点),先求"通",再求"对"。

1. 训练配置

class Config: model_type = "vit_b" # vit_b / vit_l / vit_h checkpoint = "checkpoints/sam_vit_b.pth" lr = 1e-4 batch_size = 2 epochs = 30 target_size = 1024 # ResizeLongestSide 的边长

2. 数据加载:COCO → 图像 + 提示 + 真值掩码

import cv2, torch from pycocotools.coco import COCO from pycocotools import mask as rle_mask from segment_anything.utils.transforms import ResizeLongestSide class DomainDataset(torch.utils.data.Dataset): def __init__(self, ann_file, img_dir, size=1024): self.coco = COCO(ann_file) self.img_dir = img_dir self.t = ResizeLongestSide(size) def __len__(self): return len(self.coco.imgs) def __getitem__(self, i): img_id = list(self.coco.imgs)[i] info = self.coco.loadImgs(img_id)[0] img = cv2.imread(f"{self.img_dir}/{info['file_name']}", cv2.IMREAD_COLOR) img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img = torch.from_numpy(self.t.apply_image(img)).permute(2, 0, 1).float() ann = self.coco.loadAnns(self.coco.getAnnIds(img_id))[0] x, y, w, h = ann["bbox"] point = torch.tensor([[x + w / 2, y + h / 2]]) # 正提示点 label = torch.tensor([1.0]) gt = torch.from_numpy(rle_mask.decode(ann["segmentation"])).float()[None] return img, (point, label), gt

3. 训练循环:冻结编码器,只练提示与解码

SAM 的图像编码器预训练得非常充分,第一版训练直接把它冻住,只更新提示编码器和掩码解码器,省显存还稳:

from segment_anything import sam_model_registry def train(cfg): model = sam_model_registrycfg.model_type.cuda() for p in model.image_encoder.parameters(): p.requires_grad = False # 冻结,第一版别动它 loader = torch.utils.data.DataLoader( DomainDataset("data/labels/train.json", "data/images"), batch_size=cfg.batch_size, shuffle=True) opt = torch.optim.AdamW(model.mask_decoder.parameters(), lr=cfg.lr) loss_fn = torch.nn.BCEWithLogitsLoss() for epoch in range(cfg.epochs): for images, (pts, labels), gt in loader: images, pts, labels, gt = map(lambda t: t.cuda(), (images, pts, labels, gt)) with torch.no_grad(): img_embed = model.image_encoder(images) # 冻结,不算梯度 img_pe = model.prompt_encoder.get_dense_pe() sparse_pe, dense_pe = model.prompt_encoder(points=(pts, labels)) masks, _ = model.mask_decoder( img_embed, img_pe, sparse_pe, dense_pe, multimask_output=False) loss = loss_fn(masks, gt) opt.zero_grad(); loss.backward(); opt.step() print(f"epoch {epoch:02d} loss {loss.item():.4f}") if (epoch + 1) % 10 == 0: torch.save(model.state_dict(), f"checkpoints/ft_{epoch+1:02d}.pth")

跑完几轮后,如果 loss 平稳下降、验证指标在涨,恭喜,骨架通了。剩下的都是"调"的功夫。

从能跑到跑得好:分层微调策略与超参数调节

一次放开所有参数是新手最容易踩的坑。推荐的节奏是分层解冻:先让下游模块学会"怎么接你的领域信号",验证曲线稳定后再回头微调编码器。

超参数别迷信"最优值",记住"起始值 + 往哪个方向调"就够了:

超参数起始值症状与调整方向优先关注
学习率1e-4(解冻编码器后 1e-5)震荡/发散 → 减半;学不动 → 翻倍最高
批量大小2 ~ 8(vit_b)显存不足 → 减半并开混合精度;太慢 → 加梯度累积
训练轮数30验证指标先升后降 → 提前停在最高点
权重衰减1e-4过拟合明显 → 升到 1e-3
输入边长1024小目标漏检 → 升到 1280,显存换精度

它真的学会了吗:评估指标白话版与性能对比

光看 loss 下降不放心,得用验证集说话。四个常用指标,用一句话说清各自防什么:

  • mIoU(平均交并比):预测和真值的重叠部分占并集的比例。最主流的综合指标,越大越好。
  • Dice 系数:对"小目标"更敏感的交并比变体。产线上的小裂纹、影像里的微病灶,看它比 mIoU 更诚实。
  • Precision(精确率):预测出来的像素里有多少是真的——防"多切了"。
  • Recall(召回率):真目标里有多少被切到了——防"漏切了"。

Precision 高 Recall 低 = 模型保守漏切;反过来 = 激进乱扩。两个都低才是真没学会。

微调前后的典型变化(示意数据,实际幅度取决于你的领域与数据量):

规格预训练 mIoU微调后 mIoU相对提升单图推理耗时
vit_b0.750.88+17%~45ms
vit_l0.780.90+15%~80ms
vit_h0.810.92+13%~130ms

规律很明显:预训练起点越高,微调空间越小。小数据场景下,vit_b 微调常常是性价比之王。

微调后的自动预测掩码效果(同一张图上的多目标叠加):

把微调后的模型用起来:ONNX 导出与推理加速

ONNX 导出

仓库自带导出脚本 scripts/export_onnx_model.py,支持单掩码输出和动态量化:

python scripts/export_onnx_model.py \ --checkpoint checkpoints/ft_30.pth \ --model-type vit_b \ --output sam_decoder.onnx \ --return-single-mask \ --quantize-out sam_decoder_quant.onnx

两个实用参数:

  • --return-single-mask:只输出最优掩码,高清图上能明显省时间;
  • --quantize-out:动态量化,CPU 推理提速明显,精度损失通常可忽略。

注意 ONNX 导出的是"提示编码器 + 掩码解码器"这部分,图像编码器仍留在 PyTorch 侧,这恰好是加速的关键。

推理缓存:一张图只编码一次

SAM 的设计红利就是"编码器算一次,提示随便点"。生产环境里同一张图往往要跑很多提示,务必复用图像嵌入:

from segment_anything import SamPredictor predictor = SamPredictor(model) predictor.set_image(img) # 图像嵌入在这里算一次并缓存 for box in detected_boxes: # 对每个检测框反复提示 masks, scores, _ = predictor.predict( point_coords=[], point_labels=[], box=box, multimask_output=False)

部署前的自查清单:

  • 微调 checkpoint 已保存到checkpoints/并验证可加载
  • 验证集指标(mIoU/Dice)已记录,留作上线基线
  • ONNX 单掩码版本导出成功,量化后精度复测通过
  • 同一图像的重复推理走缓存,不重复跑编码器
  • 显存不足时启用混合精度推理

踩过的坑:SAM 模型训练常见问题排查

现象第一嫌疑怎么处理
loss 原地不动把所有参数都冻住了 / 学习率为 0 或过小打印p.requires_grad检查;跑一次学习率扫描
训练 loss 降、验证指标涨不动数据量太少,过拟合了加增强、早停,或混入部分通用数据防遗忘
显存 OOM解冻编码器 + batch 过大先减 batch、开 AMP;仍不行再考虑梯度检查点
掩码边缘锯齿、细节糊输入分辨率低 / 真值掩码本身粗糙输入边长提到 1280;用多边形重标关键样本
小目标 Recall 一直上不去提示点落在了目标外检查提示点构造逻辑;对小目标改用框提示
微调后通用场景反而变差领域数据占比过高训练时按比例混入通用图片(如 1:3)

收尾:关键收获与延伸方向

回顾一下这条路径,四步走:

  1. 选对体型:显存紧张从 vit_b 起步,数据多再上 vit_l / vit_h
  2. 喂对数据:COCO 标注 + bbox 中心点提示,标注质量决定上限
  3. 分层训练:先冻编码器练下游,验证稳定后再解冻、降学习率
  4. 闭环验证:mIoU/Dice 看综合,Precision/Recall 判断"多切还是漏切"

接下来可以继续探索的方向:

  • 负提示点(把背景点也喂给模型)进一步提升边界精度
  • 接入自动标注流水线:让微调后的模型给新图打初标,人工只做修正,滚雪球扩数据
  • 尝试知识蒸馏,把 vit_h 的教师模型压成 vit_b 级学生模型
  • 跟进社区中分辨率更强的新一代分割模型,评估是否值得迁移你的微调经验

微调没有银弹,数据、策略、验证三件事做好,SAM 在你领域里的表现会有肉眼可见的变化。动手跑通第一个 epoch,比读完十篇教程更有用。

【免费下载链接】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),仅供参考

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

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

立即咨询