尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

TransUnet实现DRIVE视网膜血管分割:混合架构与迁移学习实战

TransUnet实现DRIVE视网膜血管分割:混合架构与迁移学习实战 简介基于TransUnet的DRIVE视网膜血管分割实战资源面向医学图像分割方向的开发者与研究者包含完整代码与DRIVE数据集0背景、1前景标注可帮助读者从训练、评估到推理快速跑通分割流程。压缩包共76个文件以Python脚本py、数据集图片png和编译缓存pyc为主并附带README、依赖说明等文件整体仅7.87MB目录结构按训练、评估、预测等模块划分便于对照学习与二次开发。代码注释详细训练脚本会输出训练集与验证集的loss、IoU曲线、学习率衰减曲线、训练日志及数据集可视化图像评估脚本可计算测试集的IoU、Recall、Precision、像素准确率预测脚本则能生成GT与GTimage掩膜图方便逐张检验分割效果。结合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 TransUnetDRIVE 数据怎么变成 224×224 的 patch2.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 能整除 16patch 划分没有余数如果你机器显存够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, size224, augmentFalse): 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), dtypenp.float32) pad_label np.zeros((s, s), dtypenp.float32) pad_mask np.zeros((s, s), dtypenp.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格式很多网盘转存的版本变成了.jpgImage.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 维度 768head 数 12中间的 MLP 用 GELU 激活。Decoder 侧把 Transformer 输出的最后一层和指定层一般是第 3、6、9、12 层的特征挑选出来通过 UpSample 块逐级恢复到 224×224 分辨率每级上采样时把 CNN 对应 stage 的输出做 Concatenate最终接一个 1×1 卷积把通道数压成 1sigmoid 输出概率图。如果你用网上最常见的 vi_t 实现要注意设置img_size224, patch_size16, in_chans1024, embed_dim768这里的in_chans不是输入图像通道数而是 ResNet 输出的 1024。很多人在这里踩坑把in_chans写成 3patch 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.02decoder 通道[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_lossBCE 提供逐像素的稳定梯度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, reductionnone) bce (bce * mask).sum() / mask.sum() dice dice_loss(pred, target, mask) return 0.5 * bce 0.5 * dice逻辑说明bce用的reductionnone是为了手动乘 mask只统计 FOV 内部的误差分母是mask.sum()而不是 batch 内像素总数避免黑边区域占比太高把 loss 稀释。DiceLoss 分母加了1e-6防止某张图血管区域为空导致除零。两个 loss 各占 0.5这个比例对小目标分割基本不会出问题。如果你发现训练时 loss 曲线剧烈震荡可以把 BCE 权重提到 0.7Dice 保持 0.3梯度会更平滑。3.3 优化器与学习率调度小数据集的收敛节奏优化器用 AdamW 而不是 SGD原因是 Transformer 部分对学习率很敏感AdamW 的逐参数自适应能天然处理 ViT 和 CNN 骨干的尺度差异。学习率设1e-4weight 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_lr1e-4, wd1e-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_020, T_mult2, eta_min5e-6)参数说明T_020表示每 20 个 epoch 一个退火周期T_mult2表示下一个周期长度翻倍这样前 20 个 epoch 用较激进的下降探索后 40 个 epoch 用更细的步长收敛。eta_min5e-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_weight1.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(fepoch{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 参数总量约 90MDRIVE 训练集只有 20 张图模型有足够能力把训练样本的全部纹理细节背下来。Transformer 的全局注意力机制特别容易记住样本特有模式比如某张图的视盘轮廓、光照分布。加上数据增强只用了翻转和亮度扰动根本没形成有效的正则化。解决加载 ImageNet 预训练权重是第一步但光这样还不够。我在训练时会在 Transformer Encoder 的 MLP 层后加 Dropout 0.15并给 CNN 骨干的 BatchNorm 层设置track_running_statsFalse不这个方向是错的BatchNorm 统计量还是要保留。实际有效的三个手段是提高增强强度加随机弹性形变和局部擦除、在验证集上做 early stoppingpatience 设 20、把 batch size 调小到 4 并增加 Dropout。还有一个细节把 ResNet 骨干的最后一层 frozen 住requires_gradFalse只训练前面三层和 Transformer 部分能让模型少背很多纹理噪声。4.2 上采样输出尺寸与标签尺寸不一致现象模型 forward 输出 shape 是[B, 1, 224, 224]但数据集返回的标签是[B, 224, 224]的二维图你会说这很简单unsqueeze 一下不就完了。真正的坑在输入尺寸不是 224 时模型输出 223 或 225loss 函数直接报错 shape mismatch。原因TransUnet 的 Decoder 每一级上采样用的nn.Upsample(scale_factor2)如果输入 ResNet 前不是 16 的倍数比如 565 直接缩放成 224 没问题但如果你用自己的图 600×400 先 padding 到 600×600 再 Resize 到 240×240中间的 240 不是 16 的倍数ResNet 下采样是整除 floor 操作上采样回来就少了几个像素。解决预处理阶段统一保证img_size % 16 0这是最省事的办法。如果模型已经训了一半才发现尺寸问题可以用F.interpolate(pred, sizelabel.shape[-2:], modebilinear)把预测图 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 你就知道没加载成功但如果你用了strictFalse缺失参数会被随机初始化Transformer 的 attention 层 QKV 矩阵随机初始化后输出经过 softmax 接近均匀分布Sigmoid 后概率集中在 0.4~0.6 左右。另一个常见原因是 BatchNorm 用了track_running_statsFalse训练不稳定时统计量漂移推理时归一化完全失效。解决加载预训练时先打印缺失参数列表和无关参数列表确认 CNN 骨干和 Transformer 的权重真的进去了。我用的是 HuggingFacetimm里预训练的 ResNetV2-101 作为 backbone注意这里不要写外链也不要去给具体下载地址只说用自己 torchvision 或 timm 能拿到的预训练模型即可。拿到模型后把model.encoder.load_state_dict(checkpoint_encoder, strictTrue)逐个模块加载别图省事整个模型strictFalse。推理前在验证集上跑几个 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(mask0, 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之后无论怎么调学习率、加 epochDice 就是卡住不动。原因这是混合架构的典型瓶颈。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_size224, stride160, devicecuda): model.eval() h, w img.shape[:2] prob_map np.zeros((h, w), dtypenp.float32) count_map np.zeros((h, w), dtypenp.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:ypatch_size, x:xpatch_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:ypatch_size, x:xpatch_size] prob count_map[y:ypatch_size, x:xpatch_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:ypatch_size, x:xpatch_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:ypatch_size, x:xpatch_size] prob count_map[y:ypatch_size, x:xpatch_size] 1.0 prob_map np.divide(prob_map, count_map, outnp.zeros_like(prob_map), wherecount_map0) return prob_map参数说明stride160意味着重叠 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, cmapReds, alpha0.5)叠加把视盘区域边缘用黄色画一个圈能一眼看出模型在哪个解剖结构上翻车。经验是 DRIVE 视盘附近和血管分叉密集区永远是重灾区如果这两个区域效果差优先检查预处理而不是换模型。评估可视化跑完之后把最优阈值和对应 Dice 记录到实验表格里每个模型跑 3 次取均值再下结论不要拿单次结果对外汇报。希望这些参数设定和踩坑整理能帮到你少走我当年走过的那些弯路。本文还有配套的精品资源点击获取
返回列表