SAM2微调实战:从数据标注到LoRA训练全流程
简介这份资源面向希望基于SAM2框架训练自有数据的深度学习开发者与研究者聚焦于将特定数据集封装为模型可识别格式这一关键环节。压缩包内共2个Python脚本文件整体约5KB分别承担通用数据集创建与针对LabPicsV1的自定义封装任务覆盖数据预处理、格式统一、归一化及旋转裁剪等数据增强操作并涉及训练集与验证集的划分逻辑。已有310人学习下载适合具备一定Python与深度学习基础、需要快速搭建SAM2训练数据管线的读者参考。通过这两个脚本读者可以理解如何将原始图片数据整理成符合SAM2输入要求的dataset结构掌握数据增强在防止过拟合、提升鲁棒性方面的具体落地方式并借鉴针对特定数据集进行适配封装的工程思路从而减少从零编写数据加载代码的试错成本。1. SAM2 训练自己的数据从标注到微调一条能跑通的路径SAM2Segment Anything Model 2在图像和视频分割上的零样本能力确实惊艳但真把它往自己的业务数据上一放翻车是常态医学影像里边界模糊的病灶、工业质检里反光金属表面的划痕、遥感图里细碎的道路零样本推理出来的 mask 要么糊成一片要么漏得离谱。原因不复杂——SAM2 的预训练数据以自然场景为主你的领域分布跟它差得远。想让 SAM2 真正干活就得拿自己的数据做微调。这篇笔记讲的就是 SAM2 训练自己的数据这条链路数据怎么标、格式怎么转、环境怎么搭、参数怎么调、显存不够怎么办、训完怎么验证。适合已经跑过 SAM2 推理、手里有一批领域图像或视频、准备动手微调的工程师如果你连推理都还没跑通建议先把官方 demo 走一遍再回来。2. 先搞清楚 SAM2 微调到底在训什么2.1 SAM2 的三个组件与可训练部分SAM2 的结构可以拆成三块图像编码器Image Encoder、提示编码器Prompt Encoder、掩码解码器Mask Decoder。图像编码器通常是个 ViT 主干把输入图压成特征提示编码器处理点、框、掩码这类 prompt掩码解码器负责融合特征和 prompt输出分割结果。视频场景下还多了一个记忆注意力模块用来跨帧传播 mask。微调时你面对的第一个决策是训哪部分。全量微调图像编码器显存开销极大一张 1024×1024 的图过 ViT-L 主干激活值就能吃掉十几 G。常见做法是冻结图像编码器只训提示编码器和掩码解码器这样显存需求能降一个量级收敛也快。如果你的领域跟自然图像差异特别大比如灰度医学图、多光谱遥感冻结主干可能不够那就考虑解冻最后几个 transformer block 做部分微调或者用 LoRA 挂在主干上。我一般会先跑冻结主干的版本看验证集 IoU 能不能到可用水平不行再逐步解冻。提示冻结主干时图像编码器只做一次前向可以把它的输出缓存下来复用训练速度能快好几倍。但缓存会占磁盘数据量大时权衡一下。2.2 数据格式从标注到 SAM2 能吃的输入SAM2 官方训练代码期望的数据格式核心是每张图配一个 mask 以及对应的 prompt 信息。实际落地时你的原始标注可能是 COCO 多边形、LabelMe JSON、或者干脆是二值 PNG。不管哪种最终都要转成「图像 二值掩码 提示点/框」的组合。这里有个容易忽略的点SAM2 是 promptable 的训练时喂进去的 prompt 质量直接影响模型学到的东西。如果你只用框做 prompt模型就偏向框引导的分割如果混入点模型对点引导的鲁棒性会更好。我的习惯是训练时随机在 mask 内部采样正点、在外部采样负点同时保留框让模型见到多种 prompt 组合。下面是一个把二值 PNG 掩码转成 SAM2 训练样本的脚本骨架import numpy as np from PIL import Image import torch def mask_to_prompt_points(mask, num_pos3, num_neg3): 从二值 mask 中采样正负提示点 ys, xs np.where(mask 0) if len(xs) 0: return None, None # 正点在 mask 内部随机采 pos_idx np.random.choice(len(xs), sizemin(num_pos, len(xs)), replaceFalse) pos_points np.stack([xs[pos_idx], ys[pos_idx]], axis1) # (N, 2) x,y # 负点在 mask 外随机采简单做法是在全图随机取再过滤 h, w mask.shape neg_points [] while len(neg_points) num_neg: x, y np.random.randint(0, w), np.random.randint(0, h) if mask[y, x] 0: neg_points.append([x, y]) neg_points np.array(neg_points) return pos_points, neg_points def build_sample(image_path, mask_path): img np.array(Image.open(image_path).convert(RGB)) mask np.array(Image.open(mask_path).convert(L)) mask (mask 127).astype(np.uint8) # 二值化阈值按你的标注调整 pos_pts, neg_pts mask_to_prompt_points(mask) if pos_pts is None: return None # 框 prompt由 mask 外接矩形得到 ys, xs np.where(mask 0) box np.array([xs.min(), ys.min(), xs.max(), ys.max()]) return { image: img, mask: mask, pos_points: pos_pts, neg_points: neg_pts, box: box, }逻辑说明mask_to_prompt_points负责从掩码里采正负点正点直接取 mask 内像素坐标负点在 mask 外随机采。build_sample把图像、掩码、点、框打包成一个样本。参数上num_pos和num_neg控制每次采几个点太少模型学不到多样 prompt太多会拖慢训练我一般正负各 3 到 5 个起步。二值化阈值 127 是个经验值如果你的标注边缘有抗锯齿可以调高到 200 减少噪声。2.3 训练集、验证集、测试集怎么切分割任务的数据切分比分类更讲究因为同一场景的相邻帧或同一病灶的多张切片如果被切到不同集合验证指标会虚高。视频数据要按视频切不能按帧随机切医学影像要按病人切不能按切片切。常见做法是 7:1.5:1.5 或 8:1:1小数据集可以 6:2:2。切完之后验证集和测试集要人工过一遍确认没有跟训练集同源的样本混进去。这一步偷懒后面指标好看但上线就崩血泪经验。3. 环境搭建与最小训练闭环3.1 依赖安装与显存预估SAM2 官方仓库依赖 PyTorch、torchvision以及一些辅助库。安装时最容易踩的坑是 CUDA 版本和 PyTorch 版本对不上导致编译扩展失败。我的习惯是先确定 CUDA 版本再去 PyTorch 官网找对应命令不要直接pip install torch了事。# 以 CUDA 12.1 为例先装匹配的 PyTorch pip install torch torchvision --index-url https://download.pytorch.org/whl/cu121 # 再装 SAM2 及其依赖按官方仓库的 requirements 走 pip install -e .[train]显存预估方面冻结图像编码器、只训解码器时batch size 1、1024×1024 输入大约需要 8 到 12 G 显存如果解冻主干同样配置可能飙到 24 G 以上。显存不够时优先降输入分辨率SAM2 支持 512 输入但小目标会受影响再考虑梯度累积模拟大 batch。3.2 配置文件里必须改的几个参数SAM2 训练配置通常是个 YAML 或 Python dict里面参数一大堆但真正影响你这次微调的就那么几个。下面这张表是我每次都会检查的参数作用建议值备注lr学习率1e-4 ~ 5e-5只训解码器可以大一点解冻主干要降batch_size批大小1 ~ 4受显存限制配合梯度累积num_epochs训练轮数20 ~ 50小数据集容易过拟合早停freeze_image_encoder是否冻结主干True先 True 跑通再考虑解冻prompt_typeprompt 类型pointbox混合 prompt 泛化更好resolution输入分辨率1024降分辨率省显存但伤小目标学习率是最玄学的参数。只训解码器时1e-4 通常能收敛如果 loss 震荡厉害降到 5e-5。解冻主干后主干的学习率要比解码器低一个量级否则预训练权重被冲垮模型直接退化。3.3 跑通第一个 epoch命令与日志观察配置改好后启动训练python training/train.py \ --config configs/sam2_finetune.yaml \ --data_root /path/to/your/dataset \ --output_dir ./runs/exp1 \ --batch_size 2 \ --num_epochs 30启动后重点看三样东西loss 曲线、显存占用、数据加载速度。loss 在前几百步下降是正常的如果一直平着不动检查学习率是不是太小或者数据没喂进去。显存占用用nvidia-smi盯着如果接近上限就降 batch 或分辨率。数据加载速度看日志里的data_time如果比batch_time还大说明 dataloader 的num_workers设小了或者磁盘 IO 是瓶颈。注意第一个 epoch 不要急着看指标先确认 loss 在降、显存没爆、没有 NaN。模型训练报 NaN 是常见问题多半是学习率太大或者数据里有异常值比如全黑的 mask、坐标越界。4. 避坑与排查那些让我重跑训练集的坑4.1 现象loss 正常下降但验证集 IoU 极低原因最常见的是数据泄漏或标注格式错位。比如图像和 mask 文件名对不上代码按排序读取时错位了或者 mask 的 0/1 反了模型学的是背景。还有一种隐蔽情况验证集的 prompt 采样方式和训练集不一致训练时用点框验证时只用框指标自然掉。解决先肉眼可视化几个训练样本把图像、mask、prompt 点叠在一起看。确认无误后检查验证集的 prompt 生成逻辑是否和训练一致。我一般会写个小脚本把 dataloader 吐出来的第一个 batch 存成图这是排查数据问题的后悔药。4.2 现象训练到一半 loss 突然爆掉原因学习率过大、梯度爆炸或者某个 batch 里出现了尺寸异常的数据。SAM2 对输入尺寸敏感如果混入了非 1024 的图而没做 resize位置编码会对不上直接产生异常梯度。解决加梯度裁剪clip_grad_norm设 1.0并在 dataloader 里强制 resize 到统一尺寸。如果已经爆了从最近的 checkpoint 恢复把学习率降一半再跑。4.3 现象显存明明够却报 OOM原因PyTorch 的缓存分配器会预留显存nvidia-smi看到的占用可能比实际需求高。另外如果验证阶段没加torch.no_grad()验证也会建计算图显存翻倍。解决验证循环包torch.no_grad()训练前设torch.cuda.empty_cache()。如果还是 OOM用torch.cuda.memory_summary()看是谁在占显存通常是某个中间变量没释放。4.4 现象小目标分割效果差大目标还行原因SAM2 的图像编码器下采样倍率较高小目标在特征图上只剩几个像素解码器再强也无力回天。这是架构层面的限制不是调参能完全解决的。解决提高输入分辨率或者对小目标区域做裁剪后单独训练一个模型。另一个思路是在 prompt 里多给小目标正点引导解码器关注该区域。如果业务里小目标占比高可能得考虑换更适合的架构SAM2 不是万能药。4.5 现象视频分割时 mask 跨帧抖动原因记忆注意力模块对运动剧烈或遮挡的场景处理不好帧间特征关联断了。训练时如果只喂单帧模型没学到时序一致性。解决训练数据里加入连续帧样本让记忆模块见到时序关系。推理时可以用后处理做时序平滑比如对 mask 做滑动窗口投票。这个坑在视频场景里特别常见单帧指标好看不代表视频可用。5. 进阶用 LoRA 和缓存把微调成本压下来5.1 LoRA 挂在图像编码器上的最小改法全量微调图像编码器成本太高LoRA 是个性价比很高的替代。思路是在 ViT 的注意力层里插入低秩矩阵只训这两个小矩阵主干权重冻住。这样可训练参数能降到原来的百分之几显存和训练时间都大幅下降。import torch.nn as nn class LoRALinear(nn.Module): def __init__(self, original_linear, rank8, alpha16): super().__init__() self.original original_linear self.rank rank self.alpha alpha in_dim original_linear.in_features out_dim original_linear.out_features # 低秩矩阵 A 和 B self.lora_A nn.Parameter(torch.zeros(in_dim, rank)) self.lora_B nn.Parameter(torch.zeros(rank, out_dim)) nn.init.kaiming_uniform_(self.lora_A, a5**0.5) nn.init.zeros_(self.lora_B) # B 初始为 0保证初始时 LoRA 分支不改变输出 def forward(self, x): base self.original(x) lora (x self.lora_A) self.lora_B * (self.alpha / self.rank) return base lora逻辑说明LoRALinear包住原来的线性层前向时把原输出和 LoRA 分支相加。lora_B初始化为 0保证训练开始时模型行为和原模型一致不会一上来就破坏预训练特征。rank控制低秩维度8 或 16 是常用值alpha是缩放系数一般设成 rank 的两倍。替换时遍历图像编码器的注意力层把qkv或proj线性层换成LoRALinear即可。5.2 缓存图像特征加速训练冻结图像编码器时每张图的特征其实是不变的可以提前算好存磁盘训练时直接读特征跳过主干前向。这样单步训练时间能降一半以上尤其适合数据集不大但想多跑几轮的场景。# 预计算特征并保存 import os def cache_features(model, dataloader, cache_dir): os.makedirs(cache_dir, exist_okTrue) model.image_encoder.eval() with torch.no_grad(): for i, batch in enumerate(dataloader): imgs batch[image].cuda() feats model.image_encoder(imgs) torch.save(feats.cpu(), os.path.join(cache_dir, f{i}.pt))逻辑说明把图像编码器切到 eval 模式关掉梯度逐 batch 算特征存下来。训练时 dataloader 直接返回缓存的特征和对应的 mask、prompt。注意缓存的特征要跟图像一一对应文件名或索引别搞乱否则又是数据错位的老坑。缓存占磁盘1024 分辨率下每张图特征大概几 MB数据集大时算一下磁盘够不够。5.3 验证微调是否真的有效训完之后别只看 loss要做三组对比零样本 SAM2、微调后的 SAM2、以及一个简单的基线比如用传统方法或小 U-Net。在同一测试集上算 IoU 和边界 F1。如果微调后比零样本提升明显说明方向对了如果提升微弱检查是不是数据量太少或领域差异不够大。我习惯把预测结果可视化几十张肉眼看边界质量指标有时候会骗人尤其是小目标被平均掉的时候。微调 SAM2 这件事值不值得做取决于你的领域跟自然图像的差距和标注数据的量。差距大、数据有几百张以上微调收益通常明显差距小、数据就几十张可能调 prompt 和后处理更划算。我自己踩过最深的坑是急着上全量微调结果显存爆了三天最后发现冻结主干加 LoRA 效果差不多还省一半时间。先跑最小闭环再逐步加码这个习惯帮我省了很多重跑训练集的时间。希望帮到你。本文还有配套的精品资源点击获取