简介:面向需要利用自定义数据训练Sam2模型的机器学习开发者,这份资源提供了完整的数据封装示例,聚焦从原始数据集到模型可训练格式的转化流程。资源共2个文件,均为Python脚本,压缩包约5KB。两个脚本分别承担通用数据集创建与预处理逻辑,以及针对LabPicsV1数据集的定制封装,涵盖数据清洗、格式统一、旋转裁剪等增强操作,并明确划分训练集与验证集,便于直接套用或改造到自身任务中。已有310人学习下载。通过阅读脚本,读者可快速理解Sam2训练数据的组织方式,掌握封装、数据增强及训练/验证划分的工程技巧,避免从零踩坑,适合具备一定Python基础、希望深入Sam2微调或扩展应用的开发者。整体流程从数据加载、增强到批量划分均有体现,能为后续模型微调提供直接参考。
1. sam2训练自己的数据:从一个边界不齐的掩码到能用的模型
做了几年分割模型,我最深的感受是:像sam2这种交互式分割模型,真正磨人的不是跑通训练脚本,而是让模型在你自己的数据上稳定出边界。你拿官方权重推理一张猫图效果惊艳,但换成工厂缺陷、遥感地物、医学切片之后,画面里的边缘立刻变得敷衍。所谓sam2训练自己的数据,本质上就是拿少量带标注或者半标注的样本,微调那一套已经很强的 prompt encoder 和 mask decoder,让它学会你场景里的尺度、纹理和边界习惯。它跟跑yolov5训练自己的数据集的节奏很像,但多出一个“提示”维度——点、框、mask都参与训练。这篇笔记我按数据准备、训练脚本、避坑、验证这条链来写,适合手里已有几十张标注图但被训练细节卡住的人。
2. 数据准备:把零散图片组织成 Sam2 能直接读的形态
2.1 先定目录和标注文件结构,别信“自动识别”
我一开始跑通 befor 的时候图省事,直接把图片塞进一个文件夹,想靠 Sam2 自己读路径,结果折腾半天发现它的训练入口需要一份标注索引文件。常见做法是做成类似 COCO 但更扁的结构:一个根目录里放 images 和 annotations,再配一份 JSON 索引。JSON 里每条记录要包含图像路径、类别、以及以多边形或 mask 文件路径形式存在的标注。
我这里用最小可跑的结构来举例:
sam2_finetune/ ├── data/ │ ├── images/ │ │ ├── img_001.jpg │ │ ├── img_002.jpg │ └── annotations/ │ ├── img_001.png │ ├── img_002.png └── train_index.json{ "annotations": [ { "image_path": "data/images/img_001.jpg", "seg_path": "data/annotations/img_001.png", "category": "defect" }, { "image_path": "data/images/img_002.jpg", "seg_path": "data/annotations/img_002.png", "category": "defect" } ] }这样写的好处是后续无论是转成 video 帧序列还是按 batch 读取都很直接。image_path和seg_path用相对路径,方便在不同机器之间迁移,不用改绝对路径。类别字段先留着,San2 本身不强制类别语义,但后续如果你想加 prompt 类别条件,这就是扩展口。
2.2 标注没到像素级?让 Sam2 自己先出一版粗糙 mask
这是 Sam2 区别于传统分割训练的地方:你不必先手绘完整多边形。我处理一批工业零件数据时,先用 Sam2 的交互式分割在每张图上点几个关键点,让模型把目标大概框出来,再手动把明显错误的边界删掉。这里的关键是点选的位置要覆盖目标的两端,比如长条缺陷要分别在头部、中断、尾部各点一次,否则 mask 容易只包住中间段。
from sam2.build_sam import build_sam2 from sam2.sam2_image_predictor import SAM2ImagePredictor predictor = SAM2ImagePredictor(build_sam2( config_file="configs/sam2.1/sam2.1_hiera_base_plus.yaml", ckpt_path="checkpoints/sam2.1_hiera_base_plus.pt")) image = cv2.imread("data/images/img_001.jpg") predictor.set_image(image) point_coords = np.array([[120, 80], [240, 160]], dtype=np.float32) point_labels = np.array([1, 1], dtype=np.int32) masks, _, _ = predictor.predict( point_coords=point_coords, point_labels=point_labels, multimask_output=True, ) best_mask = masks[0] cv2.imwrite("data/annotations/img_001.png", (best_mask * 255).astype("uint8"))这段代码做的事情是先加载 Sam2 的 base_plus 权重,然后对单张图做一次基于两个正点的交互预测。注意multimask_output=True会返回多个候选 mask,我习惯取第一个,但更稳的做法是检查masks里的 score 数组,选分数最高的。best_mask存成 8 位灰度 PNG 时,Sam2 输出是 float 的 0~1 矩阵,记得乘 255 再转 uint8,否则就成了全黑图。
2.3 训练前做一次数据体检:先过滤,再进训练管线
很多人在 Sam2 上翻车不是因为网络结构,而是因为标注里混进了大面积只有几像素的小碎片。分割训练对这类噪声很敏感。我一般会在训练前跑一遍面积和边缘检查,统计每张 mask 的像素占比,把小于全图 0.5% 的目标单独拎出来人工确认。
import cv2 import numpy as np import json index = json.load(open("train_index.json")) for item in index["annotations"]: mask = cv2.imread(item["seg_path"], 0) > 0 ratio = mask.mean() contours, _ = cv2.findContours( (mask * 255).astype("uint8"), cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE ) small_parts = [c for c in contours if cv2.contourArea(c) < 50] if ratio < 0.005 or len(small_parts) > 0: print("需要复查:", item["image_path"], "目标占比", round(ratio, 4))这个脚本不改变任何数据,只负责把有问题的样本列出来。ratio小于 0.005 的目标在训练时对 loss 的贡献极小,模型学不到边界特征;small_parts超过 0 则说明标注里有碎片噪声,真要大面积标注可以用形态学开运算先清一遍。这一步做好,后续训练省下的时间远大于你在这里花的十分钟。
3. 训练脚本拆解:在官方 hf_segment 基础上改出自己的 Sam2
3.1 安装和版本配对:这几个包最容易踩坑
训练 Sam2 不是pip install sam2就结束的。它会依赖hydra、timm、opencv,而且不同 PyTorch 版本对 checkpoint 里的权重键名有影响。我用的组合是 PyTorch 2.1 + CUDA 11.8,Sam2 仓库切到sam2.1分支,因为它的权重和配置文件名是对应的。安装命令通常是这样:
git clone git@github.com:facebookresearch/sam2.git cd sam2 pip install -e . pip install opencv-python timm hydra-core注意pip install -e .会把包安装成编辑模式,源码改一下就直接生效,方便调试,但也会有缓存问题——如果你改过 Sam2 源码却感觉没生效,先检查是不是装成了非编辑模式。另外权重文件记得放在checkpoints/目录下,因为配置文件里写的是相对路径,文件放错位置会在加载时报“找不到 checkpoint”。
3.2 最小训练脚本逐段注释:冻结 encoder,只训解码器
Sam2 参数量很大,完全端到端微调不是不行,但成本高,收敛也慢。常见做法是冻结 image encoder,只训练 mask decoder 和 prompt encoder,这样显存压力小,几十张图也能跑出可用效果。
import torch from sam2.build_sam import build_sam2 from sam2.sam2_image_predictor import SAM2ImagePredictor from torch.utils.data import Dataset, DataLoader import cv2, json, numpy as np config_file = "configs/sam2.1/sam2.1_hiera_base_plus.yaml" checkpoint = "checkpoints/sam2.1_hiera_base_plus.pt" sam2_model = build_sam2(config_file, checkpoint, device="cuda") torch.cuda.empty_cache() for name, param in sam2_model.image_encoder.named_parameters(): param.requires_grad = False trainable = [p for p in sam2_model.parameters() if p.requires_grad] optimizer = torch.optim.AdamW(trainable, lr=1e-4, weight_decay=1e-4) class SegDataset(Dataset): def __init__(self, index_file): with open(index_file, "r") as f: self.data = json.load(f)["annotations"] def __len__(self): return len(self.data) def __getitem__(self, idx): item = self.data[idx] img = cv2.imread(item["image_path"]) img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) mask = cv2.imread(item["seg_path"], 0) > 0 img = torch.from_numpy(img.transpose(2, 0, 1)).float() / 255.0 mask = torch.from_numpy(mask).float() return img, mask def focal_dice_loss(pred, target): pred = torch.sigmoid(pred) focal = -target * torch.log(pred + 1e-6) - (1 - target) * torch.log(1 - pred + 1e-6) focal = focal.mean() inter = (pred * target).sum() dice = 1 - (2 * inter + 1) / (pred.sum() + target.sum() + 1) return focal + dice dataloader = DataLoader(SegDataset("train_index.json"), batch_size=4, shuffle=True) for epoch in range(30): epoch_loss = 0 for img, mask in dataloader: img, mask = img.cuda(), mask.cuda() with torch.no_grad(): image_embeddings = sam2_model.image_encoder(img) pred = sam2_model.mask_decoder( image_embeddings, sam2_model.prompt_encoder.get_dense_pe(), multimask_output=False, ) loss = focal_dice_loss(pred["masks"].squeeze(1), mask) optimizer.zero_grad() loss.backward() optimizer.step() epoch_loss += loss.item() print(f"epoch {epoch}, loss {epoch_loss / len(dataloader):.4f}")逐段解释一下。冻结image_encoder后,每次 forward 它的输出是固定的,所以我在循环里用with torch.no_grad()包住image_encoder,然后把计算好的 embedding 送给mask_decoder,这样可以省掉 encoder 的反向传播开销。get_dense_pe()返回稠密位置编码,它是 prompt encoder 的内部状态,作为 decoder 的条件输入。multimask_output=False表示只输出单 mask,因为我们做的是语义分割而不是交互式多候选。loss 用 focal 加 dice 的组合,focal 负责处理正负样本不平衡,dice 负责边界区域的梯度。
3.3 超参数对照:照着这个范围调,别上来就莽
我刚跑的时候把 batch 拉到 8,结果 24G 显存的卡直接 OOM。后来总结了一套相对安全的参数区间,列在下面。
| 参数 | 安全区间 | 说明 |
|---|---|---|
| batch_size | 2~8 | 取决于输入分辨率,1024 以上建议 2~4 |
| 输入分辨率 | 512~1024 | 统一缩放到 1024 训练,推理时保持同尺寸 |
| learning rate | 5e-5 ~ 3e-4 | 冻结 encoder 时用 1e-4 起调比较稳 |
| epochs | 20~50 | 数据量小就 30 轮上下,看 loss 是否平台期 |
| 混合精度 | amp fp16 | 能省一半显存,但 loss 出现 nan 时关掉排查 |
batch_size的取值要跟分辨率一起看,不是单纯越大越好。我一般先用 512 分辨率跑通,再切到 1024 看收益。多数情况下 512 已经够做业务验证,1024 是给最终模型留的余量。
4. 避坑与常见问题:第一次跑 Sam2 训练最容易翻车的四个点
4.1 现象:OOM,而且把 batch 降到 1 还是炸
如果你把 batch 都降到 1 还爆显存,问题基本不在 batch,而在分辨率。我遇到过一张 4000×3000 的遥感图直接塞进去,embedding 特征图大得离谱,显存瞬间吃满。解决方法是先做 resize 到 1024 以内,而且训练和推理要保持同一套预处理。mmrotate训练 DOTA 数据集时大家都有过类似经验,遥感图直接训是走不通的。
4.2 现象:loss 在下降,但输出的 mask 全是黑的
这是最典型的“看起来在训练,其实没学到东西”的情况。原因多半是标注 mask 是 0/255 的灰度图,而代码里用了阈值 127 之前直接转 float,被误判成全零。我之前写过一个转换脚本,忘了> 0的判断,结果 target 全为 0,dice loss 反而一路往下掉。解决方法是跑数据体检脚本,统计训练集里 mask 的非零像素比例,确保正样本占比不是 0。
4.3 现象:权重加载时报 key 不匹配,或者尺寸对不上
这个坑主要出在换了配置,比如你下载的是 sam2.1 权重,但配置文件仍指向旧的 sam2_hiera_base_plus。Sam2 和 Sam2.1 的 decoder 结构不完全一致,load_state_dict自然会报错。解决方法是严格匹配权重与配置文件的版本前缀,比如sam2.1_hiera_base_plus.yaml对应sam2.1_hiera_base_plus.pt,同时检查 checkpoint 是完整模型还是只存了 decoder 的增量权重。
4.4 现象:单图预测效果好,一到视频帧序列就崩
Sam2 同时支持图像和视频,但视频训练走的是另一套数据逻辑,需要帧序列和帧间 mask 对应关系。如果只做图像交互分割,训练时务必用SAM2ImagePredictor而不是SAM2VideoPredictor。我见过有人拿 video predictor 去读单张图,输出的 mask 带时间维,后处理直接错乱。记住一个原则:图像微调用 image predictor,视频微调才用 video predictor,两者不混用。
5. 验证与导出:别被单张图的可视化结果骗了
5.1 写一个验证脚本,在未见过的图上算边界 IoU
单看几张训练集的预测图,永远都是好的。尤其是 Sam2 这种强交互模型,你在推理时点选的位置跟训练时相近,效果自然好,但换个角度点选可能就崩了。所以我习惯在验证集上同时算 mask IoU 和边界 IoU。边界 IoU 更能反映 Sam2 的边界敏感度。
from skimage.metrics import adapted_rand_error import numpy as np def boundary_iou(pred, gt, dilation=2): from scipy.ndimage import binary_dilation pred_b = binary_dilation(pred, iterations=dilation) gt_b = binary_dilation(gt, iterations=dilation) inter = np.logical_and(pred_b, gt_b).sum() union = np.logical_or(pred_b, gt_b).sum() return inter / union # pred_mask: model output, gt_mask: ground truth print("boundary IoU:", round(boundary_iou(pred_mask, gt_mask), 4))dilation=2表示对边界做两次膨胀,膨胀范围越大,对边界偏移的容忍度越高。如果边界 IoU 明显低于 mask IoU,说明模型内部区域分得不错,但边界不贴,这时优先增强训练数据里边缘清晰的样本,而不是继续加训练轮数。
5.2 导出:把训练好的权重存成部署可用的格式
训练完的模型权重是一个完整的 Sam2 状态字典,部署时如果只做交互分割,可以只保留image_encoder、prompt_encoder、mask_decoder三个子模块导出为 TorchScript,这样体积小、加载快。另一种做法是直接存成state_dict,推理时再加载构建模型,但会依赖原来的 Python 类定义。
torch.save({ "image_encoder": sam2_model.image_encoder.state_dict(), "prompt_encoder": sam2_model.prompt_encoder.state_dict(), "mask_decoder": sam2_model.mask_decoder.state_dict(), }, "sam2_finetuned_weights.pth")这里只存子模块,不存完整模型,加载时先build_sam2构建原始结构,再用load_state_dict把三个子模块灌回去。这么做的好处是部署时可以不依赖训练脚本里的自定义 loss 函数。
6. 进阶技巧:把交互式点选变成你的数据生产引擎
6.1 用点选迭代修正低质量 mask
训练数据不足时,与其花一晚上人工抠图,不如用已经微调过的模型配合人工点选来扩数据。每张图先随机点三个点生成 mask,人工只看边界,哪条边不对就在哪边补一个点。这样一张图的标注时间能从十分钟压到两分钟以内。补出来的新 mask 再回填到训练集,迭代两三轮,模型会越来越懂你场景里的边界。
6.2 微调数据量参考:几十张也能起步
很多人卡在“没有数据集就不敢训练”,实际上 Sam2 微调对数据量要求没那么高。
| 数据规模 | 推荐策略 |
|---|---|
| 20~50 张 | 冻结 encoder,只训 decoder,重点把边界学准 |
| 50~200 张 | 可以放开部分 encoder 层,或者加 LoRA 微调视觉主干 |
| 200 张以上 | 端到端小学习率训练,数据增强要跟上 |
数据量少时不要贪多,把lr降到 5e-5,防止过拟合。有人还把这种思路类比成 lora 训练,本质上都是用小数据微调大模型的关键分支,视觉这边没有 LoRA 那么流行的现成方案,但直接调低学习率同样有效。
6.3 检查点保存与恢复训练的习惯
从那以后,我每次训练都会强制走一遍“每 5 轮存一个 checkpint + 记录验证集边界 IoU”的流程。不是为了别的,就是防止跑了一宿之后发现 loss 曲线已经平台期,却拿不出一个可回退的中间版本。训练中断也不用从头再跑,把--resume指向最近的 checkpoint 就行。希望帮到你。
本文还有配套的精品资源,点击获取