TransUnet实战:DRIVE视网膜血管分割的完整复现与避坑指南

发布时间:2026/9/28 12:03:28
TransUnet实战:DRIVE视网膜血管分割的完整复现与避坑指南 简介面向医学图像分割的入门者与进阶开发者这份视网膜血管分割实战资源基于Transformer与U-Net结合的分割模型并配有完整代码和公开数据集可直接复现实验也能迁移到自定义数据上训练。压缩包共76个文件以Python脚本、PNG图像和说明文档为主整体约7.87MB按训练、评估、推理等模块组织查找方便。已有323人学习下载。配套的训练脚本会生成训练集与验证集的损失曲线、交并比曲线、学习率衰减曲线、训练日志以及数据集可视化图像评估脚本可计算测试集的交并比、召回率、精确率、像素准确率等指标推理脚本用于生成掩膜图像并输出真实标注以及标注叠加原图的结果。代码注释详尽按说明文档操作即可轻松运行适合快速掌握该模型在医学图像分割任务中的完整流程。1. TransUnet 在 DRIVE 上做血管分割先接受一个反直觉结论头一次用 TransUnet 跑 DRIVE 视网膜血管分割的人十有八九会以为自己是来刷新高分的结果第一轮验证 Dice 停在 0.75 上下反而不如一个老到的 UNet。没错TransUnet 自己也会翻车。但我依然建议你把这条路线走一遍不是因为它一定让指标上天而是它能让你亲手体会“CNN 局部特征 Transformer 全局建模”在医学小目标分割里的真实手感。DRIVE 只有 20 张训练图标签还是细线血管这种数据规模下搞出一个能稳定复现的训练、评估流程才是本文想交付的东西。本文读者是已经会跑 UNet想试 Transformer 变体的人以及被导师扔来一句“用 TransUnet 试一下”的硕士生。我不会贴在网盘里给你一堆压缩包而是把工程上最关键的代码和参数写给你看照抄能跑然后自己去改。2. TransUnet 凭什么能切视网膜血管从编码器 token 化到跳跃连接的取舍2.1 小数据集 细血管为什么照样需要全局上下文DRIVE 里的视网膜血管图人眼看就是暗红背景下的一堆粗细不一的亮色线条血管细到只有几个像素分叉处更是断断续续。传统 UNet 的感受野有限越靠深层才能看到越大范围所以在这种细线性结构上容易把分叉处切断。CNN 擅长抓局部边缘但血管走向是连续弯曲的局部信息再强也补不了一条被断成两截的细线。TransUnet 的关键是在 CNN 编码器后面接一个 Transformer Encoder。它把整张图的特征压缩成一级一级的 token让每个位置都能和全图其他位置算注意力。血管虽然在局部很细但整体走向有强一致性Transformer 能把“这根血管从视野边缘一路延伸到中心”这种全局依赖学出来。对 DRIVE 这种几百乘几百的图token 序列不算长计算量可控效果也明显比单靠局部卷积好。这里有一个容易被忽略的点DRIVE 训练集非常小20 张图不足以端到端训练一个 Transformer。所以实际可行的做法是先用 ImageNet 预训练的 ResNet50 做特征提取器把 Transformer 当作一个“全局关系增强模块”。这样参数大部分已经有了靠谱初始值Transformer 只负责在医学图像特有的纹理上做微调。换句话说TransUnet 在 DRIVE 上的胜出不是胜在凭空训练一个 Transformer而是胜在“预训练 CNN 提特征 全局注意力补长程依赖”的组合。2.2 编码器 token 化和解码器级联上采样的关键维度用代码讲清维度流转我在自己复现时习惯先把数据流写在注释里再动手写模块。一个精简的 TransUnet 结构输入是(B, 3, H, W)输出是(B, 1, H, W)中间维度大致如下# 输入形状 (B, 3, H, W)这里以 256x256 为例 # 用 ResNet50 作为编码器 # stem : (B, 64, 128, 128) # layer1: (B, 256, 64, 64) # layer2: (B, 512, 32, 32) # layer3: (B, 1024, 16, 16) B, C, h, w f3.shape emb self.embed_linear(f3) # (B, embed_dim, 16, 16) emb emb.flatten(2).transpose(1, 2) # (B, 256, embed_dim) # 给 256 个位置的 token 加上可学习位置编码 emb emb self.pos_embed[:, :256, :] emb self.transformer(emb) # (B, 256, embed_dim) t emb.transpose(1, 2).reshape(B, self.embed_dim, h, w) # 还原成 16x16 特征图这段代码解释了为什么说 TransUnet 是“混合架构”它没有把整张图直接压成一长串序列而是先用 CNN 降采样到 1/16 分辨率再把每个空间位置变成一个 token。256 个 token 对应原图中 16x16 的空间网格每个 token 都拥有原图 16x16 区域的原始感受野经过 Transformer 后又能和其余 255 个 token 交换信息。解码器这边常规做法是从这个(B, embed_dim, 16, 16)特征开始做两次上采样并在上采样后把 layer2、layer1 的编码器特征拼回来。为什么拼跳跃连接至关重要因为 Transformer 的输出天然偏向全局结构对小血管的边缘细节不敏感。layer1 和 layer2 的编码器特征里保留了原始分辨率下的锐利边缘把它们拼接回来相当于给解码器提供了一个“细节急救包”。如果只靠 Transformer 特征上采样输出的血管边界会糊成一片Dice 会很难看。2.3 TransUnet 和 UNet、纯 Transformer 的选型边界我在不同项目里对比过三类模型它们的分工大致如下模型优势劣势适合场景UNet结构简单收敛快跳跃连接直达细节缺乏全局建模长距离断裂处表现一般数据量小、任务通道多、上线求稳纯 Transformer 分割全局依赖强大模型刷分潜力高小数据过拟合严重训练不稳定数据量充足、算力充裕、科研刷点TransUnet兼顾局部细节与全局结构适合细线目标参数多训练需要预训练权重踩坑点多医学小样本、需要论文对比实验DRIVE 上的血管分割本质是“局部结构强、全局走向弱”的任务。TransUnet 的注意力机制能把断线接上但代价是训练时要满足几个前提输入尺寸不能太小否则 token 太少全局建模没意义、必须用预训练 CNN 初始化、解码器的跳跃连接不能省。如果你只是想快速产出一个能上线的模型UNet 仍是最稳的起点但如果你想做对比实验、发论文或者验证“Transformer 在特定医学任务上有用”TransUnet 是值得投入的方向。3. DRIVE 数据集预处理与环境搭建把官方文件夹变成能直接训练的样本3.1 拿到 DRIVE 先看文件结构别让扩展名和调色板坑死你DRIVE 数据集的官方结构是 training 和 test 两个大目录每个目录里都有 images、mask 和手工标注文件夹。我第一次跑的时候没仔细看扩展名直接用 OpenCV 的imread读标签结果读出来全是灰灰白白的乱图血管边缘还错位后来才发现是 GIF 调色板问题。所以拿到数据集的第一件事不是写模型而是用脚本把每个子目录的文件格式和尺寸摸一遍。我一般先在项目根目录建一个清单脚本把每张图的路径、尺寸、通道数、像素取值打印出来tree DRIVE输出大致长这样DRIVE/ ├── training/ │ ├── images/ # *.tif彩色眼底图 │ ├── mask/ # *.gifFOV 有效区域 │ ├── 1st_manual/ # *.gif第一标注者结果 │ └── 2nd_manual/ # *.gif第二标注者结果 └── test/ ├── images/ ├── mask/ └── 1st_manual/看明白结构后自己建一个dataset.py把文件名统一改成image_01.tif这种方式也行但不必重命名。关键是把标签和 mask 都转成 0/1 的 numpy 数组。多提一句DRIVE 官方评估时只在 mask 覆盖的眼底区域算指标所以 mask 必须原尺寸保存不要随意裁剪丢掉边界。3.2 图片、标签、FOV mask 的三路读取与归一化一份能落地的 DataLoader我最常用的 DRIVE 预处理函数不长但每个参数都会影响训练稳定性。下面这段可以直接放进dataloader.pyimport cv2 import numpy as np from PIL import Image def load_drive_triplet(img_path, lbl_path, mask_path, size(256, 256)): # 彩色眼底图用 OpenCV 读后续要转 RGB img cv2.imread(img_path, cv2.IMREAD_COLOR) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) img cv2.resize(img, size, interpolationcv2.INTER_AREA) # 标签和 mask 用 PIL 读GIF/PNG 的调色板坑更少 lbl Image.open(lbl_path).convert(L) lbl lbl.resize(size, Image.NEAREST) lbl np.array(lbl) lbl (lbl 0).astype(np.uint8) mask Image.open(mask_path).convert(L) mask mask.resize(size, Image.NEAREST) mask np.array(mask) mask (mask 0).astype(np.uint8) # 归一化到 [-1, 1] img img.astype(np.float32) / 127.5 - 1.0 return img, lbl, mask这里两个参数直接决定了你后面翻不翻车第一lbl和mask的 resize 插值必须用Image.NEAREST不能用双线性。因为双线性会把 0 和 255 插出中间灰色二值标签就变成灰度图Dice 损失会算得莫名其妙。第二mask在训练时不是用来遮挡输入的而是用来告诉损失函数“哪些地方可以不承担责任”。DRIVE 图像四周是黑色背景背景占比太大如果和血管区域一样计算损失模型会倾向把整张图都预测成背景因为这样 BCE 损失已经很低。所以 DataLoader 每次要返回img, lbl, mask三元组。3.3 环境依赖和显存估算torch timm 的安装顺序与最小启动TransUnet 这类模型吃显存主要不在卷积而在 Transformer 的注意力矩阵。如果输入 512x512token 序列长度是 1024多头注意力的中间张量会非常占地。我的建议是从 256x256 起步batch size 不要一上来就拉满。经验上 256x256、batch size 4 大约需要 6~8GB 显存如果你只有 4GB 显卡老老实实batch_size2 混合精度。环境依赖最简单的一组是conda create -n transunet python3.9 -y conda activate transunet pip install torch2.1.0 torchvision0.16.0 --index-url https://download.pytorch.org/whl/cu118 pip install timm opencv-python pillow matplotlib tqdm说明一下为什么需要timm很多实验代码会用 timm 里的 ResNet 结构做 backbone它的timm.create_model可以直接加载预训练权重。但如果你想完全自己写 ResNet50不用 timm 也完全可以。torchvision的resnet50(weightsResNet50_Weights.IMAGENET1K_V1)也能做同样的事。装完环境后不要急着训先跑一句import torch from simple_transunet import SimpleTransUNet net SimpleTransUNet(img_size256) x torch.randn(2, 3, 256, 256) print(net(x).shape) # 期望 (2, 1, 256, 256)这一步能同时验证模型定义、显存占用和输入输出维度是否一致。如果这里就炸后面训练脚本根本没有必要往下调。4. 训练 TransUnet一套能跑 100 epoch 的脚本和一组稳的参数量级4.1 先写一个能跑的 SimpleTransUNet 模型保留核心结构去掉花活网上能搜到的 TransUnet 实现往往很长几十个类互相嵌套第一次看容易劝退。我实际落地时会把它压缩成一个能跑通的最小版本ResNet50 编码器一个 Transformer Encoder解码器全靠双线性上采样拼接。下面这个类足够你训练和验证import torch.nn as nn import torch.nn.functional as F import torchvision class SimpleTransUNet(nn.Module): def __init__(self, embed_dim768, num_heads12, depth6, img_size256): super().__init__() resnet torchvision.models.resnet50(weightsNone) self.stem nn.Sequential( resnet.conv1, resnet.bn1, resnet.relu, resnet.maxpool ) self.layer1 resnet.layer1 # 输出通道 256, 下采样到 1/4 self.layer2 resnet.layer2 # 输出通道 512, 1/8 self.layer3 resnet.layer3 # 输出通道 1024, 1/16 self.embed_linear nn.Conv2d(1024, embed_dim, kernel_size1) self.pos_embed nn.Parameter( torch.zeros(1, (img_size // 16) ** 2, embed_dim) ) encoder_layer nn.TransformerEncoderLayer( d_modelembed_dim, nheadnum_heads, dim_feedforwardembed_dim * 4, dropout0.1, activationgelu, batch_firstTrue ) self.transformer nn.TransformerEncoder(encoder_layer, num_layersdepth) self.up1 DecoderBlock(embed_dim 512, 256) self.up2 DecoderBlock(256 256, 64) self.head nn.Conv2d(64, 1, kernel_size1) def forward(self, x): x self.stem(x) f1 self.layer1(x) # (B, 256, H/4, W/4) f2 self.layer2(f1) # (B, 512, H/8, W/8) f3 self.layer3(f2) # (B, 1024, H/16, W/16) B, C, h, w f3.shape emb self.embed_linear(f3) emb emb.flatten(2).transpose(1, 2) # (B, hw, embed_dim) emb emb self.pos_embed[:, :h * w, :] emb self.transformer(emb) t emb.transpose(1, 2).reshape(B, -1, h, w) x self.up1(t, f2) # 通道 768512 - 256分辨率到 1/8 x self.up2(x, f1) # 通道 256256 - 64分辨率到 1/4 x F.interpolate(x, scale_factor4, modebilinear, align_cornersFalse) x self.head(x) return x class DecoderBlock(nn.Module): def __init__(self, in_c, out_c): super().__init__() self.up nn.Upsample(scale_factor2, modebilinear, align_cornersFalse) self.conv nn.Sequential( nn.Conv2d(in_c, out_c, kernel_size3, padding1), nn.BatchNorm2d(out_c), nn.ReLU(inplaceTrue), ) def forward(self, x, skip): x self.up(x) x F.interpolate(x, sizeskip.shape[-2:], modebilinear, align_cornersFalse) x torch.cat([x, skip], dim1) return self.conv(x)这个模型结构上不是完整版 TransUnet但核心思想完全一致CNN 提取局部特征Transformer 做全局建模解码器通过跳跃连接恢复细节。embed_dim768、depth6是我在 256x256 输入下常用的一组参数显存压力比官方默认的 12 层小一半。如果你发现训练不稳定可以先把depth降成 4num_heads降到 8后面再慢慢加。4.2 训练循环与混合损失Dice 和 BCE 怎么分摊背景压力训练血管分割最怕背景主导梯度。DRIVE 里血管像素可能只占整张图的 8%~10%如果只用 BCE网络学到的就是把所有像素输出为 0。所以我会把 Dice Loss 和 BCE 按 1:1 混合。Dice Loss 天然对正负样本比例不敏感BCE 则能提供更平滑的梯度两者互补。class MixedLoss(nn.Module): def __init__(self, smooth1.0): super().__init__() self.bce nn.BCEWithLogitsLoss(reductionmean) self.smooth smooth def dice_per_image(self, pred, target, mask): pred pred.sigmoid() # 只保留 FOV 内区域 pred pred * mask target target * mask inter (pred * target).sum(dim(2, 3)) union pred.sum(dim(2, 3)) target.sum(dim(2, 3)) dice (2.0 * inter self.smooth) / (union self.smooth) return dice.mean() def forward(self, logit, target, mask): if mask is not None: loss_pixel self.bce(logit.view(-1), target.view(-1)) # 全图范围 else: loss_pixel self.bce(logit.view(-1), target.view(-1)) loss_dice 1.0 - self.dice_per_image(logit, target.float(), mask.float()) return loss_pixel loss_dice注意我这里 BCE 仍然在全图范围计算Dice 只统计 mask 内。这样做的原因很简单BCE 的梯度足够平滑即使背景多它也能给出一个稳定的学习信号Dice 作为主要评价目标只关注 FOV 有效区域避免黑边干扰。如果你的显存允许也可以把 BCE 也改成 mask 内但那样会丢失一部分对背景的约束训练早期可能更容易输出全 0。训练循环本身不复杂关键是每次迭代要拿到img, lbl, maskfor epoch in range(epochs): model.train() for img, lbl, mask in loader: img, lbl, mask img.cuda(), lbl.float().cuda(), mask.float().cuda() logit model(img) loss loss_fn(logit, lbl, mask) optimizer.zero_grad() loss.backward() optimizer.step()别小看这个循环我第一次跑的时候因为lbl没转 float在 Dice 计算里被自动转成整数结果梯度直接断掉损失停在原地不动。建议在每个 batch 开头加一句断言assert lbl.max() 1.0避免标签变成 255 这种隐性问题。4.3 学习率、输入尺寸、batch size、epoch 的配法参数范围对照表超参数这个东西在医学分割里几乎是玄学但一组保守参数一定能让你先跑通再优化。我总结的参考范围如下参数推荐值说明输入尺寸256x256显存小token 序列短训练快batch size4~8显存 6~10GB 都能承受epoch 数100~150DRIVE 小训练集每轮只有 20 张要跑够轮数优化器AdamW权重衰减建议 1e-4学习率1e-4 ~ 5e-5太大容易崩太小收敛极慢学习率调度ReduceLROnPlateau监测验证 Dice连续 10 轮不升就降半混合损失比例Dice:BCE 1:1可按效果微调为什么建议ReduceLROnPlateau而不是 Cosine因为数据集太小训练曲线抖动很大固定的余弦衰减很容易在抖动中错过好的局部极小点。Plateau 策略更稳也更符合“看验证集表现”的工程习惯。4.4 训练日志里哪些数字值得信Loss 在降不代表血管在长我见过不少朋友看到训练 Loss 从 0.5 降到 0.2 就以为成功了结果验证 Dice 才 0.4。为什么因为 BCE 的主要贡献来自背景像素。即便血管全预测错只要背景预测对BCE 也能降得很好看。Dice Loss 因为分母是预测和标签的总和对全背景输出会压在 0 附近所以一定要在训练日志里同时打印dice和loss。建议每个 epoch 跑完验证集后打印这四列epoch、train_loss、val_dice、val_bce。观察val_dice是否在稳定上升而不是看train_loss。另外训练早期如果val_dice一直不涨不是模型不行很可能是标签或者 mask 读错了。先用单张图做一次推理把logit.sigmoid()画出来看看再决定要不要调参数这比改十轮学习率都有效。5. TransUnet 实战避坑5 个让人血压升高的常见问题与排查记录5.1 现象预训练权重加载直接报错权重尺寸不匹配这个问题几乎是 100% 会遇到的。ResNet50 的state_dict里有fc.weight和fc.bias而我们的SimpleTransUNet只取了layer1、layer2、layer3没有分类头直接load_state_dict一定报错。原因Torchvision 预训练权重是给“图像分类”任务用的多了一整个全连接头这部分没法塞进分割模型。解决加载时删掉不需要的 key或者设置strictFalse。我的习惯是先把分类头删掉再加载from torchvision.models import ResNet50_Weights resnet_state torchvision.models.resnet50( weightsResNet50_Weights.IMAGENET1K_V1 ).state_dict() for k in [fc.weight, fc.bias]: resnet_state.pop(k, None) bs SimpleTransUNet.__new__(SimpleTransUNet) # 实际项目中请正常初始化模型 model SimpleTransUNet() model.load_state_dict(resnet_state, strictFalse)严格来说这个代码片段和SimpleTransUNet的结构不一定完全对得上因为 ResNet 的 stem 卷积名可能带conv1而我们的定义是resnet.conv1具体 key 名要以model.state_dict()为准。加载完权重后强烈建议逐层打印统计值确认 backbone 的参数真的被更新了而不是静默失败。5.2 现象训练的 Dice 一直很低输出图全黑现象很典型训练 30 个 epoch验证 Dice 始终在 0.2 以下可视化预测图全黑。首先检查是不是标签在预处理里被(lbl 0).astype(np.uint8)转成了 0/1但模型输出没有经过 sigmoid你直接拿logit画图当然全黑。原因BCEWithLogitsLoss 期望模型输出 logits而你在可视化时忘了加 sigmoid。另一个隐藏原因是 Dice Loss 的 smoothing 设得不对默认smooth是 1但如果数据是 float 且像素值在 0~255会导致分母和分子差异巨大。解决先跑一次前向打印raw_logit.min(),raw_logit.max()和raw_logit.sigmoid().min/max。如果 sigmoid 后概率分布全在 0.1 以下说明模型还没有学会任何模式如果分布正常但全黑那一定是可视化时阈值定错了。血管分割不要固定用 0.5 作为阈值因为 DRIVE 血管细模型输出往往偏保守0.3 或 0.4 更合适。这个放到最后一部分验证脚本里再说。5.3 现象显存不够OOMOOM 在 TransUnet 里非常常见因为 Transformer 的中间状态随序列长度平方增长。256x256 输入时序列长度 256注意力矩阵是 256x256这个数字还好一旦换到 512x512序列长度变 1024注意力矩阵是 1024x1024显存直接爆炸。原因不是模型参数太多而是中间激活太大。解决第一选择是batch_size2第二选择是输入尺寸降到 224 或 200。再不行用混合精度scaler torch.cuda.amp.GradScaler() for img, lbl, mask in loader: with torch.autocast(cuda): logit model(img) loss loss_fn(logit, lbl, mask) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update()混合精度能省将近一半显存而且新卡对 fp16 支持很好几乎不丢精度。如果你还想省显存可以把dim_feedforward从 768x4 降到 768x2Transformer 层数从 6 降到 4。代价是全局建模能力变弱但训练会快很多。5.4 现象验证指标虚高原因是用整图计算了背景区域有一种“假成功”特别迷惑人验证 Dice 高达 0.9但看测试预测图血管根本没分割出来全图都预测成背景。这种虚高通常发生在你拿整张图作为 Dice 分母时因为背景区域占了 90% 以上只要模型全输出背景Dice 就已经接近 0.8。原因DRIVE 官方评估只在眼底 FOV 内计算指标而你用自己的指标在整图上算背景把自己“评分”刷上去了。解决所有验证和测试指标都只统计mask 0的区域。做法很简单在 Dice 计算函数里先predict predict * masktarget target * mask然后再算交集和并集。别再拿全图算指标那除了骗自己没有任何参考价值。5.5 现象DataLoader 加载 mask 时png 打开是灰色但血管边缘错位DRIVE 的标签和 mask 是 GIF 格式有些教程代码会先cv2.imread再转灰度结果 OpenCV 对 GIF 的配色表处理有问题读出来血管位置和原图对不上。你会在可视化时看到血管像“描边错位”一样叠在另一层。原因OpenCV 的imread对索引色格式支持不强GIF 调色板被当成 BGR 通道读导致灰度出来不是真正的标注值。解决统一用 PIL 读标签和 mask不要用 OpenCV。PIL 的Image.open(...).convert(L)会正确解析 GIF 的索引色。如果你坚持用 OpenCV也要先Image.open(...).convert(L)再转成 numpy不要直接调cv2.imread。6. 用 DRIVE 的 FOV mask 算官方指标再输出预测图最后的验证脚本6.1 只统计 FOV 内的 Dice 和 IoU才是 DRIVE 的官方算法训练结束后真正能写进论文的指标是 mask 内 Dice 和 IoU。我用下面这个函数验证每一张测试图def compute_metrics(logit, lbl, mask, threshold0.5): prob logit.sigmoid() prob prob * mask lbl lbl * mask pred (prob threshold).float() inter (pred * lbl).sum(dim(2, 3)) union (pred lbl).sum(dim(2, 3)) - inter dice (2 * inter 1e-6) / (inter * 2 union 1e-6) iou (inter 1e-6) / (union 1e-6) return dice.mean().item(), iou.mean().item()注意这里prob和lbl都先乘了mask所以背景黑边不会参与计算。这个函数输出的是一个 batch 内的均值如果你要对 20 张测试图分别出结果建议一张图一个单独调用再把结果存成 CSV。这样你做论文表格时可以直接用。6.2 把概率图导出成 PNG顺便做一次自适应阈值对比我复现时养成的习惯是每轮实验结束把测试集预测图导出成 PNG用肉眼看一遍。模型分数再高如果预测图上血管断成一节一节分数也只是“看起来很厉害”。import cv2 import numpy as np def save_prediction(logit, mask, dst): prob logit.sigmoid().cpu().numpy()[0, 0] mask mask.cpu().numpy()[0, 0] prob prob * mask # 二值化 pred (prob 0.4).astype(np.uint8) * 255 cv2.imwrite(dst, pred)这里的阈值 0.4 不是定死的我一般会同时输出 0.3、0.4、0.5 三张对比哪张血管分叉点最完整。有的项目里用 Otsu 自动阈值效果也不错但视网膜血管的灰度分布太不均衡Otsu 经常把阈值抬得太高导致细血管断裂。如果你的目标是发论文用固定阈值并写清楚阈值设置是更稳妥的做法。6.3 我的复现习惯固定随机种子并记录每次实验的配置最后分享一个让我少走很多弯路的习惯每次训练前固定随机种子并把参数自动写进一个配置文件。DRIVE 就 20 张训练图随机种子不同训练结果波动很大。不固定种子你分不清某次提升是改结构带来的还是纯运气。import random import numpy as np import torch def set_seed(seed42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed)这个习惯救过我一次之前调参调了一周发现同样的参数在不同机器上跑出两个结果最后发现是 DataLoader 的 shuffle 顺序不同导致验证集指标差异。固定种子之后每次改动才有可比性。做算法实验的底线就是可复现TransUnet 这种训练曲线抖动大的模型尤其需要这个后悔药。我能给你的最好建议是把“只看 Dice”改掉改成“Dice IoU 预测图三样一起看”。模型输出全黑但 Dice 很高这种坑我踩过希望你这次能绕过。希望帮到你。本文还有配套的精品资源点击获取