ARTICLE DETAIL

建站实战干货

来自一线的建站与推广经验沉淀,每一条都经过真实交付验证。

6 步跑通 SAM 微调:从自有数据集到 ONNX 部署的完整实践

2026/8/30 9:00:00 拓冰建站 浏览量
6 步跑通 SAM 微调:从自有数据集到 ONNX 部署的完整实践 6 步跑通 SAM 微调从自有数据集到 ONNX 部署的完整实践【免费下载链接】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打开 SAMSegment Anything的演示页面点几下鼠标掩码就出来了。可换成你自己的医疗片、工业图或卫星图边界就开始发飘、漏边——预训练模型的通用在垂直领域往往不够用。这时候正确的动作不是从头训练而是在自己标注的数据上做一次 SAM 微调。下面的内容全部对齐 segment-anything 仓库的真实接口跟着做就能跑。读完你可以拿到COCO 标注怎么写、怎么变成模型能吃的提示点带最小示例与仓库Sam.forward接口一致的最小训练循环分阶段解冻策略和一组可直接起步的超参ONNX 导出命令与推理加速 checklist 快速认识三个部件各干什么先看全局再动手才不会迷路。SAM 由重和轻两部分组成图像编码器只跑一次提示编码器掩码解码器每次预测都跑——微调时冻结谁、训练谁就决定了显存和速度。三档规格配置来自 segment_anything/build_sam.py型号编码器维度层数 / 头数输入尺寸参数量级什么时候选它vit_hdefault128032 / 161024~636M精度优先显存够vit_l102424 / 161024~308M精度与速度折中vit_b76812 / 121024~91M微调起步首选微调建议一律从vit_b开始收敛快、单卡即可跑。动手前环境与目录这一步只花十分钟但能避免后面九成ModuleNotFoundError。# 1) 环境python3.8PyTorch1.7强烈建议带 CUDA conda create -n sam python3.9 -y conda activate sam pip install torch torchvision # 2) 克隆并安装项目 git clone https://gitcode.com/GitHub_Trending/se/segment-anything.git cd segment-anything pip install -e . # 3) 可选依赖RLE 输出 / ONNX 导出 / 示例 notebook 需要 pip install opencv-python pycocotools matplotlib onnx onnxruntimesegment-anything/ ├── segment_anything/ # 核心推理代码 │ ├── modeling/ # 图像编码器 / 提示编码器 / 掩码解码器 │ ├── utils/ # ResizeLongestSide、RLE 编解码、ONNX 封装 │ ├── build_sam.py # vit_h/l/b 构建函数 sam_model_registry │ ├── predictor.py # SamPredictorset_image 一次predict 多次 │ └── automatic_mask_generator.py # 全图自动出掩码 ├── notebooks/ # predictor / amg / onnx 三个用法示例 ├── scripts/ # amg.py 命令行生成、export_onnx_model.py 导出 └── demo/ # 浏览器里跑 ONNX 掩码解码的网页演示下载对应型号的 checkpoint 后sam_model_registryvit_b就能加载后面所有代码都基于这两个入口。 喂数据COCO 标注怎么写微调的质量上限由标注决定而 COCO 格式是 SAM 生态的原生格式——自动掩码生成的输出segment_anything/automatic_mask_generator.py本身就是 COCO RLE标注与推理格式天然对齐。最小可运行示例segmentation.counts为 RLE 编码用pycocotools.mask.decode还原成二值图{ images: [{id: 1, file_name: 0001.jpg, width: 1024, height: 768}], annotations: [{ id: 1, image_id: 1, category_id: 1, bbox: [120, 88, 340, 460], area: 156400, segmentation: {size: [768, 1024], counts: RLE...}, iscrowd: 0 }], categories: [{id: 1, name: weld_bead}] }注意bbox是 XYWH左上角 x、y 宽高segmentation.size是 [高, 宽]两者别写反。标注完成后用预训练模型跑一遍 scripts/amg.py 做自动预标注人工只修边界效率能翻几倍数据增强对微调的效果分级SAM 官方训练只用翻转领域数据可适度加码增强手段建议幅度解决什么问题效果水平/垂直翻转必开方向鲁棒性零成本⭐⭐⭐⭐⭐随机裁剪保留 0.8~1.0 面积尺度变化⭐⭐⭐⭐随机旋转±30°摆位变化注意掩码同步转⭐⭐⭐亮度/对比度±20%光照波动⭐⭐⭐高斯噪声σ≈0.01成像噪声⭐⭐跑起来配置、数据集、训练循环这是核心环节。原则不动仓库代码只写一个外部训练脚本调用它现有的forward接口。超参起步值ConfigCFG dict( model_typevit_b, # 先用小模型验证流水线 checkpointsam_vit_b_01ec64.pth, # 预训练权重路径 long_side1024, # ResizeLongestSide 对齐编码器输入 lr1e-4, weight_decay1e-4, epochs50, warmup3, freeze_encoderTrue, # 阶段一先冻结图像编码器 batch_size4, )写一个 SAM 能读的数据集关键就两件事图像走ResizeLongestSide提示坐标走apply_coords两者共用同一变换坐标才不会错位。import cv2, torch from torch.utils.data import Dataset from pycocotools.coco import COCO from segment_anything.utils.transforms import ResizeLongestSide class DomainDataset(Dataset): def __init__(self, ann, img_dir, long_side1024): self.coco, self.dir COCO(ann), img_dir self.ids list(self.coco.imgs) self.tr ResizeLongestSide(long_side) # 长边缩到 1024 def __len__(self): return len(self.ids) def __getitem__(self, i): info self.coco.loadImgs(self.ids[i])[0] img cv2.cvtColor(cv2.imread(f{self.dir}/{info[file_name]}), cv2.COLOR_BGR2RGB) pts, lbls self._sample_from_gt(info) # 前景取中心点、背景取框外点 return { image: torch.as_tensor(self.tr.apply_image(img)).permute(2,0,1), original_size: img.shape[:2], point_coords: self.tr.apply_coords(pts, img.shape[:2]), point_labels: lbls, gt_256: self._decode_to_256(info), # RLE 解码并缩放到 256x256 }_sample_from_gt和_decode_to_256各三五行前者从掩码质心采样前景点、在框外采样背景点后者用mask_utils.decode还原掩码再插值到 256×256与解码器输出同尺寸。最小训练循环Sam.forward接收的是字典列表键名见 segment_anything/modeling/sam.py 文档字符串损失直接在 256 低分辨率 logits 上算省掉一次上采样。from segment_anything import sam_model_registry import torch.nn as nn sam sam_model_registryvit_b.cuda() if CFG[freeze_encoder]: for p in sam.image_encoder.parameters(): p.requires_grad_(False) opt torch.optim.AdamW(filter(lambda p: p.requires_grad, sam.parameters()), lrCFG[lr], weight_decayCFG[weight_decay]) bce nn.BCEWithLogitsLoss() for epoch in range(CFG[epochs]): for batch in train_loader: out sam(batched_input[batch], multimask_outputFalse)[0] loss bce(out[low_res_logits], batch[gt_256]) # 256x256 低分空间 opt.zero_grad(); loss.backward(); opt.step() val_iou evaluate(sam, val_loader) # 见下节 torch.save(sam.state_dict(), fruns/epoch{epoch}.pth)调得动先轻后重地解冻一步到位全参数微调90% 的概率会崩编码器学到的通用特征被少量领域数据冲掉。正确顺序是先训轻的再动重的参数起步值可用区间调错代价学习率1e-4解码器/ 1e-5编码器1e-5 ~ 1e-3⭐⭐⭐⭐⭐ 过大即发散批量大小42~16⭐⭐⭐ 受显存约束权重衰减1e-41e-5 ~ 1e-3⭐⭐⭐热身轮数32~10⭐⭐总轮数30~50配早停30~100⭐⭐⭐ 过多过拟合看得准怎么量化变好了评估别只看 loss——loss 降了掩码照样可能发虚。用验证集逐掩码算 IoU再平均才是能跟基线比较的数iou (pred gt).sum() / (pred | gt).sum() # 逐掩码原图分辨率 dice 2 * (pred gt).sum() / (pred.sum() gt.sum())参考量级具体以你的验证集实测为准型号预训练基线 IoU领域微调后相对提升单帧耗时A10Gvit_b0.720.8619%~45msvit_l0.750.8918%~78msvit_h0.790.9216%~125ms小模型提升空间最大这也是先用 vit_b 验证、再用大模型吃精度的量化依据。用得爽ONNX 导出与推理加速SAM 官方导出思路很巧妙重的图像编码器留在 PyTorch 里只跑一次导出到 ONNX 的只是轻量的提示编码器掩码解码器这样浏览器、边缘设备都能跑解码端参考 notebooks/onnx_model_example.ipynb 和 demo/。python scripts/export_onnx_model.py --checkpoint runs/best.pth \ --model-type vit_b --output sam_decoder.onnx \ --opset 17 --quantize-out sam_decoder_q8.onnx --gelu-approximate几个实战要点同一张图多次提示时用SamPredictor.set_image只编码一次predict反复调用跨进程传输可改用get_image_embedding()直接拿 64×64×256 特征高分辨率图只出单掩码时加--return-single-mask省掉多掩码上采样开销量化版sam_decoder_q8.onnx体积和延迟都能再降一档适合 CPU/移动端部署前过一遍这个清单验证集 IoU 达标且已记录基线checkpoint 用sam_model_registry重新加载跑 3 张图核对一致ONNX 导出后用onnxruntime.InferenceSession冒烟一次脚本自带对比 PyTorch 与 ONNX 输出掩码的 IoU 差异 1%记录单帧耗时作为线上压测基线 卡住了排障速查现象大概率原因处理loss 卡在原地或 NaN学习率过大降到 1e-5加 2~3 轮热身掩码整体偏移、贴着 1024 边界提示坐标没换算到输入空间坐标务必经tr.apply_coords图像经tr.apply_image二者用同一 transform训练 OOMbatch 内含全尺寸张量缩小批量开torch.amp混合精度掩码又多又碎过滤阈值太松pred_iou_thresh提至 0.88、stability_score_thresh提至 0.95scripts/amg.py 同名参数ONNX 端运行慢 / 算子缺失opset 低、GELU 走 erf--opset 17 --gelu-approximate保存 RLE 直接报错缺 pycocotoolspip install pycocotools回头看架构分重和轻编码器一次、解码器每次微调就按这个成本结构分配算力数据链路只有两条线图像走ResizeLongestSide坐标走apply_coords对齐了才谈得上训练解冻顺序是稳定性的来源先解码端、再编码器、后全参每阶段都有验证门槛部署端只导解码器图像特征可缓存可传输这是 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),仅供参考