左前降支(Left Anterior Descending Artery,LAD)是冠状动脉里最容易出问题、也最让分割算法头疼的一段血管。说它容易出问题,是因为冠脉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 Transformer(ViT)正是这样做的:图像先切成 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_size=7): """ 简化版 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_size=kernel_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(dim=1, keepdim=True) / (C ** 0.5) # attn_logits: (B, 1, k*k, L) # 在邻域维度上做 softmax attn = torch.softmax(attn_logits, dim=2) # 聚合 value out = (attn * v).sum(dim=2) # (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_size=7, mlp_ratio=4.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,下采样用 stride=2 的卷积,解码器用转置卷积。这个网络可以直接跑 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_channels=1, out_channels=2, base_dim=32, depths=(2, 2, 4), kernel_size=7): super().__init__() self.encoder1 = nn.Sequential( nn.Conv2d(in_channels, base_dim, 3, padding=1), 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, stride=2) 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, stride=2) 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, stride=2) 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, stride=2) 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, smooth=1e-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, dim=1) if target.dim() == 3: target_onehot = F.one_hot( target, num_classes=probs.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(dim=1) 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_channels=1, out_channels=2).to(device) loss_fn = DiceCE() optimizer = torch.optim.AdamW(model.parameters(), lr=3e-4) # dataloader 需要根据实际数据格式补齐 for epoch in range(100): avg_loss = train_one_epoch( model, dataloader, optimizer, loss_fn, device ) print(f"epoch {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