尧图建网站 尧图建网站 YAOTU WEB BUILD 免费咨询
ARTICLE DETAIL

资讯详情

深耕网站建设与建站编程的一线实战洞察。

SAM(Segment Anything)微调实战:从自己的数据集到 ONNX 部署的 8 步完整教程

SAM(Segment Anything)微调实战:从自己的数据集到 ONNX 部署的 8 步完整教程 SAMSegment Anything微调实战从自己的数据集到 ONNX 部署的 8 步完整教程【免费下载链接】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 AnythingSAM拿到医疗影像、工业检测这类领域里跑大概率会撞上同一堵墙自然图上点一下就能出漂亮掩码换到你的业务数据上边缘开始漂移、细目标漏检、背景里颜色相近的物体被误切。通用预训练权重覆盖不到你的数据分布此时唯一靠谱的解法就是用自己的数据集做 SAM 微调。这篇文章带你完整走通这个流程从选版本、整理 COCO 标注到读懂三模块结构、写训练循环、让损失收敛、用 mIoU 验收最后导出 ONNX 上线每一步都有可直接照做的命令和代码骨架。读完本文你将获得✅ 选对 SAM 版本搭好可复现的微调环境✅ 把自己的数据集整理成模型能直接消费的 COCO 标注格式✅ 读懂三模块源码结构写出能跑通的 SAM 微调训练循环✅ 用分层微调策略让训练稳定收敛并用 mIoU 与 TensorBoard 验收✅ 导出微调后的 ONNX 模型用动态量化降低部署体积一、先选版本SAM 三个版本怎么选本节目标在写一行训练代码之前先确定用哪个骨干、怎么把环境搭起来。SAM 有三个骨干尺寸参数越少越省显存。对于微调新手vit_b 是默认推荐显存压力小、收敛快效果差距在绝大多数垂直场景里可以接受。版本键骨干配置embed_dim × depth参数量*适用场景vit_b768 × 12约 91M快速迭代、显存有限微调首选vit_l1024 × 24约 308M精度与速度平衡vit_hdefault1280 × 32约 636M高精度任务默认版本*示例值仅供参考。三个版本都由 segment_anything/build_sam.py 中的sam_model_registry注册通过sam_model_registry[vit_b]这类键名实例化。环境安装仓库要求python3.8、pytorch1.7强烈建议装 CUDA 版本# 创建环境 conda create -n sam_finetune python3.9 -y conda activate sam_finetune # 克隆仓库并本地安装 git clone https://gitcode.com/GitHub_Trending/se/segment-anything cd segment-anything pip install -e . # 可选依赖掩码后处理、COCO 格式、ONNX 导出 pip install opencv-python pycocotools matplotlib onnxruntime onnx建议在仓库外建一个训练工程目录这样规划sam_finetune/ ├── configs/ # 训练配置 ├── data/ │ ├── images/ # 你的数据集图片 │ └── coco/ # COCO 标注 json ├── scripts/ # 数据集类、训练循环、评估脚本 └── outputs/ # 检查点、日志、导出模型二、整理数据用 COCO 格式准备你自己的数据集本节目标把原始标注统一成 COCO 格式让pycocotools能直接读训练时还能顺手采样提示点。SAM 官方数据集本身就是 COCO RLE 格式所以你的自定义数据集照抄这个格式最省事。一个标注文件长这样节选{ images: [{ id: 1, file_name: img_001.jpg, width: 1024, height: 768 }], annotations: [{ id: 1, image_id: 1, category_id: 1, bbox: [x, y, w, h], area: 12345, segmentation: { counts: RLE编码, size: [768, 1024] }, iscrowd: 0 }], categories: [{ id: 1, name: target_object }] }微调 SAM 时标注的掩码用来算损失bbox 则用来采样提示点——这正好对上SamPredictor的输入习惯。数据增强不必复杂方向不变性和尺度不变性收益最大增强方法参数范围适用场景效果随机裁剪保留 0.8~1.0尺度不变性⭐⭐⭐⭐⭐随机旋转±30°方向不变性⭐⭐⭐⭐水平翻转p0.5通用⭐⭐⭐⭐亮度/对比度抖动±20%光照变化⭐⭐⭐高斯噪声σ0.01抗噪能力⭐⭐注意增强图像时必须同步变换 bbox 和掩码否则提示点会指错位置训练目标也会错位。三、读懂模型再写微调代码本节目标先弄清三模块各管什么再照真实接口写数据集类和训练循环避免写出和源码对不上的代码。SAM 的前向流程是三段式图像编码器把整张图压成 64×64 的特征嵌入一次算完、可复用提示编码器把点/框/掩码转成向量轻量掩码解码器结合两者输出低分辨率掩码和 IoU 质量分。图像输入前必须经过ResizeLongestSide(1024)缩放、归一化pixel_mean/pixel_std提示坐标也要用apply_coords同步变换到输入帧——这些工具都在 segment_anything/utils/transforms.py训练时必须和推理用同一套否则掩码会整体漂移。数据集类骨架接口与源码保持一致import cv2, torch from torch.utils.data import Dataset from pycocotools.coco import COCO from segment_anything.utils.transforms import ResizeLongestSide class CustomSAMDataset(Dataset): 读 COCO 标注输出图像与点提示。 def __init__(self, ann_file, image_dir, target_length1024): self.coco COCO(ann_file) self.image_dir image_dir self.image_ids list(self.coco.imgs.keys()) self.transform ResizeLongestSide(target_length) # 与官方推理一致 def __len__(self): return len(self.image_ids) def __getitem__(self, idx): img_id self.image_ids[idx] img cv2.imread(f{self.image_dir}/{self.coco.imgs[img_id][file_name]}) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) # SAM 要求 RGB x torch.as_tensor(self.transform.apply_image(img)) x x.permute(2, 0, 1).contiguous() # HWC - CHW anns self.coco.loadAnns(self.coco.getAnnIds(imgIdsimg_id)) ax, ay, aw, ah anns[0][bbox] # 框中心采样前景点 pt self.transform.apply_coords([[ax aw/2, ay ah/2]], img.shape[:2]) return { image: x, original_size: img.shape[:2], point_coords: torch.as_tensor(pt, dtypetorch.float), point_labels: torch.tensor([1]), ann_ids: [a[id] for a in anns], }训练循环骨架。关键坑仓库里Sam.forward带了torch.no_grad()直接调它不会有梯度所以要按 segment_anything/predictor.py 的写法手动串三个子模块import torch from segment_anything import sam_model_registry sam sam_model_registryvit_b sam.train() # 第一阶段只训解码器图像编码器冻结后续解冻见第四节 optimizer torch.optim.AdamW( list(sam.mask_decoder.parameters()) list(sam.prompt_encoder.parameters()), lr1e-4, weight_decay1e-4) def train_one_epoch(): total 0.0 for batch in train_loader: images batch[image].to(device) feats sam.image_encoder(sam.preprocess(images)) # 1) 图像嵌入 sparse_emb, dense_emb sam.prompt_encoder( # 2) 编码提示 points(batch[point_coords].to(device), batch[point_labels].to(device)), boxesNone, masksNone) low_res_masks, iou_pred sam.mask_decoder( # 3) 解码掩码 image_embeddingsfeats, image_pesam.prompt_encoder.get_dense_pe(), sparse_prompt_embeddingssparse_emb, dense_prompt_embeddingsdense_emb, multimask_outputFalse) loss mask_loss(low_res_masks, build_gt_logits(batch)) # 4) 对 GT 算 BCE optimizer.zero_grad(); loss.backward(); optimizer.step() total loss.item() return total / len(train_loader)四、让训练收敛分层微调与超参数调优本节目标给出一条可复现的调优路径和一组起步超参数少碰壁。策略上先冻结图像编码器、只训提示编码器和掩码解码器解码器参数量小收敛快验证 mIoU 达标后再解冻编码器把学习率降一个数量级整体微调。这条路径的好处是显存需求低、不容易破坏预训练特征。超参数起步值影响程度用 ⭐ 标注按你的数据量在范围内搜索超参数推荐值调整范围影响程度学习率1e-4编码器解冻后 1e-51e-5 ~ 1e-3⭐⭐⭐⭐⭐批量大小42 ~ 16⭐⭐⭐权重衰减1e-41e-5 ~ 1e-3⭐⭐⭐⭐训练轮数3010 ~ 50⭐⭐⭐学习率调度CosineAnnealing—⭐⭐⭐提示点采样框中心 1 个负样本多框/多点⭐⭐⭐⭐提示点采样方式对收敛影响很大只用前景点模型容易把整块区域都当成目标在背景加 1~2 个负样本点通常能明显改善边界。五、验收效果评估指标与训练监控本节目标用 mIoU 客观量化微调收益并用 TensorBoard 跟踪训练过程。评估直接复用SamPredictor的推理接口和训练用同一套ResizeLongestSide变换import numpy as np, torch from segment_anything import SamPredictor def compute_iou(pred, gt): inter (pred gt).sum() return inter / max((pred | gt).sum(), 1) torch.no_grad() def evaluate(sam, val_loader): sam.eval() predictor SamPredictor(sam) ious [] for batch in val_loader: img batch[image].numpy()[0] predictor.set_image(img, image_formatRGB) masks, scores, _ predictor.predict( point_coordsbatch[point_coords][0].numpy(), point_labelsbatch[point_labels][0].numpy(), multimask_outputFalse) gt decode_gt_mask(batch) # 把 GT 标注解码成二值掩码 ious.append(compute_iou(masks[0], gt)) return float(np.mean(ious))同时算 Dice2×交并和与 Precision/Recall防止只盯 IoU 漏掉边界细节。下表为一次典型微调的效果对照示例值仅供参考实际取决于你的领域数据模型微调前 mIoU微调后 mIoU单图解码耗时*vit_b0.620.81约 45msvit_l0.660.85约 78msvit_h0.690.88约 125ms监控方面用SummaryWriter把train/loss、val/mIoU、当前学习率写进 TensorBoard每个 epoch 记一次。看曲线的重点有两个训练损失下降但验证 mIoU 走平通常该加强增强或早停了两者同时抖动多半是学习率偏高。六、把模型用起来SAM ONNX 导出与推理优化本节目标把微调成果导出成可移植的 ONNX 模型并做两个立竿见影的推理优化。仓库的设计是图像编码器留在 PyTorch 骨干里算一次就够只把提示编码器 掩码解码器导出成 ONNX这样浏览器、移动端等任何 ONNX 运行时都能做掩码预测。导出命令直接复用仓库脚本 scripts/export_onnx_model.py# 导出微调后的轻量部分prompt_encoder mask_decoder python scripts/export_onnx_model.py \ --checkpoint sam_vit_b_finetuned.pth \ --model-type vit_b \ --output sam_finetuned_decoder.onnx \ --quantize-out sam_finetuned_decoder_quant.onnx # 可选动态量化压缩体积推理侧两个优化点缓存图像嵌入同一张图换提示重预测时别重跑编码器。SamPredictor.set_image已经把嵌入缓存在self.features里predict可以反复调用自建服务时也可以自己按图 hash 建一层get_image_embedding()缓存。高分辨率图只回传单个掩码导出时加--return-single-mask避免多掩码上采样拖慢高分辨率输出。七、踩坑速查常见问题怎么解决本节目标把微调 SAM 最容易翻车的几处一次性列清楚遇到按表查。问题现象可能原因解决方案反向传播损失为 0 / 无梯度直接调了Sam.forward它带torch.no_grad()手动串image_encoder→prompt_encoder→mask_decoder见第三节代码掩码整体偏移、指哪错哪提示坐标没经过ResizeLongestSide.apply_coords变换训练与推理统一用官方 transform 同步坐标加载检查点 shape 报错微调版与原版骨干尺寸不一致用--model-type与检查点同版本实例化显存爆掉批量过大或用了 vit_h降到 vit_b、减 batch、梯度累积验证 mIoU 不涨标注噪声大 / 只有前景点清洗标注、加负样本点、加增强ONNX 导出后 erf 报错运行时 GELU 实现问题导出时加--gelu-approximate换 tanh 近似性能优化 checklist训练开启混合精度AMP数据加载开num_workers与pin_memory导出 ONNX 后跑一遍onnxruntime校验前向输出用--quantize-out动态量化压缩部署体积服务端缓存图像嵌入同一张图多次提示只编码一次八、总结你拿到了什么选型意识vit_b 起步、按需升 vit_l/vit_h显存和精度可以量化权衡。数据能力COCO 标注 提示点采样让任何领域数据集都能喂给 SAM 微调。训练手艺绕开no_grad坑手动串三模块分层冻结让收敛又快又稳。部署闭环ONNX 只导解码器 图像嵌入缓存微调成果能低成本上线。下一步建议换points_per_side参数试 自动掩码生成器评估微调对无提示场景的提升参考 notebooks/onnx_model_example.ipynb 把导出模型接进你自己的 Web 服务关注 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),仅供参考
返回列表