☰
SAM2与UNet融合实战:图像分割边界精修与mIoU提升指南
2026/10/11 17:03:48 网站建设 项目流程

简介:本资源面向计算机视觉方向的开发者与图像分割学习者,提供一套将SAM2与UNet结合的高精度分割算法完整项目源码,适合具备一定深度学习基础、希望深入理解分割模型融合思路的中高级读者参考实践。压缩包共82个文件,约999KB,以29个py源码文件为核心,辅以40个pyc编译文件、4个yaml配置、3个sh脚本及pyd、cu、drawio、jpg、md等类型,涵盖模型构建、数据集处理、训练与评估等模块,结构清晰便于按需查阅。项目包含SAM2UNet主模型、图像与视频预测器、自动掩码生成器及多套SAM2配置,并配有训练、测试、评估脚本与流程图,方便读者快速复现实验、理解模型融合细节与调参思路。目前已有119人学习,适合作为分割算法实战与二次开发的参考案例。

1. 分割算法落地:SAM2 与 UNet 结合到底解决了什么痛点

做过图像分割的工程师大多有过这种体验:UNet 在固定域数据上表现稳定,换一批光照、尺度、背景复杂度不同的图,边缘就开始糊,小目标直接漏。SAM2 的出现让零样本分割能力上了一个台阶,但它对提示(prompt)高度依赖,点、框、掩码给得不好,结果就飘。把两者拼起来,用 UNet 做粗定位和语义约束,用 SAM2 做精细边界回归,是当前工业界比较务实的一条路线。这个方案适合谁?适合手里有几百到几千张标注图、需要在医疗影像、工业质检、遥感或电商抠图场景里把 mIoU 从 0.7 推到 0.85 以上的团队。它不要求你从头训一个基础模型,但要求你理解两套网络的输出怎么对齐、损失怎么配、推理时显存怎么控。下面按「先立住原理、再跑通最小闭环、最后调参避坑」的顺序拆开讲,源码结构也会在中间章节给出可抄的骨架。

2. SAM2 与 UNet 的分工逻辑:为什么不是二选一

2.1 两套网络各自擅长什么、短板在哪

UNet 的核心优势在于编码器-解码器结构加跳跃连接,能把浅层纹理和深层语义拼在一起,对训练域内的类别边界非常敏感。但它的泛化能力受限于标注数据的分布,遇到训练集里没出现过的形状、遮挡或低对比度区域,分割结果往往出现「语义对、边界错」的情况。SAM2 则相反,它在海量数据上预训练过,具备强零样本边界感知能力,给它一个粗略的框或点,它能还你一条相当干净的边缘。但 SAM2 不懂你的类别语义,它不知道「这个区域是病灶还是正常组织」,也没有类别标签输出。所以常见做法是:UNet 负责出类别概率图和粗掩码,SAM2 负责在粗掩码附近做边界精修。这样既保留了 UNet 的语义判别力,又借了 SAM2 的边界先验。

2.2 融合位置的选择:像素级、特征级还是决策级

融合位置决定了实现复杂度和最终收益。像素级融合最直接:UNet 输出粗掩码,二值化后取连通域外接框,作为 SAM2 的 box prompt,SAM2 输出精细掩码,再与 UNet 的类别图做逐像素加权。特征级融合需要把 SAM2 的图像编码器特征和 UNet 解码器特征做对齐拼接,对显存和训练技巧要求更高,适合数据量充足、追求极致指标的团队。决策级融合则是两套结果做投票或 CRF 后处理,实现最快但提升有限。我一般推荐从像素级入手,原因是改动小、可解释、容易回退。下面给一个像素级融合的最小推理流程,先跑通再谈优化。

import torch import torch.nn.functional as F from segment_anything import sam2_model_registry, SamPredictor # 假设 unet 已加载并 eval,sam2 使用官方注册的 tiny 或 base 版本 unet = load_unet(checkpoint="unet_best.pth").eval().cuda() sam2 = sam2_model_registry["sam2_hiera_b+"](checkpoint="sam2_hiera_b+.pt").cuda().eval() predictor = SamPredictor(sam2) def fuse_predict(image_tensor, unet, predictor, box_pad=8, mask_thr=0.5): # image_tensor: 1x3xHxW, 已归一化 with torch.no_grad(): coarse_logit = unet(image_tensor) # 1xCxHxW coarse_mask = (coarse_logit.argmax(1) > 0).float() # 1xHxW 二值前景 # 取最大连通域的外接框作为 prompt ys, xs = torch.where(coarse_mask[0] > 0) if len(xs) == 0: return coarse_mask box = torch.tensor([xs.min()-box_pad, ys.min()-box_pad, xs.max()+box_pad, ys.max()+box_pad]).cpu().numpy() # SAM2 需要 RGB numpy 输入 img_np = (image_tensor[0].permute(1,2,0).cpu().numpy() * 255).astype("uint8") predictor.set_image(img_np) fine_mask, _, _ = predictor.predict(box=box[None, :], multimask_output=False) fine_mask = torch.from_numpy(fine_mask[0]).float().cuda() # 像素级加权:UNet 语义概率与 SAM2 边界掩码相乘 prob = F.softmax(coarse_logit, dim=1)[:, 1] # 前景概率 fused = prob * fine_mask + prob * (1 - fine_mask) * 0.3 return (fused > mask_thr).float()

这段代码的关键参数有三个:box_pad控制外接框外扩像素,太小会切掉边界,太大会引入背景干扰,通常取 5 到 15;mask_thr是最终二值化阈值,0.5 是起点,类别不平衡时往 0.6 到 0.7 调;multimask_output=False表示只取 SAM2 的最高分掩码,如果目标内部有孔洞或分离区域,可以设为 True 再按面积筛选。逻辑上先让 UNet 出粗掩码,再用粗掩码的包围盒去提示 SAM2,最后把 UNet 的语义概率和 SAM2 的边界掩码做加权融合,既保留了类别信息,又让边缘更贴真实轮廓。

3. 从零跑通训练闭环:数据、损失与两阶段调度

3.1 数据准备与标注格式对齐

这套方案对数据的要求比纯 UNet 略高,因为 SAM2 的 prompt 质量依赖粗掩码的准确性。常见做法是准备三份内容:原始图像、像素级类别掩码、以及由掩码自动生成的边界框文件。掩码格式建议用单通道 PNG,像素值就是类别 id,0 为背景。如果原始标注是 COCO 多边形,先转成掩码再统一尺寸。下面这个脚本把 VOC 风格的 XML 或 COCO json 转成训练用的掩码和框,注意类别映射要固定,否则 UNet 输出通道和后续融合会对不上。

import os, json, numpy as np from PIL import Image, ImageDraw def coco_to_mask(coco_json, img_dir, out_mask_dir, out_box_dir, class_map): os.makedirs(out_mask_dir, exist_ok=True) os.makedirs(out_box_dir, exist_ok=True) data = json.load(open(coco_json)) img_info = {im["id"]: im for im in data["images"]} boxes = {} for ann in data["annotations"]: img_id = ann["image_id"] info = img_info[img_id] mask_path = os.path.join(out_mask_dir, info["file_name"].replace(".jpg", ".png")) if not os.path.exists(mask_path): Image.new("L", (info["width"], info["height"]), 0).save(mask_path) mask = Image.open(mask_path) draw = ImageDraw.Draw(mask) for seg in ann["segmentation"]: poly = [(seg[i], seg[i+1]) for i in range(0, len(seg), 2)] draw.polygon(poly, fill=class_map[ann["category_id"]]) mask.save(mask_path) x, y, w, h = ann["bbox"] boxes.setdefault(info["file_name"], []).append([x, y, x+w, y+h]) for fname, bxs in boxes.items(): np.save(os.path.join(out_box_dir, fname.replace(".jpg", ".npy")), np.array(bxs))

class_map把 COCO 的 category_id 映射到 1 到 N 的连续整数,背景固定为 0。out_box_dir里存的框会在第二阶段作为 SAM2 的 prompt 监督信号。注意多边形转掩码时 PIL 的draw.polygon对自相交多边形会填充异常,遇到这种情况先用 shapely 做 buffer(0) 清理。

3.2 损失函数配置:Dice、CE 与边界损失的配比

UNet 分支的损失不能只用交叉熵,否则小目标会被背景淹没。我一般用 CE 加 Dice 再加一个边界加权项。CE 负责像素分类,Dice 拉正样本召回,边界项用 Laplacian 或形态学梯度提取边缘区域,给边缘像素更高权重。配比上,CE 权重 1.0,Dice 权重 1.0,边界项 0.5 起步。如果验证集上边缘 mIoU 明显低于区域 mIoU,把边界项提到 1.0。SAM2 分支在训练时通常冻结图像编码器,只微调 prompt 编码器和掩码解码器,学习率设成 UNet 的十分之一,避免破坏预训练边界先验。

import torch.nn as nn import torch.nn.functional as F class ComboLoss(nn.Module): def __init__(self, ce_w=1.0, dice_w=1.0, edge_w=0.5): super().__init__() self.ce_w, self.dice_w, self.edge_w = ce_w, dice_w, edge_w self.ce = nn.CrossEntropyLoss() def edge_map(self, mask): # 简单形态学梯度:膨胀减腐蚀 k = torch.ones(1,1,3,3, device=mask.device) dil = F.max_pool2d(mask.float(), 3, 1, 1) ero = -F.max_pool2d(-mask.float(), 3, 1, 1) return (dil - ero).clamp(0,1) def forward(self, logit, target): ce_loss = self.ce(logit, target) prob = F.softmax(logit, dim=1)[:, 1] tgt = (target > 0).float() inter = (prob * tgt).sum() dice_loss = 1 - (2*inter + 1e-6) / (prob.sum() + tgt.sum() + 1e-6) edge = self.edge_map(tgt.unsqueeze(1)).squeeze(1) edge_loss = (F.binary_cross_entropy(prob.clamp(1e-6,1-1e-6), tgt, reduction="none") * edge).mean() return self.ce_w*ce_loss + self.dice_w*dice_loss + self.edge_w*edge_loss

edge_map用最大池化和反向最大池化近似膨胀腐蚀,比调 OpenCV 更省事且可导。edge_loss只对边缘像素算 BCE,权重由edge_w控制。如果训练初期 loss 震荡,先把edge_w设为 0,等 CE 和 Dice 稳定后再加回来。

3.3 两阶段训练调度与显存控制

直接端到端训 UNet 加 SAM2 对显存要求很高,常见做法是两阶段。第一阶段只训 UNet,用 ComboLoss,batch size 能开多大开多大,直到验证集 mIoU 不再涨。第二阶段冻结 UNet 编码器,只微调解码器和 SAM2 的 prompt 相关模块,此时把 UNet 输出的粗掩码转成框,作为 SAM2 的 prompt 输入,损失只算 SAM2 掩码输出与真值的 Dice。显存不够时用梯度累积,累积步数 4 到 8,同时把 SAM2 图像编码器换成 tiny 版本。推理阶段可以只保留 UNet 加 SAM2 解码器,图像编码器输出缓存一次即可,避免每张图重复编码。

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

4.1 现象:融合后边缘反而比纯 UNet 更毛糙

原因通常是 SAM2 的 prompt 框给得太大,把背景纹理也框进去了,SAM2 在框内做二分类时把背景误判成前景。解决方法是收紧box_pad,或者用 UNet 概率图做一次阈值过滤,只保留概率大于 0.7 的连通域再取框。另一个可能是 SAM2 输入图像没有做正确的归一化,SAM2 期望 RGB 0 到 255 的 uint8,如果传了 0 到 1 的 float,边界会整体偏移。

4.2 现象:小目标在融合结果里直接消失

UNet 对小目标的粗掩码可能只有几个像素,取外接框后 SAM2 的 prompt 太小,掩码解码器输出全零。解决方法是设一个最小框尺寸,比如宽高都不小于 16 像素,不够就按中心点扩展。同时检查 UNet 的损失里 Dice 权重是否太低,小目标被 CE 淹没了。可以在采样时对含小目标的图做 oversampling,或者在损失里给小目标像素额外权重。

4.3 现象:训练 loss 正常但验证集 mIoU 卡在 0.6 不涨

先看数据里有没有类别不平衡或标注噪声。SAM2 对噪声标注很敏感,如果粗掩码本身在边界处抖动,SAM2 会放大这种抖动。常见做法是用形态学开闭运算平滑一下 UNet 输出的粗掩码再取框,或者用 CRF 做一次后处理。另外检查第二阶段学习率,如果和第一阶段一样大,SAM2 的预训练权重会被快速破坏,边界先验丢失,表现反而不如纯 UNet。

4.4 现象:推理速度从 30 FPS 掉到 3 FPS

SAM2 图像编码器是主要瓶颈,尤其 hiera large 版本。如果业务对实时性有要求,换 tiny 或 base 版本,或者把图像编码器做成 TensorRT 引擎。另一个隐藏开销是每张图都重新set_image,如果视频流相邻帧差异小,可以每隔几帧才更新一次图像嵌入,中间帧复用。UNet 侧用半精度推理也能省不少时间,但注意 SAM2 的某些算子对 fp16 支持不完整,需要逐层验证。

4.5 现象:换一批新数据后 SAM2 分支输出全空

这通常是因为新数据的图像均值方差和训练域差异大,UNet 粗掩码本身就不准,导致 prompt 框落在背景上。解决方法是先在新域上做少量微调,或者用无监督域适应方法对齐特征。如果没法微调,退化成纯 UNet 推理,至少保证语义结果可用。另一个检查点是 SAM2 的输入尺寸,它内部会 resize 到 1024,如果原图长宽比极端,resize 后目标变形,prompt 框坐标要按比例映射回去。

5. 进阶技巧:用掩码质量打分做自适应融合

跑通基础融合后,真正拉开差距的是「什么时候信 SAM2、什么时候信 UNet」。我一般会算一个掩码质量分,综合 UNet 前景概率均值、SAM2 输出掩码的稳定性和两者 IoU。如果 UNet 概率均值高且 SAM2 掩码与粗掩码 IoU 大于 0.7,就按 0.7 比 0.3 加权偏向 SAM2;如果 IoU 低于 0.4,说明两者分歧大,回退到 UNet 结果并标记该图待人工复核。下面这个打分函数可以直接嵌到推理流程里。

def mask_quality(unet_prob, sam_mask, coarse_mask): # unet_prob: HxW 前景概率, sam_mask/coarse_mask: HxW 二值 conf = unet_prob.mean().item() inter = ((sam_mask > 0) & (coarse_mask > 0)).sum().item() union = ((sam_mask > 0) | (coarse_mask > 0)).sum().item() + 1e-6 iou = inter / union # 稳定性:SAM2 掩码面积与粗掩码面积比,偏离 1 太多说明不稳定 area_ratio = sam_mask.sum().item() / (coarse_mask.sum().item() + 1e-6) stability = 1 - min(abs(area_ratio - 1), 1) score = 0.4*conf + 0.4*iou + 0.2*stability return score, iou def adaptive_fuse(unet_prob, sam_mask, coarse_mask, thr=0.6): score, iou = mask_quality(unet_prob, sam_mask, coarse_mask) if score > thr and iou > 0.7: return 0.7*sam_mask + 0.3*coarse_mask elif iou < 0.4: return coarse_mask # 分歧大,回退 else: return 0.5*sam_mask + 0.5*coarse_mask

conf是 UNet 对前景的自信程度,iou衡量两个掩码的一致性,stability惩罚面积突变。阈值thr和 IoU 分界可以根据验证集画 PR 曲线来定,我通常把thr设在 0.55 到 0.65 之间。这套自适应策略在工业质检数据上把误检率压了将近三成,代价是每张图多算一次 IoU,开销可以忽略。

验证方法上,不要只看整体 mIoU,要分区域看:边界带(真值边缘外扩 3 像素)的 mIoU、小目标(面积小于 32×32)的 mIoU、以及不同光照子集的 mIoU。如果边界带提升明显但小目标下降,说明融合权重对小目标不友好,回到 4.2 去调最小框尺寸。我自己的习惯是每次改完融合参数,固定跑一遍这三个子集,记录成表格,避免被整体指标骗了。这套方案值不值得投入?如果你手里的数据标注质量可控、业务对边界精度有硬要求、且能接受推理时多一个 SAM2 编码器的开销,它比从头设计一个新网络要稳得多。希望帮到你。

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

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

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

立即咨询