SE-Res Block U型卷积神经网络:乳腺癌靶区分割实战指南
简介这份PDF文献面向医学影像与深度学习方向的研究者、放疗物理师及医工交叉专业学生聚焦乳腺癌保乳术后放疗中临床靶区与危及器官的自动分割难题。研究在传统U-net基础上引入残差单元与SE注意力单元构建SE-Res Block U型卷积神经网络以482例患者CT图像及结构信息为样本对临床靶区、心脏、左右肺和脊髓进行分割并采用戴斯相似系数与豪斯多夫距离评估精度为理解深度学习在放疗勾画中的应用提供了完整实验范式。资源包内仅含1个PDF文件约2.06MB即该篇正式发表的期刊论文全文涵盖研究背景、方法设计、结果数据与结论讨论便于精读与引用。目前已有422人学习下载。读者可从中获取网络结构改进思路、分割评价指标的具体数值以及小样本与微小体积目标预测不足的局限分析对开展医学图像分割建模与放疗自动勾画研究具有参考价值。1. SE-Res Block U型卷积神经网络做乳腺癌靶区分割为什么临床医生开始盯上这套方案放疗科物理师最头疼的环节之一就是对着 CT 一层一层勾画乳腺癌靶区和危及器官。一个熟练的医生勾完一例乳腺病例靶区加心脏、肺、脊髓这些危及器官少说也要四十分钟到一个小时遇到术后解剖结构改变或者体型特殊的病例时间还得翻倍。更麻烦的是不同医生之间勾画差异不小靶区边界差个几毫米剂量分布就跟着变。这几年 U-net 在医学图像分割里几乎成了默认选项但直接把原始 U-net 拿来跑乳腺癌 CT效果往往不够看——靶区边界模糊、小体积危及器官漏检、不同扫描协议下泛化差都是血泪经验。SE-Res Block U型卷积神经网络这套组合核心思路就是在 U-net 的编码器里嵌入残差块加通道注意力让网络自己学会“哪些特征通道对靶区更重要”从而在乳腺癌临床靶区与危及器官自动分割上把精度往上推一截。这篇不是讲论文是讲如果你要复现或者落地这套方案从数据准备、网络搭建、训练调参到推理部署每一步该怎么做、参数怎么设、哪里容易翻车。适合有一定深度学习基础、想在医学图像分割方向做落地的工程师和临床科研人员。2. 从 U-net 到 SE-Res Block为什么原始结构在乳腺 CT 上不够用2.1 U-net 在乳腺癌分割任务里的三个硬伤U-net 的编码器-解码器加跳跃连接结构在细胞、肺结节这类边界相对清晰的场景里表现很好。但乳腺癌靶区分割有几个特殊之处第一临床靶区CTV和计划靶区PTV的边界不像器官轮廓那样有明确的解剖分界很多时候依赖医生对浸润范围的判断边界灰度变化平缓第二乳腺 CT 里心脏、肺、脊髓这些危及器官的对比度差异大小体积结构比如脊髓在低分辨率特征图上容易丢第三不同扫描设备、不同层厚、有无对比剂导致数据分布差异明显。原始 U-net 的编码器用普通卷积堆叠每层提取的特征通道权重是固定的网络没法根据输入内容动态调整哪些通道更重要。结果就是靶区边界处特征响应弱小器官在深层特征图上被背景淹没。常见做法是在编码器里加残差连接缓解梯度消失但残差块本身不解决通道选择问题。2.2 SE-Res Block 到底改了什么SE-Res Block 是残差块和 Squeeze-and-Excitation 模块的结合。残差块负责让梯度跨层流动SE 模块负责通道注意力。具体来说SE 模块对每个残差块的输出特征图做两步操作Squeeze 阶段用全局平均池化把每个通道的空间信息压成一个标量相当于给每个通道算一个“全局得分”Excitation 阶段用两个全连接层学习通道间的非线性关系输出每个通道的权重系数再乘回原特征图。用公式描述就是给定特征图 $U \in \mathbb{R}^{H \times W \times C}$先做全局平均池化得到 $z_c \frac{1}{H \times W}\sum_{i1}^{H}\sum_{j1}^{W} u_c(i,j)$然后经过 $s \sigma(W_2 \delta(W_1 z))$ 得到通道权重最后 $\tilde{U} s \cdot U$。其中 $\delta$ 是 ReLU$\sigma$ 是 Sigmoid$W_1$ 和 $W_2$ 是两个全连接层中间有个降维比例 $r$通常取 16。在乳腺癌分割任务里这个机制的价值在于网络可以自动学习到“靶区边界处的通道”和“小器官对应的通道”应该给更高权重而背景纹理通道被抑制。这比单纯堆卷积层或者加注意力模块更轻量参数量增加很少。2.3 把 SE-Res Block 嵌进 U-net 的具体位置不是随便塞进去就行。我一般会按下面的原则放编码器每个下采样阶段后的两个 3×3 卷积替换成 SE-Res Block保持特征图尺寸不变瓶颈层用两个 SE-Res Block 堆叠因为这里通道数最多通道注意力收益最大解码器的上采样卷积后也加一个 SE-Res Block帮助恢复空间细节时重新校准通道跳跃连接处不做 SE因为跳跃连接传递的是浅层高分辨率特征加注意力反而可能引入噪声。这样改下来整体参数量比原始 U-net 增加大约 5% 到 8%但推理时间增加不到 10%在单张 12GB 显存的卡上跑 512×512 的 CT 切片完全没问题。3. 数据准备与预处理乳腺 CT 分割的脏活累活3.1 数据来源与标注格式乳腺癌放疗 CT 数据通常来自医院放疗科的计划系统格式多为 DICOM。标注一般由高年资医生在 CT 上勾画靶区CTV、PTV和危及器官心脏、双肺、脊髓、食管、甲状腺等。如果你拿到的标注是 RTSTRUCT 文件需要先转成掩膜。常见做法是用pydicom读 DICOM 序列用rt-utils或platipy解析 RTSTRUCT把每个结构的轮廓转成与 CT 同尺寸的二值掩膜。这里有个坑不同结构的 HU 值范围差异大CT 的窗宽窗位设置会影响可视化但不影响原始 HU。预处理时不要直接对原始 HU 做归一化而是先做窗宽窗位截断。乳腺软组织窗通常取 WL40, WW400肺窗取 WL-600, WW1500。但分割任务里我一般统一用软组织窗截断到 [-200, 300] HU然后归一化到 [0,1]这样靶区和大部分危及器官都能保留。3.2 重采样与各向同性处理乳腺 CT 层厚常见 2.5mm 或 3mm层内像素间距 0.8~1.2mm各向异性明显。直接送进网络会导致 Z 方向信息损失。标准做法是重采样到各向同性比如 1mm×1mm×1mm 或者 1.5mm×1.5mm×1.5mm。用SimpleITK的ResampleImageFilter插值方式选线性掩膜用最近邻。import SimpleITK as sitk def resample_image(image, target_spacing(1.0, 1.0, 1.0), is_labelFalse): original_spacing image.GetSpacing() original_size image.GetSize() target_size [ int(round(original_size[i] * original_spacing[i] / target_spacing[i])) for i in range(3) ] resampler sitk.ResampleImageFilter() resampler.SetSize(target_size) resampler.SetOutputSpacing(target_spacing) resampler.SetOutputOrigin(image.GetOrigin()) resampler.SetOutputDirection(image.GetDirection()) if is_label: resampler.SetInterpolator(sitk.sitkNearestNeighbor) else: resampler.SetInterpolator(sitk.sitkLinear) return resampler.Execute(image)这段代码的关键参数是target_spacing我一般设 (1.0, 1.0, 1.0)。如果显存紧张可以放宽到 (1.5, 1.5, 1.5)但靶区边界精度会下降约 2% 到 3% Dice。is_labelTrue时用最近邻插值避免掩膜出现小数标签。3.3 数据增强与类别不平衡处理乳腺癌分割里靶区和危及器官的体素占比通常不到 5%背景占 95% 以上。直接训练会导致网络偏向预测背景。常见做法是在损失函数里用 Dice Loss 加交叉熵的加权组合Dice Loss 对类别不平衡不敏感同时在数据增强时多做随机旋转±15°、随机缩放0.9~1.1、随机弹性形变增加正样本的多样性。我一般还会做一个操作对包含靶区的切片做 oversampling让每个 batch 里至少有一半切片包含靶区或危及器官。这样比单纯调损失权重更直接。4. 网络搭建与训练SE-Res Block U型网络的代码级实现4.1 SE-Res Block 的 PyTorch 实现import torch import torch.nn as nn class SEBlock(nn.Module): def __init__(self, channels, reduction16): super(SEBlock, self).__init__() self.avg_pool nn.AdaptiveAvgPool3d(1) self.fc nn.Sequential( nn.Linear(channels, channels // reduction, biasFalse), nn.ReLU(inplaceTrue), nn.Linear(channels // reduction, channels, biasFalse), nn.Sigmoid() ) def forward(self, x): b, c, _, _, _ x.size() y self.avg_pool(x).view(b, c) y self.fc(y).view(b, c, 1, 1, 1) return x * y class SEResBlock(nn.Module): def __init__(self, in_channels, out_channels, stride1, reduction16): super(SEResBlock, self).__init__() self.conv1 nn.Conv3d(in_channels, out_channels, kernel_size3, stridestride, padding1, biasFalse) self.bn1 nn.BatchNorm3d(out_channels) self.conv2 nn.Conv3d(out_channels, out_channels, kernel_size3, stride1, padding1, biasFalse) self.bn2 nn.BatchNorm3d(out_channels) self.se SEBlock(out_channels, reduction) self.relu nn.ReLU(inplaceTrue) self.shortcut nn.Sequential() if stride ! 1 or in_channels ! out_channels: self.shortcut nn.Sequential( nn.Conv3d(in_channels, out_channels, kernel_size1, stridestride, biasFalse), nn.BatchNorm3d(out_channels) ) def forward(self, x): residual self.shortcut(x) out self.relu(self.bn1(self.conv1(x))) out self.bn2(self.conv2(out)) out self.se(out) out residual return self.relu(out)reduction16是 SE 模块的降维比例通道数少于 16 时建议改成 8 或 4否则全连接层降维后信息损失太大。stride在编码器下采样时设为 2解码器里设为 1。残差分支的shortcut在通道数或尺寸变化时用 1×1 卷积对齐。4.2 U 型主干网络组装class SEResUNet(nn.Module): def __init__(self, in_channels1, num_classes8, base_channels32): super(SEResUNet, self).__init__() # 编码器 self.enc1 SEResBlock(in_channels, base_channels) self.enc2 SEResBlock(base_channels, base_channels * 2, stride2) self.enc3 SEResBlock(base_channels * 2, base_channels * 4, stride2) self.enc4 SEResBlock(base_channels * 4, base_channels * 8, stride2) # 瓶颈层 self.bottleneck nn.Sequential( SEResBlock(base_channels * 8, base_channels * 16, stride2), SEResBlock(base_channels * 16, base_channels * 16) ) # 解码器 self.up4 nn.ConvTranspose3d(base_channels * 16, base_channels * 8, kernel_size2, stride2) self.dec4 SEResBlock(base_channels * 16, base_channels * 8) self.up3 nn.ConvTranspose3d(base_channels * 8, base_channels * 4, kernel_size2, stride2) self.dec3 SEResBlock(base_channels * 8, base_channels * 4) self.up2 nn.ConvTranspose3d(base_channels * 4, base_channels * 2, kernel_size2, stride2) self.dec2 SEResBlock(base_channels * 4, base_channels * 2) self.up1 nn.ConvTranspose3d(base_channels * 2, base_channels, kernel_size2, stride2) self.dec1 SEResBlock(base_channels * 2, base_channels) # 输出层 self.out_conv nn.Conv3d(base_channels, num_classes, kernel_size1) def forward(self, x): e1 self.enc1(x) e2 self.enc2(e1) e3 self.enc3(e2) e4 self.enc4(e3) b self.bottleneck(e4) d4 self.dec4(torch.cat([self.up4(b), e4], dim1)) d3 self.dec3(torch.cat([self.up3(d4), e3], dim1)) d2 self.dec2(torch.cat([self.up2(d3), e2], dim1)) d1 self.dec1(torch.cat([self.up1(d2), e1], dim1)) return self.out_conv(d1)base_channels32是起点显存够可以加到 48 或 64但参数量会翻倍。num_classes根据你的标签数设比如靶区加 7 个危及器官就是 8。输入通道in_channels1因为 CT 是单通道灰度。4.3 损失函数与训练参数import torch.nn.functional as F class DiceCELoss(nn.Module): def __init__(self, weight_dice0.7, weight_ce0.3): super(DiceCELoss, self).__init__() self.weight_dice weight_dice self.weight_ce weight_ce def forward(self, pred, target): # pred: [B, C, D, H, W], target: [B, D, H, W] ce_loss F.cross_entropy(pred, target) pred_soft F.softmax(pred, dim1) target_onehot F.one_hot(target, num_classespred.shape[1]) target_onehot target_onehot.permute(0, 4, 1, 2, 3).float() dice_loss 0.0 for c in range(1, pred.shape[1]): # 跳过背景 intersection (pred_soft[:, c] * target_onehot[:, c]).sum() union pred_soft[:, c].sum() target_onehot[:, c].sum() dice_loss 1 - (2 * intersection 1e-5) / (union 1e-5) dice_loss / (pred.shape[1] - 1) return self.weight_dice * dice_loss self.weight_ce * ce_lossweight_dice0.7和weight_ce0.3是我在乳腺数据上试出来的比例Dice 占主导能缓解类别不平衡交叉熵保留梯度稳定性。优化器用 AdamW初始学习率 1e-4权重衰减 1e-5。学习率调度用 CosineAnnealingLR周期设 100 个 epoch。Batch size 根据显存单卡 12GB 跑 3D 数据一般设 2 到 4配合梯度累积模拟大 batch。训练时每 10 个 epoch 在验证集上算一次 Dice 和 Hausdorff 距离保存验证 Dice 最高的模型。如果验证损失连续 20 个 epoch 不降就提前停。5. 避坑与排查乳腺癌分割落地时最容易翻车的五个地方5.1 现象训练 Dice 很高推理时靶区边界偏移严重原因数据增强里的随机旋转和缩放对掩膜用了线性插值导致边界出现小数标签训练时网络学到模糊边界。解决掩膜增强必须用最近邻插值或者先增强图像再根据增强参数同步变换掩膜不要对掩膜单独做插值。5.2 现象小体积危及器官如脊髓Dice 始终低于 0.5原因脊髓在 CT 上体积小下采样四次后特征图上的响应几乎消失。解决在损失函数里给每个类别加权重小器官权重设大一点或者在解码器最后两层加辅助输出用深监督让浅层特征也参与小器官预测。5.3 现象不同医院数据混训后模型在某一批数据上性能骤降原因不同扫描协议的 HU 分布和噪声水平不同BatchNorm 统计量被某一批数据主导。解决用 GroupNorm 替代 BatchNorm或者做域自适应比如在训练时对输入做直方图匹配把不同来源的 CT 对齐到同一灰度分布。5.4 现象显存溢出batch size 只能设 1原因3D 卷积对显存消耗大SE-Res Block 里的全连接层在通道数大时也会占显存。解决用混合精度训练把部分卷积换成深度可分离卷积或者把输入 patch 从 128×128×128 降到 96×96×96。如果还不行用梯度检查点牺牲 20% 速度换显存。5.5 现象推理结果出现孤立小连通域靶区不连续原因网络在边界处预测概率波动阈值化后产生碎片。解决后处理用连通域分析保留最大连通域或者用形态学闭运算填补小孔。但注意不要过度平滑否则靶区边界会被侵蚀。6. 进阶技巧用测试时增强和模型集成把 Dice 再推两个点训练完一个 SE-Res Block U型网络如果验证 Dice 卡在 0.85 左右上不去可以试两个技巧。第一个是测试时增强对同一张 CT分别做原始、水平翻转、旋转 90°、旋转 180° 四种变换每种都跑一次推理然后把 softmax 概率图反变换回原空间取平均。这个操作不需要重新训练推理时间变成四倍但 Dice 通常能涨 1 到 2 个点。注意翻转和旋转后的概率图要正确逆变换否则叠加位置对不上。第二个是模型集成用不同的随机种子训练 3 到 5 个 SE-Res Block U型网络推理时把它们的概率图平均。集成对边界模糊的靶区特别有效因为不同模型在边界处的错误模式不完全相关平均后能互相纠正。代价是推理显存和时间的线性增长如果部署环境紧张可以只集成两个模型收益也能有 1 个点左右。还有一个我常用的技巧在推理阶段把输入 CT 的窗宽窗位做两种设置一种软组织窗一种宽窗分别推理后取平均。这样网络能看到不同对比度下的结构信息对小器官尤其有用。这个技巧在脊髓和食管分割上效果明显Dice 能涨 3 到 5 个点。最后说一个验证方法不要只看整体 Dice要按结构分别算 Dice 和 95% Hausdorff 距离。靶区的 Hausdorff 距离比 Dice 更能反映边界偏移临床上靶区边界差 3mm 就可能影响剂量覆盖。我一般会把每个结构的 Dice 和 HD95 列成表逐项对比哪一项掉得厉害就回去查对应的数据增强和损失权重。这套方案我从数据清洗到推理部署跑通大概花了两周中间在掩膜插值和 BatchNorm 上翻过车。如果你也在做乳腺癌靶区自动分割建议先把数据预处理和损失函数调稳再动网络结构。希望帮到你。本文还有配套的精品资源点击获取