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_h | 636M | 1280 维 × 32 层 | 数据多、精度要求高时用 |
| vit_l | 308M | 1024 维 × 24 层 | 精度与显存的折中 |
| vit_b | 91M | 768 维 × 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), gt3. 训练循环:冻结编码器,只练提示与解码
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_b | 0.75 | 0.88 | +17% | ~45ms |
| vit_l | 0.78 | 0.90 | +15% | ~80ms |
| vit_h | 0.81 | 0.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) |
收尾:关键收获与延伸方向
回顾一下这条路径,四步走:
- 选对体型:显存紧张从 vit_b 起步,数据多再上 vit_l / vit_h
- 喂对数据:COCO 标注 + bbox 中心点提示,标注质量决定上限
- 分层训练:先冻编码器练下游,验证稳定后再解冻、降学习率
- 闭环验证: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),仅供参考