
简介基于U-Net的脑肿瘤分割完整代码项目面向深度学习与医学影像交叉领域的入门者提供自动分割方案解决MRI/CT图像中肿瘤区域识别与边界定位问题。压缩包内共2000个文件以1989个tif格式脑部影像数据为主另有7个Python脚本分别对应模型构建、数据集处理、训练与测试环节以及1张网络结构示意图整体约291.31MB。代码采用U-Net对称编解码结构配合自定义Dataset完成数据加载训练脚本默认20个epoch可快速体验完整训练与Dice/IoU评估流程包内含pyc缓存便于直接运行。已有1017人学习下载这套可复现的示例能帮助读者深入理解U-Net原理并迁移到其他医学分割任务中文件组织清晰适合作为首个医学图像分割实战项目。 如果你第一次拿到BraTS数据集看着满屏的.nii.gz文件大概率会愣住几秒。不用怀疑我当初也是这样。脑肿瘤分割这个任务入门门槛不在模型本身而在于数据怎么读、预处理怎么做、损失函数怎么设——这些细节决定了你跑出来的到底是一坨没法看的噪声还是一个能用的分割结果。这篇文章就用一套完整的UNet代码把从数据加载到模型训练再到评估的每个环节捋一遍代码全部可以直接复制跑通。我默认读者会用PyTorch并且对卷积网络有基础认知但没真正上手过医学图像分割。实验环境是PyTorch 2.x CUDA 11.8单卡训练如果你只有CPU也可以跑通只是时间会长一些。1. 脑肿瘤分割为什么首选UNet这个结构1.1 脑肿瘤在影像上的特殊性脑肿瘤分割和常规的语义分割有一个非常明显的差异肿瘤区域的边界极其不规则而且不同模态的MRI图像上肿瘤不同子区域的对比度完全不同。BraTS数据集提供四种模态T1、T1ce、T2、FLAIR。这四张图像不是简单的重复它们各自能看清的组织不一样。T1适合看解剖结构T1ce打药后增强区域通常对应活跃的肿瘤核心T2对水肿区域敏感FLAIR则能抑制脑脊液信号让肿瘤周围的水肿更明显。诊断时医生也是四张图来回对照所以模型输入端最好把这四个模态当成四个通道一起喂进去让网络自己学着融合这些信息。肿瘤本身还有一个特点大小差异极大。有的病灶只有几个像素有的则占据大半个脑半球。加上肿瘤和正常脑组织的灰度值在很多区域非常接近边缘在视觉上几乎没有明确分界这对分割模型的细节捕捉能力要求很高。1.2 UNet的编码器-解码器结构为什么能应对UNet的核心思路就是对称的编码器-解码器结构中间用跳跃连接把同尺度的特征直接拼起来。编码器一路下采样把图像逐步降成高语义、低分辨率的特征图解码器再一路往上采样把低分辨率的语义信息逐步恢复成高分辨率的空间细节。跳跃连接是整个设计的灵魂。如果没有跳跃连接解码器只能靠来自最底层的特征图去恢复细节边缘信息早就丢光了。加入跳跃连接后编码器每一层的边缘纹理信息都能直接传到解码器的对应层网络在做像素级分类时既能看到全局语义又能参考局部细节肿瘤边界的分割精度明显提升。对比经典的FCNFCN虽然也有跨层相加但结构上远不如UNet的U型对称设计来得彻底。在实际效果上同样数据量下UNet的收敛速度和小目标分割能力都更好一些。1.3 为什么先用2D UNet而不是直接上3D很多人一上来就问BraTS是三维体数据为什么不用3D UNet我的回答是如果你的目标是先把流程跑通、把baseline做出来2D UNet是最合适的选择。3D模型理论上能利用层间空间信息效果通常会更好但显存占用和数据预处理复杂度也成倍上升。一个全分辨率3D体数据直接喂进去在普通单卡上根本跑不动必须做patch采样、混合精度、梯度累积等一堆操作。2D UNet可以把整个体数据切成切片按普通图像分割的方式训练入门成本低很多而且效果绝对不差——在BraTS上一个调参得当的2D UNetDice也能做到0.85以上。提示先把2D流程跑通再考虑3D升级。这是最稳妥的路线不要一开始就挑战高难度。2. 完整代码分段拆解数据、模型、损失与评估2.1 BraTS数据读取与预处理BraTS的原始数据是NIfTI格式需要用nibabel库读取。一个典型的病例目录下有flair、t1、t1ce、t2四个模态的MRI图像外加一个seg标签文件。标签里的数值含义是0为背景1为坏死核心2为水肿4为增强肿瘤。数据读取的代码分两步先读四模态合并成多通道体数据再做预处理。import nibabel as nib import numpy as np def load_patient(data_dir): modalities [flair, t1, t1ce, t2] images [] for m in modalities: img nib.load(f{data_dir}/{m}.nii.gz).get_fdata() images.append(img) vol np.stack(images, axis-1) # (H, W, D, 4) seg nib.load(f{data_dir}/seg.nii.gz).get_fdata() return vol, seg预处理的重点有两个一是归一化二是提取切片。归一化不能对所有体素直接算全局均值和方差因为背景区域全是0会严重拉低有效统计量。正确做法是只对非零体素做z-score归一化每个模态独立计算def normalize_modality(data): mask data 0 if mask.sum() 0: return data mean data[mask].mean() std data[mask].std() normalized np.zeros_like(data) normalized[mask] (data[mask] - mean) / (std 1e-8) return normalized然后沿着z轴切2D切片。这里有个非常关键的过滤条件我只保留包含肿瘤标签的切片因为整包数据里大量切片全是背景喂给网络只会浪费时间。如果你想增加负样本防止误检可以按比例混入一些无肿瘤切片但baseline阶段先跳过。def extract_2d_slices(vol, seg): slices, labels [], [] for i in range(vol.shape[2]): if (seg[:, :, i] 0).sum() 0: slices.append(vol[:, :, i]) labels.append(seg[:, :, i]) x np.transpose(np.array(slices), (0, 3, 1, 2)).astype(np.float32) y np.array(labels).astype(np.int64) return x, y2.2 UNet模型定义下面这个UNet实现是PyTorch里非常经典的写法结构清晰直接贴在模型文件里就能用。import torch import torch.nn as nn class DoubleConv(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.conv nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue) ) def forward(self, x): return self.conv(x) class Down(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.mpconv nn.Sequential( nn.MaxPool2d(2), DoubleConv(in_ch, out_ch) ) def forward(self, x): return self.mpconv(x) class Up(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.up nn.ConvTranspose2d(in_ch, in_ch // 2, 2, stride2) self.conv DoubleConv(in_ch, out_ch) def forward(self, x1, x2): x1 self.up(x1) x torch.cat([x2, x1], dim1) return self.conv(x) class UNet(nn.Module): def __init__(self, n_channels, n_classes): super().__init__() self.inc DoubleConv(n_channels, 64) self.down1 Down(64, 128) self.down2 Down(128, 256) self.down3 Down(256, 512) self.down4 Down(512, 512) self.up1 Up(512, 256) self.up2 Up(256, 128) self.up3 Up(128, 64) self.up4 Up(64, 64) self.outc nn.Conv2d(64, n_classes, 1) def forward(self, x): x1 self.inc(x) x2 self.down1(x1) x3 self.down2(x2) x4 self.down3(x3) x5 self.down4(x4) x self.up1(x5, x4) x self.up2(x, x3) x self.up3(x, x2) x self.up4(x, x1) return self.outc(x)这里我稍微解释一下设计细节。DoubleConv在每个尺度上做两层3x3卷积加BatchNorm加ReLU作用是提取当前尺度的特征。Down模块先做2x2最大池化降采样再继续卷积感受野逐步扩大。Up模块用转置卷积做上采样关键一步是torch.cat([x2, x1], dim1)把编码器同尺度的特征直接拼在通道维度上这就是跳跃连接的落地方式。最后用1x1卷积把64通道映射到类别数。实例化模型时输入通道数设为4对应四个模态输出通道数按你的标签类别数来。如果只做二分类肿瘤vs背景n_classes2如果做三分类整个肿瘤、核心、增强需要额外处理标签。2.3 Dice Loss与评估函数脑肿瘤分割最常用的损失函数是Dice Loss而不是交叉熵。原因是肿瘤区域在整张图里占的比例很低背景像素往往占了95%以上普通交叉熵会被背景主导模型倾向于把所有像素都预测成背景。Dice系数的本质是衡量两个集合的重叠程度对类别不平衡天然不敏感。Dice Loss就是1减去Dice系数import torch.nn.functional as F class DiceLoss(nn.Module): def __init__(self, n_classes, smooth1e-5): super().__init__() self.n_classes n_classes self.smooth smooth def forward(self, pred, target): # pred: (B, C, H, W) 未归一化logits # target: (B, H, W) 整型标签 pred F.softmax(pred, dim1) target_onehot F.one_hot(target, self.n_classes).permute(0, 3, 1, 2).float() intersection (pred * target_onehot).sum(dim(2, 3)) dice (2 * intersection self.smooth) / ( pred.sum(dim(2, 3)) target_onehot.sum(dim(2, 3)) self.smooth ) return 1 - dice.mean()评估的时候我们用的是Dice系数和上面的损失正好相反越大越好def dice_score(pred, target, eps1e-6): # pred和target都是二值数组 intersection (pred * target).sum() return (2 * intersection eps) / (pred.sum() target.sum() eps)3. 训练过程中最容易踩的坑以及我最终的配置3.1 切片尺寸与Patch策略BraTS原始图像尺寸通常是240x240x155直接整图喂进UNet显存会很紧张而且很多边缘区域全是背景浪费算力。实际使用中我采用两种方式处理要么Center Crop到固定尺寸比如224x224要么随机裁剪Patch比如128x128。Center Crop简单粗暴缺点是会丢掉脑壳边缘的部分信息。Patch训练更灵活每次随机裁一块包含肿瘤的区域相当于做了数据增强还能缓解显存压力。我最终用的是随机128x128 Patch训练推理阶段用Overlap-tile策略滑窗时让patch之间重叠一部分重叠区域取平均这样能有效避免输出边界出现接缝伪影。注意UNet下采样四次输入尺寸必须是16的倍数否则特征图尺寸对不上转置卷积上采样后无法精确恢复原分辨率。128、192、224、256这些尺寸都没问题。3.2 类别不平衡与标签重映射BraTS原始标签有4个值直接做4分类问题不大但很多baseline会先转成二分类整个肿瘤区域统一标为1背景标为0。这样做的好处是任务简单收敛快Dice指标更容易做高。如果你要做更细的分割就需要把标签从0、1、2、4重映射为0、1、2、3这步别忘了def remap_label(seg): seg seg.copy() seg[seg 4] 3 return seg二分类模式下Dice Loss的n_classes设为2即可。3.3 归一化细节与数据增强不同MRI机器的采集参数不一样图像灰度范围差异很大。所以测试时对每个病例的每个模态也要单独做z-score归一化而不是用训练集的统计量硬套。我的经验是每个模态分开算均值和方差比所有模态一起算效果更稳因为T1和T2的灰度分布本来就不是一个量级。数据增强我用的是随机翻转水平和垂直、随机旋转±15度、随机缩放。弹性形变对医学图像效果不错但会增加训练时间我一般先不加等基线跑通了再尝试。3.4 模型状态字典加载的常见问题训练过程中经常遇到的情况是改了一下网络结构比如把输入通道从3改成4然后直接加载以前的checkpoint结果报出一堆size mismatch错误。这个就是很多人困惑的“加载器和checkpoint模型是否通用”问题。答案很简单不通用。load_state_dict要求每个参数的名字和shape完全一致哪怕只改了一层卷积的输入通道数整个模型能加载的参数也会对不上。正确做法是checkpoint torch.load(unet.pth) model.load_state_dict(checkpoint, strictFalse)strictFalse可以忽略缺失和不匹配的参数但你得心里有数哪些层被跳过了。如果第一层输入通道改了加载进来的结果其实只保留了后面的层第一层相当于从头训练。这种情况我建议直接改第一层卷积权重的加载逻辑# 第一层Conv2d(3, 64, 3)加载到Conv2d(4, 64, 3)时 new_conv nn.Conv2d(4, 64, 3, padding1) with torch.no_grad(): new_conv.weight[:, :3] old_conv.weight # 复用原权重 new_conv.weight[:, 3] old_conv.weight[:, 0] # 新增通道用已有权重初始化改完再加载新通道就有了一个合理的初始状态训练起来会稳定很多。4. 评估指标、改进方向与进一步优化4.1 BraTS的评估指标怎么看BraTS挑战赛官方评估通常用Dice系数和Hausdorff距离两个指标。Dice衡量重叠度直观易懂Hausdorff距离衡量边界误差能捕捉到Dice反映不出来的“肿块的形状完全不对但重叠率还行”的情况。如果只做二分类baseline说实话只用Dice就够了。但心里要清楚Dice到0.85以上之后肉眼观察分割结果往往比Dice数值更能说明问题。很多情况下Dice提升0.02看上去的差别已经很大了。4.2 简单有效的改进在跳跃连接里加注意力如果你想让2D UNet再进一步我建议先试Attention U-Net里的注意力门控Attention Gate。思路非常简单在把编码器特征拼到解码器之前先用一个注意力系数过滤掉与目标无关的区域让网络更关注肿瘤附近的信息。核心代码就是这种门控结构class AttentionGate(nn.Module): def __init__(self, in_ch, g_ch): super().__init__() self.Wg nn.Conv2d(g_ch, in_ch, 1) self.Wx nn.Conv2d(in_ch, in_ch, 1) self.psi nn.Conv2d(in_ch, 1, 1) self.relu nn.ReLU(inplaceTrue) self.sigmoid nn.Sigmoid() def forward(self, x, g): # x: 编码器特征, g: 解码器上采样后的特征 g1 self.Wg(g) x1 self.Wx(x) attn self.sigmoid(self.psi(self.relu(x1 g1))) return x * attn在跳跃连接拼接前套一个这样的门控参数量增加不多但效果往往比单纯堆深度来得明显。4.3 进一步优化方向当2D UNet的baseline稳定之后可以考虑的方向包括深度监督Deep Supervision在解码器每个尺度都加一个1x1卷积输出对多尺度输出分别计算损失能缓解梯度消失加速收敛。混合精度训练用PyTorch的torch.cuda.amp显存占用能下降40%左右训练速度也更快。3D UNet在2D流程跑通后再切换到3D充分使用体数据的空间连续性Dice通常还能再涨2到3个点。4.4 推理阶段的Overlap-tile策略原版UNet论文里有个很实用的技巧推理时不要直接把整张大图一次性输入而是用滑动窗口做带重叠的patch推理。窗口和窗口之间重叠一部分重叠区域的预测概率取平均。这样可以避免单个patch边缘信息不完整导致的预测误差对脑肿瘤这种边界精细的任务尤其明显。我实测下来推理patch边缘的伪影确实会减少。代价只是推理时间变长一些但换来的分割质量提升是实打实的。5. 关于这套代码的补充说明整套代码跑下来我最想强调的一点是不要迷信花哨的网络结构先把这个最基础的UNet调稳把数据流程打通再逐步追加改进。我实际用过很多结构最后发现基础UNet配合好的预处理、合适的损失函数和增强策略效果已经能超过很多加了各种花活的模型。最后分享一个实操中的小技巧训练时每隔几个epoch把预测结果可视化一次保存成图片。这个非常有用因为loss曲线只能告诉你模型在变好还是变坏但看不出具体哪里分割得不对。可视化能让你一眼看出模型是不是把整个脑部都包进去了、有没有漏掉边缘的小病灶很多时候比盯着Dice数字更直观。调试模型这种事多看图少猜效率会高很多。本文还有配套的精品资源点击获取