尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

SAM2自定义数据训练全攻略:从格式转换到避坑实践

SAM2自定义数据训练全攻略:从格式转换到避坑实践 简介针对需要训练SAM2自定义数据的开发者这份资源提供了一套清晰的数据封装与训练流程参考。压缩包共2个Python脚本大小仅5KB分别承担通用数据集创建和针对LabPicsV1数据集的定制封装完整覆盖数据预处理、格式统一、数据增强以及训练集与验证集划分等关键环节。已有310人学习下载适合具备一定Python基础、希望快速上手SAM2自定义训练的入门与中级用户。通过阅读这两个脚本读者可了解SAM2框架对输入数据的具体要求掌握将原始图片数据转换为模型可识别格式的操作方法并能在自己的项目中直接借鉴数据增强策略和数据集划分逻辑从而减少环境配置与代码调试的试错成本提升模型训练的效率和泛化能力。围绕LabPicsV1展开的封装体现了对特定领域数据特性的针对性适配整个流程虽小但链路完整对理解数据驱动训练下从原始样本到可用训练集的完整过程很有帮助也可作为后续扩展更多数据集时的参照模板。脚本命名和结构清晰便于对照学习关键流程。1. 通用分割权重在你数据上失灵sam2 训练自己的数据到底要改什么用过 Meta 的 SAM2 做业务分割的人大概率经历过这个瞬间官方权重在演示图上点哪里切哪里换到自己的产线缺陷图、遥感影像或者病理切片上要么点中了目标却框出一大片背景要么自动分割模式下目标直接被跳过。原因不是 SAM2 效果不行而是它的通用权重见过的是 SA-V 里那批自然图像和目标分布你数据里的纹理、尺度、目标形态和背景噪声它没吃过。所以“sam2训练自己的数据”这句话背后其实是一条完整的微调链路把自有标注转换成它认得的格式改训练配置跑通采样和 loss再解决显存与过拟合问题。这篇笔记就按这个顺序把每一步的参数、命令和容易翻车的地方写清楚适合手里已经有标注、想把 SAM2 稳定用在自己场景上的工程师。2. 把标注整理成 sam2 能吃的数据格式COCO JSON 与掩码体检2.1 为什么微调绕不开 COCO 格式以及和 YOLO 格式的差别很多人是从 YOLO 那条路过来的手里有 yolov5 或 yolov8 训练自己的数据集时留下的 txt 标注、XML 或者 labelme 的 JSON。下意识会觉得 SAM2 也能直接吃这些文件实际不行。常见做法是先把所有标注统一成 COCO 格式的 JSON一个 JSON 文件里包含 images、annotations、categories 三块annotations 里的 segmentation 字段要么是多边形点集要么是 RLE 编码掩码。SAM2 的训练管线读取的就是这种结构加载后会把 segmentation 还原成与图片等大的 mask 张量再送到训练流程里。这里最容易搞混的是YOLO 的 txt 给的是归一化中心点和宽高SAM2 的 COCO segmentation 给的是像素坐标系下的多边形顶点mmrotate 训练 DOTA 数据集时用的旋转框格式也和 SAM2 要的稠密掩码不是一回事。也就是说即使你之前的任务用过 COCO 检测格式也需要确认 segmentation 是真掩码或多边形而不只是 bbox。如果你只有检测框得先用分割模型或标注工具补出掩码如果你有 HRSC2016 这类带旋转框的数据也要先把旋转多边形展开成普通多边形顶点再写进 COCO JSON。数据格式这步偷懒后面训练会一直出现 mask 与图片对不上的问题而且是黑匣子式的报错不好查。2.2 用脚本把已有标注转成有效的掩码并做体检无论标注原来存在哪里最终落到磁盘上的建议是一个 images 目录放原图一个 annotations 目录放 COCO JSON。转换脚本的核心工作是读原标注、按图片尺寸把多边形或掩码还原、过滤空标注、再写入新 JSON。下面是一个针对 labelme 多边形标注转 COCO 的参考脚本可以直接改路径用import json import numpy as np from pathlib import Path from pycocotools import mask as mask_utils def labelme_to_coco(img_dir, label_dir, out_json): images, annotations [], [] ann_id 1 for idx, img_path in enumerate(sorted(Path(img_dir).glob(*.jpg))): label_path Path(label_dir) / (img_path.stem .json) if not label_path.exists(): continue with open(label_path) as f: label json.load(f) h, w label[imageHeight], label[imageWidth] images.append({id: idx, file_name: img_path.name, height: h, width: w}) for shape in label[shapes]: pts np.array(shape[points], dtypenp.float32).flatten().tolist() if len(pts) 6: # 至少 3 个点 continue # 用多边形生成 RLE避免抽样后顶点数过多 rle mask_utils.frPyObjects([pts], h, w) mask mask_utils.decode(rle) area int(mask.sum()) if area 0: continue annotations.append({ id: ann_id, image_id: idx, category_id: 1, segmentation: [pts], area: area, bbox: mask_utils.toBbox(rle).tolist(), iscrowd: 0 }) ann_id 1 coco {images: images, annotations: annotations, categories: [{id: 1, name: object}]} with open(out_json, w) as f: json.dump(coco, f) print(fimages: {len(images)}, annotations: {len(annotations)})逻辑说明逐张图片读取对应 labelme 文件把每个 shape 的顶点转成 COCO 的多边形表示frPyObjects会把多边形编码成 RLEmask_utils.decode再还原成像素掩码这样既过滤掉顶点数不足的异常标注也能顺手算出面积和 bbox。需要注意的是SAM2 的类别编号从 1 开始背景隐式作为 0所以这里统一用category_id: 1。参数说明frPyObjects的第二个和第三个参数是图片高度、宽度顺序不能反否则 RLE 解码出来是旋转 90 度的掩码训练时 mask 与图像错位bbox要的是[x, y, w, h]浮点列表mask_utils.toBbox返回的就是这个结构。如果标注里有多类别建议先在 categories 里维护 id 映射再在循环里按 shape 的 label 查表不要直接在 JSON 里写字符串类别名。2.3 目录组织与 train/val 切分单独留一个“交互式检验集”数据转换只是第一步目录组织直接决定后续脚本和配置好不好写。我一般会建一个 datasets 根目录下面按任务名分子目录每个子目录里固定三个位置train.json、val.json、images。val.json 不只用来算 mIoU还要从里面再抽一小批图片单独放着不参与任何训练和指标统计专门留给训练结束后的交互式点击测试因为 SAM2 这类模型最终的验收方式是“人点一下看 mask 贴不贴”。如果你手上有多个来源的数据比如一部分是公开数据集加自己的业务数据建议先把公开集合并进训练集再在 val 里混入一部分同分布的业务数据。原因是 SAM2 在自家 SA-V 上已经很强你训练的目标是让它吸收新域的纹理和分割模式val 里如果全是新域数据loss 回落会很好看但拿到没见过的同域数据上可能立刻打回原形。train/val 的切分比例常见是 8:2 或 9:1数据量少于 500 张时建议用 9:1val 里每类目标至少要有 20 个实例否则评估时方差太大看不出真实水平。切分脚本很简单按图片 id 做随机划分确保同一张图不会同时出现在 train 和 val 的 JSON 里。tree datasets/my_data/ ├── train.json ├── val.json └── images/ ├── img_001.jpg ├── img_002.jpg └── ...目录里的 train.json 和 val.json 由上面转换脚本生成images 只放原图。如果你用的是视频数据目录组织稍微不同会在第 4 章单独说。到这一步新手最容易在“JSON 里 image_id 和 file_name 对不上”上踩坑所以准备一个校验脚本把 train/val 里每个 image_id 对应的图片路径存在性检查一遍缺文件的提前补不要等到训练中途 DataLoader 报错再回头找。3. 修改 sam2 训练配置从权重路径到数据路径的三处硬改动3.1 配置文件里必须改的键模型权重、数据路径与类别数SAM2 的训练入口是仓库里的training/train.py它用 Hydra 管理配置默认会加载configs/sam2.1_hiera_b.yaml这类模型配置文件。不建议直接改仓库自带的 yaml复制一份放到自己的配置目录再在命令行用-c指向它。需要改的第一个位置是模型权重路径model.checkpoint指向你下载好的官方 SAM2 权重文件比如 sam2.1_hiera_base_plus.pt 或 sam2.1_hiera_large.pt这个路径必须是绝对路径或者相对于仓库根目录的路径。第二个必改位置是数据集配置。SAM2 的 data 配置块长下面这样核心是data.dataset.train.json_path和data.dataset.train.img_dir测试时还要带上 val 路径data: dataset: train: _target_: sam2.data.datasets.ImageVideoDataset img_dir: ./datasets/my_data/images json_path: ./datasets/my_data/train.json num_frames: 1 val: _target_: sam2.data.datasets.ImageVideoDataset img_dir: ./datasets/my_data/images json_path: ./datasets/my_data/val.json num_frames: 1 num_workers: 8逻辑说明_target_指定 SAM2 内部的数据集类num_frames: 1表示只按单帧图像读不启用视频帧序列。如果 json 里只有图片没有视频帧却把num_frames设成大于 1DataLoader 会按索引取不到对应帧出现奇怪的越界报错。data.num_workers在 Linux 上建议 8 或 16Windows 下先降到 2否则多进程加载会频繁报错。另外类别数在 SAM2 的配置里不是直接写num_classes而是通过model.num_frames、model.decoder等结构隐式定义你只要保证 COCO JSON 里的category_id是从 1 开始递增且 categories 的长度与真实类别数一致即可。3.2 启动训练的命令与单卡适配配置改完启动训练的命令长这样python training/train.py \ -c configs/my_sam2_bplus.yaml \ data.dataset.train.json_path./datasets/my_data/train.json \ data.dataset.train.img_dir./datasets/my_data/images \ data.dataset.val.json_path./datasets/my_data/val.json \ data.dataset.val.img_dir./datasets/my_data/images \ model.checkpoint./checkpoints/sam2.1_hiera_base_plus.pt \ learning_parameters.per_gpu_batch_size2 \ learning_parameters.max_epoch200逻辑说明-c后跟的是你自定义的 yaml 路径命令行里冒号后面的覆盖项能在不改 yaml 的情况下临时调参数这对跑消融实验很有用。per_gpu_batch_size是每张卡上的批次大小不是全局批次大小SAM2 配置里默认写在learning_parameters块下多卡时全局批量 单卡批次 × 卡数。max_epoch按数据集大小给500 张图以内的任务 100-200 epoch 常见数据量超过 2000 张可以减到 50-80。单卡用户最容易卡在显存上。sam2.1_hiera_base_plus 是较小的权重12GB 显存能跑 batch_size2、输入图短边 1024如果用的是 large 权重或图片分辨率更高先把 batch_size 降到 1或 2再不行就改配置里的data.dataset.eval_img_size和data.dataset.input_img_size把长边限制在 1024 或 768。SAM2 内部会把输入图 pad 到 14 的倍数所以你传给 DataLoader 的图不需要自己预先裁剪只要控制原始图片尺寸不要太大就行。3.3 超参数怎么给小数据集别盲目抄大模型的批次大小训练自己的数据时最不该做的是直接把 SAM2 预训练阶段的超参搬运过来。预训练 SAM2 用的是几十万张图和很大的 batch你的数据可能只有几百张全局 batch 太大反而让模型在少量 epoch 内看过重复样本记忆而不是泛化。小数据集上我一般会这样给学习率用 1e-4 到 2e-4 之间配合 AdamW如果验证集 mIoU 波动很大降到 5e-5 重跑权重衰减用 0.05这是 AdamW 的常见设置。learning_parameters.lr指的是 backbone 和 decoder 的共享学习率SAM2 配置支持分别设置backbone_lr_multiplier和decoder_lr_multiplier常见做法是让 decoder 学得快一点backbone 学得慢一点防止破坏预训练图像特征。冻结骨干网络是另一个可控的选项训练命令里加上learning_parameters.backbone_lr_multiplier0.0这意味着 backbone 完全不更新只训练 mask decoder 和 memory 相关模块。显存占用明显下降训练速度也快不少前提是你的任务和自然图像没有剧烈偏移。像医学影像、遥感这类跟自然图像差距大的场景还是让 backbone 以很低的倍率跟着学一般取 0.1 到 0.3。4. 处理两类容易翻车的数据视频标注与交互式点击数据4.1 视频数据用 SA-V 结构组织帧间同一 object_idSAM2 最强的地方在视频分割这也是它与 SAM 最大的区别。但如果你的数据是视频组织方式跟单帧图像完全不同。SAM2 的视频数据沿用 SA-V 格式images 目录下按视频片段分文件夹annotations 里每个标注都带video_id同一目标在连续帧中必须保持同一个instance_id这样训练时才能让模型学到“跨帧保持对象一致性”。如果你从开源数据集转过来比如想把 HRSC2016 这类遥感视频或 DOTA 的帧序列喂给 SAM2要注意这些数据集往往只有静态框没有跨帧 id需要先做视频追踪标注不然模型学到的是“每一帧独立分割”直接把 memory 机制学废了。视频配置里num_frames要改成大于 1 的值常见有 2 或 4。不要一口气拉到 8显存会迅速爆掉而且帧间目标形变大时模型反而学不到稳定的特征。我吃过这个亏第一次跑视频数据时把num_frames设成 8A100 上 batch 只能设 1训练速度慢到怀疑人生后来发现 4 帧和 8 帧的结果几乎没差别。另外视频训练时 val 集抽帧方式最好和训练一致从同一个片段里按相同间隔抽帧不要训练用连续帧、验证用随机帧那样评估出来的数字是虚高的。4.2 没有交互点击数据时用 GT 掩码自造 prompt交互式分割是 SAM2 的招牌能力训练时它需要 prompt 输入常见的是点或框。问题来了你手上的标注只有完整掩码没有人工点击记录怎么办常见做法是在训练管线里把 GT 掩码转成合成 prompt从掩码的连通域内部随机采样一个或多个点作为正样本点再归一化到 [0,1] 区间送给模型。SAM2 的数据加载代码里就内置了类似的 prompt 生成逻辑你只需要在配置里打开data.dataset.use_gt_prompts并设置采点数量训练时每个 batch 都会随机生成不同的 prompt相当于免费做了数据增强。import numpy as np def sample_prompt_from_mask(mask, num_points1): ys, xs np.where(mask 0) if len(xs) 0: return None idx np.random.choice(len(xs), sizenum_points, replaceFalse) h, w mask.shape points np.stack([xs[idx] / w, ys[idx] / h], axis1).astype(np.float32) return points.tolist()逻辑说明函数输入是像素掩码和需要的点数量输出的是归一化坐标列表。SAM2 内部要求 prompt 坐标是归一化到图像宽高的浮点值顺序是[x, y]先列后行很多从 OpenCV 转过来的工程师会习惯写成[y, x]这个顺序错了模型推理时完全无法收敛。replaceFalse保证同一个点不会被重复采样。参数说明num_points对最终效果影响很大。1 个点时模型偏向于从边界内某个点向四周扩散适合目标比较紧凑的场景3 到 5 个点时模型更稳定但对噪声更敏感。训练时建议每张图 2 到 3 个点推理时只用 1 个点就能工作这样模型在“单点丢进来也能切出完整目标”的能力上会更强。5. 训练 sam2 的避坑清单显存、分辨率与 loss 不降的排查5.1 显存 OOM先检查 pad 和批次而不是急着换大卡现象训练刚开始第一个 step 直接抛出 CUDA out of memory连 eval 都跑不起来。这时很多人第一反应是换 80G 大卡其实多数情况下不是模型太大而是输入图被放得太大或者批次没降下来。SAM2 的 Hiera backbone 对分辨率非常敏感长边接近 2048 的遥感图在 batch4 下40G 显存都不够。原因除了输入图尺寸和批次还有一个隐藏变量是 pad 到 14 的倍数。SAM2 的 patch size 是 14输入任意的H × W会被 pad 到ceil(H/14)*14 × ceil(W/14)*14如果原图是 1001×1001pad 后变成 1008×1008。看似只多了几像素但配合注意力机制中间张量会成倍放大。解决先用一条命令看当前数据集的原始尺寸分布把过长边的图在训练前统一 resize 到 1024 或 768并调低per_gpu_batch_size。如果图像数量大可以启用梯度累积即用较小 batch 积累多步后再更新一次参数效果接近大批次但显存友好得多。5.2 掩码在图中占比过低loss 不降的一个隐蔽原因现象训练跑了 50 个 epochloss 曲线一直在高位震荡mIoU 几乎为 0。检查数据和代码都没发现明显错误最终发现是掩码目标区域只占整张图的千分之几比如工业缺陷检测里一块小瑕疵在 4000×3000 的图上只有几十个像素。原因SAM2 虽然用 focal loss 和 dice loss 的混合对正负样本不平衡有一定承受力但当正样本区域占比低于 0.1% 时dice loss 的梯度被背景主导模型学到的是“全输出背景”的局部最优解。解决不要直接拿整张大图训练先把包含目标的外接框裁剪出来扩边 30 像素后 resize 到 1024再送进训练管线。裁剪后的掩码占比通常能提高到 5% 以上模型才真正学到目标的结构。如果目标太稀疏可以先用目标检测器定位再对每个框做分割这条路线和 mmrotate 训练 DOTA 数据的做法类似先检测再分割而不是让一个大模型端到端硬啃。5.3 验证指标虚高评估方式配不上交互式使用场景现象训练结束后的 val mIoU 到了 0.85看起来非常理想但部署时用户点了一下目标中心SAM2 却给了个大范围的错误掩码。原因val 集的评估用的是“给定 GT 掩码附近的 prompt 再预测整个掩码”而真实使用里用户点的位置不可控可能落在目标边缘、遮挡区域甚至背景里。这种情况下模型没见过分布外的 prompt表现自然崩。解决单独构造一个交互式验证集每张图上标注 5 到 10 个不同的点击点包括目标中心、边界内侧、目标与背景交界处跑一遍推理记录 mask 与 GT 的 IoU然后取中位数而不是平均数。平均数会被“正好点中中心”的样本拉高中位数更能代表真实手感。SAM2 的评估脚本里本来就有类似的 point prompt 采样逻辑只是默认采点偏保守建议把采样点数量加大并加入边界点。5.4 过拟合到背景纹理训练数据太少时的典型症状现象训练集 mIoU 持续上升val 集从第 30 个 epoch 开始不涨反跌预测结果里出现大量与目标颜色相近的背景碎片。原因数据量太少模型把背景纹理当成了判别特征典型的过拟合。解决分三步走第一步数据增强打开水平翻转和颜色抖动并适当加大随机裁剪比例第二步调整学习率把 lr 从 1e-4 降到 5e-5同时提高 weight_decay 到 0.1第三步如果增强后还是过拟合说明数据量确实不够回到 2.3 节提到的方向补充数据优先补有目标形变和背景变化的样本而不是相同场景的重复帧。对于视频数据还有一个专用技巧抽帧时不要集中在视频前几秒均匀抽完整段视频否则同一目标的相同姿态反复出现加剧记忆效应。5.5 配置里改了参数没生效被 Hydra 覆盖层骗了现象在命令行里补了 learning_parameters.max_epoch500训练日志里显示的 epoch 数却是默认值 200所有参数覆盖都像没生效一样。原因SAM2 的 Hydra 配置是多层组合的命令行覆盖优先级最高但如果配置文件的 defaults 列表里有 override hydra/job_logging 这类设置可能把命令行参数重新覆盖回去。解决启动训练前打印当前生效的配置确认参数真的被读到了。命令是 python training/train.py -c configs/my_sam2_bplus.yaml --cfg job会输出完整的解析结果在日志里搜索 max_epoch如果还是默认值说明你的自定义 yaml 里 defaults 的顺序写错了。另外注意命令行覆盖项里点号分隔的 key 必须和 yaml 里的层级完全一致比如 learning_parameters.per_gpu_batch_size少写一个层级就会静默失败训练照跑但参数用的是默认值。6. 验证与进阶用“点一点”实测再用 LoRA 控制显存6.1 加载 checkpoint 做交互式验证的最小脚本训练结束后的第一个动作不要看 tensorboard 里的曲线直接加载 checkpoint 跑一个手动点击脚本感受一下真实的手感。下面是一个最小验证脚本import torch from PIL import Image import numpy as np from sam2.build_sam import build_sam2 from sam2.sam2_image_predictor import SAM2ImagePredictor checkpoint ./checkpoints/my_finetuned.pt model_cfg configs/sam2.1_hiera_b.yaml predictor SAM2ImagePredictor(build_sam2(model_cfg, checkpoint)) image np.array(Image.open(./test_imgs/defect_001.jpg).convert(RGB)) predictor.set_image(image) click_point np.array([[300, 250]]) # [x, y] click_label np.array([1]) # 1 表示正样本点 masks, scores, _ predictor.predict( point_coordsclick_point, point_labelsclick_label, multimask_outputTrue ) best masks[np.argmax(scores)] print(best mask shape:, best.shape, score:, scores.max())逻辑说明SAM2ImagePredictor是对外提供的推理封装set_image会先跑一次图像编码器后续多次点击不需要重复编码验证交互式体验时必须先调用它。predict里的point_coords是像素坐标不需要归一化预测器内部会处理。multimask_outputTrue会让模型同时输出多个候选掩码取分数最高的那个作为交互结果这也是 SAM2 的推荐用法。这个脚本同时可以用来验证 5.3 节提到的边界点击场景把点击点分别放在目标中心、靠近边缘、重叠区域看掩码质量变化。如果中心点击效果好、边缘点击明显崩说明训练时的 prompt 采样太保守回到 4.2 节调大num_points并加入边界点重训一轮。6.2 小样本下的进阶方向冻结 backbone 与 LoRA如果你的显存瓶颈始终卡在 backbone 上且数据量只有一两百张可以尝试在 SAM2 上做 LoRA。做法是冻结原始权重在 Hiera 的 attention 层插入低秩适配器只训练低秩矩阵。这种方法在 SAM2 社区里已有不少实践配合 6.1 的验证脚本能在 8GB 显存上完成 base_plus 的微调。实际操作时一般只对 image encoder 的 query、key、value 投影层加 LoRArank 取 16 到 32超过 32 收益急剧下降反而增加过拟合风险。训练完成后可以做一个简单评估同一张测试图分别用官方权重、全量微调权重和 LoRA 权重跑一次点击分割对比三者的掩码质量和显存峰值。你会发现 LoRA 未必比全量微调差很多尤其在数据分布集中、目标形态相对固定的业务场景下LoRA 甚至可能因为正则化效果而更稳。我的习惯是把 LoRA 权重单独保存输出一个“sam2_lora_weights.pt”和 base checkpoint 分离这样以后换数据只要重新训练 LoRA 权重就行不需要复制整个大模型。这条工作流比较适合团队里模型多、显存少、迭代快的场景。最后补一句个人习惯每次调完参数都会把“训练命令 配置文件 关键日志片段”存在同一个实验目录里下次改数据或调参时直接对比比记笔记可靠得多。希望帮到你。本文还有配套的精品资源点击获取
返回列表