ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

SAM2与UNet融合实战:图像分割边界精修与mIoU提升指南

SAM2与UNet融合实战:图像分割边界精修与mIoU提升指南 简介本资源面向计算机视觉方向的开发者与图像分割学习者提供一套将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 promptSAM2 输出精细掩码再与 UNet 的类别图做逐像素加权。特征级融合需要把 SAM2 的图像编码器特征和 UNet 解码器特征做对齐拼接对显存和训练技巧要求更高适合数据量充足、追求极致指标的团队。决策级融合则是两套结果做投票或 CRF 后处理实现最快但提升有限。我一般推荐从像素级入手原因是改动小、可解释、容易回退。下面给一个像素级融合的最小推理流程先跑通再谈优化。import torch import torch.nn.functional as F from segment_anything import sam2_model_registry, SamPredictor # 假设 unet 已加载并 evalsam2 使用官方注册的 tiny 或 base 版本 unet load_unet(checkpointunet_best.pth).eval().cuda() sam2 sam2_model_registry[sam2_hiera_b](checkpointsam2_hiera_b.pt).cuda().eval() predictor SamPredictor(sam2) def fuse_predict(image_tensor, unet, predictor, box_pad8, mask_thr0.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(boxbox[None, :], multimask_outputFalse) fine_mask torch.from_numpy(fine_mask[0]).float().cuda() # 像素级加权UNet 语义概率与 SAM2 边界掩码相乘 prob F.softmax(coarse_logit, dim1)[:, 1] # 前景概率 fused prob * fine_mask prob * (1 - fine_mask) * 0.3 return (fused mask_thr).float()这段代码的关键参数有三个box_pad控制外接框外扩像素太小会切掉边界太大会引入背景干扰通常取 5 到 15mask_thr是最终二值化阈值0.5 是起点类别不平衡时往 0.6 到 0.7 调multimask_outputFalse表示只取 SAM2 的最高分掩码如果目标内部有孔洞或分离区域可以设为 True 再按面积筛选。逻辑上先让 UNet 出粗掩码再用粗掩码的包围盒去提示 SAM2最后把 UNet 的语义概率和 SAM2 的边界掩码做加权融合既保留了类别信息又让边缘更贴真实轮廓。3. 从零跑通训练闭环数据、损失与两阶段调度3.1 数据准备与标注格式对齐这套方案对数据的要求比纯 UNet 略高因为 SAM2 的 prompt 质量依赖粗掩码的准确性。常见做法是准备三份内容原始图像、像素级类别掩码、以及由掩码自动生成的边界框文件。掩码格式建议用单通道 PNG像素值就是类别 id0 为背景。如果原始标注是 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_okTrue) os.makedirs(out_box_dir, exist_okTrue) 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[i1]) for i in range(0, len(seg), 2)] draw.polygon(poly, fillclass_map[ann[category_id]]) mask.save(mask_path) x, y, w, h ann[bbox] boxes.setdefault(info[file_name], []).append([x, y, xw, yh]) 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.0Dice 权重 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_w1.0, dice_w1.0, edge_w0.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, devicemask.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, dim1)[:, 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, reductionnone) * edge).mean() return self.ce_w*ce_loss self.dice_w*dice_loss self.edge_w*edge_lossedge_map用最大池化和反向最大池化近似膨胀腐蚀比调 OpenCV 更省事且可导。edge_loss只对边缘像素算 BCE权重由edge_w控制。如果训练初期 loss 震荡先把edge_w设为 0等 CE 和 Dice 稳定后再加回来。3.3 两阶段训练调度与显存控制直接端到端训 UNet 加 SAM2 对显存要求很高常见做法是两阶段。第一阶段只训 UNet用 ComboLossbatch 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 FPSSAM2 图像编码器是主要瓶颈尤其 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, thr0.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_maskconf是 UNet 对前景的自信程度iou衡量两个掩码的一致性stability惩罚面积突变。阈值thr和 IoU 分界可以根据验证集画 PR 曲线来定我通常把thr设在 0.55 到 0.65 之间。这套自适应策略在工业质检数据上把误检率压了将近三成代价是每张图多算一次 IoU开销可以忽略。验证方法上不要只看整体 mIoU要分区域看边界带真值边缘外扩 3 像素的 mIoU、小目标面积小于 32×32的 mIoU、以及不同光照子集的 mIoU。如果边界带提升明显但小目标下降说明融合权重对小目标不友好回到 4.2 去调最小框尺寸。我自己的习惯是每次改完融合参数固定跑一遍这三个子集记录成表格避免被整体指标骗了。这套方案值不值得投入如果你手里的数据标注质量可控、业务对边界精度有硬要求、且能接受推理时多一个 SAM2 编码器的开销它比从头设计一个新网络要稳得多。希望帮到你。本文还有配套的精品资源点击获取
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进