SAM2+UNet高精度图像分割:从选型到TTA实战

发布时间:2026/10/11 20:36:53
SAM2+UNet高精度图像分割:从选型到TTA实战
简介这份资源面向计算机视觉方向的学习者与算法工程师聚焦图像分割任务提供一套将SAM2与UNet结合的高精度分割算法完整项目源码适合希望深入理解分割模型结构、动手复现并用于课程设计或工程实战的中高级开发者。压缩包共82个文件约999KB以29个Python源码文件为核心涵盖模型构建、数据集处理、训练与评估脚本另有4个YAML配置文件、3个Shell脚本及1个CUDA算子文件辅以少量编译产物与说明文档整体结构清晰、便于按模块阅读。项目围绕SAM2UNet.py组织主干网络配套train、test、eval等入口脚本并附带模型结构图与README说明方便读者快速理清数据流与训练流程。目前已有119人学习可作为分割算法入门到进阶的参考范例帮助读者掌握从配置到推理的完整实现思路。1. 从一次抠图翻车说起SAM2 加 UNet 这套组合到底在做什么上个月帮朋友处理一批商品图背景是那种半透明的磨砂玻璃边缘还带反光。我一开始图省事直接拿现成的分割模型跑结果玻璃边缘全糊成一团反光区域被当成背景切掉了。后来换成 SAM2 做粗定位、UNet 做精细边缘回归同一批图 IoU 从 0.71 拉到 0.89。这就是这套「SAM2 UNet 高精度图像分割算法」项目源码想解决的问题用 SAM2 的强泛化能力圈出目标大致范围再用 UNet 的编码器-解码器结构把边缘、细小结构一点点抠回来。它适合谁做电商抠图、医疗影像预处理、遥感地物提取、工业质检的从业者只要你的场景对边缘精度有要求又不想从零标注几万张 mask这套思路就能直接抄。源码包里是完整的训练和推理脚本不是只给一个模型权重让你猜怎么用。下面我按「先搞懂为什么这么搭 → 再动手跑通 → 最后避开我踩过的坑」的顺序拆一遍。2. 为什么是 SAM2 加 UNet选型逻辑与网络结构拆解2.1 SAM2 负责「找得到」UNet 负责「切得准」SAM2 的核心是提示驱动的分割你给一个点、一个框或者一个粗略 mask它就能把目标区域推出来。它的优势在于零样本泛化——没见过的东西也能分个大概。但问题也在这SAM2 的输出分辨率受限于它的 mask decoder边缘往往偏软遇到细长结构比如电线、血管、裂缝容易断。UNet 正好补这个短板。它的 skip connection 把浅层的高分辨率特征直接送到解码器边缘和纹理信息不会在降采样里丢掉。所以这套组合的分工很明确SAM2 的输出当作先验 mask和原图一起送进 UNet让 UNet 去学「SAM2 哪里分对了、哪里分错了、边缘该往哪收」。常见做法是两阶段训练第一阶段冻结 SAM2只训 UNet第二阶段把 SAM2 的 image encoder 用低学习率解冻联合微调。我一般会把第二阶段的学习率设成第一阶段的十分之一不然 SAM2 的预训练权重容易被带偏。2.2 网络结构先验 mask 怎么喂进 UNet很多人第一次搭这个结构会卡在一个地方SAM2 输出的 mask 是单通道概率图直接和 RGB 三通道拼在一起变成四通道输入UNet 的第一层卷积核维度要改。但更稳的做法是把 mask 先做一次归一化和二值化再作为额外的条件通道而不是简单 concat。import torch import torch.nn as nn class SAM2UNet(nn.Module): def __init__(self, sam2_encoder, unet_decoder, mask_channels1): super().__init__() self.sam2_encoder sam2_encoder # 冻结或低学习率微调 self.unet_decoder unet_decoder # 把先验 mask 映射到和图像特征同维度 self.mask_proj nn.Sequential( nn.Conv2d(mask_channels, 64, kernel_size3, padding1), nn.BatchNorm2d(64), nn.ReLU(inplaceTrue) ) def forward(self, image, prior_mask): # image: [B,3,H,W], prior_mask: [B,1,H,W] img_feat self.sam2_encoder(image) # SAM2 图像编码特征 mask_feat self.mask_proj(prior_mask) # 先验 mask 投影 fused img_feat mask_feat # 逐元素相加避免通道爆炸 out self.unet_decoder(fused) return out这段代码的关键在mask_proj和fused这两步。mask_proj把单通道 mask 升到 64 维和 SAM2 编码器输出的浅层特征对齐用加法而不是 concat是因为 concat 会让后续卷积核参数量翻倍小数据集上更容易过拟合。prior_mask在训练时用 SAM2 的预测结果推理时也用同一套流程保证训练和推理一致。参数上注意两点mask_proj里的 BatchNorm 在 batch size 小于 8 时统计量不稳建议换成 GroupNormfused之后如果显存吃紧可以在 UNet decoder 的第一层之前加一个 1x1 卷积降维。2.3 损失函数Dice 和 BCE 怎么配比分割任务里单独用 BCE 会让小目标被淹没单独用 Dice 在训练初期梯度不稳。这套源码里用的是 Dice BCE 加权权重比一般设 0.6:0.4 到 0.8:0.2 之间。我的经验是目标占比小于 5% 时 Dice 权重要拉到 0.8 以上否则模型会倾向于全预测背景。class DiceBCELoss(nn.Module): def __init__(self, dice_weight0.7): super().__init__() self.dice_weight dice_weight self.bce nn.BCEWithLogitsLoss() def forward(self, pred, target): bce_loss self.bce(pred, target) pred_sigmoid torch.sigmoid(pred) intersection (pred_sigmoid * target).sum() dice_loss 1 - (2 * intersection 1e-6) / (pred_sigmoid.sum() target.sum() 1e-6) return self.dice_weight * dice_loss (1 - self.dice_weight) * bce_loss1e-6是平滑项防止分母为零。dice_weight这个参数没有绝对最优值我一般先在验证集上跑三组0.5、0.7、0.9看 IoU 曲线选收敛最稳的那个。3. 把源码跑起来环境、数据、训练三步走3.1 环境依赖与版本对齐这套源码依赖 PyTorch、SAM2 官方实现和几个图像处理库。SAM2 对 PyTorch 版本比较敏感常见做法是锁 2.1 以上、CUDA 11.8 或 12.1。我踩过的坑是 torchvision 版本和 PyTorch 不匹配导致 SAM2 加载权重时报unexpected key查了半天以为是模型结构问题其实是版本错位。conda create -n sam2unet python3.10 -y conda activate sam2unet pip install torch2.1.2 torchvision0.16.2 --index-url https://download.pytorch.org/whl/cu118 pip install opencv-python pillow numpy tqdm tensorboard # SAM2 按官方仓库方式安装注意不要混用 pip 和源码安装装完之后先跑一个最小验证加载 SAM2 权重随便找张图给个点看能不能出 mask。这一步过了再往下走不然训练脚本报错你分不清是环境问题还是代码问题。3.2 数据组织图像和 mask 的命名要对齐源码默认的数据集结构是images/和masks/两个文件夹文件名一一对应。mask 是单通道 PNG前景 255、背景 0。如果你的原始标注是 JSON 或者 COCO 格式需要先转成这种配对结构。import os import cv2 import numpy as np def convert_coco_to_mask(coco_json, image_dir, output_mask_dir): # 解析 COCO 标注按 image_id 聚合多边形 # 每个 image_id 生成一张同名 PNG mask os.makedirs(output_mask_dir, exist_okTrue) # ... 解析逻辑略核心是 cv2.fillPoly for img_id, anns in annotations.items(): file_name image_info[img_id][file_name] h, w image_info[img_id][height], image_info[img_id][width] mask np.zeros((h, w), dtypenp.uint8) for ann in anns: for seg in ann[segmentation]: pts np.array(seg).reshape(-1, 2).astype(np.int32) cv2.fillPoly(mask, [pts], 255) cv2.imwrite(os.path.join(output_mask_dir, file_name.replace(.jpg, .png)), mask)转换完一定要抽查几张确认 mask 和原图对齐。我遇到过标注里有多边形自相交fillPoly出来的 mask 中间有空洞训练时模型学到的就是错的。另外注意文件名后缀统一jpg和png混用会导致 DataLoader 找不到文件。3.3 训练脚本参数怎么改源码的train.py里几个关键参数batch_size、lr、epochs、sam2_freeze_epochs。sam2_freeze_epochs控制前多少轮冻结 SAM2一般设总 epoch 的 1/3。lr在冻结阶段可以设 1e-3解冻后降到 1e-4。python train.py \ --data_root ./datasets/mydata \ --batch_size 8 \ --lr 1e-3 \ --epochs 60 \ --sam2_freeze_epochs 20 \ --dice_weight 0.7 \ --img_size 512 \ --save_dir ./checkpointsimg_size设 512 是精度和显存的平衡点设 1024 边缘更细但 batch size 要降到 2 以下。save_dir里会存每个 epoch 的权重和 TensorBoard 日志训练时盯着验证集 IoU如果连续 10 轮不涨就说明该调学习率或者加数据了。4. 避坑与排查我在这套流程里翻过的五次车4.1 现象训练 loss 正常下降但验证 IoU 一直 0.3 左右原因通常是先验 mask 和原图没对齐。SAM2 推理时如果用了 resizemask 回到原图尺寸时会有偏移UNet 学到的是错位的先验。解决方法是统一在 SAM2 推理阶段就保持原图尺寸或者在 resize 后把 mask 用最近邻插值还原不要用双线性。4.2 现象推理时边缘出现锯齿状抖动这是 UNet 解码器最后一层上采样用了转置卷积导致的棋盘效应。换成双线性插值加 3x3 卷积或者用 PixelShuffle边缘会平滑很多。源码里默认是转置卷积我一般会手动改掉。4.3 现象显存溢出batch size 降到 1 还是 OOMSAM2 的 image encoder 在 1024 分辨率下很吃显存。除了降img_size还可以用梯度累积batch_size2累积 4 次等效于 8。另外检查有没有在训练循环里保留计算图比如把 loss 累加到一个列表里忘了 detach。4.4 现象SAM2 权重加载后提示 missing keys大概率是 SAM2 版本和权重不匹配。官方有 sam2_hiera_tiny、small、base_plus、large 几个规格源码里默认用 base_plus你下的权重如果是 tiny 就会缺 key。确认sam2_config和权重文件对应别混用。4.5 现象小目标分割效果差大目标还行除了调 Dice 权重还要检查 UNet 的深监督有没有开。源码里如果没加深监督可以在解码器每个阶段加一个辅助 loss让浅层也直接学边缘。另外img_size太小会让小目标在降采样后只剩几个像素适当放大输入分辨率比改网络结构更直接。5. 进阶技巧用测试时增强把 IoU 再抬两个点训练跑通之后如果还想压榨精度测试时增强TTA是性价比最高的手段。做法很简单推理时对同一张图做水平翻转、垂直翻转、多尺度缩放把多次预测的 mask 平均后再二值化。这套源码的inference.py里预留了 TTA 开关但默认是关的。def tta_predict(model, image, prior_mask, scales[1.0, 1.25, 1.5]): preds [] for scale in scales: h, w image.shape[-2:] new_h, new_w int(h * scale), int(w * scale) img_s torch.nn.functional.interpolate(image, size(new_h, new_w), modebilinear) mask_s torch.nn.functional.interpolate(prior_mask, size(new_h, new_w), modenearest) with torch.no_grad(): out model(img_s, mask_s) out torch.nn.functional.interpolate(out, size(h, w), modebilinear) preds.append(torch.sigmoid(out)) # 水平翻转 out_f model(torch.flip(img_s, dims[-1]), torch.flip(mask_s, dims[-1])) out_f torch.flip(out_f, dims[-1]) out_f torch.nn.functional.interpolate(out_f, size(h, w), modebilinear) preds.append(torch.sigmoid(out_f)) return torch.stack(preds).mean(dim0)scales不要设太多三档足够再多收益递减还拖慢推理。翻转增强对对称目标有效如果你的目标有明确方向性比如文字垂直翻转反而会掉点建议只保留水平翻转。平均之后用 0.5 做阈值二值化如果验证集上最优阈值不是 0.5可以在验证集上扫一遍 0.3 到 0.7选 IoU 最高的那个存下来。还有一个容易被忽略的点TTA 之后 mask 边缘会变模糊如果下游任务对边缘锐度有要求可以再做一次导向滤波或者用原图做联合双边滤波。我一般会在 TTA 之后加一步cv2.ximgproc.guidedFilter边缘贴合度会更好。从那以后我每次跑分割任务都会先把 TTA 开关打开跑一遍验证集确认增益再决定要不要在推理服务里常驻。这套 SAM2 UNet 的源码我前后在三个项目里复用最大的感受是先验 mask 的质量决定上限UNet 的结构决定能不能摸到上限而 TTA 和损失权重这些细节决定你离上限还差几个点。希望帮到你。本文还有配套的精品资源点击获取