
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 AnythingSAM 的三大件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 python3.9 -y conda activate sam-ft # 2. PyTorch按自己的 CUDA 版本选这里给通用写法 pip install torch torchvision # 3. Segment Anything 本体 pip install githttps://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 JSONLabelMe、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, size1024): 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_sizecfg.batch_size, shuffleTrue) opt torch.optim.AdamW(model.mask_decoder.parameters(), lrcfg.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_outputFalse) loss loss_fn(masks, gt) opt.zero_grad(); loss.backward(); opt.step() print(fepoch {epoch:02d} loss {loss.item():.4f}) if (epoch 1) % 10 0: torch.save(model.state_dict(), fcheckpoints/ft_{epoch1:02d}.pth)跑完几轮后如果 loss 平稳下降、验证指标在涨恭喜骨架通了。剩下的都是调的功夫。从能跑到跑得好分层微调策略与超参数调节一次放开所有参数是新手最容易踩的坑。推荐的节奏是分层解冻先让下游模块学会怎么接你的领域信号验证曲线稳定后再回头微调编码器。超参数别迷信最优值记住起始值 往哪个方向调就够了超参数起始值症状与调整方向优先关注学习率1e-4解冻编码器后 1e-5震荡/发散 → 减半学不动 → 翻倍最高批量大小2 ~ 8vit_b显存不足 → 减半并开混合精度太慢 → 加梯度累积中训练轮数30验证指标先升后降 → 提前停在最高点中权重衰减1e-4过拟合明显 → 升到 1e-3低输入边长1024小目标漏检 → 升到 1280显存换精度中它真的学会了吗评估指标白话版与性能对比光看 loss 下降不放心得用验证集说话。四个常用指标用一句话说清各自防什么mIoU平均交并比预测和真值的重叠部分占并集的比例。最主流的综合指标越大越好。Dice 系数对小目标更敏感的交并比变体。产线上的小裂纹、影像里的微病灶看它比 mIoU 更诚实。Precision精确率预测出来的像素里有多少是真的——防多切了。Recall召回率真目标里有多少被切到了——防漏切了。Precision 高 Recall 低 模型保守漏切反过来 激进乱扩。两个都低才是真没学会。微调前后的典型变化示意数据实际幅度取决于你的领域与数据量规格预训练 mIoU微调后 mIoU相对提升单图推理耗时vit_b0.750.8817%~45msvit_l0.780.9015%~80msvit_h0.810.9213%~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[], boxbox, multimask_outputFalse)部署前的自查清单微调 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),仅供参考