简介:面向医学图像处理研究者与深度学习开发者的Unet改进方案资源,通过引入SAM提示框机制针对性提升息肉肿瘤语义分割精度。资源包内含2000个文件,以1992张jpeg医学图像数据集为主体,附带5个py源码文件、2个txt说明及1个readme文档,整体约263.61MB,覆盖数据标注、模型定义到推理验证的完整链路。改进思路是在Unet基础上嵌入SAM注意力模块,训练阶段自动生成边界框提示以聚焦肿瘤区域,同时提供带UI界面的推理脚本,支持手动框选进行交互式分割。数据集包含原始图像与对应分割标签,可直接用于模型训练和性能评估;源码涵盖网络结构、训练流程、推理逻辑及界面设计等关键模块,便于复现论文实验或开展二次开发。目前已有773人学习下载,适合希望快速上手息肉分割任务并深入理解改进机制的中高级研究者。
1. 当 U-Net 被息肉边缘绊倒:SAM 提示框能补上什么
用 SAM 提示框改进 U-Net,做息肉肿瘤语义分割,核心动作就是一句话:在标准 unet 的输入上多给一个目标框提示,让解码器在框的约束下去抠息肉边缘,而不是在整张肠镜图里漫无目的地扫。这个改进对两类人最值钱:一是被小息肉和模糊边缘反复折磨的医学图像算法工程师,框能把网络注意力直接钉在病灶区域;二是刚入门、想拿开源数据集跑通 unet 改进全流程的研究生,这套东西从数据准备到源码改造都够具体,不需要自己发明网络结构。下面按“为什么要这样加、数据集怎么做、源码怎么改、坑在哪、怎么验证收益”的顺序,把一套能复现的方案完整拆开讲。
2. 提示框为什么管用:SAM 先验与 U-Net 的融合选型
2.1 息肉分割里 U-Net 的真实瓶颈:边缘模糊和背景混杂
标准 U-Net 在语义分割任务里表现稳定,但设计上并没有“知道”哪个区域是目标。编码器连续下采样四到五次,特征图分辨率降到输入的 1/16、1/32,这对小目标极不友好。息肉里直径小于 6mm 的平坦型病灶,经过四层池化后往往只剩几个像素的响应强度;跳跃连接虽然把浅层细节传回解码器,但传的是全局语境,不是位置先验。
内镜图像更麻烦:息肉和周围黏膜在 RGB 颜色上非常接近,有些平坦型息肉只比背景亮 8~10 个灰度级。我实际跑 Kvasir 息肉数据时,基线 U-Net 收敛后 DSC 能到 0.80 附近,但把 GT 边缘膨胀 5 像素再算边界召回,只有 0.6 左右。问题不在网络深度,而在于模型不知道目标大概在哪,把大量参数用在了抑制背景噪声上。边界被吞进去一块是常事,医生回放时一眼就能看出来。
2.2 SAM 的提示框先验:box prompt 把“目标在哪”写进模型
SAM(Segment Anything Model)的训练逻辑是同时接受图像和 prompt,输出 prompt 对应的目标掩码。prompt 分为点、框、掩码三类,其中框提示(box prompt)非常契合医学分割场景:它提供一个空间范围,告诉分割器“目标在这个框内,别到外面去找”。
在息肉场景里,我不会把 SAM 整个当成分割器直接用,而是借它的两个能力:
- 图像编码器:用 ViT 提取全图特征,输出 256 通道的语义向量,可以当作一个强力的全局特征提取器;
- 提示编码器:把框编码成可学习 token,在解码时通过 attention 与图像特征交互,产生“目标边界应该长什么样”的先验。
点提示虽然更精细,但对噪声敏感,点在息肉边缘或血管上时容易把模型引导到错误区域;框提示虽然粗,但容错高。医学图像从业者画一个框只需要一次鼠标拖拽,比逐点画掩码便宜得多,也更接近真实工作流中的交互习惯。
2.3 特征拼接、热图注入还是损失约束:三条融合路径的取舍
把 SAM 框提示加进 U-Net,工程上有三条常见路线,我把它们拉成一个对比表再逐个说明。
| 融合方式 | 改造成本 | 显存开销 | 推理时是否依赖 SAM | 对边缘增益 |
|---|---|---|---|---|
| 热图通道注入 | 低 | 低 | 不依赖 | 中等 |
| SAM 特征拼接 | 中 | 高 | 依赖 | 较高 |
| 损失级约束 | 低 | 无 | 不依赖 | 中低 |
热图通道注入:把提示框画成一张二维高斯热图,作为额外输入通道拼到 U-Net 第一层。工程改动最小,网络自行学习框内目标和框外抑制的关系。缺点是如果框画得太大,热图接近全 1,提示作用接近于零。
特征拼接:真正把 SAM 图像编码器输出的 256 维特征在编码器出口和 U-Net 特征 concat,全局语义更丰富。缺点是训练和推理都要跑一次 SAM 图像编码器,12G 显存卡上 batch size 会被压得很小。
损失级约束:只在训练时用框对 ROI 内外做梯度加权,框不出现在推理前向里。它不改特征表示,只改误差信号的权重,对边界模糊问题的改善有限。
我最终采用“热图通道 + SAM 特征在编码器出口拼接”的组合:热图负责框级别的硬约束,SAM 特征负责语义层面的软先验。如果部署端无法接受 SAM 带来的额外推理开销,可以在推理时把 SAM 分支裁掉,模型退化成“热图引导 U-Net”,DSC 会降一点,但比纯 U-Net 还是稳。
2.4 先跑通哪条路:我建议的最小可行组合
刚接触这个方向的人,我一般建议按三步走,每一步都能独立出指标,不会把“数据集错了”和“融合方法不行”混在一起。
第一步,只加热图通道,把数据管线和训练循环跑通,确认模型代码没问题;第二步,接入 SAM 图像特征拼接,验证增益是否超过 0.5 个点;第三步,如果边界召回提升不明显,再往损失约束那边加 ROI 加权,把框的影响进一步放大。这个顺序也方便你做消融,最后汇报“每加一个模块涨多少点”才站得住。
3. 数据集准备:息肉公开集转成“框提示+掩码”训练样本
3.1 选 Kvasir-SEG 还是 CVC-ClinicDB:标注与划分差异
息肉语义分割的开源数据集里,最常用的是 Kvasir-SEG 和 CVC-ClinicDB。Kvasir-SEG 有千张级别带掩码的内镜图像,CVC-ClinicDB 也有六百张左右,二者都提供原图和二值掩码,掩码是黑白 PNG,目标区域为白色,背景为黑色。
两个数据集在标注粒度上有差异:Kvasir-SEG 的掩码覆盖范围偏保守,边缘留白较多;CVC-ClinicDB 的标注更贴近息肉实际边界。如果混在一起训练,模型会被迫在两个标注风格之间折中,边缘指标反而不稳定。我的习惯是先用 Kvasir-SEG 做训练和验证,把整套流程跑通,再用 CVC-ClinicDB 做跨数据集泛化测试,而不是最早就把两个混在一起。
划分方面没有统一官方规则,常见做法是随机 8:2 分成 train/val,或者沿用某个公开论文的划分。注意一点:息肉数据里同一患者可能有多帧图像,随机划分可能让同患者帧同时出现在训练和验证里,导致验证指标虚高。稳妥做法是先按患者或序列分组,再在组级别上划分。
3.2 从掩码自动生成提示框:代码与关键参数
训练时我们拿得到 GT 掩码,所以提示框可以自动生成,不需要人工标注。核心函数是把二值掩码的最小外接矩形提取出来,再按参数做外扩和扰动。
import json import cv2 import numpy as np from pathlib import Path def mask_to_box(mask, pad=8, jitter=0, min_area=100): """从GT掩码生成SAM提示框,返回[x1, y1, x2, y2] pad: 框外扩像素,给边缘留余量 jitter: 随机扰动像素数,用于模拟检测器框不准的情况 min_area: 过滤过小的掩码,防止噪声干扰 """ if mask.sum() < min_area: return None ys, xs = np.where(mask > 0) x1, x2 = int(xs.min()), int(xs.max()) y1, y2 = int(ys.min()), int(ys.max()) if jitter > 0: # 在原始边界的基础上随机外扩或收缩,模拟真实框的误差 x1 = max(0, x1 + np.random.randint(-jitter, jitter)) y1 = max(0, y1 + np.random.randint(-jitter, jitter)) x2 = min(mask.shape[1], x2 + np.random.randint(-jitter, jitter)) y2 = min(mask.shape[0], y2 + np.random.randint(-jitter, jitter)) return [max(0, x1 - pad), max(0, y1 - pad), min(mask.shape[1], x2 + pad), min(mask.shape[0], y2 + pad)] def generate_boxes(data_dir): """遍历数据集,为每个训练样本生成提示框,存入JSON""" records = [] for mask_path in Path(data_dir).glob("*_mask.png"): mask = cv2.imread(str(mask_path), cv2.IMREAD_GRAYSCALE) box = mask_to_box(mask, pad=8, jitter=0) if box is None: continue records.append({ "image": str(mask_path).replace("_mask.png", ".png"), "mask": str(mask_path), "box": box, }) return records逻辑说明:mask_to_box先从掩码中找到所有前景像素的坐标范围,再对四个边界分别做外扩和扰动。外扩的物理意义是让网络不要把注意力全压在紧边界上,jitter 的意义是模拟实际检测器输出的框不会像 GT 框那么精确——训练时你的框是抖动的,测试时即使框偏了,模型也不会因为“从没见过歪框”而崩掉。
参数说明:pad=8在 512×512 输入尺度下外扩约 1.5%,属于安全范围;jitter建议训练时取 8~16,不要在生成 JSON 时做,因为离线固定后等于没做数据增强;min_area=100用于过滤掉掩码面积太小的样本,这类样本通常是标注噪声。
3.3 目录组织与双分支增强:SAM 和 U-Net 输入要对齐
训练时我推荐用在线生成热图,而不是把热图存成文件。理由很简单:数据增强会改变图像尺寸和位置,预生成热图容易错位,而在线生成可以保证热图和增强后的图像严格对齐。
import albumentations as A import numpy as np def get_train_augment(): """图像、掩码、框三同步增强,保证提示框和息肉位置严格一致""" return A.Compose([ A.RandomResizedCrop((512, 512), scale=(0.8, 1.0), p=0.6), A.HorizontalFlip(p=0.5), A.RandomRotate90(p=0.5), A.OneOf([ A.MotionBlur(blur_limit=3), A.GaussNoise(var_limit=(10.0, 40.0)), ], p=0.3), ], bbox_params=A.BboxParams( format="pascal_voc", # [x_min, y_min, x_max, y_max] min_visibility=0.5, # 增强后框可见面积低于50%则丢弃 label_fields=["labels"], )) def box_to_heatmap(box, size, sigma=12): """把提示框编码成2D高斯热图,作为额外输入通道""" h, w = size x1, y1, x2, y2 = [int(v) for v in box] cx, cy = (x1 + x2) // 2, (y1 + y2) // 2 yy, xx = np.mgrid[0:h, 0:w] heat = np.exp(-((xx - cx) ** 2 + (yy - cy) ** 2) / (2 * sigma ** 2)) # 框内区域直接置1,框外按高斯衰减 heat[(xx >= x1) & (xx <= x2) & (yy >= y1) & (yy <= y2)] = 1.0 return heat.astype(np.float32) def load_sample(image, mask, box, aug=None): """加载一个样本并生成热图提示通道""" sample = {"image": image, "mask": mask, "bboxes": [box], "labels": ["polyp"]} if aug is not None: sample = aug(**sample) h, w = sample["image"].shape[:2] heat = box_to_heatmap(sample["bboxes"][0], (h, w), sigma=12) x = np.concatenate([sample["image"], heat[..., None]], axis=-1) return x, sample["mask"]逻辑说明:get_train_augment用 albumentations 同时处理图像、掩码和边界框。关键在于bbox_params的存在——如果你自己写增强只翻转图像不翻转框,热图就会和息肉位置错开,训练出来模型的效果直接下一个档位。
参数说明:sigma=12控制高斯衰减的强度,太小则热图只在框中心附近有响应,太大则整个框区域衰减不明显;min_visibility=0.5防止增强后框被裁掉大半时仍然保留,那会让模型收到一个基本无意义的提示。还有一个细节:RandomResizedCrop 会把框缩放,所以box_to_heatmap必须在增强之后执行,用增强后的尺寸和框坐标。
4. 源码落地:SAM 框提示如何接进 U-Net
4.1 最小改造:输入多一个通道,编码器出口拼 SAM 特征
模型结构上,我推荐的改造不是重写 U-Net,而是只动两处:第一处是输入层,把原来的 3 通道改成 4 通道(图像 3 通道 + 热图 1 通道);第二处是编码器出口,把 SAM 图像特征投影后和 U-Net 特征 concat 再送入解码器。
import torch import torch.nn as nn import torch.nn.functional as F class SAMGuidedUNet(nn.Module): def __init__(self, sam_image_encoder, unet_backbone): super().__init__() self.sam_encoder = sam_image_encoder # 冻结 SAM 图像编码器,不参与梯度更新 for p in self.sam_encoder.parameters(): p.requires_grad = False self.unet = unet_backbone # 替换 U-Net 第一层:3 通道输入 -> 4 通道(多一个热图提示) self.first_conv = nn.Conv2d(3 + 1, 64, kernel_size=3, padding=1, bias=False) # SAM 特征投影层:把 256 维压到和 U-Net 编码器特征一致 self.sam_proj = nn.Sequential( nn.Conv2d(256, 64, kernel_size=1), nn.BatchNorm2d(64), nn.ReLU(inplace=True), ) # 融合层:把投影后的 SAM 特征和 U-Net 特征 concat 后压缩回 64 维 self.fuse_conv = nn.Sequential( nn.Conv2d(64 + 64, 64, kernel_size=1), nn.BatchNorm2d(64), nn.ReLU(inplace=True), ) def forward(self, x, heat): # x: [B,3,H,W], heat: [B,1,H,W] x_in = torch.cat([x, heat], dim=1) x_in = self.first_conv(x_in) # 走 U-Net 编码器,拿到 [B,64,H/16,W/16] 的特征 encoder_features = self.unet.encode(x_in) # 假设 unet 已拆成 encode/decode # SAM 图像编码器,输入原图 x,输出全局特征 with torch.no_grad(): sam_feat = self.sam_encoder(x) # [B,256,H/16,W/16] sam_feat = self.sam_proj(sam_feat) # 特征级融合 fused = self.fuse_conv(torch.cat([encoder_features, sam_feat], dim=1)) return self.unet.decode(fused)逻辑说明:SAM 图像编码器在整个网络中扮演“全局语义提供者”,它吃的是没有热图的原始图像,因为提示框信息已经通过热图通道走 U-Net 本身了,不需要再给 SAM 一份。torch.no_grad()包裹 SAM 前向,一是省显存,二是保证冻结参数的 BatchNorm 和 Dropout 不会在训练时产生错误统计。U-Net 的编码器部分需要拆成encode/decode两个方法;如果你用的 segmentation_models_pytorch 这类库,可以只拿 encoder 和 decoder 拼装,不影响整体结构。
参数说明:sam_proj用 1×1 卷积把 256 维压缩到 64 维,是为了和 U-Net 编码器出口的特征维度对齐。fuse_conv同样用 1×1 卷积做跨通道融合,避免引入大卷积核带来额外计算量。这里有个改进空间:如果显存允许,可以把 SAM 特征和多层跳跃连接分别融合,而不是只在编码器出口做一次;实践下来只提升边缘适度,但显存开销明显增加。
4.2 训练循环:冻结 SAM、混合精度与 Dice+BCE 损失
训练时的核心约束是:SAM 已经冻结,优化器只更新 U-Net 参数;损失函数用 Dice 和 BCE 的组合,而不是单一交叉熵。Dice 对小目标更友好,但收敛不稳定,所以要和 BCE 混合。
optimizer = torch.optim.AdamW( [p for p in model.parameters() if p.requires_grad], lr=1e-4, weight_decay=1e-4, ) scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=20) scaler = torch.cuda.amp.GradScaler() criterion = CombinedDiceBCE(alpha=0.5, smooth=1e-5) for epoch in range(20): model.train() # 推理前必须把 SAM 切到 eval 模式,否则冻结层统计量错乱 model.sam_encoder.eval() for images, heats, targets in train_loader: images, heats, targets = images.cuda(), heats.cuda(), targets.cuda() optimizer.zero_grad() with torch.cuda.amp.autocast(): preds = model(images, heats) loss = criterion(preds, targets) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() scheduler.step() class CombinedDiceBCE(nn.Module): def __init__(self, alpha=0.5, smooth=1e-5): super().__init__() self.alpha = alpha self.smooth = smooth self.bce = nn.BCEWithLogitsLoss() def forward(self, preds, targets): # preds 是 logits,先算 Dice p = torch.sigmoid(preds) intersection = (p * targets).sum(dim=(2, 3)) dice = (2.0 * intersection + self.smooth) / ( p.sum(dim=(2, 3)) + targets.sum(dim=(2, 3)) + self.smooth ) dice_loss = 1.0 - dice.mean() bce_loss = self.bce(preds, targets) return self.alpha * dice_loss + (1.0 - self.alpha) * bce_loss逻辑说明:CombinedDiceBCE返回的是损失值,所以 Dice 用1 - dice的形式。alpha=0.5表示两者各占一半,实际使用时如果小目标多,把 alpha 提到 0.7 会更稳定。优化器只收集requires_grad为 True 的参数,SAM 冻结参数不会进优化器,省内存也避免误更新。
参数说明:lr=1e-4对预训练 U-Net 是安全的起步值,如果 backbone 没预训练则要降到 3e-5。T_max=20要和训练轮数一致,否则余弦退火还没走完训练就结束了。amp 混合精度下GradScaler必须全程参与,不能用autocast之后直接loss.backward(),否则 fp16 梯度下溢会直接埋下 NaN 的雷。
4.3 推理管线:手动框和检测器框都怎么变成提示
推理时的框来源通常有两个:医生手动在界面上画框,或者上游检测器(例如 YOLO)输出的检测框。无论哪种,都需要先把原图尺寸的框缩放到模型输入尺寸,再转成热图。这一步最容易被忽略:很多人直接拿原图坐标去box_to_heatmap,当模型输入经过 resize 后,框就错位了。
def infer_polyp(model, image, box_xyxy, input_size=512, sigma=12): """ image: 原始输入图像,HWC 格式 box_xyxy: 原始图像坐标系下的提示框,[x1, y1, x2, y2] """ orig_h, orig_w = image.shape[:2] # 1. 坐标缩放:从原图坐标系映射到模型输入坐标系 scale_x = input_size / orig_w scale_y = input_size / orig_h box = [ int(box_xyxy[0] * scale_x), int(box_xyxy[1] * scale_y), int(box_xyxy[2] * scale_x), int(box_xyxy[3] * scale_y), ] # 2. 图像本身也要做等比例或拉伸 resize resized = cv2.resize(image, (input_size, input_size)) # 3. 生成热图并与图像拼接 heat = box_to_heatmap(box, (input_size, input_size), sigma=sigma) x = np.concatenate([resized, heat[..., None]], axis=-1) # 4. 前向推理 x_tensor = torch.from_numpy(x).permute(2, 0, 1).unsqueeze(0).float().cuda() with torch.no_grad(): logit = model(x_tensor, x_tensor[:, 3:4, :, :]) mask = torch.sigmoid(logit).squeeze().cpu().numpy() return mask, box逻辑说明:x_tensor[:, 3:4, :, :]从拼接后的张量里把热图通道单独切出来传给模型,和训练时的传参方式保持一致。如果模型 forward 签名是(x, heat),就必须把热图单独传入,不能只传拼接后的张量。
参数说明:input_size=512要和训练时的增强输出尺寸一致。如果你训练时用的 512,推理时为了性能临时改成 320,模型会面临两套分布差异,效果掉得比你想的快。sigma保持和训练一致,不要推理时随手调成 20,那会让模型看到的提示从“框内确认目标”变成“整个画面都是目标”。
5. 避坑:U-Net+SAM 提示框训练中的翻车现场
5.1 现象:加框之后 DSC 反而掉了 2 个点
加了 SAM 提示框,验证集指标不升反降,这是我见过最多的翻车现场。
原因基本有两个:一是训练时热图来自 GT 框,测试时手动框或检测器框和 GT 框偏差太大,模型把框当成了“作弊信号”,正确框给对结果,歪框就直接崩;二是数据划分时同一患者的多帧图像同时出现在训练和验证里,验证集虚高,加框之后过拟合又放大了这种虚高。
解决:训练时对框加 jitter 并设置一定概率随机丢弃(比如 15% 的样本不加热图通道,强制模型在缺失提示时也能兜底)。测试时不要只报随机划分的指标,单独留一组和训练集无重合患者的数据来做泛化验证。SAM 提示框合适的增益区间是 DSC 提 0.5~1.5 个点、边界召回提 3~5 个点,超出这个范围先怀疑数据泄漏和框来源不一致。
5.2 现象:混合精度训练到第 10 个 epoch 突然 NaN
损失降到 0.6 左右突然变 NaN,而且每次翻车的位置还不一样,这是典型的梯度下溢问题。
原因是 Dice loss 的分母包含了p.sum + targets.sum + smooth,当一张图里息肉面积特别小时,分母接近 1e-5,Dice 梯度在 fp16 下直接下溢变成无穷大。
解决:smooth从 1e-6 提到 1e-5,并且在CombinedDiceBCE里对p做torch.clamp(p, min=1e-4, max=1-1e-4),让梯度始终落在一个安全的数值区间。如果还不稳定,就把小目标比例高的样本踢出 amp 的 autocast 范围,或者干脆不用混合精度,12G 显存下 512×512 输入 batch size 2 也能跑。
5.3 现象:小息肉边缘被模型“吃掉”一圈
模型整体能定位到息肉,但预测掩码总是比 GT 小一圈,尤其 10mm 以下的小息肉很明显。
原因是提示框 padding 太紧,框紧贴目标边界之后,模型在解码器里会把靠近框边缘的像素往背景推;另一个原因是热图sigma设置太大,框内高斯过饱和,中心区域和边缘区域在热图上没有梯度差异,模型学不到边界空间关系。
解决:把pad外扩到 12~16 像素,sigma保持在训练和推理一致,不要用 0 padding 的框去硬训。如果小息肉占比高,还可以在增强里对图像做小幅随机缩放,让同一大小的息肉出现不同的相对尺寸,模型对尺度变化的鲁棒性会好一些。
5.4 现象:冻结 SAM 后训练 loss 平稳、验证指标剧烈震荡
训练集 loss 一路向下,验证集 DSC 忽高忽低,每两个 epoch 波动超过 3 个点。
原因多数是冻结的 SAM 分支没有切到 eval 模式。SAM 图像编码器里有 LayerNorm 和 GELU,虽然不像 BatchNorm 那样依赖 batch 统计量,但在训练模式下某些实现会保留 Dropout 路径;更常见的是 U-Net 主干里的 BatchNorm 还在用训练时的滑动统计,导致推理阶段特征分布漂移。
解决:在训练循环里显式加一行model.sam_encoder.eval(),并且对 U-Net 里的 BatchNorm 也做一次显式检查——如果分类头是预训练权重带来的,第一轮先把 BatchNorm 的track_running_stats冻结,等 loss 稳定后再放开。这个操作看起来小,但对验证集指标的影响经常大于融合方式本身。
5.5 现象:可视化时热图框和息肉完全错开
训练时看不出问题,一可视化发现热图的框画在了背景上,息肉在框外面。
原因是数据增强没有同步边界框。常见于自己手写 transform 或者从语义分割代码里拷来一套只变换 image 和 mask 的 pipeline,翻转、旋转只处理了图像,没有处理 bbox。
解决:用 albumentations 的BboxParams统一管理,或者把增强限定为水平翻转和旋转 90 度这两种不改变框坐标系的变换。如果你坚持手写,就在 transform 函数里对 image、mask、box 三个对象使用同一个随机种子,不要分三次调np.random。
6. 用 DSC、IoU 和边界召回验证这次改进值不值
评估部分建议固定三件事:数据集划分、输入尺寸、指标口径。数据集划分用 8:2 或者按患者组拆分都可以,但一定要单独留一组跨数据集测试;输入尺寸训练和测试必须一致,否则模型看到的特征分布变了,指标对比没有意义;指标除了 DSC 和 IoU,一定要算边界召回,因为 SAM 提示框带来的主要收益就在边缘区域。
边界召回的计算方式是:把 GT 掩码做形态学膨胀减去腐蚀,得到一个宽度约 10 像素的边缘带,然后算预测掩码在这个边缘带内的召回率。这个指标比整体 DSC 更敏感,能直接反映“框把模型注意力钉在边缘上”的效果。消融实验建议跑三组:纯 U-Net、U-Net+热图、U-Net+热图+SAM 特征拼接。我自己的习惯是先跑纯 U-Net 拿到基线,再逐项叠加,每次只改一个变量。如果加 SAM 之后边界召回不升反降,不要怀疑论文里写的增益是玄学,先回第 5 章逐条排查框来源、增强同步和归一化差异。
进阶一点的做法是“迭代框精修”:第一轮用基线 U-Net 的输出掩码生成伪框,第二轮把伪框作为提示框输入带 SAM 的模型,模型输出更好的掩码后再更新伪框,循环两次。这个做法在跨数据集泛化测试里经常比直接训练更有用,因为它把“框不准”这件事暴露在训练阶段,模型被迫学会校正偏移的提示。
我印象最深的一次翻车就是第一版实验把 GT 框直接当提示框训练和评估,指标漂亮得吓人,换到检测器给的框立刻掉了四个点。之后我所有实验都强制要求:训练加 jitter,评估用手动框或检测器框至少跑一组。这个习惯帮我避开了大量无效的“假改进”。希望帮到你。
本文还有配套的精品资源,点击获取