
左前降支Left Anterior Descending ArteryLAD是冠状动脉里最容易出问题、也最让分割算法头疼的一段血管。说它容易出问题是因为冠脉CTA影像里它走行最长、分支最多而且一路贴着心室表面既要穿过心肌又要绕开静脉和相邻心腔的阴影说它让算法头疼是因为它整体纤细、弯曲管腔体素占比极低在三维体数据里往往只有极少的前景体素而背景体素却有上亿个。用普通 3D U-Net 能分出大致轮廓但细小分支和血管连续性经常出问题直接上全局 Transformer显存又先撑不住。这篇文章要讨论的是“Neighborhood Attention Transformer”这一组合如何用于 LAD 的 3D 分割增强。核心判断是在细长管状结构分割里真正有价值的不是“看得更全”的全局注意力而是“知道该往哪里看”的局部注意力。读完这篇文章你可以理解邻域注意力与全局注意力、窗口注意力的差异了解这类网络在 3D 医学分割中的典型结构设计并拿到一份可以直接跑通的 PyTorch 语义实现和训练验证思路。需要提前说明的是本文不臆造论文里的具体实验数值和超参数重点是把方案原理、适用场景、工程实现和常见坑讲清楚。真正复现时请以原文实验设置和你的设备情况为准。1. 这篇文章真正要解决的问题LAD 的分割结果之所以重要不只是为了“把血管画出来”。在冠心病诊疗流程里LAD 的三维几何信息直接影响三种下游任务冠脉狭窄程度的量化评估。医生需要知道管腔在哪一段变窄、窄了多少这依赖于准确的管腔边界。血流动力学模拟例如基于 CTA 的 FFR 计算。模拟结果对血管中心线走向、分叉角度和管腔截面积非常敏感血管分割误差会直接放大到压力降的计算里。介入手术规划。支架长度、球囊直径、是否覆盖分叉病变都需要三维血管形态作为参考。所以LAD 分割不是一个“学术玩具”而是有明确临床价值的任务。但也是典型的高难度任务它的难点可以归纳成四点目标小。LAD 管腔直径通常只有几毫米体素占比极低正负样本极度不平衡。形状细长且弯曲。整个血管像一条空间曲线分割网络必须保持连续性中间断一段就是致命错误。对比度不稳定。钙化斑块、支架、心腔造影剂残留都会让局部外观变化很大。三维数据大。CTA 体数据往往达到数百万甚至上亿体素全局注意力在这种规模下几乎不可行。如果我们把目光放在“用什么网络结构”上过去常用的 3D U-Net 属于卷积路线优点是局部建模好、显存可控缺点是没有显式的长程依赖建模后来 Transformer 路线进入医学图像分割大家又发现全局注意力在 3D 场景里内存爆炸、小目标上容易学偏。该往哪边走就成了一个很实际的问题。而邻域注意力 Transformer 提供了一条值得尝试的中间路线它保留 Transformer 的动态聚合能力同时用“只看邻域”的方式把计算复杂度压下来。这篇文章后面所有内容都是围绕这条路线展开的。2. Neighborhood Attention 与 Transformer 核心概念2.1 自注意力机制回顾Transformer 的核心是自注意力。对于输入特征 $X \in R^{N \times C}$先通过三个线性变换得到 Query、Key、Value$$Q XW_Q, \quad K XW_K, \quad V XW_V$$然后计算注意力权重$$Attention(Q, K, V) softmax(\frac{QK^T}{\sqrt{d_k}})V$$在视觉任务里如果 $N$ 是整幅图像的所有 patch 数量这就是全局注意力。Vision TransformerViT正是这样做的图像先切成 patch再经过多层全局自注意力提取特征。从数学上看每个位置的输出是所有位置的加权和权重由内容相似度决定。这个设计的好处是动态、自适应用关系建模理论上可以捕获任意距离的依赖。2.2 全局注意力在 3D 医学图像中的三个短板第一个短板是显存和计算量。注意力矩阵大小是 $N \times N$在 3D 数据里 $N$ 很容易达到几十万甚至上百万$N^2$ 直接不可接受。即使只做一次注意力计算现代 GPU 的显存也撑不住。第二个短板是优化困难。全局注意力让每个位置都要和所有位置交互对数据和训练策略很敏感。医学图像里背景体素占绝对多数全局注意力很容易把大量权重分配给背景细小血管反而得不到足够关注。第三个短板是归纳偏置弱。Vision Transformer 刚被提出时数据效率不如 CNN必须靠大规模预训练或强数据增强。医学图像数据集普遍较小直接上全局注意力容易过拟合。所以一个很自然的想法是不学全部位置只学邻域位置。这正好是 Neighborhood Attention 的思路。2.3 从窗口注意力到邻域注意力讲 Neighborhood Attention邻域注意力之前先看它常用的两个参照系卷积。卷积核在局部窗口内滑动每个位置只聚合固定邻域内的信息位置之间共享权重。它天然适合局部结构但权重是静态共享的不能根据输入内容动态调整。Swin Transformer 的窗口注意力。把图像分成不重叠的窗口窗口内做全局注意力再通过移位窗口让信息跨窗口流动。它解决了全局注意力的计算问题但引入了窗口划分和掩码实现复杂度较高。Neighborhood Attention 的思路和这两者都不一样。它让每个 query 只关注自己周围某个半径内的 key相当于给每个位置都做一个“跟随自身移动的窗口”。注意三个要点窗口是滑动的不需要把图像切成固定 block。每个位置关注的邻域大小由 kernel_size 决定例如 7x7。不同位置的邻域可以重叠且位置越近天然越容易产生交互。这个思路在不同任务里都证明了有效性。把 Transformer 的“全局交互”改成“局部交互”换来的是线性计算复杂度同时保留了注意力的动态权重特性。2.4 三种注意力机制对比机制每个位置的关注范围计算复杂度实现复杂度代表方法全局注意力所有位置O(N^2)低ViT、原始 Transformer窗口注意力固定窗口内所有位置O(N * k^2)中Swin Transformer邻域注意力以当前位置为中心的邻域O(N * k^2)低至中Neighborhood Attention Transformer从工程角度看邻域注意力和窗口注意力复杂度差不多但少了窗口划分和掩码逻辑更适合向下游分割任务扩展。对于 LAD 这类局部结构高度重要的任务这个“限制注意力范围”的改动意义甚至比“换一个更大的 Transformer”更关键。3. 为什么邻域注意力适合 LAD 这类细长结构前面已经提到了计算复杂度这一节从医学图像本身的特点再多说几步因为这决定了网络设计时应该怎么分配计算量。3.1 血管形态是强局部连续的结构LAD 从冠脉开口出发沿前室间沟向下走行。相邻体素之间在灰度、位置、走向上高度相关血管壁的连续性也是局部的。也就是说一个体素是否属于管腔最有判别力的信息几乎都来自它周围的小邻域而不是远隔几十毫米的其他切片。全局注意力在这种任务里反而显得“浪费”它拿大量参数和显存去建模远处背景之间的相关性对血管局部细节的贡献有限。邻域注意力把计算集中到局部更贴合管状结构的形态先验。3.2 小目标分割需要更高分辨率CTA 里的 LAD 管腔直径可能只有 3 到 5 毫米转换成体素往往只占几十个像素宽度。要保住这些细节网络不能为了省显存把分辨率降得太狠。全局注意力受限于 $N^2$ 复杂度通常在低分辨率上使用邻域注意力是线性的可以配合更高分辨率特征图使用。这意味着同样的显存预算下邻域注意力能让你保留更多细节这对细血管分割是实打实的收益。3.3 分辨率灵活性和平移等变性Swin Transformer 的窗口注意力要求特征图尺寸能被窗口大小整除多尺度设计需要额外处理。邻域注意力没有这种约束只要 padding 做对任意尺寸都可以计算平移等变性和卷积更接近在分割任务里可以更灵活地组合不同尺度特征。但要强调一点邻域注意力不等于没有长期依赖。单个邻域注意力层的感受野是有限的长期依赖靠的是层层堆叠和多尺度下采样。所以在设计网络时不能只放一个邻域注意力层就期待它捕获长距离信息必须配合编码器-解码器结构和跳跃连接。3.4 局限在哪里邻域注意力也有代价。如果 kernel_size 太小网络只看局部容易丢失大范围上下文比如心脏整体位置信息、血管远端与心尖的关系如果 kernel_size 太大计算量又会上升和全局注意力的差距缩小。实际使用中一般通过不同 stage 使用不同 kernel_size或者混合 3D 卷积和邻域注意力来平衡。所以更稳妥的判断是邻域注意力适合作为 3D 分割网络里的“主力注意力机制”但最好是和卷积、和下采样结构配合而不是彻底替换一切。4. 面向 LAD 3D 分割的网络架构设计拆解4.1 总体结构一个典型的“基于邻域注意力 Transformer 的 3D 分割网络”通常采用编码器-解码器结构这一点和 3D U-Net、UNETR、Swin UNETR 是共通的。大致如下编码器先做 patch embedding 或卷积 stem再经过多级下采样每个 stage 由若干邻域注意力 Transformer 块组成。瓶颈最深层继续堆叠邻域注意力块保持全局抽象特征。解码器逐级上采样恢复分辨率并通过跳跃连接把编码器的细节特征传给解码器。输出头最后用 1x1x1 卷积输出每个体素的类别 logits通常包含背景和 LAD 两类。这个结构和 3D U-Net 几乎一模一样区别只在编码器内部的“基本块”从卷积块换成了邻域注意力 Transformer 块。4.2 Stem 与 Patch Embedding网络最开始需要把原始体数据变成特征图。常见做法有两种卷积 stem第一层用 stride 为 2 的 3D 卷积直接降采样并增加通道。Patch Embedding把相邻 patch 拉平并通过线性映射投影到特征维度类似 ViT 的做法。在医学图像里卷积 stem 通常更稳定。因为原始 CTA 数据噪声多、灰度范围大卷积能先做一个局部平滑和特征抽取再交给注意力块。4.3 邻域注意力 Transformer 块一个标准的邻域注意力块包含两个子层邻域注意力层对每个位置只和它周围 kernel_size 范围内的位置计算注意力。多层感知机MLP对每个位置的通道维度做非线性变换。每个子层前面有 LayerNorm后面有残差连接。这个结构和标准 Transformer 块几乎一样只把全局注意力替换成邻域注意力。用公式表示$$z z NeighborhoodAttention(LayerNorm(z))$$ $$z z MLP(LayerNorm(z))$$实际实现时3D 特征图的维度是 (B, C, D, H, W)而 LayerNorm 通常作用于 (N, L, C) 布局因此需要调整维度顺序。4.4 多尺度下采样与跳跃连接多尺度是血管分割的关键。LAD 在近段、中段、远段直径差别很大不同尺度的特征关注的信息不一样浅层高分辨率特征关注血管壁边缘、管腔边界。深层低分辨率特征关注血管整体走向、分叉关系、与心腔的相对位置。所以编码器通常包含 3 到 4 个下采样 stage每个 stage 后分辨率减半通道数翻倍。解码器通过上采样逐步恢复分辨率并把编码器对应的特征通过跳跃连接合并进来帮助恢复细节。4.5 输出头与损失函数输出头通常是一个 1x1x1 卷积把特征通道映射到类别数。由于 LAD 是一个二类分割问题logits 一般输出 2 个通道背景和 LAD。损失函数建议用 Dice Loss 和交叉熵损失的组合。Dice Loss 天然关注前景体素对类不平衡友好交叉熵损失提供更平滑的梯度有利于稳定训练。另外还可以加一层深监督让网络在多个解码器层级上都计算损失能明显加速收敛。5. 环境准备与数据前置条件5.1 硬件环境这个任务的内存和显存压力比较大。一份完整 CTA 体数据通常是 512x512x300 左右即使裁剪到感兴趣区域也要几百兆。建议准备GPU 显存不低于 12GB理想是 24GB 以上。内存不低于 32GB。存储空间预留训练数据和模型权重。如果设备有限可以把训练 patch 减小到 96x96x64 之类的尺寸通过滑动窗口推理来完整分割体数据。5.2 软件环境建议使用以下环境Linux 或 Windows 都可以Linux 更适合长时间训练。Python 3.9 或更高版本。PyTorch 2.x 或更高版本版本请以实际项目为准本文演示通用 API。NumPy、SimpleITK 或 nibabel 用于读取 NIfTI 等医学图像格式。MONAI 可以按需安装它提供了很多医学图像预处理和评估组件但不是必须。5.3 数据组织医学图像分割项目一般建议统一使用 NIfTI 格式目录结构如下data/ ├── imagesTr/ │ ├── case_001.nii.gz │ └── ... ├── labelsTr/ │ ├── case_001.nii.gz │ └── ... ├── imagesVal/ │ └── ... └── labelsVal/ └── ...预处理时注意几个关键点统一体素间距spacing。CTA 各向异性明显统一到接近各向同性的 spacing 可以提升分割一致性。灰度归一化。建议对每个样本计算窗宽窗位把 CT 值裁剪到合适范围再做 z-score 归一化。裁剪感兴趣区域。如果只关心 LAD可以先裁剪到包含冠状动脉的区域节省显存。标注复核。血管标注容易出现标注员不一致训练前务必人工抽检。6. 核心代码实现邻域注意力与简易分割网络这一节给出可以直接运行的 PyTorch 语义实现。为了便于演示下面代码以 2D 版本为例说明原理扩展到 3D 时将卷积替换为 3D 卷积把邻域注意力算子换成支持 3D 的版本即可。核心思想完全一致。6.1 邻域注意力简化实现# 文件路径models/neighbor_attention.py import torch import torch.nn as nn import torch.nn.functional as F def neighborhood_attention(x, kernel_size7): 简化版 2D 邻域注意力实现。 每个位置只与周围 kernel_size x kernel_size 邻域内的位置计算注意力。 参数: x: (B, C, H, W) kernel_size: 邻域大小应为奇数 返回: out: (B, C, H, W) B, C, H, W x.shape pad kernel_size // 2 q x # 作为 query # 对 x 做 padding然后用 unfold 取出每个位置的邻域块 x_pad F.pad(x, [pad, pad, pad, pad]) windows F.unfold(x_pad, kernel_sizekernel_size) # windows: (B, C * k * k, L)其中 L H * W B, CK, L windows.shape k windows.view(B, C, kernel_size * kernel_size, L) v k # 这里 key 和 value 来自同一特征图 # 计算注意力分数query 与邻域内每个 key 做点积 q_flat q.view(B, C, L).unsqueeze(2) # (B, C, 1, L) attn_logits (q_flat * k).sum(dim1, keepdimTrue) / (C ** 0.5) # attn_logits: (B, 1, k*k, L) # 在邻域维度上做 softmax attn torch.softmax(attn_logits, dim2) # 聚合 value out (attn * v).sum(dim2) # (B, C, L) out out.view(B, C, H, W) return out这个实现是为了展示原理没有加入相对位置编码。真实项目里Neighborhood Attention 通常会配合相对位置偏置效果会更好。另外这个实现用 unfold 取邻域的效率一般生产环境建议使用更高效的 CUDA 算子例如 NATTEN 加速库。6.2 邻域注意力 Transformer 块# 文件路径models/nat_block.py import torch import torch.nn as nn from models.neighbor_attention import neighborhood_attention class NeighborhoodAttentionBlock(nn.Module): 标准 Transformer 块把全局注意力替换为邻域注意力。 包含 LayerNorm NeighborAttention 残差、LayerNorm MLP 残差。 def __init__(self, dim, kernel_size7, mlp_ratio4.0): super().__init__() self.norm1 nn.LayerNorm(dim) self.norm2 nn.LayerNorm(dim) self.kernel_size kernel_size self.mlp nn.Sequential( nn.Linear(dim, int(dim * mlp_ratio)), nn.GELU(), nn.Linear(int(dim * mlp_ratio), dim), ) def forward(self, x): # x: (B, C, H, W) B, C, H, W x.shape # LayerNorm 作用在通道维先调整维度再还原 x_norm self.norm1(x.permute(0, 2, 3, 1)).permute(0, 3, 1, 2) attn_out neighborhood_attention(x_norm, self.kernel_size) x x attn_out x_norm self.norm2(x.permute(0, 2, 3, 1)).permute(0, 3, 1, 2) x x self.mlp(x_norm.permute(0, 2, 3, 1)).permute(0, 3, 1, 2) return x这里真正容易踩坑的地方是维度顺序。PyTorch 的 LayerNorm 默认作用于最后一维而卷积特征图是 (B, C, H, W)所以必须把通道维换到最后一维再做 LayerNorm。如果忘了这一步训练时大概率直接报错或者内存悄悄涨得飞快。6.3 简易 U 型分割网络下面是一个极简的 U 型网络编码器用多个 NeighborhoodAttentionBlock下采样用 stride2 的卷积解码器用转置卷积。这个网络可以直接跑 2D 分割验证思路是否成立。# 文件路径models/simple_nat_unet.py import torch import torch.nn as nn from models.nat_block import NeighborhoodAttentionBlock class SimpleNATUNet(nn.Module): 简易 2D 邻域注意力分割网络。 用于验证 Neighborhood Attention 在分割任务上的效果。 3D 版本需要把卷积替换为 3D 卷积并适配注意力算子。 def __init__(self, in_channels1, out_channels2, base_dim32, depths(2, 2, 4), kernel_size7): super().__init__() self.encoder1 nn.Sequential( nn.Conv2d(in_channels, base_dim, 3, padding1), nn.GELU(), ) self.blocks1 nn.ModuleList([ NeighborhoodAttentionBlock(base_dim, kernel_size) for _ in range(depths[0]) ]) self.down1 nn.Conv2d(base_dim, base_dim * 2, 2, stride2) self.blocks2 nn.ModuleList([ NeighborhoodAttentionBlock(base_dim * 2, kernel_size) for _ in range(depths[1]) ]) self.down2 nn.Conv2d(base_dim * 2, base_dim * 4, 2, stride2) self.blocks3 nn.ModuleList([ NeighborhoodAttentionBlock(base_dim * 4, kernel_size) for _ in range(depths[2]) ]) self.up2 nn.ConvTranspose2d(base_dim * 4, base_dim * 2, 2, stride2) self.blocks4 nn.ModuleList([ NeighborhoodAttentionBlock(base_dim * 2, kernel_size) for _ in range(depths[1]) ]) self.up1 nn.ConvTranspose2d(base_dim * 2, base_dim, 2, stride2) self.blocks5 nn.ModuleList([ NeighborhoodAttentionBlock(base_dim, kernel_size) for _ in range(depths[0]) ]) self.head nn.Conv2d(base_dim, out_channels, 1) def forward(self, x): # 编码器 x1 self.encoder1(x) for block in self.blocks1: x1 block(x1) x2 self.down1(x1) for block in self.blocks2: x2 block(x2) x3 self.down2(x2) for block in self.blocks3: x3 block(x3) # 解码器 x self.up2(x3) x x x2 for block in self.blocks4: x block(x) x self.up1(x) x x x1 for block in self.blocks5: x block(x) return self.head(x)这里的跳跃连接用了最简单的直接相加。实际项目中也可以像 U-Net 那样在通道维拼接。对于 LAD 这种小目标拼接往往能保留更多浅层细节效果会更好。6.4 损失函数Dice Loss 与交叉熵组合# 文件路径losses/dice_ce.py import torch import torch.nn as nn import torch.nn.functional as F class DiceCE(nn.Module): Dice Loss CrossEntropy Loss 的组合损失适合小目标分割。 def __init__(self, smooth1e-6): super().__init__() self.smooth smooth def forward(self, logits, target): # logits: (B, C, H, W) # target: (B, H, W) 或 one-hot (B, C, H, W) probs torch.softmax(logits, dim1) if target.dim() 3: target_onehot F.one_hot( target, num_classesprobs.shape[1] ).permute(0, 3, 1, 2).float() else: target_onehot target.float() # 只计算前景类别跳过背景对血管这类小目标更友好 dice 0.0 for c in range(1, probs.shape[1]): p probs[:, c] t target_onehot[:, c] inter (p * t).sum() dice (2.0 * inter self.smooth) / (p.sum() t.sum() self.smooth) dice dice / (probs.shape[1] - 1) ce_target target if target.dim() 3 else target.argmax(dim1) ce F.cross_entropy(logits, ce_target) return dice * 0.5 ce * 0.5这类损失函数设计对血管分割影响很大。只用交叉熵网络容易偏向背景只用 Dice训练早期容易不稳定。二者按 0.5 和 0.5 组合是一个常用起点实际可以根据验证集表现调整权重。6.5 训练循环骨架# 文件路径train_seg.py import torch from models.simple_nat_unet import SimpleNATUNet from losses.dice_ce import DiceCE def train_one_epoch(model, dataloader, optimizer, loss_fn, device): model.train() total_loss 0.0 for batch in dataloader: image batch[image].to(device) # (B, 1, H, W) label batch[label].to(device) # (B, H, W) logits model(image) # (B, C, H, W) loss loss_fn(logits, label) optimizer.zero_grad() loss.backward() optimizer.step() total_loss loss.item() return total_loss / max(len(dataloader), 1) if __name__ __main__: device torch.device(cuda if torch.cuda.is_available() else cpu) model SimpleNATUNet(in_channels1, out_channels2).to(device) loss_fn DiceCE() optimizer torch.optim.AdamW(model.parameters(), lr3e-4) # dataloader 需要根据实际数据格式补齐 for epoch in range(100): avg_loss train_one_epoch( model, dataloader, optimizer, loss_fn, device ) print(fepoch {epoch:03d} loss {avg_loss:.4f})在完整项目里这个骨架需要补充验证集评估、学习率调度、模型保存和日志记录。但作为最小验证这样已经足够判断“邻域注意力能不能跑通”。7. 模型训练、验证与结果评估7.1 数据划分与增强建议把病例按患者维度划分避免同一患者的不同序列同时出现在训练和验证集里。常见比例是训练 70%、验证 15%、测试 15%。数据增强对血管分割的提升非常大。常用的有随机翻转。随机旋转角度不宜过大。随机缩放。弹性形变可以模拟血管走行变化。灰度扰动模拟不同 CT 设备差异。需要谨慎的是血管结构对几何变形比较敏感弹性形变强度过大会让标注和图像错位反而伤害精度。7.2 训练配置示例以下是 yaml 格式的示例配置实际数值需要根据设备和数据调整# 文件路径configs/lad_nat.yaml model: name: SimpleNATUNet in_channels: 1 out_channels: 2 base_dim: 48 depths: [2, 2, 6] kernel_size: 7 data: patch_size: [128, 128, 64] spacing: [0.25, 0.25, 0.5] normalization: zscore augmentation: flip: true rotation_degree: 15 intensity_shift: 0.1 training: optimizer: AdamW lr: 3.0e-4 weight_decay: 1.0e-4 scheduler: cosine epochs: 300