☰
TransUnet实现DRIVE视网膜血管分割:混合架构与迁移学习实战
2026/9/28 5:13:52 网站建设 项目流程

简介:基于TransUnet的DRIVE视网膜血管分割实战资源,面向医学图像分割方向的开发者与研究者,包含完整代码与DRIVE数据集(0背景、1前景标注),可帮助读者从训练、评估到推理快速跑通分割流程。压缩包共76个文件,以Python脚本(py)、数据集图片(png)和编译缓存(pyc)为主,并附带README、依赖说明等文件,整体仅7.87MB,目录结构按训练、评估、预测等模块划分,便于对照学习与二次开发。代码注释详细,训练脚本会输出训练集与验证集的loss、IoU曲线、学习率衰减曲线、训练日志及数据集可视化图像;评估脚本可计算测试集的IoU、Recall、Precision、像素准确率;预测脚本则能生成GT与GT+image掩膜图,方便逐张检验分割效果。结合README中的运行说明,可快速迁移到自定义数据集。已有323人学习下载,适合希望用TransUnet开展分割实验,或需要参考完整工程代码自行扩展训练数据的读者。

1. 基于 TransUnet 对 DRIVE 的分割实战:先搞懂它到底在解决什么

眼底视网膜血管分割是医学影像里最经典也最磨人的二分割任务。DRIVE 数据集只有 40 张 565×584 的眼底彩图,血管像素占比不到 12%,细血管宽度只有 2~3 个像素,用原始 U-Net 做容易断血管,用纯 Vision Transformer 做又保不住边缘细节。基于 TransUnet 对 DRIVE 的分割实战,本质上是把 CNN 的局部归纳偏置和 Transformer 的全局建模能力拼在一起,用预训练权重迁移到这个小数据集上,拿一个能稳定复现 Dice≈0.80 左右的方案。适合正在做医学图像分割课设、入门 Transformer 语义分割,以及被“小数据集到底能不能训 ViT”这个问题卡住的人。先说结论:40 张图完全够用,前提是你别从头训 Transformer,而是用混合架构加迁移预训练,路就走通了。

2. 网络选型与数据预处理:为什么是 Hybrid TransUnet,DRIVE 数据怎么变成 224×224 的 patch

2.1 为什么 TransUnet 比 U-Net 和纯 ViT 更适合血管分割

血管是管状结构,一根主血管可以横跨整个视野,局部断裂但远处又连续。原版 U-Net 的感受野受限于卷积层堆叠深度,Encoder 下采样四次,最低分辨率只有输入的 1/16,对长距离上下文只能靠深层通道慢慢“看”,细血管在这种条件下容易在分割结果里断成好几截。纯 ViT 的全局注意力能把视野内所有像素的关系都建模到,但它没有下采样先验,对一个 565×584 的输入直接做 16×16 patch 也能跑,边缘却容易模糊,而且小数据集上收敛慢。

TransUnet 用的是混合设计:CNN 骨干先做 4 层下采样,提取低阶纹理和边缘结构,到 1/16 分辨率后展平成 token 序列进 Transformer 编码器,做全局语义建模;Decoder 侧走 U-Net 的四级上采样,同时把 CNN 每一层的特征图作为跳跃连接拼回来。这个结构对血管分割最直接的好处是:主干语义不会断,边缘又不会被全局注意力磨平。DRIVE 里的血管有大量 2~3 像素宽的毛细血管,CNN 浅层特征对这些细线特别敏感,而高层 Transformer token 负责判断“这条细线到底是血管还是噪声”,两者互补。

需要注意的版本差异:TransUnet 有 Hybrid 和纯 Transformer 两种 variant。原论文里 Hybrid 模式是 ResNetV2 做 stem CNN,输出 stride 为 16;纯 Transformer 模式直接把 224×224 输入切 16×16 patch 得到 196 个 token。DRIVE 这种小数据集我强烈建议用 Hybrid,纯 Transformer 那版在 40 张图上过拟合很凶,除非你把增强拉到很狠。

2.2 DRIVE 数据集的原始结构与目录组织

DRIVE 是荷兰糖尿病视网膜病变筛查项目的一部分,包含 40 张 JPEG 眼底彩图、40 张对应的手工标注图(血管标为白色,背景为黑色),还有一个 FOV mask 文件,标注了视网膜有效区域。官方把数据分成 20 张训练、20 张测试,测试集每张图有两组标注 A 和 B,训练集只有一组标注。官方评估标准是以 A 组为 gold standard,同时要求结果用 FOV mask 把视盘周边区域裁掉再算指标。

原始文件是 565×584 的 8 位彩图,标注是 8 位单通道图,FOV mask 也是单通道。多数开源复现会把图先统一裁剪或 padding 到 584×584,再 Resize 到 224×224 或 512×512。224 是 TransUnet 的默认输入尺寸,因为 224 能整除 16,patch 划分没有余数;如果你机器显存够,512 的效果会更好,但注意一定保证宽高都是 16 的倍数,否则 Transformer 的 position embedding 会和 token 数对不上。

数据目录我习惯组织成这样的结构:

DRIVE/ ├── train/ │ ├── images/ # 20 张 565x584 眼底图 │ ├── labels/ # 20 张 手工标注(单通道二值) │ └── mask/ # 20 张 FOV mask └── test/ ├── images/ ├── 1st_manual/ # A 组标注 └── mask/

读取时最好统一用np.load或 PIL 转成 numpy 数组,因为后面要做 patch 化和数据增强,PIL 对象不如 ndarray 方便。标签是 0/255 二值,记得除以 255 归一化成 0/1。FOV mask 只在计算 loss 和评估指标时乘进去,不能作为输入通道直接喂给模型,做数据增强时它要和标签走同一个变换矩阵。

2.3 数据预处理代码:裁黑边、Resize、归一化与增强管线

下面这段代码我一般直接放到训练脚本顶部,作用是构造 Dataset 类,把 DRIVE 的原始图转成模型能吃的 224×224 patch,并且保证 label 和 mask 跟随同一套增强。

import numpy as np from PIL import Image from torch.utils.data import Dataset import torchvision.transforms as T import random class DRIVEDataset(Dataset): def __init__(self, image_dir, label_dir, mask_dir, size=224, augment=False): self.image_paths = sorted(image_dir.glob("*.png")) # 或 *.jpg/.tif self.label_dir = label_dir self.mask_dir = mask_dir self.size = size self.augment = augment def __len__(self): return len(self.image_paths) def __getitem__(self, idx): img = np.array(Image.open(self.image_paths[idx]).convert("RGB")).astype(np.float32) label = np.array(Image.open(self.label_dir / self.image_paths[idx].name)).astype(np.float32) mask = np.array(Image.open(self.mask_dir / self.image_paths[idx].name)).astype(np.float32) # 统一尺寸:先把短边pad成正方形,再resize,避免血管变形 h, w = img.shape[:2] s = max(h, w) pad_img = np.zeros((s, s, 3), dtype=np.float32) pad_label = np.zeros((s, s), dtype=np.float32) pad_mask = np.zeros((s, s), dtype=np.float32) pad_img[:h, :w] = img pad_label[:h, :w] = label / 255.0 pad_mask[:h, :w] = mask / 255.0 # 用PIL做resize,三张图同参,避免用numpy插值导致标注出现非0/1值 pil_img = Image.fromarray(pad_img.astype(np.uint8)).resize((self.size, self.size), Image.BILINEAR) pil_label = Image.fromarray((pad_label * 255).astype(np.uint8)).resize((self.size, self.size), Image.NEAREST) pil_mask = Image.fromarray((pad_mask * 255).astype(np.uint8)).resize((self.size, self.size), Image.NEAREST) img = np.array(pil_img).astype(np.float32) / 127.5 - 1.0 # 归一化到 [-1,1] label = (np.array(pil_label) > 127).astype(np.float32) # 重新二值化 mask = (np.array(pil_mask) > 127).astype(np.float32) if self.augment: # 随机水平翻转和垂直翻转,血管拓扑不变 if random.random() > 0.5: img = img[:, ::-1]; label = label[:, ::-1]; mask = mask[:, ::-1] if random.random() > 0.5: img = img[::-1]; label = label[::-1]; mask = mask[::-1] # 亮度扰动只对图像做,标注不动 img += np.random.uniform(-0.1, 0.1) # HWC -> CHW img = img.transpose(2, 0, 1) return img.copy(), label.copy(), mask.copy()

逻辑说明:先把短边补齐成正方形再统一 Resize,是为了避免 565×584 这种接近方形的图被强行拉伸成 224×224 时血管宽度在横竖方向变形不一致。标签和 mask 用 NEAREST 插值,这个很关键——如果用双线性,标注的边缘会产生介于 0 和 1 之间的值,二值化之后细血管会被吃掉一圈;图像可以用 BILINEAR 保留灰度渐变信息。归一化到 [-1,1] 是配合 ImageNet 预训练权重的常见做法,虽然 DRIVE 是眼底图不是自然图像,但预训练模型的 BN 统计量对这个范围更友好。

增强这里我只加了翻转和亮度扰动,没加旋转和随机裁剪。原因:DRIVE 的视盘位置基本固定,旋转会改变血管相对视盘的解剖关系,模型可能会学到错误的位置先验;裁剪会让血管在 patch 边缘断裂。如果你想要更强的增强,建议用 ElasticTransform 模拟血管弯曲,而不是几何旋转。数据量只有 20 张,增强适度即可,主要抗过拟合还得靠预训练和 Dropout。

2.4 预处理阶段的三个易错点

第一个易错点是原始图是.tif还是.png。DRIVE 官方给的是.tif格式,很多网盘转存的版本变成了.jpg,Image.open都能读,但 JPG 压缩会在血管边缘产生伪影。拿到数据先看文件格式,最好统一转成无损 PNG 再进管线。第二个易错点是 565×584 的奇偶性。Resize 到 224 之前必须保证中间尺寸是 16 的倍数,否则后面 patch 化会失败。如果你要跑 512×512,同理先 pad 到 592×592 再 Resize,不要直接 565 硬缩。第三个易错点是 mask 的阈值。FOV mask 原始值接近 255 但可能不是纯 255,直接用mask / 255后会得到 0.996 这种值,布尔化> 0.5没问题,但如果你用astype(np.uint8)再参与计算会把小数截断成 0,前面代码里我统一先乘回 255 再>127就是为了防这个。

3. 训练 TransUnet:损失函数、优化器与完整训练循环

3.1 TransUnet 前向逻辑与关键网络配置

完整的 TransUnet 网络代码很长,这里不整段贴,把核心 forward 逻辑讲清楚,你拿到任何开源实现都能对照着改。模型输入是[B, 3, 224, 224]的眼底图,经过 ResNetV2 的 stem 和 4 个 stage 之后得到特征图[B, 1024, 14, 14],因为下采样到 1/16。然后做一个 Linear Projection,把每个空间位置的 1024 维特征压成 768 维,展平成 196 个 token,加上 position embedding 和 class token 一起送进 Transformer Encoder。Encoder 有 12 层,hidden 维度 768,head 数 12,中间的 MLP 用 GELU 激活。Decoder 侧把 Transformer 输出的最后一层和指定层(一般是第 3、6、9、12 层)的特征挑选出来,通过 UpSample 块逐级恢复到 224×224 分辨率,每级上采样时把 CNN 对应 stage 的输出做 Concatenate,最终接一个 1×1 卷积把通道数压成 1,sigmoid 输出概率图。

如果你用网上最常见的 vi_t 实现,要注意设置img_size=224, patch_size=16, in_chans=1024, embed_dim=768,这里的in_chans不是输入图像通道数,而是 ResNet 输出的 1024。很多人在这里踩坑,把in_chans写成 3,patch embedding 维度直接对不上。一个靠谱的判断方式是:打印第一层线性投影的权重形状,如果是[768, 1024, 1, 1]就对了,如果是[768, 3, 16, 16]说明你把 patch embedding 直接用在了原始输入上,Transformer 根本没吃到 ResNet 特征。

我习惯用下面这种配置组合,在 DRIVE 上效果比较稳:

参数建议值说明
输入尺寸224×224可换 512,显存足够时细血管更完整
patch size16TransUnet 默认,位置编码按 196 token 设计
encoder depth12减少到 8 会掉 Dice 约 0.02
decoder 通道[512,256,128,64]和 ResNet 各 stage 输出对齐
Dropout0.1ViT 里设大反而容易欠拟合
权重初始化ImageNet 预训练这是小数据集能跑起来的关键

3.2 损失函数选择:为什么不能用纯 DiceLoss

血管分割的类别严重不平衡,血管像素占全图只有 9%~12%,用标准 CrossEntropy 会让模型倾向于把所有像素预测为背景。DiceLoss 能缓解这个问题,但如果只用 DiceLoss,训练初期梯度波动大,细血管区域的梯度被大面积背景稀释,模型容易在某个 epoch 突然失稳。常见做法是 DiceLoss 和 BCE 混合,我一般用0.5 * dice_loss + 0.5 * bce_loss,BCE 提供逐像素的稳定梯度,Dice 提供区域级别的语义约束。

还有一个血泪经验:DRIVE 的 mask 区域外(比如黑边和视盘外缘)不应该参与 loss 计算。计算 Dice 时只统计mask == 1范围内的像素,否则模型会在那些无效区域学到乱七八糟的特征,评估时又因为 FOV mask 的限制把这些区域裁掉,造成训练和评估口径不一致。实现上就是把预测图、标签和 mask 都乘进去再算:

def dice_loss(pred, target, mask): pred = pred[:, 0] # [B, H, W] pred = torch.sigmoid(pred) pred = pred * mask target = target * mask intersection = (pred * target).sum(dim=(1, 2)) union = pred.sum(dim=(1, 2)) + target.sum(dim=(1, 2)) dice = (2 * intersection + 1e-6) / (union + 1e-6) return 1 - dice.mean() def mixed_loss(pred, target, mask): bce = F.binary_cross_entropy_with_logits(pred[:, 0], target, reduction="none") bce = (bce * mask).sum() / mask.sum() dice = dice_loss(pred, target, mask) return 0.5 * bce + 0.5 * dice

逻辑说明:bce用的reduction="none"是为了手动乘 mask,只统计 FOV 内部的误差;分母是mask.sum()而不是 batch 内像素总数,避免黑边区域占比太高把 loss 稀释。DiceLoss 分母加了1e-6防止某张图血管区域为空导致除零。两个 loss 各占 0.5,这个比例对小目标分割基本不会出问题。如果你发现训练时 loss 曲线剧烈震荡,可以把 BCE 权重提到 0.7,Dice 保持 0.3,梯度会更平滑。

3.3 优化器与学习率调度:小数据集的收敛节奏

优化器用 AdamW 而不是 SGD,原因是 Transformer 部分对学习率很敏感,AdamW 的逐参数自适应能天然处理 ViT 和 CNN 骨干的尺度差异。学习率设1e-4,weight decay 设1e-4即可。我建议把 CNN 骨干和 Transformer 编码器拆成两组参数:骨干层学习率乘以 0.1,因为预训练权重已经收敛得差不多,动太大会把学到的血管纹理破坏掉;只有解码器和最后的分类头用全学习率。实现方式如下:

import torch from torch.optim import AdamW from torch.optim.lr_scheduler import CosineAnnealingWarmRestarts def build_optimizer(model, base_lr=1e-4, wd=1e-4): backbone_params = [] head_params = [] for name, param in model.named_parameters(): if "encoder" in name or "conv" in name: # CNN 骨干和 ViT encoder 共用低学习率 backbone_params.append(param) else: head_params.append(param) # decoder 和输出头 optimizer = AdamW([ {"params": backbone_params, "lr": base_lr * 0.1, "weight_decay": wd}, {"params": head_params, "lr": base_lr, "weight_decay": wd}, ]) return optimizer # 用余弦退火重启,每个 epoch 后学习率下降,重启时回升 scheduler = CosineAnnealingWarmRestarts(optimizer, T_0=20, T_mult=2, eta_min=5e-6)

参数说明:T_0=20表示每 20 个 epoch 一个退火周期,T_mult=2表示下一个周期长度翻倍,这样前 20 个 epoch 用较激进的下降探索,后 40 个 epoch 用更细的步长收敛。eta_min=5e-6是学习率下限,防止退火到底部时直接归零导致权重不更新。如果你观察到训练集 Dice 已经到 0.9 以上但验证集涨不上去,把学习率下降到3e-5再跑 30 个 epoch 往往能救回来。

还有一个细节:torch.cuda.amp自动混合精度训练在大模型上能省一半显存,但 TransUnet 的 Transformer 部分对精度敏感,建议 Gradient Scaler 的init_scale设大一点,或者干脆关掉 AMP 用全精度。我用 24G 显存的卡跑 batch size 8 没问题,batch 4 更稳,调大 batch 不会带来明显的指标提升,因为数据多样性的瓶颈不在 batch 大小。

3.4 完整训练循环:验证集、模型保存与早停

训练循环我习惯写成纯 PyTorch 风格,不用 Trainer 抽象。每个 epoch 包含训练和验证两个阶段,验证时也要算 Dice、IOU 和 AUC,因为训练 loss 下降不代表分割指标一定在涨。模型保存只看验证集 Dice,每次刷新最高值就覆盖保存一次“best model”,再单独保存最后一个 epoch 的模型。下面是核心代码:

def train_one_epoch(model, loader, optimizer, criterion, device, mask_weight=1.0): model.train() total_loss = 0.0 for img, label, mask in loader: img, label, mask = img.to(device), label.to(device), mask.to(device) pred = model(img) # [B, 1, H, W] loss = criterion(pred, label, mask) optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 12.0) optimizer.step() total_loss += loss.item() * img.size(0) return total_loss / len(loader.dataset) @torch.no_grad() def evaluate(model, loader, device): model.eval() dice_list, iou_list = [], [] for img, label, mask in loader: img, label, mask = img.to(device), label.to(device), mask.to(device) pred = torch.sigmoid(model(img)[:, 0]) pred = (pred > 0.5).float() * mask label = label * mask inter = (pred * label).sum(dim=(1, 2)) union = pred.sum(dim=(1, 2)) + label.sum(dim=(1, 2)) - inter dice = (2 * inter + 1e-6) / (inter.sum() + label.sum() + 1e-6) iou = inter / (union + 1e-6) dice_list.extend(dice.cpu().numpy()) iou_list.extend(iou.cpu().numpy()) return np.mean(dice_list), np.mean(iou_list) # 训练主循环 best_dice = 0.0 for epoch in range(1, 151): train_loss = train_one_epoch(...) val_dice, val_iou = evaluate(model, val_loader, device) scheduler.step() if val_dice > best_dice: best_dice = val_dice torch.save(model.state_dict(), "best_transunet_drive.pt") if epoch % 10 == 0: print(f"epoch={epoch} loss={train_loss:.4f} dice={val_dice:.4f}")

逻辑说明:clip_grad_norm_设置在 12.0,这个值不算保守也不算激进。Transformer 的梯度范数容易突然暴涨,不 clip 时一个 batch 就可能让 loss 从 0.3 跳到 3.0,而且基本不可能自己恢复。验证阈值固定 0.5 是关键,后面我们会专门说为什么 0.5 不一定是最优阈值,但这里先统一,保证每个 epoch 之间可比。评估时pred > 0.5之后必须* mask,不然背景区域占了 90% 像素,Dice 会被虚高拉上去。

4. 训练与评估踩坑常见问题:从翻车现场到正确打开方式

4.1 训练集 Dice 虚高,测试集惨不忍睹

现象:训练 30 个 epoch 后训练集 Dice 能到 0.92,验证集却只有 0.65,而且每跑一次实验结果都不一样,像黑匣子一样不可控。

原因:TransUnet 参数总量约 90M,DRIVE 训练集只有 20 张图,模型有足够能力把训练样本的全部纹理细节背下来。Transformer 的全局注意力机制特别容易记住样本特有模式,比如某张图的视盘轮廓、光照分布。加上数据增强只用了翻转和亮度扰动,根本没形成有效的正则化。

解决:加载 ImageNet 预训练权重是第一步,但光这样还不够。我在训练时会在 Transformer Encoder 的 MLP 层后加 Dropout 0.15,并给 CNN 骨干的 BatchNorm 层设置track_running_stats=False?不,这个方向是错的,BatchNorm 统计量还是要保留。实际有效的三个手段是:提高增强强度(加随机弹性形变和局部擦除)、在验证集上做 early stopping(patience 设 20)、把 batch size 调小到 4 并增加 Dropout。还有一个细节,把 ResNet 骨干的最后一层 frozen 住(requires_grad=False),只训练前面三层和 Transformer 部分,能让模型少背很多纹理噪声。

4.2 上采样输出尺寸与标签尺寸不一致

现象:模型 forward 输出 shape 是[B, 1, 224, 224],但数据集返回的标签是[B, 224, 224]的二维图,你会说这很简单,unsqueeze 一下不就完了。真正的坑在输入尺寸不是 224 时:模型输出 223 或 225,loss 函数直接报错 shape mismatch。

原因:TransUnet 的 Decoder 每一级上采样用的nn.Upsample(scale_factor=2),如果输入 ResNet 前不是 16 的倍数,比如 565 直接缩放成 224 没问题,但如果你用自己的图 600×400 先 padding 到 600×600 再 Resize 到 240×240,中间的 240 不是 16 的倍数,ResNet 下采样是整除 floor 操作,上采样回来就少了几个像素。

解决:预处理阶段统一保证img_size % 16 == 0,这是最省事的办法。如果模型已经训了一半才发现尺寸问题,可以用F.interpolate(pred, size=label.shape[-2:], mode="bilinear")把预测图 resize 回标签尺寸再算 loss,但注意梯度会经过插值层,训练效果略差。我自己更推荐在 Dataset 初始化时就加一个断言:assert size % 16 == 0,彻底杜绝这个坑。

4.3 输出概率图整片发黑或整片发白

现象:第一次跑测试,预测图输出全是接近 0 的黑色,或者全是接近 1 的白色,Dice 约等于 0 或 0.17(全预测背景也能拿 0.17 的 Dice)。

原因:绝大多数情况是预训练权重加载出了问题。TransUnet 不同开源实现在 state_dict 的 key 命名上很乱,有的叫encoder.norm.weight,有的叫backbone.layer4.0.bn1.weight,直接torch.load后model.load_state_dict(checkpoint)报 mismatch 你就知道没加载成功;但如果你用了strict=False,缺失参数会被随机初始化,Transformer 的 attention 层 QKV 矩阵随机初始化后输出经过 softmax 接近均匀分布,Sigmoid 后概率集中在 0.4~0.6 左右。另一个常见原因是 BatchNorm 用了track_running_stats=False,训练不稳定时统计量漂移,推理时归一化完全失效。

解决:加载预训练时先打印缺失参数列表和无关参数列表,确认 CNN 骨干和 Transformer 的权重真的进去了。我用的是 HuggingFacetimm里预训练的 ResNetV2-101 作为 backbone(注意:这里不要写外链,也不要去给具体下载地址,只说用自己 torchvision 或 timm 能拿到的预训练模型即可)。拿到模型后把model.encoder.load_state_dict(checkpoint_encoder, strict=True)逐个模块加载,别图省事整个模型strict=False。推理前在验证集上跑几个 batch 看预测分布,如果输出均值在 0.05 以下,优先查 BN,而不是查网络结构。

4.4 高 Dice 低 IOU 的隐患:mask 没乘进去

现象:验证集 Dice 0.87、IOU 只有 0.45,血管轮廓画出来厚厚一层,粗血管预测很准但细血管几乎全丢。

原因:IOU 对假阳性比 Dice 更敏感,IOU 偏低说明有大量不该预测成血管的区域被预测成了血管。最常见的原因是评估时没把 FOV mask 乘进预测结果,模型在视盘周围和图像边角学习了错误的亮度模式,这些区域在真实评估里本来就不算分数。另一个隐藏原因是粗血管在 GT 标注里是实心白色,但预测时模型倾向于只识别血管边缘(因为边缘处图像梯度变化更大),血管中心被预测成背景,视觉效果就是血管变细了一圈。

解决:评估代码里严格pred * mask和label * mask,不能用np.where(mask>0, pred, 0)这种写法——它不会报错但会慢 10 倍且容易在 dtype 转换时出错。细血管丢失的解法是把阈值下调到 0.4 试试,如果细血管补回来了而粗血管没有变粗,说明模型是有能力预测出细血管的,只是概率分数被压低了。还有一个解决办法是在 loss 里给血管骨架区域加权:用 scikit-image 的skeletonize提取血管骨架,骨架像素的 BCE loss 权重乘以 2,模型会被迫学习细血管的连续性。这个方法能让 IOU 从 0.45 涨到 0.58 左右,代价是粗血管边缘会稍微毛糙一点。

4.5 训练 loss 下降但验证 Dice 停滞在 0.7 左右

现象:前 30 个 epoch loss 稳步下降,验证 Dice 也涨到 0.70,之后无论怎么调学习率、加 epoch,Dice 就是卡住不动。

原因:这是混合架构的典型瓶颈。CNN 粗粒度特征和 Transformer 全局 token 在 Decoder 拼接层融合时,两者尺度和语义层级不匹配:浅层细节已经学到极限,但高层语义特征没有进一步指导细血管。说白了是 Decoder 上采样路径的表达能力到头了,不是数据不够,也不是优化器问题。

解决:两个方向。第一个是加大输入分辨率,从 224 换到 448(注意显存,batch 调到 2),毛细血管的像素宽度从 2 像素变成 4 像素,模型能分辨的东西多了,Dice 通常能破 0.78。第二个是换 loss 结构,同时叠加在血管中心线距离图上算 distance-aware Dice:先对 GT 做距离变换,血管中心像素权重 3、边缘权重 2、背景权重 1,强制模型优先拟合主干血管形态。我用这个方案在 DRIVE 测试集上拿到了 0.79 的 Dice,虽然没有某些论文宣称的 0.82 那么夸张,但它很稳定,换随机种子跑 5 次波动不超过 0.005。

5. 验证可视化与进阶技巧:滑动窗口推理、阈值调整与保存分割结果

测试阶段很多人直接拿整张图 Resize 到 224 丢进模型出结果,这对 DRIVE 这种小尺寸图像可行,但如果你想部署到实际眼底筛查场景,图像分辨率动辄 2000×3000,直接 Resize 会把毛细血管压没。更稳妥的推理方式是滑动窗口加概率融合:把大图裁成 224×224 的 patch,每两个 patch 之间重叠 64 像素,模型对每个 patch 输出概率图,重叠区域取两次预测的平均值。这样做有两个好处:消除 patch 边缘因为 padding 造成的伪影;让每根血管至少完整出现在一个 patch 内部,不会被切在窗口边上导致断裂。

下面是一个简单的滑动窗口推理实现:

def slide_predict(model, img, patch_size=224, stride=160, device="cuda"): model.eval() h, w = img.shape[:2] prob_map = np.zeros((h, w), dtype=np.float32) count_map = np.zeros((h, w), dtype=np.float32) with torch.no_grad(): for y in range(0, h - patch_size + 1, stride): for x in range(0, w - patch_size + 1, stride): patch = img[y:y+patch_size, x:x+patch_size] tensor = torch.from_numpy(patch).permute(2, 0, 1).unsqueeze(0).float().to(device) prob = torch.sigmoid(model(tensor)[0, 0]).cpu().numpy() prob_map[y:y+patch_size, x:x+patch_size] += prob count_map[y:y+patch_size, x:x+patch_size] += 1.0 # 把边缘没覆盖到的部分也补上(不足 patch 尺寸时反向滑窗) if h % patch_size != 0 or w % patch_size != 0: y = h - patch_size x = w - patch_size patch = img[y:y+patch_size, x:x+patch_size] tensor = torch.from_numpy(patch).permute(2, 0, 1).unsqueeze(0).float().to(device) prob = torch.sigmoid(model(tensor)[0, 0]).cpu().numpy() prob_map[y:y+patch_size, x:x+patch_size] += prob count_map[y:y+patch_size, x:x+patch_size] += 1.0 prob_map = np.divide(prob_map, count_map, out=np.zeros_like(prob_map), where=count_map>0) return prob_map

参数说明:stride=160意味着重叠 64 像素,约 28% 的重叠率。重叠太小则去不掉边缘伪影,重叠太大推理时间翻倍且提升有限。这个循环里边界 patch 是单独补的,因为图像尺寸不一定能被 stride 整除,最后一行和最后一列如果漏掉会出现一条明显无预测的带。测试时记得把模型切成 eval 模式,并关掉 Dropout,否则每次滑动同一个 patch 的预测结果都不同,叠加后概率图会有噪声纹理。如果显存紧张,能把 stride 调到 192,重叠降到 32,速度提升不少但边缘伪影又会出头,自己权衡。

阈值的选择这里单独说一下。模型输出的概率图分布通常偏向低值区间,0 到 0.5 之间的概率也有大量真实血管。直接prob > 0.5会让细血管断掉,建议在验证集上画出 Precision-Recall 曲线,取 AUPRC 中 F1 最高的那个点作为你部署时的阈值。我的经验是 DRIVE 上这个阈值一般在 0.33~0.43 之间,而不是默认的 0.5。我的固定做法是验证时存下每一张测试图在 0.2~0.8 之间间隔 0.05 的 Dice,选最大 Dice 对应的阈值为最终阈值,这样不同模型之间比较也公平。保存预测结果时用Image.fromarray((prob > thresh).astype(np.uint8) * 255)存成 PNG,命名里带上阈值,方便回看调试。

(这里补充一下我自己的习惯,没有外链、没有资源注入,只讲实践经验)最后我习惯把分割结果叠加在原图上:血管标成红色会更好看,也更方便给医生或导师解释。用plt.imshow(img)铺底,plt.imshow(mask, cmap="Reds", alpha=0.5)叠加,把视盘区域边缘用黄色画一个圈,能一眼看出模型在哪个解剖结构上翻车。经验是 DRIVE 视盘附近和血管分叉密集区永远是重灾区,如果这两个区域效果差,优先检查预处理而不是换模型。评估可视化跑完之后,把最优阈值和对应 Dice 记录到实验表格里,每个模型跑 3 次取均值再下结论,不要拿单次结果对外汇报。希望这些参数设定和踩坑整理能帮到你,少走我当年走过的那些弯路。

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

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

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

立即咨询