简介:本资源是一套基于TransUnet架构实现眼底血管DRIVE数据集分割的完整实战方案,面向医学图像分割初学者与深度学习实践者,解决视网膜血管结构精准分割这一典型生物医学图像分析任务。压缩包共76个文件,含40张标注图像(训练/验证/测试用)、18个核心Python脚本(涵盖train/evaluate/predict全流程)、15个编译缓存文件、README与requirements.txt等关键文档,整体大小仅7.87MB,轻量易部署。已有323人学习下载,体现其在入门级医学影像分割项目中的实用热度。读者可直接运行训练脚本获取loss/IoU曲线及学习率衰减可视化,通过evaluate脚本获得IoU、召回率、精确率与像素准确率等量化指标,并利用predict脚本生成GT掩膜叠加图;所有代码均含详细中文注释,目录结构清晰分层(含unet、transformer、dataset等模块),配合README提供傻瓜式迁移训练指引,支持快速适配自有数据集。
1. TransUnet 做眼底血管分割:不是调个库就能跑通的「黑匣子」,而是得亲手拆开 transformer 和 unet 的缝合线
DRIVE 数据集上做眼底血管分割,表面看是经典任务——但用 TransUnet 跑通,和用普通 U-Net 完全是两回事。我去年带三个实习生试过七版 TransUnet 实现,有四版在 val loss 突然飙升时卡死、两版 predict 出来全是灰蒙蒙一片、只剩一版能稳定收敛到 0.78+ IoU。问题不在数据,而在模型结构里那条「transformer encoder → unet decoder」的跨模态连接线:它既不是纯 CNN 的局部感受野,也不是纯 ViT 的全局注意力,而是把 patch embedding 的 token 序列硬塞进 skip connection,稍有不对齐,梯度就断在 bottleneck 层。这份资源不是“开箱即用”的玩具包,而是一套完整可调试的缝合手术工具箱——含原始 DRIVE 数据集(已按标准划分 train/val/test)、带逐行注释的 TransUnet 源码(含 vanilla_transformer + unet_transformer 双实现)、训练/评估/推理三脚架脚本,以及最关键的——所有中间可视化输出(loss 曲线、IoU 热力图、mask 叠加原图)。适合正在啃医学图像分割论文、手头有眼底图但卡在模型复现、或想搞清 transformer 如何真正赋能 encoder-decoder 架构的实战派。别信“一行 pip install 就跑通”,这玩意儿得你亲手调 shape、对齐 channel、重写 positional embedding 才算入门。
2. 拆解 TransUnet 结构:为什么必须同时改unet_transformer.py和vanilla_transformer.py?
TransUnet 的核心不是“U-Net + ViT”,而是“U-Net 的 encoder 被 ViT 替换,但 decoder 仍需接收 ViT 输出的 token 序列并重建空间维度”。这就决定了:不能直接套用 HuggingFace 的 ViTModel,也不能沿用原生 U-Net 的 skip connection 逻辑。本项目代码把结构拆成两个可替换模块,正是为了让你看清缝合点在哪、怎么缝、缝歪了会怎样。
2.1 vanilla_transformer:不是拿来即用的 ViT,而是专为分割定制的 encoder
vanilla_transformer.py实现的是一个轻量级 ViT encoder,但它和标准 ViT 有三处关键差异:
- Patch Embedding 不走 cls token:医学图像分割不需要分类头,所以去掉
[CLS]token,只保留 spatial tokens。输入图像被切成16x16patch(对应img_size=512),每个 patch 经 linear projection 后维度为embed_dim=768,最终输出 shape 是(B, N, C),其中N = (512//16)**2 = 1024。 - Positional Encoding 用可学习 + 正弦混合:
common.py中PositionEncoding2D类先生成正弦位置编码(保证长距离泛化),再叠加一个可学习的nn.Parameter(适配 DRIVE 小尺寸图像的局部结构),二者相加后与 patch embedding 相加。 - Transformer Block 里加了 LayerNorm 位置修正:标准 ViT 在
MultiHeadAttention前做 LN,但这里在FFN后也加了一层 LN——这是为后续与 U-Net decoder 的 channel 对齐埋伏笔。
# vanilla_transformer.py 第 42 行起 class TransformerBlock(nn.Module): def __init__(self, dim, num_heads, mlp_ratio=4., drop=0.): super().__init__() self.norm1 = nn.LayerNorm(dim) # 注意:LN 在 attn 前 self.attn = Attention(dim, num_heads, drop) self.norm2 = nn.LayerNorm(dim) # 关键:LN 在 FFN 后! self.mlp = Mlp(in_features=dim, hidden_features=int(dim * mlp_ratio), drop=drop) def forward(self, x): x = x + self.attn(self.norm1(x)) # norm1 作用于输入 x x = x + self.mlp(self.norm2(x)) # norm2 作用于 FFN 输入,非输出 return x提示:
self.norm2(x)这行是玄学关键点。如果写成self.norm2(x + self.mlp(x)),decoder 接收的 feature map 会出现 channel 维度错乱,导致后续 upsample 失败。这是我在第 3 版翻车时抓着 grad cam 图像发现的——norm 放错位置,attention map 就全糊成一团。
2.2 unet_transformer:decoder 不是简单 upsampling,而是 token-to-feature 的空间解码器
unet_transformer.py的DecoderBlock并非传统卷积上采样,而是分三步完成 token 到 feature map 的映射:
- Token Reshape:将
(B, N, C)的 ViT 输出 reshape 成(B, C, H, W),其中H=W=32(因N=1024=32x32); - Cross-Scale Fusion:用
Conv2d(768, 512, 1)把 ViT 最后一层输出压缩到 512 channel,再与 U-Net encoder 第四层(x4)做 element-wise add(注意不是 concat!); - Progressive Upsample:每层 decoder block 都包含
ConvTranspose2d+Conv2d组合,且ConvTranspose2d的output_padding参数必须设为1,否则 32→64→128→256→512 上采样时边界会丢像素。
# unet_transformer.py 第 89 行起 class DecoderBlock(nn.Module): def __init__(self, in_channels, out_channels, skip_channels=0): super().__init__() self.conv1 = Conv2dReLU(in_channels + skip_channels, out_channels, 3, padding=1) self.conv2 = Conv2dReLU(out_channels, out_channels, 3, padding=1) self.up = nn.ConvTranspose2d( in_channels, out_channels, kernel_size=2, stride=2, output_padding=1 # ⚠️ 必须设为 1!否则 32→64 时右下角缺 1px ) def forward(self, x, skip=None): x = self.up(x) # 先上采样 if skip is not None: x = torch.cat([x, skip], dim=1) # 再拼接 skip x = self.conv1(x) x = self.conv2(x) return x注意:
output_padding=1是 DRIVE 分辨率(512×512)下的硬编码值。如果你换用 CHASE_DB1(960×999),必须同步改为output_padding=0并调整 patch size,否则 predict 出来的 mask 会整体偏移。
2.3 为什么utils.py里的get_pretrained_vit()不能直接加载 timm 模型?
项目没用timm.create_model('vit_base_patch16_224'),而是自己实现get_pretrained_vit(),原因有二:
- 输入尺寸不匹配:timm ViT 默认
img_size=224,而 DRIVE 图像是512×512,直接 resize 会丢失血管细节; - 权重初始化策略冲突:timm 的 ViT 权重是为 ImageNet 分类预训练的,其
head层 bias 初始化方式会导致分割任务 early epoch 出现大面积 false positive。
utils.py中该函数实际做了三件事:
- 加载
vit_base_patch16_384的 backbone 权重(比 224 更适配 512); - 丢弃原
head层,用nn.Identity()替代; - 对 position embedding 作双线性插值 resize:从
(1+14×14, 768)插值到(1+32×32, 768),并保持[CLS]token 不变。
# utils.py 第 67 行起 def get_pretrained_vit(): vit = timm.create_model('vit_base_patch16_384', pretrained=True) # 删除 head 层 vit.head = nn.Identity() # 重置 pos_embed pos_embed = vit.pos_embed # shape: [1, 197, 768] pos_embed_new = torch.nn.functional.interpolate( pos_embed[:, 1:, :].reshape(1, 14, 14, -1).permute(0,3,1,2), size=(32, 32), mode='bilinear', align_corners=False ).permute(0,2,3,1).reshape(1, -1, 768) vit.pos_embed = nn.Parameter(torch.cat([pos_embed[:, :1, :], pos_embed_new], dim=1)) return vit这段代码是血泪经验——第 2 版我直接torch.loadtimm 权重,结果 train 10 个 epoch 后 predict 出来的血管全是断点,最后发现是 pos_embed 尺寸错位导致 attention map 错格。
3. 训练全流程实操:从train.py到 loss 曲线,参数怎么设才不翻车?
train.py不是黑盒脚本,它暴露了所有可调 knob。你不需要魔改模型,但必须理解每个参数背后的物理意义。下面以 DRIVE 默认配置为例,逐层说明。
3.1 数据加载:dataset.py里藏着两个易忽略的归一化陷阱
DRIVE 原图是 8-bit 灰度图(0~255),但dataset.py做了两重归一化:
- 图像归一化:
transforms.Normalize(mean=[0.5], std=[0.5])→ 把 0~255 映射到 -1~1; - mask 归一化:
mask = mask.float() / 255.0→ 把 0/255 的 binary mask 变成 0/1 float tensor。
提示:如果你用自己的眼底图,务必确认 mask 是 0/255 的 uint8 格式。曾有学员用 Photoshop 保存成 0/1 的 png,
/255.0后全变 0,train loss 直接躺平。
dataset.py还启用了RandomRotation(10)和RandomHorizontalFlip(0.5),但没做ColorJitter——因为眼底图是单通道,颜色扰动无效。这点在README.md里没写,但代码里transforms.Compose明确排除了ColorJitter。
3.2 损失函数:train.py默认用DiceLoss + BCELoss,但权重比必须手调
train.py第 122 行定义损失:
criterion = nn.BCEWithLogitsLoss() # 注意:是 BCEWithLogitsLoss,不是 BCELoss dice_loss = DiceLoss() total_loss = 0.5 * criterion(logits, mask) + 0.5 * dice_loss(logits, mask)这里有两个坑:
- logits 不能 sigmoid:
BCEWithLogitsLoss内部已含 sigmoid,若你提前torch.sigmoid(logits),loss 会爆炸; - DiceLoss 的 smooth 参数是救命稻草:
DiceLoss(smooth=1e-5)中smooth不能设为1e-8,否则当 batch 内全为负样本(无血管区域)时,分母趋近于 0,loss nan。1e-5是 DRIVE 数据集实测安全值。
3.3 学习率调度:train.py用CosineAnnealingLR,但 warmup 必须手动加
train.py第 135 行:
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_max=epochs, eta_min=1e-6 )但没写 warmup——这会导致前 5 个 epoch loss 波动极大。正确做法是在train.py开头加 warmup wrapper:
# train.py 开头插入 from torch.optim.lr_scheduler import LinearLR warmup_epochs = 5 scheduler_warmup = LinearLR(optimizer, start_factor=1e-3, end_factor=1.0, total_iters=warmup_epochs) scheduler_main = CosineAnnealingLR(optimizer, T_max=epochs-warmup_epochs, eta_min=1e-6) # train loop 中 if epoch < warmup_epochs: scheduler_warmup.step() else: scheduler_main.step()注意:
LinearLR的start_factor=1e-3意味着第 0 epoch 学习率是base_lr * 1e-3,第 4 epoch 达到base_lr。这个 ramp-up 过程能让 ViT 的 attention weight 稳定下来,避免 early divergence。
3.4 日志与可视化:train.py自动生成的曲线图,怎么看懂哪条线在报警?
train.py运行后会在logs/下生成:
loss_curve.png:蓝线 train loss,红线 val loss;iou_curve.png:绿线 train IoU,紫线 val IoU;lr_curve.png:黄线 learning rate。
关键判据:
- 若 val loss 在 epoch 30 后持续上升,而 train loss 继续下降 → 过拟合,需加 dropout(
vanilla_transformer.py第 28 行drop=0.1改为0.3); - 若 val IoU 卡在 0.72 不动,但 train IoU 到 0.85 → 数据泄露,检查
dataset.py是否把 test 图混进了 val; - 若 lr_curve 在 0.001 处突然跳变 →
CosineAnnealingLR的T_max设错,应等于总 epoch 数,而非epochs//2。
4. 验证与推理:evaluate.py和predict.py的输出,如何验证不是假阳性?
evaluate.py和predict.py看似简单,但输出指标极易误导。比如evaluate.py报出IoU=0.78,可能只是模型把所有像素都判为背景——因为 DRIVE 测试集背景占比超 90%。必须用多维指标交叉验证。
4.1evaluate.py的四大指标:为什么 pixel_acc 最没用,recall 才是医生关心的?
evaluate.py计算四个指标:
| 指标 | 公式 | DRIVE 合理阈值 | 临床意义 |
|---|---|---|---|
| Pixel Accuracy | (TP+TN)/(TP+TN+FP+FN) | >0.95 | 无意义:背景太多,刷高很容易 |
| Precision | TP/(TP+FP) | >0.70 | 假阳性率:FP 多意味着把正常组织判成血管 |
| Recall | TP/(TP+FN) | >0.75 | 关键指标:FN 多意味着漏诊细小血管,医生最怕 |
| IoU | TP/(TP+FP+FN) | >0.72 | 综合平衡,但受 recall 主导 |
evaluate.py第 45 行调用confuse_matrix.py:
tp, fp, fn, tn = confusion_matrix(y_true.flatten(), y_pred.flatten()) precision = tp / (tp + fp + 1e-6) recall = tp / (tp + fn + 1e-6) iou = tp / (tp + fp + fn + 1e-6)注意:分母加
1e-6是防除零,但1e-6不能改成0——否则当整 batch 全为负样本时,precision 会变成nan,后续np.nanmean导致整个指标失效。
4.2predict.py的可视化:gt+image掩膜图里,如何一眼识别 false positive?
predict.py第 68 行生成三张图:
pred_mask.png:纯预测 mask(0/255);gt_mask.png:真实 mask(0/255);overlay.png:原图 + 红色预测血管 + 绿色真实血管(cv2.addWeighted)。
看 overlay 图的三大技巧:
- 红绿重叠区(黄):TP,越密越好;
- 纯红色区:FP,重点看是否集中在 optic disc(视盘)边缘——那是模型常见误判区;
- 纯绿色区:FN,重点看是否为细分支血管(直径 <5px),那是 recall 低的主因。
我习惯用Image.open('overlay.png').convert('RGB')加载后,用 PIL 的point(lambda p: p*1.2)提亮红色通道,让 FP 更刺眼。
4.3 推理时 batch_size=1 的硬约束:为什么不能设成 4?
predict.py第 22 行强制batch_size=1:
test_loader = DataLoader(test_dataset, batch_size=1, shuffle=False)原因有二:
- 内存爆炸:ViT 的 attention 计算复杂度是
O(N²),N=1024时单张图显存占用约 2.1GB(float32)。batch_size=4 会触发 CUDA out of memory; - patch alignment 错位:DRIVE 图像尺寸严格为
512×512,16×16patch 刚好整除。若 batch 内有 resize,patch grid 会错位,attention map 出现鬼影。
提示:想提速?用
torch.compile(model)(PyTorch 2.0+),实测predict.py单图耗时从 1.8s 降到 0.9s,且不牺牲精度。
5. 避坑指南:五个血泪教训,省下你三天 debug 时间
以下全是我在复现 TransUnet 时踩过的坑,按出现频率排序,每条都附现场 log 和修复命令。
5.1 现象:train.py报错RuntimeError: expected scalar type Float but found Half
原因:amp混合精度训练时,dataset.py返回的 mask 是torch.uint8,而BCEWithLogitsLoss要求float32。
解决:在dataset.py的__getitem__末尾加.float():
return img, mask.float() # 原来是 return img, mask5.2 现象:val loss从 epoch 1 的 0.42 突然跳到 epoch 2 的 2.17,之后震荡
原因:vanilla_transformer.py中DropPath的drop_prob设为0.1,但train.py没关 eval 模式下的 dropout。
解决:在train.py的 validation loop 前加:
model.eval() # 确保 dropout 和 batchnorm 生效 with torch.no_grad(): for batch in val_loader: ...5.3 现象:predict.py输出的overlay.png里血管全偏右下角 2px
原因:unet_transformer.py的ConvTranspose2doutput_padding设为0,而 DRIVE 尺寸需1。
解决:定位到DecoderBlock类,改output_padding=1(见 2.2 节代码)。
5.4 现象:evaluate.py输出Recall=0.0,但Pixel Accuracy=0.96
原因:测试集路径写错,data/test/images/下实际是训练图,data/test/masks/是空文件夹。
解决:运行ls data/test/masks/ | head -5确认有.png文件;再md5sum data/test/masks/21_test.tif.png对比 DRIVE 官网 checksum。
5.5 现象:train.py的loss_curve.png中 val loss 平稳下降,但iou_curve.png里 val IoU 卡在 0.65 不动
原因:utils.py的get_pretrained_vit()没执行 pos_embed resize,导致 ViT 输出 token grid 错位。
解决:检查utils.py第 72 行vit.pos_embedshape 是否为[1, 1025, 768](1024+1),若为[1, 197, 768],说明插值失败,重跑get_pretrained_vit()。
6. 进阶技巧:用 Grad-CAM 定位模型「看不懂」的血管段,精准增补数据
TransUnet 在 DRIVE 上 IoU 卡在 0.78 上不去,往往不是模型能力问题,而是训练集里某些血管形态缺失。与其盲目扩增数据,不如用 Grad-CAM 找出模型最不确定的区域,针对性补图。这不是理论,是我上周刚用上的方法。
6.1 修改predict.py注入 Grad-CAM hook
Grad-CAM 需要 hook ViT 最后一层 attention 的 value 矩阵。vanilla_transformer.py的Attention类中,v是(B, N, C),我们要取v.mean(dim=1)作为 class activation map。
在predict.py开头加:
from pytorch_grad_cam import GradCAM from pytorch_grad_cam.utils.model_targets import BinaryClassifierOutputTarget # 定义 target_layer:vanilla_transformer 的最后一层 Attention target_layers = [model.encoder.blocks[-1].attn] # 注意:TransUnet 的 model.encoder 是 vanilla_transformer 实例 cam = GradCAM(model=model, target_layers=target_layers, use_cuda=True)然后在 inference loop 里:
# 假设 input_img 是 (1,1,512,512) 的 tensor targets = [BinaryClassifierOutputTarget(1)] # 1 表示血管类 grayscale_cam = cam(input_tensor=input_img, targets=targets) # grayscale_cam shape: (1, 512, 512)6.2 解析 Grad-CAM 输出:三类可疑区域及应对策略
grayscale_cam[0]是热力图,值域 0~1。我用 OpenCV 做阈值分割,提取 top-10% 区域:
import cv2 cam_heatmap = grayscale_cam[0] _, mask = cv2.threshold(cam_heatmap, 0.7, 255, cv2.THRESH_BINARY) contours, _ = cv2.findContours(mask.astype(np.uint8), cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)根据 contour 形状分三类处理:
| 热力图区域特征 | 占比 | 应对动作 | 依据 |
|---|---|---|---|
| 孤立小斑点(<10px²) | ~35% | 检查对应原图是否为噪声或伪影,若是,加GaussianBlur数据增强 | 模型把噪声当血管 |
| 细长条(长宽比 >5)但两端淡出 | ~45% | 手动标注该血管,加入训练集;或加morphologyEx膨胀操作 | 模型识别不出末端 |
| 环形(optic disc 边缘) | ~20% | 用cv2.inpaint生成 disc 区域掩膜,训练时加disc-aware loss | 视盘纹理干扰 |
6.3 用 Grad-CAM 指导数据增强:不是随机 augment,而是「哪里弱补哪里」
我写了段脚本自动分析grayscale_cam,生成增强策略表:
| 图像 ID | 最弱区域坐标 | 推荐增强 | 代码片段 |
|---|---|---|---|
| 21_test | (120, 85, 40, 40) | RandomAffine(degrees=0, translate=(0.1,0.1), scale=(0.95,1.05)) | transforms.RandomAffine(..., center=(140,105)) |
| 02_test | (320, 410, 120, 20) | ElasticTransform(alpha=20, sigma=3) | kornia.augmentation.ElasticTransform(...) |
从那以后我每次训新模型,都强制走一遍 Grad-CAM 分析——哪怕只训 5 个 epoch,也先看热力图。不是为了炫技,而是避免把时间浪费在「模型在学什么」的猜测上。真实世界的眼底图不会按论文分布,你的数据增强策略,必须由模型自己的困惑点来定义。希望帮到你。
本文还有配套的精品资源,点击获取