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仓库的最低要求见 READMEPyTorch torchvision≥ 1.7建议带 CUDA编码器是 12 层以上 ViTCPU 上训不动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 truthannotations.category_id类别 id方便分类别统计指标有一个坑要提前避开bbox 是提示segmentation 是答案两者必须来自同一条标注且 RLE 解码后要和图像尺寸对齐否则训练数字看着对、实际在学错位。数据增强微调集通常只有几千张图增强是用来防过拟合的。按需取用不用全上增强手段参数范围什么时候用随机翻转p0.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, epochs30, lr1e-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], lrlr) 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~8A1001024 输入的编码器很吃显存不够就减半并用梯度累积补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),仅供参考