ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

U-Net与Attention U-Net医学图像分割实战:从数据到训练避坑全解析

U-Net与Attention U-Net医学图像分割实战:从数据到训练避坑全解析 简介一套面向医学图像分割任务的开源实现基于U-Net与Attention U-Net双架构适用于CT等医学影像的语义分割。资源覆盖数据处理、模型搭建、训练评估与预测全流程dataset模块支持自定义路径、尺寸、随机翻转及CT窗宽窗位增强model模块实现标准U-Net和带注意力门控的变体跳跃连接融合多层特征train模块集成学习率余弦衰减、AdamW优化器及Dice、IoU、混淆矩阵等多类指标并自动保存训练日志与曲线predict模块加载模型输出原图叠加掩码结果utils提供设备检测、可视化等工具。包内含14个文件以Python源码为主5个py脚本另有7个pyc编译文件、依赖清单和README说明压缩包仅16KB轻量易部署。已有118人学习浏览既适合医学图像处理初学者快速理解分割框架也便于中高级研究者扩展多类别实验。1. 医学图像分割绕不开的基线U-Net 和 Attention U-Net 这套代码能直接跑通做医学图像分割的人对 U-Net 基本都有感情。不管后来出现了 Transformer、SAM 这类大模型真到了处理 CT 序列、器官勾画、病灶提取这类任务时U-Net 和它的注意力变体 Attention U-Net 依然是最靠谱的起点。这套代码最大的价值不是给你看一个孤立的网络结构而是把数据处理、模型定义、训练评估、预测可视化整个链路串起来了解压之后改改路径就能跑适合刚开始接触分割任务、又不想从零造轮子的人。我拆这套资源的时候第一感觉是它的文件划分非常“工程化”dataset.py 管数据怎么读、怎么增强、怎么把灰度值映射成类别标签model.py 里同时给了 U-Net 和 Attention U-Net 两种结构train.py 把训练、验证、指标统计、模型保存全部串好predict.py 负责加载权重、跑单张图像并可视化。作者把训练日志按 JSON 格式落盘连混淆矩阵的更新逻辑都写进了 utils.py这对后续做多类别评估非常有帮助。适合谁用如果你正在做 CT 图像的器官分割、病灶区域提取或者你想在上采样结构里加注意力机制但不确定怎么改跳跃连接这套代码是一个可以直接拿来对比学习的参考实现。2. 数据处理与工具函数dataset.py 和 utils.py 里的关键逻辑2.1 数据加载自定义路径、格式与尺寸的读图方案医学图像分割的第一步永远是数据读入。这个 dataset.py 没有依赖特定的医学图像库而是走通用图像读取路线支持你自定义图像路径、掩码路径、图像格式png、jpg、bmp 等和输入尺寸。这样做的好处是你不用非得把数据转成某个特定格式只要把原图和标签图放在对应目录下就能直接喂给模型。# dataset.py 中核心的数据读取逻辑简化示意 class MedicalDataset(Dataset): def __init__(self, img_dir, mask_dir, img_size(256, 256), transformNone): self.img_paths sorted(glob.glob(os.path.join(img_dir, *.png))) self.mask_paths sorted(glob.glob(os.path.join(mask_dir, *.png))) self.img_size img_size self.transform transform def __len__(self): return len(self.img_paths) def __getitem__(self, idx): image cv2.imread(self.img_paths[idx]) image cv2.cvtColor(image, cv2.COLOR_BGR2RGB) mask cv2.imread(self.mask_paths[idx], cv2.IMREAD_GRAYSCALE) image cv2.resize(image, self.img_size, interpolationcv2.INTER_LINEAR) mask cv2.resize(mask, self.img_size, interpolationcv2.INTER_NEAREST) # 灰度值映射将掩码中的原始灰度值映射为类别索引 mask self._apply_gray_mapping(mask) image torch.from_numpy(image.transpose(2, 0, 1)).float() / 255.0 mask torch.from_numpy(mask).long() return image, mask这里有一个关键细节掩码在 resize 时用的是cv2.INTER_NEAREST也就是最近邻插值。原因很简单分割标签是离散类别如果用双线性插值去缩放掩码边缘部分会生成介于两个类别之间的新灰度值这些值映射回类别时会产生错误标签。常见做法是图像用线性插值保持平滑掩码一律用最近邻插值保持语义。这个代码在这个点上处理得很规范。灰度值映射是 dataset.py 里另一个值得讲的逻辑。很多公开数据集的分割标签不是从 0 开始的连续整数比如背景是 0器官是 1病灶是 2但某些数据集会把标签定义成 0、85、255 这样的灰度值。代码里的_apply_gray_mapping就是干这件事的它会把输入灰度值逐个映射成[0, num_classes-1]的连续整数。你拿到自己的数据时先统计一下掩码里出现了哪几个灰度值再把这个映射关系对应上否则训练出来的模型类别索引是乱的。2.2 数据增强随机翻转与 CT 窗宽窗位调整训练医学分割模型数据增强不是“锦上添花”而是“续命手段”。这套代码实现了两个增强策略随机翻转和 CT 窗宽窗位调整。随机翻转很好理解左右翻转、上下翻转各以一定概率执行相当于把训练样本数量翻倍同时让模型对器官位置的左右偏移不敏感。# 数据增强的典型实现片段 def random_flip(image, mask, p0.5): if random.random() p: image cv2.flip(image, 1) # 水平翻转 mask cv2.flip(mask, 1) if random.random() p: image cv2.flip(image, 0) # 垂直翻转 mask cv2.flip(mask, 0) return image, mask def adjust_window(image, window_center40, window_width400): # 将CT值映射到窗宽窗位范围内提升软组织对比度 min_val window_center - window_width / 2.0 max_val window_center window_width / 2.0 image np.clip(image, min_val, max_val) image (image - min_val) / (max_val - min_val) return imageCT 窗宽窗位调整是这套代码里比较专业的部分。做过 CT 图像处理的人都知道CT 原始像素值范围通常在 -1024 到 3071 之间直接归一化会把软组织、血管、病灶的对比度压得很低模型根本学不到有效特征。临床上不同组织有不同窗宽窗位比如腹部看软组织常用窗宽 400、窗位 40看骨骼则用更大的窗宽。代码里把窗宽窗位作为可调参数放在增强流程里你要处理肺部 CT 就把窗位调到 -600 左右效果立竿见影。但要注意一点窗宽窗位调整在推理阶段也要做。如果你训练时用了窗宽窗位限定预测时没做同样的预处理输入分布就不一致分割效果会明显变差。我一般会把窗宽窗位调整写在数据读取里而不是写在随机增强里保证训练和推理走同一条处理链路。2.3 utils.py设备检测、指标计算与混淆矩阵更新utils.py 是这套代码的“后勤部”核心功能包括三类设备检测和模型初始化、Dice 与 IoU 等分割指标计算、混淆矩阵的更新与基于混淆矩阵的精确率、召回率、F1 分数计算。# utils.py 中混淆矩阵更新的简化逻辑 def update_confusion_matrix(cm, pred, target, num_classes): pred pred.view(-1) target target.view(-1) mask (target 0) (target num_classes) pred pred[mask] target target[mask] # 统计每个真实类别被预测成哪个类别的频次 counts torch.bincount(num_classes * target pred, minlengthnum_classes ** 2) cm counts.view(num_classes, num_classes).cpu().numpy() return cm这段更新逻辑看似简单但有两个隐藏的设计决策值得注意。第一它用torch.bincount一次性统计所有类别的预测-真实组合比逐像素 for 循环快得多在大尺寸图像上这个性能差异会被放大到不可忽略的程度。第二它强制过滤掉了标签值不在[0, num_classes-1]范围内的像素这正好对应了前文灰度值映射的作用——如果映射没做干净这里会把异常像素静默丢弃评估结果会虚高。指标计算部分Dice 和 IoU 的写法也需要留意边界情况。代码里通常在分子分母上加了平滑项避免预测和标签全为空时出现除零。从工程角度来说这个处理是必须的特别是处理那些接近图像边缘的小病灶时很多切片里目标区域确实为空如果没有平滑项训练过程会直接产生 NaN 损失。2.4 混淆矩阵在训练评估中的实际作用很多同学用 U-Net 训练分割模型习惯只看一个平均 Dice 就完事。这套代码把混淆矩阵写进训练流程意味着它支持按类别查看精确率、召回率、F1 分数。对于医学图像分割这非常重要类别不平衡是常态如果只看平均指标占比大的背景类别会掩盖占比极小的病灶类别的真实表现。# 基于混淆矩阵计算各类别指标 def compute_class_metrics(cm): num_classes cm.shape[0] precision np.zeros(num_classes) recall np.zeros(num_classes) f1 np.zeros(num_classes) for i in range(num_classes): tp cm[i, i] fp cm[:, i].sum() - tp fn cm[i, :].sum() - tp precision[i] tp / (tp fp 1e-6) recall[i] tp / (tp fn 1e-6) f1[i] 2 * precision[i] * recall[i] / (precision[i] recall[i] 1e-6) return precision, recall, f1我看过很多训练日志发现一个常见误区打印每个类别的 Dice 是用sklearn.metrics现成函数算的这个函数是样本级别统计和图像分割里的像素级别统计定义完全不同数值上会有偏差。这套代码自己实现了混淆矩阵更新是像素级别的和论文里报告的指标口径一致这点对复现工作非常关键。3. 模型结构model.py 中 U-Net 与 Attention U-Net 的差异3.1 标准 U-Net 的卷积块与循环卷积块设计model.py 里的标准 U-Net 遵循了经典结构编码器、解码器、跳跃连接三部分组成。编码器是四个阶段的卷积块加下采样每经过一个阶段通道数翻倍特征图尺寸减半。解码器是上采样加拼接再加卷积逐步恢复空间分辨率。# U-Net 中的卷积块与下采样模块示意 class ConvBlock(nn.Module): def __init__(self, in_channels, out_channels): super(ConvBlock, self).__init__() self.conv1 nn.Conv2d(in_channels, out_channels, kernel_size3, padding1) self.bn1 nn.BatchNorm2d(out_channels) self.relu nn.ReLU(inplaceTrue) self.conv2 nn.Conv2d(out_channels, out_channels, kernel_size3, padding1) self.bn2 nn.BatchNorm2d(out_channels) def forward(self, x): x self.relu(self.bn1(self.conv1(x))) x self.relu(self.bn2(self.conv2(x))) return x class DownBlock(nn.Module): def __init__(self, in_channels, out_channels): super(DownBlock, self).__init__() self.conv ConvBlock(in_channels, out_channels) self.pool nn.MaxPool2d(kernel_size2, stride2) def forward(self, x): x self.conv(x) return x, self.pool(x)这段代码里有几个参数值得说明。卷积核固定为 3×3、padding 为 1这是 U-Net 论文里的默认选择保证特征图尺寸在卷积前后不变。每个卷积块是“两次卷积 批归一化 ReLU”的组合批归一化放在卷积和激活之间这在训练时能显著缓解梯度消失允许你使用更大的学习率。我实际跑的时候发现如果去掉 BatchNorm同样的任务需要多训练将近一倍的 epoch 才能达到相近的指标所以这层不是可有可无的。这个实现里还有一个“循环卷积块”的设计也就是说某些卷积块内部会在相同参数下执行两次以上卷积操作。循环卷积块的作用是增大感受野同时不增加太多参数。但这种设计稍微会增加显存占用如果你在 8GB 显存以下的卡上训练可以先把循环次数调低或者直接换成普通卷积块对最终指标的影响通常不大。3.2 Attention U-Net 的注意力门控机制Attention U-Net 和标准 U-Net 的唯一区别在于解码器的每个上采样阶段前会插入一个注意力门控模块Attention Gate。它的核心思想是编码器每一层输出的特征图里并不是所有空间位置都值得被传递到解码器注意力门控会根据解码器当前的特征图计算一个空间权重图把编码器特征中与当前目标类别相关的区域放大不相关的区域抑制。# Attention Gate 的核心结构 class AttentionBlock(nn.Module): def __init__(self, enc_channels, dec_channels, out_channels): super(AttentionBlock, self).__init__() self.enc_conv nn.Conv2d(enc_channels, out_channels, kernel_size1) self.dec_conv nn.Conv2d(dec_channels, out_channels, kernel_size1) self.alpha nn.Conv2d(out_channels, 1, kernel_size1) def forward(self, enc_feat, dec_feat): g1 self.dec_conv(dec_feat) x1 self.enc_conv(enc_feat) # 相加后经过ReLU与Sigmoid得到空间注意力权重 alpha torch.sigmoid(self.alpha(F.relu(g1 x1))) return enc_feat * alpha这个注意力模块的工作流程是从解码器特征图上采样或直接使用当前尺寸通过 1×1 卷积压缩通道数与压缩后的编码器特征相加经过 ReLU 和 Sigmoid 生成一个空间权重图最后用这个权重图去乘编码器特征。权重值的范围是 0 到 1接近 1 的位置表示该空间区域与分割目标强相关接近 0 的位置是背景或无关区域。从参数角度看注意力门控增加的网络参数量非常少因为每个注意力块只有三个 1×1 卷积。这个特性决定了 Attention U-Net 不会比 U-Net 慢很多但分割精度通常会提升尤其在目标边缘和低对比度区域。做腹部 CT 肝肿瘤分割时肿瘤和周围软组织灰度值接近普通 U-Net 经常出现边缘模糊、区域过大或过小的问题加注意力门控之后边缘清晰度会有肉眼可见的改善。3.3 两种模型在 training 中的切换机制model.py 里通常会提供一个统一的创建模型入口通过参数控制返回 U-Net 还是 Attention U-Net。这个设计在代码复用上很方便你不需要改训练脚本的主体逻辑只需要在命令行参数里指定模型类型。# 模型创建函数通过 model_type 切换结构 def get_model(model_typeunet, in_channels3, num_classes2): if model_type.lower() unet: return UNet(in_channels, num_classes) elif model_type.lower() attunet: return AttentionUNet(in_channels, num_classes) else: raise ValueError(fUnknown model type: {model_type})注意这里的in_channels3默认值是 3。CT 图像通常是单通道灰度图很多人在加载数据时直接把单通道图像复制成三通道再喂给模型这样虽然能跑但浪费了第一层卷积的参数。我一般会根据数据类型修改 in_channels单通道 CT 就设 1RGB 染色病理图就设 3这样网络第一层的计算量减少 2/3训练速度明显更快。从分割能力来看在这套代码里两个模型的输出通道数都等于类别数最后一层没有上采样到头再 Softmax而是直接输出每个像素在各类别上的 logits。这个设计的好处是方便后续在训练脚本里直接用交叉熵或 Dice 损失不需要在模型内部多做一步归一化。4. 训练流程与参数配置train.py 中的优化器、学习率与评估策略4.1 AdamW 与余弦学习率衰减的配置细节train.py 里选用了 AdamW 优化器和经典 Adam 的区别在于权重衰减的单独处理。Adam 里权重衰减是直接加在梯度上AdamW 则是把权重衰减放在更新步骤之外效果是更强的正则化对医学图像这种小数据集场景非常合适。如果你之前用普通 Adam 训练分割模型总觉得过拟合换 AdamW 加上合适的 weight_decay通常能明显改善泛化。# train.py 中的优化器与学习率调度配置示意 optimizer torch.optim.AdamW(model.parameters(), lr1e-4, weight_decay1e-5) total_epochs 200 warmup_epochs 10 # 余弦退火先将学习率从很小值线性升到峰值再按余弦曲线衰减 def cosine_lr(epoch): if epoch warmup_epochs: return 0.1 * (epoch 1) / warmup_epochs # 线性预热 progress (epoch - warmup_epochs) / (total_epochs - warmup_epochs) return 0.5 * (1 math.cos(math.pi * progress)) scheduler torch.optim.lr_scheduler.LambdaLR(optimizer, lr_lambdacosine_lr)这里的学习率配置值得多讲两句。初始学习率设在 1e-4对医学图像分割是比较稳妥的起点。医学图像往往样本量少、目标区域占比小太高的学习率会让模型在训练初期就震荡甚至发散。代码里的余弦衰减加了一个 10 个 epoch 的预热阶段先让学习率从峰值的 10% 线性升上去再按余弦曲线降到接近零。这种做法在小数据集上效果非常明显相当于先用小学习率让模型找到一个大致的解区域再用大学习率在附近探索最后用衰减收敛到更平滑的局部最优。我自己的经验是这个配置不要直接搬到你自己的数据集上。如果数据集很大数千张以上可以把初始学习率调到 3e-4 到 5e-4 之间训练速度会加快如果数据集很小几百张建议维持 1e-4 甚至降到 5e-5否则过早过拟合的风险很大。权重衰减 1e-5 到 1e-4 之间通常不会对最终指标产生决定性影响但只要不是零就能起到稳定作用。4.2 训练过程中的损失记录与指标可视化train.py 不只是循环跑前向反向它还会把损失值、Dice、IoU、类别精确率/召回率/F1 全部记录下来。记录的方式是写入一个 JSON 文件配合 matplotlib 可视化训练曲线。我比较认可这个设计因为你训练一个分割模型动不动就要跑十几个小时中间如果没有任何日志落盘训练崩了连问题出在哪都不知道。# 训练日志记录的简化示例 log_data { epoch: epoch, train_loss: train_loss, val_dice: val_dice, val_iou: val_iou, class_precision: class_precision.tolist(), class_recall: class_recall.tolist(), class_f1: class_f1.tolist(), lr: scheduler.get_last_lr()[0] } with open(log_path, a) as f: f.write(json.dumps(log_data) \n)这种 JSON 逐行追加的写法的好处是每一个 epoch 的记录都是独立的一行即使训练中途异常退出你也能从最后一行看到崩溃前最新的训练状态。很多封装好的日志工具会把所有内容缓存到内存里最后一次性写入一旦崩了就全丢了。从这个细节能看出这套代码是经过实际训练验证过的不是玩具工程。可视化部分代码通常会把损失曲线、学习率衰减曲线、各类别指标变化画在一个大图里。我看这个可视化有两个用途一是确认学习率是否真的在按预期衰减二是看每个类别的 Dice 是否都在同步上升。如果某个类别的 Dice 一直不涨说明这个类别样本太少或者特征太弱需要针对性处理比如调整窗宽窗位突出该组织区域。4.3 最佳模型保存策略基于验证集 Dice 与 IoUtrain.py 的模型保存策略是“自动保存最佳模型”核心判断依据是验证集上的 Dice 或 IoU。通常的逻辑是每个 epoch 结束验证一次如果当前指标比历史最佳高就把当前模型权重保存为 best_model.pth同时记录这个最佳指标对应的 epoch 号。# 最佳模型保存逻辑 if val_dice best_dice: best_dice val_dice torch.save({ model_state_dict: model.state_dict(), epoch: epoch, val_dice: val_dice, val_iou: val_iou }, best_model.pth)这里有一个值得注意的细节保存的不是裸 state_dict而是打包了 epoch 和验证指标。这样你训练完回头看模型文件时不需要额外记录文件对应的训练进度模型文件本身就是完整的元信息存储。如果你要从最佳模型继续训练或者在多个候选模型之间做选择这种打包方式非常方便。从实操角度我建议在训练时同时保存最后一个 epoch 的模型和最佳模型。因为最佳模型是验证集上的表现最后一个 epoch 是训练集上的最终状态两者在不同的迁移场景下各有优势。有些同学只保存最佳模型部署时发现验证集指标高但实际测试效果一般回头想换最后的模型没得换只能重新训练。代码里如果不带这个功能你自己加一行也不难。4.4 损失函数的选择与类别不平衡处理思路train.py 里没有写明用的哪种损失函数但从代码结构和医学分割场景推断最常用的组合是 Dice Loss 加上带权重的交叉熵。对于小目标分割单独使用交叉熵会出现一个严重问题图像里背景像素占了大多数模型只需要把所有像素预测成背景就能得到很低的损失但分割指标会非常难看。Dice Loss 直接优化你真正关心的指标它对类别不平衡天然不那么敏感。如果你在这个框架上做实验我建议采用组合损失先用普通交叉熵训练前几个 epoch 让模型学会基本轮廓再切换到 Dice Loss 或 Dice Focal Loss 精细调优。这种两阶段策略比从头到尾用同一种损失都能拿到更高的最终指标。代码里没有强制绑定损失函数所以你可以自由替换 train.py 里的 loss 部分不需要动模型结构。5. 避坑指南数据处理、训练与预测环节的常见问题5.1 掩码灰度值与类别索引错位现象是训练可以正常进行但验证时每个类别的 Dice 都很低尤其是当数据集里的掩码灰度值是 0、128、255 这类非连续数值时训练曲线看起来正常指标却上不去。我遇到过一位同行他直接把掩码当成三通道图像读入然后用 One-Hot 编码做分割结果所有目标的边缘区域都被预测成了背景。原因在数据集处理阶段没有做灰度值映射。U-Net 输出的通道数是类别数每个类别对应一个通道标签图需要以类别索引的形式计算损失。如果掩码里是 0 和 255而模型输出只有两个通道计算交叉熵时会把 255 当作类别 255但模型只能输出 0 和 1 两个类别的预测梯度计算完全错乱。解决方法是先用 np.unique 统计掩码里的所有灰度值建立一个字典把原始灰度映射到连续整数再在 dataset.py 的_apply_gray_mapping里执行映射。训练前一定要单独可视化几张映射后的掩码图确认没有某个类别的像素被映射成错误值。5.2 自定义数据集尺寸与网络下采样倍数不匹配现象是训练过程报错显示尺寸不匹配。常见于输入图像尺寸不是 16 的整数倍。U-Net 有四个下采样阶段每次下采样 2 倍总共 16 倍。如果输入是 250×250经过四次下采样后尺寸是 16×16 左右再经过四次上采样解码器端拼接时会出现尺寸差 1 个像素的情况直接导致 concat 操作维度不匹配。原因是对 U-Net 的参数化结构理解不够。很多同学直接把数据丢进网络没有检查输入尺寸与下采样倍数的关系。解决方法是把所有图像统一 resize 到 256×256 或 512×512这类尺寸对 16 倍下采样是整除的。如果你坚持保留原始输入尺寸那需要修改网络结构里的 padding 策略但这对分割任务来说性价比不高。建议在 dataset.py 里初始化时就直接固定 img_size同时设置assert img_size[0] % 16 0 and img_size[1] % 16 0从源头上杜绝这个问题。这个检查只需要一行代码但能省下一个小时的排查时间。5.3 推理阶段忘记使用训练时的预处理流程现象是训练时验证指标很高但用 predict.py 单独预测一张图时分割结果出现大片误判或目标完全丢失。最常见的原因是窗宽窗位调整没有在预测流程里执行。你训练时每张图都做了窗宽窗位映射把 CT 值压缩到了固定范围但预测时只做了简单的归一化甚至直接读原图数据分布完全不一致。解决方法是把预处理逻辑抽成一个独立函数放在 dataset.py 或者 utils.py 里训练和预测共用这个函数。当年我在一个胰腺分割项目里踩过这个坑训练时平均 Dice 0.87预测实测只有 0.6 左右排查到最后就是预处理不一致。从那以后我所有分割项目的数据读取都强制走同一个预处理函数。5.4 显存不足与 batch size 选择的权衡现象是训练刚开始就报 CUDA out of memory。U-Net 虽然结构并不算深但跳跃连接会在解码器端保留多尺度的特征图显存占用相当可观。尤其在 512×512 输入、batch size 为 8 的条件下8GB 显存基本不够用。解决思路有三个方向第一把 batch size 降到 2 或 4对于分割任务 batch size 小一点影响不大因为 BN 层的统计量在小 batch 下略有波动但医学图像样本数少通常也在可接受范围第二把输入尺寸从 512 降到 256显存占用直接降到四分之一第三开启梯度累积每两个 batch 更新一次梯度模拟更大的 batch size。代码里默认可能没开梯度累积你自己加一个判断即可。5.5 标签类别不平衡导致模型偏向背景类现象是训练到后期损失还在下降但目标类别的 Dice 始终在 0.3 以下背景类别的精确率和召回率都是 99% 以上。这说明模型把所有像素都预测成了背景。原因在于背景像素占比可能超过 90%模型只要预测全背景就能获得极低的损失。解决方式首选 Dice Loss因为 Dice 衡量的是区域重叠程度全背景预测的 Dice 为 0损失接近 1模型无法靠“偷懒”获得低损失。也可以在 dataset.py 里做类别权重采样让网络每个 batch 看到的正样本比例更高。代码里如果用的是交叉熵建议改成class_weight [0.1, 1.0]这类配置给目标类别更大的权重效果会立刻改善。6. 用 predict.py 做单图推理与双模型效果对比验证predict.py 解决了 U-Net 这类分割模型落地时的最后一公里问题拿训练好的模型对单张图做推理并把分割掩码叠加在原图上展示。我先讲它的基本用法再讲怎么用它做 U-Net 和 Attention U-Net 的对比实验。# 用训练好的模型对单张图像进行分割 python predict.py --model_type attunet --checkpoint best_model.pth --input data/sample_ct.png --output result.pngpredict.py 加载模型权重后会先对输入图像做和训练时一致的预处理包括 resize、归一化、窗宽窗位调整然后前向推理得到每个类别的概率图取概率最大的类别作为预测结果。之后它会把预测掩码转成彩色形式并叠加在原图上方便肉眼观察分割效果是否覆盖了目标区域。# predict.py 中灰度值还原的关键逻辑 def restore_label(mask_gray, mapping_dict): # 将预测的类别索引还原为原始标签灰度值 restored np.zeros_like(mask_gray) for idx, gray_val in mapping_dict.items(): restored[mask_gray idx] gray_val return restored灰度值还原是 predict.py 里容易被忽略但非常重要的功能。训练时你把原始掩码灰度值映射成了连续类别索引预测输出自然是类别索引。如果评测工具需要和原始标签做对比必须把类别索引再映射回原始灰度值。不少同学用外部评测脚本时发现预测结果和标签对不上就是漏了这一步。要对比 U-Net 和 Attention U-Net 在同一数据集上的效果我建议用同一个 checkpoint 路径下存两个模型的训练结果分别训练两个模型用各自的 best_model.pth 分别跑同一张测试图然后并排展示预测结果。不要只看 Dice 数值要盯着边缘区域看注意力模型在低对比度组织边界通常更平滑普通 U-Net 容易出现锯齿状边缘或小虚影。实际跑下来Attention U-Net 通常比普通 U-Net 在指标上高 1% 到 3%但训练时间和显存消耗只增加了不到 10%。如果你的任务目标是器官分割这类结构相对固定的任务Attention U-Net 的收益主要体现在边界精度上如果是病灶分割这类目标很小、位置不固定的任务注意力机制的收益会更明显因为注意力门控能帮助模型把特征集中在病灶区域。对新人来说我推荐把 U-Net 作为基线先跑通全流程再切到 Attention U-Net 观察差异这样你对模型变化导致的效果差异会有更直接的认知。最后说一个我自己的习惯每拿到一个分割数据集我会先找三张典型图像跑一遍 predict.py一张目标大的、一张目标小的、一张目标边缘模糊的用这三张结果快速判断模型是不是真的学到了有效特征。如果三张都表现稳定再去做完整的验证集指标评测。从那次预处理不一致翻车以后我每改一道预处理逻辑都会强制先跑一次预测再开训练这个习惯帮我挡住了不少低级错误。希望你能把这套代码用顺少踩几个我踩过的坑。本文还有配套的精品资源点击获取
返回列表