简介:面向需要将自有图像数据接入sam2框架进行训练的开发者与算法学习者,这份压缩包提供了一套可直接运行与扩展的数据封装脚本,帮助解决从原始图片到模型可读Dataset格式的转化难题。包内共2个文件,均为Python脚本,总体积约5KB,分别承担通用训练数据集构建和针对特定数据集(如LabPicsV1)的自定义封装,内容涉及数据预处理、格式统一、数据增强以及训练/验证集划分等关键环节。目前已有310人学习下载,轻量且实用,尤其适合刚接触sam2或希望快速验证自定义数据训练流程的读者。通过阅读和改写这两个脚本,可清晰理解数据管道的搭建细节,并根据自身图像数据调整参数;两个脚本相互对照,也能帮助掌握通用封装与定制化适配两种常见写法,降低sam2自定义训练的实践门槛。
1. Sam2 能训自己的数据吗:先用一遍才知道它和 YOLO 那套训练管线差在哪
Sam2训练自己的数据,听起来就是把 YOLO 那套“准备 mask → 改配置文件 → 跑 train.py”搬过来,真上手会发现完全不是一回事。很多人拿 yolov8 训练自己的数据集很顺手,换成 sam2 训练自己的数据,第一步就卡在输入格式上:SAM2 是带 prompt 的交互式分割模型,训练时除了图像和 GT 掩膜,还要给它点坐标、边界框或者掩膜提示,缺了这些,模型根本不知道你在让它学什么。
这篇文章就是讲怎么把普通分割数据改造成 SAM2 能直接吃的格式,怎么选预训练权重和微调范围,以及训练时最容易翻车的几个点。适合两种人:手里有语义分割或实例分割掩膜,想换 SAM2 做底座的工程师;以及想验证 SAM2 在自己领域能不能比传统分割网络更稳的研究者。它不会替你把数据变多,但能让你少走两三周的歪路。
2. Sam2 训练数据准备:把普通掩膜转成 prompt 格式,自动采点
2.1 数据组织:为什么不能拿 png 目录直接开训
SAM2 的训练脚本一般按 COCO 或 SA-V 风格读取 json:images 里记录图像路径,annotations 里挂实例掩膜,掩膜以 RLE 编码存储;训练时数据加载器再从实例掩膜里动态采点。如果你直接建一个 png 文件夹然后让 train.py 去读,绝大多数版本会直接报错,或者“成功”读到但每个样本都没有点提示,后面 loss 怎么调都降不下去。
为什么这样设计?因为 SAM2 的定位是“任意分割”,它学的是条件概率 P(mask | image, prompt),而不是单纯的 P(mask | image)。只给它图像和 GT,等价于把 prompt 随机留空,模型就会退化成一个很弱的全图分割器。所以数据准备这一步,不是把掩膜换个格式存一遍,而是要把“点提示、框提示”一起编码进训练样本。这也解释了为什么很多人跑 mmsegmentation 训练 cityscapes 很熟、跑 mask2former 训练也顺,一换到 SAM2 就发现连数据校验脚本都要重写。
我一般的数据组织方式是:图像和掩膜继续按目录放,但额外生成一个 json 索引文件,里面每张图挂上所有实例掩膜、每个实例的外接框,以及一组预采样的前景点。json 本身不大,几千张图也就几十 MB,真正的重资产还是原始 PNG。这样 train.py 不用扫描目录,直接读 json 就能判断训练集和验证集怎么切。
2.2 掩膜转 RLE + 自动采前景点:转换脚本与参数说明
下面这个脚本是我常用的“普通分割掩膜 → SAM2 训练 json”转换器。它假设你的掩膜目录里每张 PNG 是单通道标签图,像素值就是类别或实例 ID;如果同一个类别在图像里有多个连通域,我会在脚本里拆开,否则一个类别会被当成一个巨型实例,小目标全被淹没。
# convert_masks_to_sam2_json.py # 作用:把普通掩膜目录(单通道 png)转成 SAM2 训练可读的 json import json import glob import cv2 import numpy as np from scipy import ndimage from pycocotools import mask as mask_util def mask_to_rle(binary_mask): # 转成 fortran order 的 uint8,pycocotools 才能正确编码 fortran_mask = np.asfortranarray(binary_mask.astype(np.uint8)) encoded = mask_util.encode(fortran_mask) return encoded["counts"].decode("utf-8"), encoded["size"] def pick_foreground_points(binary_mask, num_points=3): # 用 distance transform 选“远离边界”的像素,比随机质心更稳 dist = cv2.distanceTransform((binary_mask * 255).astype(np.uint8), cv2.DIST_L2, 3) ys, xs = np.where(dist > 1.0) # 只保留距边界至少 1px 的点 if len(xs) == 0: ys, xs = np.where(binary_mask > 0) idx = np.random.choice(len(xs), min(num_points, len(xs)), replace=False) return [[int(xs[i]), int(ys[i])] for i in idx] def split_instances(label_map, class_id=None): # 按连通域拆分。如果掩膜本身就是实例ID,class_id 可以传 None if class_id is None: mask = label_map > 0 labels, num = ndimage.label(mask) else: mask = label_map == class_id labels, num = ndimage.label(mask) instances = [] for i in range(1, num + 1): inst_mask = labels == i if inst_mask.sum() < 50: # 过滤掉面积过小的噪点 continue instances.append(inst_mask) return instances ann_id = 0 dataset = {"images": [], "annotations": []} for img_idx, mask_path in enumerate(glob.glob("masks/*.png")): label_map = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE) if label_map is None: continue h, w = label_map.shape image_id = img_idx dataset["images"].append({ "id": image_id, "file_name": mask_path.replace("masks", "images").replace(".png", ".jpg"), "width": w, "height": h, }) for inst_mask in split_instances(label_map, class_id=None): rle_counts, rle_size = mask_to_rle(inst_mask) pts = pick_foreground_points(inst_mask, num_points=3) ys, xs = np.where(inst_mask > 0) bbox = [int(xs.min()), int(ys.min()), int(xs.max() - xs.min()), int(ys.max() - ys.min())] dataset["annotations"].append({ "id": ann_id, "image_id": image_id, "category_id": 0, "bbox": bbox, "segmentation": {"size": rle_size, "counts": rle_counts}, "points": pts, # SAM2 训练时还会动态重新采样,这里是初始点 "labels": [1] * len(pts) # 1 表示前景点 }) ann_id += 1 with open("sam2_train.json", "w") as f: json.dump(dataset, f) print("done, annotations:", ann_id)逻辑说明:脚本先按连通域把一张图里的掩膜拆成实例,再用 distance transform 从掩膜内部挑远离边界的点作为前景提示。用距离变换而不是随机像素,是因为边缘附近的点容易被标注噪声干扰,训练时模型会更难收敛。参数num_points=3是初始点数量,训练加载器一般会在每个 epoch 重新随机采样更多点,所以这里给 3 个保证 json 结构完整即可。
mask_to_rle里必须用np.asfortranarray,pycocotools 对非 Fortran 序的数组会静默编出错误结果;split_instances里的面积阈值 50 要按你的数据调,如果你的目标最小面积只有几十像素,把它调到 10 以下。另外掩膜路径里用了字符串替换,前提是你的图像和掩膜目录在同级目录且文件名一一对应,实际项目里更稳的写法是直接维护一张 filename 映射表。
2.3 负样本与背景点怎么给:训练时不只“指对”还要“指开”
如果训练样本里只有前景点,模型会学会“点哪里哪里就是前景”,但遇到背景上的点击就茫然。SAM2 这类 prompt 模型真正要学的是:给一个前景点,把目标整块挖出来;给一个背景点,把误激活区域压制下去。所以数据集中不能只有正样本点,还得有标签为 0 的背景点。
常见做法是给每个实例额外采样一个负样本点,位置在该实例外接框之外、且落在图像内部。具体实现不复杂:把掩膜外扩一圈,从外扩区域里随机取一点;如果取到的点穿到了其他实例内部,就重新采样。个别极端情况是图像里几乎全是目标,负样本点采不出来,这时可以退化成“从掩膜边缘反向扩张 10 像素再取点”。
在 json 里对应的字段就是points和labels,labels 里 1 代表前景、0 代表背景。不同训练脚本对负样本的读取方式不太一样,有些会自动从背景区域补充,有些只会读你给的。为了保险,我一般把负样本点也写进 json:points里同时放前景点和背景点,labels跟着写 1 和 0。这样不管你用的那版代码是“读 json 点”还是“从掩膜动态采样”,都不会出现训练时只有正样本的尴尬。
3. 用 SAM2 官方 train 入口跑通自己的数据:模型选择与最小训练命令
3.1 预训练权重与模型选型:别一上来就挑最大号
SAM2 的预训练权重一般来自 SA-V 视频掩膜数据集,也有一部分人从 COCO 预训练权重开始接着训。视频掩膜数据让 SAM2 对运动目标和遮挡更敏感,但也会让它更依赖“时序一致性”;如果你的数据是静态图像,加载官方权重后最好先把视频相关模块关掉,再开始微调。
模型选型方面,我建议按显存和类别数来分,别一上来就上最大号。
| 模型尺寸 | 典型显卡 | 适合场景 | 训练建议 |
|---|---|---|---|
| tiny / small | 12GB 以下 | 快速验证、类别少、目标大 | 推荐先跑通流程 |
| base_plus | 16GB 左右 | 大部分业务场景 | 默认选择 |
| large / huge | 40GB 以上 | 小目标多、精度要求高 | 需要梯度累积和 bf16 |
如果你的类别数和现有预训练权重差异很大,比如从通用物体换成病理切片,backbone 前几层学到的边缘纹理还有用,但高层语义基本要重学。这时别只调 decoder,把 image encoder 的后两三层也放开训练,效果会明显好。反过来,如果只是从 20 类换到 25 类,冻结整个 encoder 只调 decoder 就够了,训起来快得多,也不容易过拟合。
3.2 最小训练命令:batch、lr、epoch 该看哪些地方
拿官方仓库里的 train.py 做底子是常见做法,不同分支的参数名略有差别,跑之前先python train.py --help核对一遍。下面的命令是我在项目里常用的最小启动方式:
python train.py \ --data_path ./sam2_train.json \ --output_dir ./sam2_finetune_output \ --model_cfg sam2.1_hiera_base_plus \ --pretrained_path ./weights/sam2.1_hiera_base_plus.pt \ --batch_size 4 \ --lr 1e-4 \ --epochs 20 \ --num_workers 4 \ --use_bf16 True参数说明:data_path指向第 2 章生成的 json;model_cfg填你选的模型配置名,不同仓库里这个名字可能是sam2_hiera_base_plus或sam2.1_hiera_base_plus,以你下载权重时附带的 cfg 文件名为准;pretrained_path千万不能省,不加载预训练权重直接随机初始化,20 个 epoch 基本学不出像样的 mask。lr=1e-4是大模型微调的常见起点,如果你的 batch size 小到 2 以下,把 lr 降到 5e-5 更稳。
很多人在这一步卡住是因为路径问题:json 里的file_name是相对路径,train.py 不一定按 json 所在目录解析,而是按当前工作目录解析。建议把所有路径都改成绝对路径,或者在 json 里直接用图片的完整路径,能省掉大量排查时间。还要注意num_workers别一次性开很大,SAM2 数据加载时会做 RLE 解码,CPU 占用不低,4 到 8 比较合适。
3.3 只调 decoder 还是连 backbone 一起调:两种微调模式怎么切
微调 SAM2 有两种常见模式:只调 prompt encoder 和 mask decoder,或者把 image encoder 的后几层也放开。第一种适合数据量小、和预训练分布差异不大的情况;第二种适合领域差异大、目标形态完全不同的情况。
如果你的训练代码直接操作 PyTorch 模型对象,冻结 encoder 只需要一段很短的逻辑:
# freeze_image_encoder.py # 只训练 decoder;backbone 全部参数不更新 for name, param in model.image_encoder.named_parameters(): param.requires_grad = False # 只放开最后两层时,把上面循环改成判断层名后缀 for name, param in model.image_encoder.named_parameters(): if "blocks.10" in name or "blocks.11" in name: param.requires_grad = True else: param.requires_grad = False逻辑说明:SAM2 的 image encoder 是 Hiera 结构,通常由多个 block 堆叠,后几层负责语义抽象,前几层保留边缘、纹理等底层特征。领域差异大时,把最后两个 block 放开训练,能让模型学新领域的高层语义,同时保住底层特征不崩。如果你只想做极轻量微调,思路接近 lora,连后两层都不动,只训练 mask decoder 和 prompt encoder。
要注意一个细节:requires_grad=False之后,优化器创建时要把该参数过滤掉,否则 PyTorch 会报“undefined gradient”或白白占用显存。过滤器写法一般是optimizer = torch.optim.AdamW(filter(lambda p: p.requires_grad, model.parameters()), lr=...)。如果用的官方 train.py 不支持指定可训练层,就需要在加载模型之后插入这段冻结逻辑,再去建优化器。
4. Sam2 训练自己的数据常见问题避坑:显存、不收敛、mask 全黑
4.1 三通道掩膜当灰度读:训练不崩但评测全错
现象:训练 loss 降得挺顺利,验证集 mIoU 也不算低,但可视化预测 mask 时,目标区域总在几个类别间“闪烁”,或者同一张图换一个输入尺寸结果差很多。
原因:标注工具导出的掩膜是 PNG,但用cv2.imread默认按三通道读。如果掩膜里类别 1、2、3 恰好渲染成 RGB 三个通道,IMREAD_GRAYSCALE丢失信息;反过来如果掩膜本来就是灰度 PNG,你却按三通道读,数据加载器又把它当三通道图像喂给网络,形状对不上后静默广播。这类问题训练时不报错,因为 PyTorch 对形状会自动处理,但语义已经乱了。
解决:统一用cv2.imread(path, cv2.IMREAD_GRAYSCALE)读掩膜,并在转换脚本里打印一下np.unique(label_map),确认像素值只有 0、1、2、3 这类类别 ID,没有 255 这种可视化颜色值。如果发现掩膜是彩色渲染的,先用 OpenCV 做颜色到 ID 的映射,别指望模型自己学回来。
4.2 训练输入里没有 prompt:loss 不降还越学越糊
现象:训练了十几个 epoch,loss 停在某个值不动,预测结果要么全黑,要么输出一堆“球状”块。把验证代码里加上一个点击点之后,效果突然变好,但训练时就是学不进去。
原因:前面强调过,SAM2 是 prompt 模型。如果数据加载器只返回 image 和 GT mask,没有构造点或框输入,模型等于在尝试从一个缺失条件里学分割,这是训练设置问题,不是模型能力问题。很多从像素分割转过来的人都会在这踩坑,因为传统分割网络根本不需要 prompt。
解决:回到数据加载逻辑,确认每个 batch 里有点坐标和标签。具体做法是在 Dataset 的__getitem__里从 GT 掩膜随机采 1 到 5 个前景点、0 到 2 个背景点,随 epoch 变化。如果你用官方 dataloader,检查 json 里points和labels字段是否被正确读取;如果你自己写 dataloader,一定把point_coords和point_labels放进返回的 dict。
4.3 显存不足的排查顺序:batch=1、梯度累积、bf16 这样组合
现象:一跑训练就CUDA out of memory,把 batch size 从 4 降到 2 还是不行,再降到 1 勉强能跑,但 loss 震荡得很厉害。
原因:SAM2 参数量不小,image encoder 的中间特征图显存占用很高。显存不足时最直接的做法不是换 tiny 模型,而是调整训练策略组合。
解决顺序:先把 batch size 降到 1;然后用梯度累积模拟 batch size 4 的效果,accumulation_steps=4时每 4 个 step 更新一次参数;再开 bf16 混合精度,显存能再省 20% 到 30%。如果还卡,关掉torch.compile,那张量显存占用也会回升但能直接跑稳。坚决不推荐的做法是一味加大num_workers,数据加载那点 CPU 显存开销根本不是瓶颈。组合后一般 16GB 卡也能跑 base_plus 模型。
4.4 小目标学不动:在指标里加 object-level recall
现象:验证集 mIoU 看还行,但可视化发现小目标经常丢,大目标边缘也粗糙,只是大目标占比高把 mIoU 拉上去了。
原因:mask 级别的 IoU 对大目标天然友好,一个 500×500 的实例多预测出 5% 面积,IoU 只掉一点;同样误差发生在 30×30 小目标上,IoU 直接掉到零点几。训练 loss 里 dice 项对大目标也有类似偏向,小目标在整个 batch 里贡献的梯度太小。
解决:评估时别只看 mIoU,加一个 object-level recall,计算每个实例的预测 mask 与 GT 的 IoU 是否超过 0.5,再按实例面积分桶统计。训练侧可以按 mask 面积给 loss 加权,小目标权重放大 1.5 到 2 倍;采 prompt 点时强制每个实例至少采一个点,而不是按面积比例随机采。这样小目标不是“没被看到”,而是“被看到后梯度至少能传到”。
4.5 视频数据误开 streaming 状态:memory bank 把显存吃干净
现象:明明是单张图像数据,训练速度却越来越慢,显存占用随 step 数线性增长,最后 OOM。
原因:SAM2 有 video streaming 分支,视频模式下每帧都会写 memory bank,保存之前帧的特征供后续帧查询。如果你没把数据按视频帧组织,却触发了 memory 相关逻辑,模型会把每一张独立图当成一个超长视频的连续帧,bank 越积越大。
解决:单图训练时显式关闭 memory bank。常见做法是在模型初始化后设置model.memory_encoder.eval()和model.memory_attention.eval(),或者直接把use_memory配置项关掉;用官方 train.py 的话,检查配置文件里有没有video或memory开关。如果你的数据本来就是视频序列,那就相反,要确保每帧的 mask 和时序 ID 都正确传入,否则 memory bank 学到的时序关联是乱的。
5. 验证与进阶:用评估结果反推该改数据还是改参数
5.1 一个最小验证脚本:点提示下的 mIoU 和 object recall
训练完后建议写一个和训练时输入完全一致的验证脚本,不要只输出一张“效果图”。下面这段逻辑我几乎每个项目都会改一改直接用:
# eval_sam2.py # 加载微调权重,给定 GT 里的一个前景点,评估 mask 预测质量 model.eval() with torch.inference_mode(): pred = model( image=images, # [B, 3, H, W] point_coords=point_coords, # [B, N, 2] point_labels=point_labels, # [B, N] multimask_output=False, ) pred_masks = pred["pred_masks"] # 多尺度输出时取最后一层 # 对每个实例算 IoU 和 object recall # object recall = count(IoU > 0.5) / count(instances)逻辑说明:point_coords里的坐标要从原图分辨率换算到模型输入分辨率,这是最常见的验证误差来源。很多人在训练时没问题,是因为训练加载器内部做了同样换算;验证脚本自己写时容易漏掉这一步。multimask_output=False让模型输出单个 mask,而不是三个候选 mask。
5.2 进阶调优顺序:embedding 缓存、多尺度与低学习率长训练
如果数据量不大,先把 image encoder 的 embedding 缓存下来,同一张图只过一次 backbone,训练时直接取缓存特征,训练速度能快好几倍。但注意,这只适用于冻结 encoder、只调 decoder 的模式;如果你把后两层也放开,缓存就失效了。
训练配置上,我建议按这个顺序调:先固定 lr=1e-4 跑 20 epoch,如果验证集 object recall 还在涨,就降低 lr 到 3e-5 再拉长到 50 到 60 epoch;如果 state 卡住不涨,再去动数据,比如增加负样本点、调整小目标权重。多数情况下,SAM2 这类大模型更适合低学习率长训练,而不是像训练小网络那样提高 lr 求快速收敛。多尺度输入也是有效的,把图像随机缩放到 768 到 1024 之间,能让模型对目标大小更鲁棒。
以前我总觉得这类大模型微调必须攒大量数据,后来在只有几百张图的场景里,把数据格式和评估口径调对,也能把目标漏检问题修掉一大半。关键不是堆数据,而是先跑通点、框、掩膜三通道的完整闭环。希望帮到你。
本文还有配套的精品资源,点击获取