ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

LoRA微调SAM实战:让通用分割大模型适配特定业务场景

LoRA微调SAM实战:让通用分割大模型适配特定业务场景 简介针对Meta发布的SAM模型缺少微调功能的问题这份资源提供了一套完整的代码方案面向希望将SAM适配到特定图像分割任务如乳腺肿块分割的研究者与开发者。代码包含训练循环、LoRA轻量微调、数据加载与预处理、推理评估等模块并配有演示notebook与可视化结果图可帮助在自定义数据上快速启动微调实验。压缩包共43个文件以Python脚本为主25个py配套13张效果对比图、少量配置文件yaml/toml、依赖锁定文件及说明文档整体仅20.72MB结构清晰便于查阅。目前已有1316人学习下载适合具备基础深度学习知识、想基于SAM开展细分场景分割工作的人群。 做CV的朋友应该对SAM不陌生2023年Meta开源的Segment Anything Model用11M掩码数据训出了一个分割大模型拿点、框、文本都能分割。但真拿它跑特定业务场景比如工业缺陷检测、医学影像里的病灶分割、遥感图像里的细小目标提取很多人会发现一个尴尬的事实效果没有Demo里那么神遇到灰度图、小目标、强噪声背景误分割和漏分割相当严重。这个项目标题是“针对任何任务微调特定 SAM 模型 - 代码”核心就干一件事把“通用分割大模型”通过少量标注数据调成“懂你这个领域的分割模型”。文章会把整个微调思路、代码实现、训练细节和踩坑记录完整过一遍重点放在LoRA微调这种性价比最高的方案上。适合有PyTorch基础、用过SAM做推理、想让它落地到自有数据集上的开发者阅读。你不需要从头读大模型源码只要按这篇文章的流程走准备一份图和对应的掩码就能跑通微调全流程。1. 微调方案的整体思路与选型拆解1.1 为什么通用分割模型到了特定场景会“失灵”SAM在SA-1B这种千万级数据集上做的预训练学到的是“通用目标”的概念比如人、车、猫、狗、道路、房子这类日常可见的东西。但特定任务的数据分布和日常图片差异很大。我举两个实际案例你会瞬间明白差距在哪一是医学CT图像器官和组织在灰度图里的边界很模糊SAM习惯用RGB纹理特征来分割换个模态之后特征空间完全不匹配二是工业场景中的表面划痕目标又细又长对比度低SAM默认生成的mask会倾向于把纹理相似的区域全部圈进去。这不是模型差而是领域偏移问题。解决方案不是让SAM十项全能而是让它学会你这个任务的“视觉经验”。微调就是把原先的通用特征和新任务的特定特征做对齐数据量不用很大几百张到几千张标注就能有肉眼可见的提升。1.2 三种主流微调方式与选型逻辑在SAM上做微调业内主流做法有三种全量微调、冻结微调、LoRA微调。很多人一开始都会直接冲全量微调觉得效果最好但实际落地时往往会翻车。我先把三个方案的核心对比如下方案可训练参数量显存需求数据量需求效果表现适用场景全量微调约6亿参数vit_h极高多卡是常态1万张以上更稳上限最高但容易过拟合和灾难性遗忘数据量充足的大团队冻结微调约400万参数只训decoder较低单卡能跑千张左右提升有限视觉特征适配不足想快速验证方案时LoRA微调约3000万参数低秩适配器低单卡8GB可跑几百张就可上手接近全量微调性价比很高小数据、小显存的标准落地选择全量微调最大的问题是SAM这种预训练模型的权重会大幅偏离原始分布你微调的数据集如果只有几百张模型很快就把原来学到的通用特征“忘”光了。冻结微调只调prompt encoder和mask decoder训练压力小但视觉主干还是原来那套旧参数遇到分布差异很大的领域比如遥感仍然力不从心。LoRA的做法是不动原模型权重只在attention层旁边挂上一组低秩矩阵训练时只更新这组小矩阵。它的巧妙之处在于既保留了SAM原本学到的通用分割能力又能通过低秩矩阵快速适配新任务。我前面提到大模型领域里LLaMA-Factory这类工具也是同样的思路加载基座模型注入LoRA适配器只训适配器权重。SAM虽然是个视觉模型但按照这套思路完全可行。2. 环境准备与数据集组织2.1 依赖安装与硬件要求我用的环境是Python 3.10、PyTorch 2.1、CUDA 11.8GPU是单张RTX 2080Ti 11GB。如果你的卡是8GB显存微调vit_b_base版本也能跑只是batch size小一点。pip install torch torchvision pip install githttps://github.com/facebookresearch/segment-anything.git pip install transformers peft opencv-python albumentations tqdm这里有个细节segment-anything官方库只负责模型结构定义和checkpoint加载训练用的数据增强和LoRA注入要额外装peft库这是HuggingFace出的参数高效微调库LoRA的实现非常成熟不必自己手写前向过程中的低秩矩阵合并。模型权重我推荐先下载vit_b的sam_vit_b_01ec64.pth375MB体型小、迭代快先把流程跑通后面需要提精度再换vit_l或vit_h。不同大小的SAM结构差异只体现在image encoder层数和特征维度上微调代码完全通用。2.2 数据集统一格式设计微调SAM最容易被忽略的就是数据格式。SAM官方训练用的数据是COCO格式但小规模微调完全没必要上COCO那一套JSON标注体系我自己整理了一个极简目录data/ ├── train/ │ ├── images/ │ │ ├── img_001.jpg │ │ └── img_002.jpg │ ├── masks/ │ │ ├── img_001.png │ │ └── img_002.png │ └── train.txt └── val/ ├── images/ ├── masks/ └── val.txt图片就是原始图像mask是8位单通道PNG背景为0目标为255。train.txt里每行写一张图片的完整路径val.txt同理。这样组织的好处是Dataset类写起来极简加载时直接用OpenCV读图片、读mask不需要解析任何标注文件。mask的二值化处理是第一个容易踩坑的地方。有些标注软件导出的mask不是0和255而是0和1或者边缘带有抗锯齿的灰度渐变。训练前必须统一做一次阈值处理我习惯在Dataset的__getitem__里顺手加一句mask (mask 127).astype(np.float32)确保训练时拿到的ground truth只有0和1。3. 核心代码实现与逐段解析3.1 加载SAM模型并冻结主网络参数微调的第一步是构建模型并加载预训练权重然后把不需要更新的参数全部冻结。这里有一个关键认知基础知识库都存在image encoder里prompt encoder也有预训练好的位置编码真正需要重点调整的是图像特征到分割掩码的映射能力。所以默认策略下我冻结image encoder和prompt encoder让mask decoder从头更新。import torch from segment_anything import sam_model_registry # 加载vit_b的预训练权重 sam sam_model_registry[vit_b](checkpoint./weights/sam_vit_b_01ec64.pth) device cuda if torch.cuda.is_available() else cpu sam.to(device) # 冻结image encoder和prompt encoder for name, param in sam.named_parameters(): if image_encoder in name or prompt_encoder in name: param.requires_grad False else: param.requires_grad True代码里requires_grad False的意思是反向传播时这些参数的梯度不计算、不更新显存占用和计算量都会大幅下降。vit_b的总参数量约9000万冻结之后真正参与训练的只剩mask decoder的约400万参数这个量级非常轻量。3.2 给Image Encoder注入LoRA适配器冻结主干只是第一步下一步是往image encoder的注意力层里注入LoRA。SAM的image encoder结构是标准的ViT内部的多头注意力层里通常有qkv这么一个完整路径peft库能自动识别并替换。from peft import LoraConfig, get_peft_model lora_config LoraConfig( r8, lora_alpha16, target_modules[qkv], lora_dropout0.1, biasnone, ) sam.image_encoder get_peft_model(sam.image_encoder, lora_config)r是低秩矩阵的秩控制可训练参数量lora_alpha是缩放系数影响LoRA对原模型的干预强度target_modules指定要注入的模块名。这里把LoRA加在qkv上本质上是让模型在计算每个token的注意力权重时多学一组任务相关的投影向量。注入完成后可以用下方代码统计可训练参数量检验是否生效trainable_params sum(p.numel() for p in sam.parameters() if p.requires_grad) total_params sum(p.numel() for p in sam.parameters()) print(fTrainable params: {trainable_params / 1e6:.2f}M / Total: {total_params / 1e6:.2f}M)我在vit_b上注入r8的LoRA后可训练参数量约为3000万远小于全量微调主要memory开销集中在特征图本身而不是梯度。实测下来11GB显存跑batch size为4没有任何压力。3.3 构造训练数据流与Box提示生成SAM的推理依赖提示prompt训练时我也需要为每张图构造至少一个提示。最稳妥的方式是直接使用mask的外接矩形框这也是最符合“任意任务”场景的输入方式。用户使用时只需要框选目标区域模型就会基于box进行分割。import cv2 import numpy as np import torch from torch.utils.data import Dataset class SAMDataset(Dataset): def __init__(self, img_dir, mask_dir, file_list, size1024): self.img_dir img_dir self.mask_dir mask_dir self.size size with open(file_list, r) as f: self.samples [line.strip() for line in f.readlines()] def __len__(self): return len(self.samples) def __getitem__(self, idx): img_path self.samples[idx] img_name img_path.split(/)[-1].replace(.jpg, .png) img cv2.imread(img_path) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) mask cv2.imread(f{self.mask_dir}/{img_name}, cv2.IMREAD_GRAYSCALE) # 二值化mask避免标注边缘杂物 mask (mask 127).astype(np.float32) # 从mask计算外接矩形框作为box prompt ys, xs np.where(mask 0) if len(xs) 0 or len(ys) 0: box np.array([0, 0, mask.shape[1], mask.shape[0]]) else: x1, x2 xs.min(), xs.max() y1, y2 ys.min(), ys.max() box np.array([x1, y1, x2, y2]) # 缩放到1024并同步调整box坐标 scale self.size / max(img.shape[0], img.shape[1]) new_w, new_h int(img.shape[1] * scale), int(img.shape[0] * scale) img cv2.resize(img, (new_w, new_h)) mask cv2.resize(mask, (new_w, new_h)) box (box * scale).astype(int) # SAM固定输入尺寸1024x1024 img_resized np.ones((self.size, self.size, 3), dtypenp.uint8) * 128 mask_resized np.zeros((self.size, self.size), dtypenp.float32) img_resized[:new_h, :new_w] img mask_resized[:new_h, :new_w] mask img_tensor torch.as_tensor(img_resized, dtypetorch.float32).permute(2, 0, 1).contiguous() img_tensor (img_tensor - 127.5) / 127.5 mask_tensor torch.as_tensor(mask_resized, dtypetorch.float32).unsqueeze(0) box_tensor torch.as_tensor(box, dtypetorch.float32).unsqueeze(0) return img_tensor, mask_tensor, box_tensor这段代码有两个细节值得说。第一缩放到1024时没有直接整体缩放成方形而是用了padding补边的方式避免图片被拉伸变形SAM原版推理也是这么处理的。第二box的坐标必须同步缩放否则提示框和图像内容错位训练出来的模型对prompt的理解就会很奇怪。3.4 训练主循环与损失函数解析训练循环的核心和前向逻辑如下。SAM的前向接口是sam(image, boxes)返回的pred_masks是未经过sigmoid的logits形状为[B, 1, 256, 256]。损失函数我用了最经典的focal loss和dice loss加权组合这也是分割任务里的黄金搭档。import torch.nn.functional as F def focal_loss(pred, target, alpha0.25, gamma2.0): pred_sigmoid pred.sigmoid() pt pred_sigmoid * target (1 - pred_sigmoid) * (1 - target) focal_weight (1 - pt).pow(gamma) bce F.binary_cross_entropy_with_logits(pred, target, reductionnone) loss focal_weight * bce loss loss.mean() return loss def dice_loss(pred, target, smooth1.0): pred_sigmoid pred.sigmoid() intersection (pred_sigmoid * target).sum() loss 1 - (2.0 * intersection smooth) / (pred_sigmoid.sum() target.sum() smooth) return lossfocal loss解决的是“目标区域小、背景区域大”的正负样本不平衡问题alpha和gamma可以控制对难样本的关注程度dice loss则是直接优化交集和并集的比值更贴近分割任务最终的评估指标。两者相加能同时兼顾像素级别和区域级别的约束。训练主循环里还要做梯度裁剪和梯度累积防止显存不足时batch size被迫减小带来的训练不稳定optimizer torch.optim.AdamW( [p for p in sam.parameters() if p.requires_grad], lr1e-4, weight_decay1e-4, ) def train_one_epoch(model, dataloader, optimizer, epoch): model.train() total_loss 0.0 for i, (images, masks, boxes) in enumerate(dataloader): images, masks, boxes images.to(device), masks.to(device), boxes.to(device) pred_masks, iou_predictions model(images, boxes) loss focal_loss(pred_masks, masks) dice_loss(pred_masks, masks) loss.backward() torch.nn.utils.clip_grad_norm_( [p for p in model.parameters() if p.requires_grad], max_norm1.0 ) optimizer.step() optimizer.zero_grad() total_loss loss.item() print(fEpoch {epoch}: loss {total_loss / len(dataloader):.4f})模型结构的前向输出和原始SAM不太一样我这里省略了prompt_encoder的中间变量实际使用时要保证model(images, boxes)内部完成了box编码和mask解码的完整链路。学习率1e-4是我在多个数据集上试出来的比较稳的值LoRA参数的更新幅度不需要太大太大容易出现loss震荡。3.5 评估与模型导出训练过程中每轮结束后我用一个简单的IoU指标评估验证集效果这个比loss更直观def evaluate(model, dataloader): model.eval() ious [] with torch.no_grad(): for images, masks, boxes in dataloader: images, masks, boxes images.to(device), masks.to(device), boxes.to(device) pred_masks, _ model(images, boxes) pred_bin (pred_masks.sigmoid() 0.5).float() intersection (pred_bin * masks).sum(dim(1, 2, 3)) union pred_bin.sum(dim(1, 2, 3)) masks.sum(dim(1, 2, 3)) - intersection iou (intersection / (union 1e-6)).mean().item() ious.append(iou) return np.mean(ious)训练完成后需要把LoRA权重合并回原模型再保存。如果直接用torch.save(sam.state_dict())会把LoRA适配器当作独立参数存下来后续加载时得额外处理。我推荐先用merge_and_unload再保存合并后的完整权重from peft import PeftModel sam.image_encoder sam.image_encoder.merge_and_unload() torch.save(sam.state_dict(), ./checkpoints/sam_vit_b_custom.pt)拿这个文件做推理时基本就是普通加载SAM权重的方式代码如下sam sam_model_registry[vit_b](checkpoint./checkpoints/sam_vit_b_custom.pt) sam.to(device) sam.eval()实际推理时只需要手动给一个box模型就能输出该区域的分割mask。到这一步“针对任何任务微调SAM”的最小闭环就算完成了。4. 高频问题排查与优化建议4.1 显存溢出LoRA微调比全量微调省显存但也不是完全没压力。如果你的显卡在8GB以下或者导入了较大的batch size很可能在训练跑到第1个迭代时直接OOM。除了调低batch size最有效的方式是开启混合精度训练。PyTorch 2.0以上版本用torch.autocast包住前向和loss计算梯度和参数保持在fp32精度feature map用fp16存储显存占用能降低40%左右。如果batch size已经小到1还不够那减输入分辨率从1024改成768对结果的影响在可接受范围。4.2 训练loss不降最常见的两个原因一个是学习率设得太大模型在最优解附近反复横跳这种情况loss曲线呈现高频震荡另一个是Dataloader里的mask没做二值化灰度mask导致loss永远不为零网络学不到确定性的监督信号。排查时先把loss和metrics打印出来看如果是loss一直维持在一个平台值不下降先检查数据是否正确然后在training loop里加一个梯度值打印如果梯度过小可能是某个模块被意外冻结了。4.3 推理时mask全黑或全白这个坑我在第一次微调时也踩过。全黑基本是输出pred_mask没经过sigmoid就做了阈值判断或者box输入的坐标类型传成了int导致SAM内部的prompt编码时位置信息错乱。全白则往往是box范围太大把整张图都包含进去了或者你对prompt_encoder做了完全错误的输入。排查步骤很简单先用原版SAM权重同一张推理如果原版正常而微调后异常那是微调过程的问题如果原版也异常那是推理代码的问题。4.4 LoRA合并后精度下降不是每次都能丝滑合并偶尔会有LoRA合并后效果反而不如合并前的情况。这可能是低秩矩阵和原权重在merge时因为精度损失导致参数漂移。遇到这种问题先确认合并操作是在eval模式下进行的然后尝试不merge直接同时保存原权重和LoRA adapter权重推理时动态加载两者。这样虽然多一个文件但能精确复现训练时的模型调试起来更方便。4.5 数据增强的经验我微调时用了albumentations库做随机翻转、随机亮度对比度调整但对SAM来说最有效的增强其实是对box边界做扰动。比如在计算外接矩形框时给x1、y1、x2、y2加上一个随机偏移量模拟用户标注不精确的场景。这样训练出来的模型对输入框的鲁棒性会明显提升这也是我在遥感目标检测和工业缺陷两个项目上对比验证过的经验。4.6 什么时候该换更大的模型vit_b微调适合快速验证流程但如果你最终的任务难度很高比如目标又小又密集或者类别数量很大建议直接上vit_l甚至vit_h。模型大的时候LoRA的rank也可以适度调大一点从8改成16或32让适配器有更强的表达能力。前提是显存和训练时间能接受我一个2.5万张图的工业质检场景从vit_b换成vit_h后IoU大概提升了4个百分点训练时间从3个小时涨到11个小时这个性价比要看你的业务需求来权衡。写在最后的实操建议这套流程我已经在医学影像、遥感道路提取、工业表面缺陷三个项目上跑通过最大的体会是整个微调过程里代码本身不是核心难点真正的关键在数据质量和评估指标的定义。我的建议是拿到新任务后先别急着写训练脚本先花一个下午手工标注二三十张图跑通一遍微调和推理确认效果能接受之后再投入人力去标注大规模数据。如果你的任务本身就是纯背景的图像比如只有单个目标那直接用freeze微调就够如果目标形态多样、背景复杂LoRA会是更稳妥的中间态。最后再分享一个小技巧微调后的模型如果在某些类上表现不佳不要盲目加数据先把这类样本的mask检查一遍很多问题是标注边缘没清理干净导致的清洗数据往往比调参更管用。本文还有配套的精品资源点击获取
返回列表