ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

SAM2训练自己的数据:从掩码到可用模型的微调实践

SAM2训练自己的数据:从掩码到可用模型的微调实践 简介面向需要利用自定义数据训练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_fileconfigs/sam2.1/sam2.1_hiera_base_plus.yaml, ckpt_pathcheckpoints/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]], dtypenp.float32) point_labels np.array([1, 1], dtypenp.int32) masks, _, _ predictor.predict( point_coordspoint_coords, point_labelspoint_labels, multimask_outputTrue, ) best_mask masks[0] cv2.imwrite(data/annotations/img_001.png, (best_mask * 255).astype(uint8))这段代码做的事情是先加载 Sam2 的 base_plus 权重然后对单张图做一次基于两个正点的交互预测。注意multimask_outputTrue会返回多个候选 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 基础上改出自己的 Sam23.1 安装和版本配对这几个包最容易踩坑训练 Sam2 不是pip install sam2就结束的。它会依赖hydra、timm、opencv而且不同 PyTorch 版本对 checkpoint 里的权重键名有影响。我用的组合是 PyTorch 2.1 CUDA 11.8Sam2 仓库切到sam2.1分支因为它的权重和配置文件名是对应的。安装命令通常是这样git clone gitgithub.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, devicecuda) 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, lr1e-4, weight_decay1e-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_size4, shuffleTrue) 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_outputFalse, ) loss focal_dice_loss(pred[masks].squeeze(1), mask) optimizer.zero_grad() loss.backward() optimizer.step() epoch_loss loss.item() print(fepoch {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_outputFalse表示只输出单 mask因为我们做的是语义分割而不是交互式多候选。loss 用 focal 加 dice 的组合focal 负责处理正负样本不平衡dice 负责边界区域的梯度。3.3 超参数对照照着这个范围调别上来就莽我刚跑的时候把 batch 拉到 8结果 24G 显存的卡直接 OOM。后来总结了一套相对安全的参数区间列在下面。参数安全区间说明batch_size2~8取决于输入分辨率1024 以上建议 2~4输入分辨率512~1024统一缩放到 1024 训练推理时保持同尺寸learning rate5e-5 ~ 3e-4冻结 encoder 时用 1e-4 起调比较稳epochs20~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 全为 0dice 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, dilation2): from scipy.ndimage import binary_dilation pred_b binary_dilation(pred, iterationsdilation) gt_b binary_dilation(gt, iterationsdilation) 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))dilation2表示对边界做两次膨胀膨胀范围越大对边界偏移的容忍度越高。如果边界 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 就行。希望帮到你。本文还有配套的精品资源点击获取
返回列表