ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

Segment Anything 自定义训练:从标注到上线的 SAM 微调实操指南

Segment Anything 自定义训练:从标注到上线的 SAM 微调实操指南 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-anythingSAM 零样本能把图里的狗切出来但你的业务要的是这批零件里哪个有划痕。通用模型不懂你的语义只能靠 Segment Anything 自定义训练SAM 微调补上用自己的数据把预训练权重推到业务域。本文按数据 → 训练 → 评估 → 上线的顺序把每一步该做什么、坑在哪讲清楚。先把 SAM 跑起来环境要求不高python ≥ 3.8PyTorch 带 CUDA。要注意一点这个仓库只提供推理代码和示例 notebook没有训练脚本——训练代码得自己写但模型构建、预处理、ONNX 导出都是现成可 import 的。安装命令就三行装完先跑通一个 inference 再谈微调git clone https://gitcode.com/GitHub_Trending/se/segment-anything cd segment-anything pip install -e . pip install opencv-python pycocotools matplotlib onnxruntime onnxSAM 由三部分组成image encoderViT参数占比超过 95%是重的部分、prompt encoder、mask decoder只有几 M 参数是轻的部分也是微调主战场。模型构建入口在 sam_model_registry有 vit_b / vit_l / vit_h 三档预处理统一走 ResizeLongestSide把长边缩到 1024。数据怎么标、怎么喂标注格式推荐 COCO理由很实际pycocotools 生态成熟而且和官方 notebook 的读取习惯一致。标注内容必须是实例级 maskRLE 或 polygon因为 SAM 推理时不直接吃 mask你得把标注转成 prompt。SAM 微调数据集准备的核心就两件事标注 prompt 化、增强时 mask 与 prompt 同步。点提示mask 内采 k 个点作正点label1mask 外采 k 个作负点label0k 取 4~8 通常够用。box 提示直接用标注的 bbox适合形状规则的场景。数据集__getitem__的伪代码看懂这四步就行img cv2.imread(path) # 原始图像 x ResizeLongestSide(1024).apply_image(img) # 模型输入 pts, labels sample_prompts(rle_ann) # 标注 → 正/负提示点 gt_mask decode_rle(rle_ann) # 损失监督用 # 返回 dict: imagex, prompts(pts,labels), gt_mask, original_size增强策略别贪多这四条性价比最高增强推荐参数备注水平 / 垂直翻转各 50%mask 和 prompt 坐标必须同步翻转颜色抖动亮度 / 对比度 ±20%只动图像标注零成本随机旋转±30° 以内点坐标要重投影实现略繁琐随机裁剪保留 ≥80% 目标需重裁标注收益一般默认不开上图是官方 automatic mask generator 的输出标注时可以参考这种效果让 SAM 先自动粗切人工修边界比从零画快得多。分层微调先解码器再碰编码器image encoder 又贵又通用所以策略是先冻结它只训 prompt encoder mask decoder等验证集 mIoU 平台期了再考虑解冻。这就是 SAM 分层微调策略的全部。别急先别调参。SAM 学习率怎么设记住一句解码器可以用常规 lr编码器必须小一个量级。基线配置如下超参起点说明learning rate解码器 1e-4 / 编码器 1e-5编码器对 lr 敏感差一个量级就会崩weight decay1e-4常规值即可batch size2~41024 输入下显存是硬约束不够就上 AMP调度warmup 5% steps cosine长训练建议加优化器AdamW搭配 weight decay训练循环给到伪代码级for epoch in range(EPOCHS): for batch in loader: masks, iou model(batch[image], batch[prompts]) # multimask_outputTrue loss bce(masks, gt) iou_reg(iou, gt_iou) # 损失组合参考论文 loss.backward(); optimizer.step() if val_plateau and phase 1: unfreeze_encoder(); set_lr(1e-5) # 进入阶段二说明一下仓库本身不带训练 loss上面是按论文思路的参考写法模型前向细节看 modeling/sam.py。 练到什么程度算好看四个数不用贴代码理解含义就行指标看什么参考线mIoU预测 mask 与标注的整体重合度主指标业务可用一般 0.85Dice对边界更敏感比 IoU 严格比 mIoU 低 5 个点以内算正常Precision / Recall差值大说明边缘没学干净过分割或欠分割两者差 0.1 要回头查数据iou 预测值模型自己估的置信度线上拿它做阈值与真实 mIoU 相关性 0.9 算校准好看 SAM 训练 mIoU 提升的曲线时前几个 epoch 涨幅大是正常现象判断收敛看验证集平台期。微调前后的典型量级示例数字用于建立预期场景零样本 mIoU微调后备注通用物体0.780.80通用域本身已强提升有限工业小目标0.610.87典型收益场景纹理单一医疗类0.550.84提示点质量影响很大 上线导出 ONNX 与缓存 embeddingSAM ONNX 部署的关键不在导出在拆分。仓库自带 export_onnx_model.py会把 SAM 拆成 image encoder / prompt encoder / mask decoder 三个 ONNX 子模型推理封装在 segment_anything/utils/onnx.py用法和 SamPredictor 对齐。拆分的好处image embedding 是个 64×64×256 的张量算一次就能反复用之后每个 prompt 只跑 mask decoder。同一张图多次提示时推理开销差一个数量级。predictor SamPredictor(sam) predictor.set_image(img) # embedding 只算这一次 for p in prompts: masks, _, _ predictor.predict(pointsp) # 后续都只走 decoder⚠️ 踩坑速查这个坑我踩过不少直接给对照表现象原因解法loss 卡住不降lr 偏大先训的模块在震荡冻结编码器解码器从 5e-5 起试微调后通用能力掉点数据少却全参微调回分层策略编码器 lr ≤1e-5开翻转后 mask 整体错位只翻了图没翻 prompt 坐标增强同时作用于 mask 与 points显存 OOMbatch4 1024 偏贪batch2 梯度累积或开 AMP多掩码排序乱只训了 mask没训 iou head损失里保留 iou 回归项收尾还能往前走哪一步SAM 微调的内核就是八个字先冻重的先训轻的数据量上去了再逐步靠近全参。数据偏少时可以借 automatic_mask_generator.py 的思路让 SAM 自动生成提示样本扩充训练对。下一站是 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),仅供参考
返回列表