☰
SAM2 结合 UNet 实现高精度图像分割:从结构设计到训练推理全解析
2026/10/11 21:47:01 网站建设 项目流程

简介:本资源面向计算机视觉方向的学习者与算法工程师,提供一套将SAM2与UNet结合的高精度图像分割项目源码,适合希望深入理解分割大模型与传统网络融合思路、并动手复现的中高级开发者。压缩包共82个文件,约999KB,以29个py源码文件为核心,辅以40个pyc编译文件、4个yaml配置、3个sh训练与评估脚本,以及pyd、cu、drawio、jpg、md等辅助文件,覆盖模型构建、数据集加载、训练与测试全流程。项目围绕SAM2UNet.py组织网络结构,包含sam2_image_predictor、automatic_mask_generator等模块,并配有train.py、eval.py、test.py及多份sam2_hiera配置,便于按需切换模型规模。目前已有119人学习。读者可从中获得完整的训练与推理代码、配置文件与脚本,理解SAM2与UNet的衔接方式,并基于现有目录结构快速开展自己的分割实验。

1. SAM2 与 UNet 合体:高精度图像分割到底在解决什么问题

做过图像分割的工程师大多经历过这种拉扯:UNet 系列在医学影像、工业质检、遥感地物这些任务上稳如老狗,边缘细节抓得准,但一旦遇到目标语义模糊、需要跨帧关联或者零样本迁移的场景,它就得靠大量标注数据重新训练;而 SAM2 这类提示式分割大模型,点一下、框一下就能出掩码,泛化能力强得离谱,可它对细粒度类别的判别又不够“懂行”,输出的是通用前景,不是你业务里定义的“病灶”“缺陷”“裂缝”。

标题里的“SAM2 结合 UNet 实现高精度图像分割”,本质上就是想把这两条路焊在一起:用 SAM2 做粗定位和候选区域生成,用 UNet 做精细语义判别和边缘回归。它适合谁?适合手里有几百到几千张标注图、显存不算宽裕、又不想从零训大模型的分割从业者。你不需要重新预训练 SAM2,也不需要把 UNet 推倒重来,核心工作在于“怎么把 SAM2 的输出变成 UNet 能吃的输入,以及怎么让两者的误差不互相放大”。这篇笔记就按这个思路,从结构设计一路写到训练、推理和踩坑,能抄的代码我尽量给全。

2. SAM2 与 UNet 的融合结构:先想清楚谁负责什么

2.1 为什么不是简单串行:SAM2 输出掩码直接喂 UNet 的三个问题

很多人第一反应是“SAM2 出掩码,UNet 再分割一遍”,听起来像级联,实际跑起来会翻车。第一个问题是 SAM2 的掩码是二值前景,没有类别通道,UNet 如果直接拿它当输入,等于把语义信息压缩成 0/1,分类头学不到“这是哪一类”。第二个问题是 SAM2 的掩码边缘偏软,尤其在低对比度区域会糊成一片,UNet 的编码器如果只看到这种软边缘,下采样几次后细节就丢了。第三个问题是显存和延迟:SAM2 的图像编码器本身就不轻,如果每张图都跑完整 SAM2 再跑 UNet,推理时间直接翻倍,工业场景很难接受。

我一般会改成“特征级融合 + 掩码提示”的双路结构:SAM2 只跑一次图像编码器,拿到多尺度特征;UNet 的主干正常走,但在跳跃连接处把 SAM2 的特征按通道拼接或注意力加权注入。这样 SAM2 提供的是“哪里可能有目标”的先验,UNet 负责“这个目标属于哪一类、边界在哪”。下面这张表是我在三个数据集上对比过的结构选型,供你参考。

融合方式显存增量推理延迟小目标 Dice适用场景
SAM2 掩码直接拼接输入低低下降 3~5%二分类、目标大
编码器特征拼接中中提升 1~2%多类别、中等目标
注意力加权注入高中高提升 2~4%小目标、边缘敏感
双路独立 + 后期融合高高提升 1~3%数据量大、显存足

选型建议:如果你做的是工业缺陷检测,缺陷往往只占几十个像素,优先选注意力加权注入;如果是遥感地物分类,目标成片出现,编码器特征拼接性价比最高。

2.2 融合模块的代码实现:一个可复现的 SAM2-UNet 主干

下面这段代码是我常用的融合模块,核心思路是把 SAM2 图像编码器的中间层特征经过 1x1 卷积对齐通道后,用空间注意力图加权到 UNet 对应尺度的跳跃连接上。代码基于 PyTorch,假设你已经能正常加载 SAM2 的 image encoder。

import torch import torch.nn as nn import torch.nn.functional as F class SAM2FeatureInjector(nn.Module): def __init__(self, sam_channels, unet_channels): super().__init__() # 1x1 卷积把 SAM2 特征通道对齐到 UNet 跳跃连接通道 self.align = nn.Conv2d(sam_channels, unet_channels, kernel_size=1) # 空间注意力:用 7x7 卷积生成单通道权重图 self.spatial_attn = nn.Sequential( nn.Conv2d(unet_channels * 2, 1, kernel_size=7, padding=3), nn.Sigmoid() ) self.bn = nn.BatchNorm2d(unet_channels) def forward(self, sam_feat, unet_skip): # sam_feat: SAM2 编码器中间层输出, unet_skip: UNet 跳跃连接特征 sam_aligned = self.align(sam_feat) # 尺寸对齐,防止下采样倍数不一致 if sam_aligned.shape[-2:] != unet_skip.shape[-2:]: sam_aligned = F.interpolate( sam_aligned, size=unet_skip.shape[-2:], mode='bilinear', align_corners=False ) # 拼接后生成空间注意力权重 concat = torch.cat([sam_aligned, unet_skip], dim=1) attn = self.spatial_attn(concat) # 加权融合:UNet 特征为主,SAM2 特征为辅 fused = unet_skip * attn + sam_aligned * (1 - attn) return self.bn(fused)

逻辑说明:align负责通道对齐,spatial_attn让网络自己学“哪些位置该信 SAM2、哪些位置该信 UNet”。参数上,sam_channels取决于你用的 SAM2 版本,常见是 256 或 1024;unet_channels对应 UNet 该层的通道数,比如 64、128、256。注意F.interpolate的align_corners=False在分割任务里更稳,别用 True,否则边缘会有半像素偏移。

2.3 训练策略:冻结 SAM2 还是联合微调

这是被问得最多的问题。我的血泪经验是:先冻结 SAM2 编码器,只训 UNet 和融合模块,等验证集 Dice 稳定后再解冻 SAM2 的最后两个 block 做小学习率微调。原因很简单,SAM2 的预训练权重是在海量数据上学的,你几百张图全量微调,灾难性遗忘几乎必然发生,表现就是训练集涨、验证集崩。

具体参数:冻结阶段学习率 1e-3,解冻阶段 SAM2 部分学习率 1e-5,UNet 部分 1e-4。优化器用 AdamW,权重衰减 1e-4。Batch size 根据显存来,8G 显存建议 2~4,配合梯度累积。损失函数用 Dice + BCE 组合,Dice 权重 0.6,BCE 权重 0.4,小目标多的话把 Dice 提到 0.7。

3. 数据准备与标注:SAM2 当预标注器怎么用才不坑

3.1 用 SAM2 生成伪标签的完整流程

如果你手里只有少量精标数据,可以用 SAM2 做半自动标注:先人工点几个点或画框,让 SAM2 生成掩码,再人工修正。这个流程能省 60% 以上的标注时间,但前提是你要控制伪标签的质量,否则噪声会污染训练集。

步骤上,我一般这样做:第一,把原始图像按 512x512 切块,重叠 64 像素,避免大图直接缩放丢细节;第二,对每块图,用 SAM2 的自动掩码生成模式跑一遍,得到候选掩码;第三,用 CLIP 或简单的颜色直方图过滤掉明显不属于目标类别的掩码;第四,人工只修正被保留的掩码,而不是从零画。下面是一个批量生成伪标签的脚本骨架。

import os import cv2 import numpy as np from sam2.build_sam import build_sam2 from sam2.automatic_mask_generator import SAM2AutomaticMaskGenerator # 加载 SAM2 模型,checkpoint 路径按你本地实际改 sam2 = build_sam2("sam2_hiera_l.yaml", "sam2_hiera_large.pt", device="cuda") mask_generator = SAM2AutomaticMaskGenerator( model=sam2, points_per_side=32, # 每边采样点数,越大越细但越慢 pred_iou_thresh=0.88, # 预测 IoU 阈值,低于此丢弃 stability_score_thresh=0.92, # 稳定性分数阈值 min_mask_region_area=100 # 最小掩码面积,过滤噪点 ) def generate_pseudo_labels(img_dir, out_dir): os.makedirs(out_dir, exist_ok=True) for fname in os.listdir(img_dir): img = cv2.imread(os.path.join(img_dir, fname)) img_rgb = cv2.cvtColor(img, cv2.COLOR_BGR2RGB) masks = mask_generator.generate(img_rgb) # 按面积排序,保留前 N 个候选 masks = sorted(masks, key=lambda x: x['area'], reverse=True)[:10] label_map = np.zeros(img.shape[:2], dtype=np.uint8) for i, m in enumerate(masks): label_map[m['segmentation']] = i + 1 cv2.imwrite(os.path.join(out_dir, fname.replace('.jpg', '.png')), label_map)

参数说明:points_per_side从 32 起步,显存不够降到 16;pred_iou_thresh和stability_score_thresh是控制伪标签质量的关键,调高会减少候选但更准,调低会引入噪声。min_mask_region_area按你的目标最小像素面积设,工业缺陷一般 50~200。

3.2 标注格式转换:从 SAM2 输出到 UNet 可读的掩码

SAM2 输出的掩码是布尔数组,UNet 训练通常需要单通道类别索引图或 one-hot。转换时最容易踩的坑是类别重叠:SAM2 可能对同一区域给出多个候选掩码,直接按顺序覆盖会导致类别错乱。我的做法是按面积从大到小排序,小的覆盖大的,因为小掩码通常是更精确的目标。转换代码如下。

import numpy as np def masks_to_label(masks, num_classes): # masks: list of dict, 每个 dict 含 'segmentation' 和 'area' # 按面积降序,先放大的,再放小的覆盖 sorted_masks = sorted(masks, key=lambda x: x['area'], reverse=True) label = np.zeros(masks[0]['segmentation'].shape, dtype=np.uint8) for idx, m in enumerate(sorted_masks): cls_id = min(idx + 1, num_classes) # 类别从 1 开始,0 为背景 label[m['segmentation']] = cls_id return label

注意num_classes要和你 UNet 输出通道一致,背景类 0 不参与 Dice 计算。如果类别数超过 255,得用 uint16,但大多数分割任务不会到这个量级。

4. 训练与推理:参数怎么设、显存怎么省

4.1 训练参数表与显存优化技巧

训练 SAM2-UNet 最现实的瓶颈是显存。SAM2 的 Hiera 编码器本身占 2~3G,UNet 再占 1~2G,加上梯度,8G 卡很容易 OOM。我常用的省显存组合是:混合精度训练(AMP)+ 梯度检查点(gradient checkpointing)+ 小 batch 梯度累积。下面这张表是我在 8G 显存下的实测配置。

参数推荐值说明
输入尺寸512x512再大显存翻倍,小目标可切块
Batch size2配合累积步数 4 等效 batch 8
精度AMP fp16显存省 30~40%,Dice 几乎不掉
梯度检查点开启SAM2 编码器部分开启,省 20% 显存
学习率1e-3(冻结)/1e-5(解冻)解冻后必须降
训练轮数80~120看验证集早停,别死磕

代码上,AMP 用torch.cuda.amp.autocast和GradScaler,梯度累积就是每 4 步optimizer.step()一次。注意 AMP 下 Dice 损失里的除法要加eps=1e-6,否则 fp16 容易出 NaN。

4.2 推理阶段:怎么把 SAM2 的提示和 UNet 的类别合起来

推理时有两种模式:一种是“提示驱动”,用户给点或框,SAM2 出候选,UNet 判类别;另一种是“全自动”,SAM2 自动掩码生成,UNet 逐候选分类。前者适合交互式标注工具,后者适合批量质检。全自动模式下,我一般会加一个后处理:对 UNet 输出的概率图做条件随机场(CRF)或简单的形态学闭运算,把 SAM2 留下的锯齿边缘磨平。

import torch import torch.nn.functional as F def inference(model, sam2_encoder, image, prompts=None): # image: 1x3xHxW, prompts: 可选,点或框 with torch.no_grad(): sam_feats = sam2_encoder(image) # 提取多尺度特征 logits = model(image, sam_feats) # UNet 融合后输出 prob = F.softmax(logits, dim=1) pred = torch.argmax(prob, dim=1) # 后处理:闭运算去毛刺 pred_np = pred.squeeze().cpu().numpy().astype('uint8') kernel = cv2.getStructuringElement(cv2.MORPH_ELLIPSE, (3, 3)) pred_np = cv2.morphologyEx(pred_np, cv2.MORPH_CLOSE, kernel) return pred_np

参数上,闭运算核大小 3x3 适合 512 输入,如果你输入更大,核可以到 5x5。别用开运算,会把小目标直接抹掉。

5. 避坑与排查:SAM2 结合 UNet 最常见的 5 个翻车现场

5.1 现象:训练 loss 正常下降,但验证集 Dice 卡在 0.6 上不去

原因:SAM2 特征和 UNet 特征尺度没对齐,融合模块学到的权重图退化成全 0 或全 1,等于没融合。解决:在融合模块后加一个nn.Identity的残差连接,让 UNet 原始特征至少能直通;同时检查F.interpolate的尺寸是否和unet_skip完全一致,差一个像素都会导致注意力图错位。

5.2 现象:推理时小目标全部漏检,大目标正常

原因:SAM2 的min_mask_region_area设太大,小目标在候选阶段就被过滤了;或者 UNet 的下采样倍数太高,小目标在编码器里已经丢光。解决:把min_mask_region_area降到 50 以下;UNet 改用轻量主干(如 MobileNetV3)并减少一次下采样,保持高分辨率特征。

5.3 现象:AMP 训练几个 epoch 后 loss 突然变 NaN

原因:Dice 损失在 fp16 下计算intersection / union时,union 可能为 0,除零导致 NaN。解决:在 Dice 计算里加eps=1e-6,并且把 Dice 损失的计算强制转成 fp32:dice_loss = dice_loss.float()。这个坑我踩过两次,每次都是训练到半夜崩掉,后悔药没得吃。

5.4 现象:SAM2 编码器解冻后,验证集指标先涨后崩

原因:学习率太大,SAM2 的预训练权重被破坏。解决:解冻部分只用 1e-5 甚至 5e-6,并且加 warmup,前 5 个 epoch 线性从 0 升到目标学习率。另外,解冻后 batch size 尽量大一点,小 batch 下 BN 统计量波动会加剧遗忘。

5.5 现象:多类别分割时,类别 1 和类别 2 的掩码互相覆盖

原因:SAM2 的候选掩码本身有重叠,转换标签时按顺序覆盖导致后处理的类别吃掉前面的。解决:用masks_to_label里按面积降序的策略,并且在 UNet 输出后加一个 softmax 前的类别互斥约束,或者用多标签 Dice 分别监督每个类别。

6. 进阶技巧:用测试时增强把 Dice 再抬 2 个点

如果你已经把上面的流程跑通,验证集 Dice 到了 0.85 左右,想再往上走,测试时增强(TTA)是性价比最高的手段。具体做法:推理时对同一张图做水平翻转、垂直翻转、90 度旋转,分别跑模型,把概率图平均后再取 argmax。这个技巧对边缘敏感的任务特别有效,因为不同视角下 SAM2 的候选掩码会有差异,平均能抵消一部分随机误差。

代码实现上,注意翻转后的概率图要翻回来再平均:

def tta_inference(model, sam2_encoder, image): probs = [] for flip in [None, 'h', 'v', 'hv']: img_aug = image if flip == 'h': img_aug = torch.flip(image, dims=[3]) elif flip == 'v': img_aug = torch.flip(image, dims=[2]) elif flip == 'hv': img_aug = torch.flip(image, dims=[2, 3]) with torch.no_grad(): sam_feats = sam2_encoder(img_aug) logits = model(img_aug, sam_feats) prob = F.softmax(logits, dim=1) if flip == 'h': prob = torch.flip(prob, dims=[3]) elif flip == 'v': prob = torch.flip(prob, dims=[2]) elif flip == 'hv': prob = torch.flip(prob, dims=[2, 3]) probs.append(prob) return torch.argmax(torch.mean(torch.stack(probs), dim=0), dim=1)

参数上,TTA 的倍数建议 4 倍,8 倍收益递减且推理时间线性增长。另外,TTA 只适合验证和最终推理,训练时别用,否则显存和时间都吃不消。

最后说个我自己的习惯:每次改完融合模块或损失函数,先拿 20 张图跑一遍过拟合测试,如果 20 张都过拟合不到 0.99 Dice,说明结构有 bug,别急着上全量数据。这个习惯帮我省了至少两周的无效训练。希望帮到你。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询