ARTICLE DETAIL

资讯详情

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

Serp-Mamba:基于蛇形扫描的视网膜血管分割新方案

Serp-Mamba:基于蛇形扫描的视网膜血管分割新方案 做眼底血管分割的同行应该都有过这种体验U-Net 这类 CNN 在 DRIVE 上能轻松把主血管分得又粗又完整但一到毛细血管、分叉点和遮挡区域就断成碎线后处理怎么连都连不回来。我去年被这个问题卡了很长一段时间试过加深网络、加注意力、换各种损失函数指标涨一点点血管拓扑却还是老样子。后来把主干换成 Mamba 结构再配合带蛇形扫描的 Serp-Mamba 方案连通性终于有了肉眼可见的改善。这篇文章不聊玄学就讲讲视网膜分割为什么是 CNN 的照妖镜、Mamba 到底补了什么短板、Serp-Mamba 的蛇形扫描究竟在扫描什么。这篇东西适合正在做医学图像分割、血管建模或者刚接触 Mamba 系列结构的人看。我会把原理尽量说人话最后也会给一套可以直接照抄的训练配置和避坑清单。1. 视网膜血管分割到底难在哪1.1 血管不是面条是树很多人第一次接触视网膜分割时都觉得这不就是一个二分类分割任务吗瞳孔照片里亮的是背景暗的是血管把像素标出来不就行了。真正上手才发现问题远没那么简单。视网膜血管是一个高度分叉的树状结构从视盘附近的主血管出发逐级变细最后变成只有几个像素宽的毛细血管。在 DRIVE 数据集里面图像分辨率是 565×584一张图里血管像素大概只占 10% 到 12%。更麻烦的是血管和背景之间的对比度在末梢区域非常低有些细血管的灰度值和背景几乎一样只靠局部像素信息根本分不出来。我用一个比较直观的比喻血管分割很像在一幅浅色水墨画里描出所有细枝。主干好认但越到末梢越靠“猜”。这个“猜”不是猜颜色深浅而是要猜这条线往哪个方向走、下一步该不该接上。换句话说分割算法不能只看每个像素像不像血管还要理解整棵血管树是怎么连起来的。这种连续性和结构性问题恰恰是传统卷积神经网络最不擅长的地方。1.2 CNN 的局部性假设在这里失效了CNN 的核心机制是卷积核它假设图像特征是局部相关的也就是离得越近的像素关系越强离得越远的像素关系可以被忽略或者由层层抽象来补偿。听起来挺合理但视网膜血管不跟你讲这个道理。一条毛细血管可能跨越几十、上百个像素中途分叉、转弯、交叉甚至和背景融合再出现。面对这种长程依赖CNN 只能用堆层数的方式把理论感受野撑大。比如一个 3×3 卷积经过 stride2 的下采样感受野的增长确实是指数级的但有效感受野远远小于理论值。关于有效感受野Luo 等人专门研究过这个问题卷积网络实际使用的有效感受野只占理论感受野的一小部分而且形状近似高斯分布中心区域贡献大边缘区域的感受野贡献几乎可以忽略。放在血管分割里这意味着网络名义上看到了 100×100 的上下文实际上注意力还是集中在中心几十个像素内。毛细血管末梢周围全是低对比度的背景局部窗口里根本没有足够线索判断它到底是不是血管。另一个更致命的点是多次下采样。UNet 这类架构通常要下采样 4 到 5 次血管主干信息保住了但细分支在下采样后的特征图里可能只占 1/32 甚至更少的像素面积。等上采样回去这些细结构早就被磨平了。可以这么理解把一张高清山水画缩小成邮票大小再放大到原尺寸细枝末节的笔触必然磨损。1.3 评价标准里藏着拓扑学的影子为什么很多 CNN 模型的指标看起来挺好实际使用却总觉得不行因为常规指标本身就不够敏感。Dice、IoU、AUC 这些都是像素级的集合相似度它们不关心分割结果里的血管是不是连成一条完整的通路。举个例子一根连续血管断成三截只要三截加起来覆盖的像素和真值差不多Dice 依然可以很高。但从医生角度看这完全不能用——断了就等于信息丢失。所以在血管分割领域大家还会看几个拓扑相关的指标比如 Betti 数误差和 clDice。Betti 数里 β0 表示连通分量个数β1 表示环形结构个数。血管树断得越多β0 误差就越大。clDice 是中心线 Dice先对预测结果和真值分别提取中心线再算中心线交叉部分的 Dice对细血管的连续性非常敏感。这些指标考察的就是“拓扑正确性”。传统 CNN 的问题在像素级指标上可能只差零点几个点但在这些拓扑指标上会差出一大截。换个说法CNN 的失败不是因为他看不清血管而是因为它不明白血管该是怎么连的。2. Mamba 凭什么打赢这场仗2.1 状态空间模型的一页纸原理Mamba 这个名词在 2023 年突然火起来核心是一类基于状态空间模型SSM的序列建模方法。很多人一听状态空间就头大我可以稍微拆一下。经典的连续状态空间模型可以写成这样的微分方程h(t) A h(t) B x(t) y(t) C h(t) D x(t)其中 x(t) 是输入信号h(t) 是隐藏状态y(t) 是输出。这里的核心思路是把输入序列逐步“读”进一个状态向量 h 里状态向量就像模型的工作记忆保存到目前为止见过的输入中有用的信息。在实际计算时要先把连续方程离散化成步进形式h_k \barA h_{k-1} \barB x_k y_k \barC h_k \barD x_k这里 \barA、\barB 等是由原参数和步长决定的离散化矩阵。每读入一个新输入 x_k模型更新一次记忆 h_k然后给出当前输出 y_k。这个过程对序列长度是线性的。处理 N 个 token 的复杂度是 O(N)而 Transformer 的注意力机制是 O(N²)。N 越大这个优势越明显。2.2 选择性扫描在选什么Mamba 真正和传统 SSM 拉开差距的地方是选择性机制。传统 SSM 的参数 A、B、C 是固定的对任何输入都用同一套规则相当于不管进来什么信息都一视同仁地记。Mamba 让 B、C 以及离散化步长 Δ 都依赖于当前输入 x_k也就是说模型可以自己决定当前这一步是记住还是忽略记住了应该往状态里写多少。这个机制在血管分割里很有价值。一路扫描过去背景占大多数真正的血管信息只有 10% 左右。如果对所有像素都用同样的权重记忆状态向量很快就被背景信息塞满。选择性机制让模型可以把关注力集中在那些“看起来像血管末端”或“血管走向有变化”的像素上关键信息能沿序列传播得更远。为了跑得动Mamba 还配了硬件感知的并行扫描算法在 GPU 上用 chunk 方式分段计算。这些工程细节不展开讲只需要知道结论Mamba 做到了线性复杂度同时保留了类似注意力的长程选择能力。2.3 为什么图像任务也能用序列模型图像本来不是一维序列但可以展平成一维去处理。把 H×W 的二维特征图按某种顺序拉成长度为 H×W 的 token 序列然后喂给 Mamba就能享受长程依赖建模的好处。问题来了顺序怎么选这太关键了。Mamba 处理的是序列数据扫描路径决定了哪些像素在序列里靠得近、哪些被拉得远。二维网格里本来相邻的两个像素如果展平顺序不好可能在序列里相隔几万步。状态向量的容量终究有限距离越远的信息越难被完整带到当前位置。早期 Vision Mamba 的方法通常做四个方向扫描左上到右下、右下到左上、右上到左下、左下到右上。通过多方向扫描让各个空间位置都有机会在序列中近距离相遇。也有方法在图像块内部先扫描再跨块扫描形成层次化序列。这些方案有效但和血管结构并不是完全对齐的。这就给 Serp-Mamba 这类专门为细长连续结构设计的扫描方式留了发挥空间。3. Serp-Mamba 的蛇形扫描黑科技拆解3.1 传统扫描方式在血管上的尴尬我先说一下传统多方向扫描的问题。假设图像是 8×8 的格子标准的逐行扫描顺序是第 1 行从左到右 第 2 行从左到右换行跳到下一行行首 第 3 行从左到右 ……这种顺序下第二行末尾的像素和第三行开头的像素在空间上其实相邻一个在最右一个在最右但在序列里中间隔了一整行距离被拉大了 8 倍。换到更大的图一行 512 个像素那么行尾和下一行行首在序列里的距离就差了 512。四方向扫描的本质是用四个方向的“硬切”来弥补这种断裂。你从左到右扫一遍再从右到左扫一遍血管在至少一个方向上会比较连续。但视网膜血管是连续弯曲的很多时候既不水平也不垂直而是斜向延伸。一条 45 度的血管在行扫描中会被切成很多小段每一段又和上一段在序列里相隔很远。模型虽然能学会跨距离建模但状态容量有限长距离传递多了信息还是会丢。3.2 蛇形扫描路径与结构对齐Serp-Mamba 的“Serp”来自 serpentine也就是蛇形。蛇形扫描的思路听上去很简单第一行从左到右第二行从右到左第三行从左到右如此反复。路径形成一条连续的蛇形曲线。示意图长这样第 1 行0 1 2 3 4 5 6 7 第 2 行15 14 13 12 11 10 9 8 第 3 行16 17 18 19 20 21 22 23 ……这种扫描方式的关键优势在于相邻两行的连接处不是从行尾跳到下一行行首而是行尾直接连接到下一行的行尾二者在二维空间里本来就是相邻的。整条序列的局部邻域和二维图像的局部邻域高度重合。这和血管有什么关系血管本质上是二维平面上的一条连续曲线沿血管走向的相邻像素在二维空间里也是相邻的。蛇形扫描保证了空间相邻性也就是保证血管的“故事线”在序列里不会被切断。对 Mamba 来说扫描是在讲一个故事故事里的上下文越连续模型理解起来就越顺。打个比方四方向扫描相当于把一张完整地图剪成四个方向的长条让模型自己拼。蛇形扫描相当于顺着地图上的小路一条条走走到头转个弯继续走路始终是连续的。3.3 Serp-Mamba 的架构是怎么组织的Serp-Mamba 并不是单纯把扫描顺序换一下就完事。我按这类思路复现时的理解是它是一个编码器-解码器结构在编码器的不同阶段用蛇形 Mamba 模块替换原来的卷积下采样模块。一个典型的 SerpMambaBlock 可以拆成四步先把二维特征图按蛇形路径展开成一维 token 序列。把 token 序列过 Mamba 状态空间层这一步负责长程依赖建模。做反序列化把序列折回二维特征图。加一层残差连接和一个轻量卷积做局部特征补充。如果只有单向蛇形模型对竖直方向的长程依赖仍然弱。因为蛇形路径在拐弯处虽然保持空间相邻但从第一行最后到第二行最后的跳跃本质上仍然等于做了一个垂直移动一次垂直移动只前进一个像素。为了缓解这个问题实际实现里往往会做两个方向的蛇形扫描一个从左到右一个从右到左再把两条序列过两个独立的 SSM 分支最后融合输出。我试过的组合里双向蛇形加一个十字交叉扫描的效果最稳。十字交叉可以补足纯水平蛇形在垂直方向连续性上的不足也让分叉点附近的信息能更快地互相接触。代码层面一个最简的蛇形展开伪代码长这样import torch def serpentine_flatten(x): # x: (B, C, H, W) B, C, H, W x.shape out [] for i in range(H): row x[:, :, i, :] # 第 i 行 if i % 2 1: row torch.flip(row, dims[-1]) # 偶数行从左到右奇数行从右到左 out.append(row) return torch.cat(out, dim-1) # (B, C, H*W) def serpentine_unflatten(seq, H, W): # 反向折回二维步骤相反 ...实际工程里不会真的用 Python 循环一行行处理因为太慢。更适合的做法是生成一个索引矩阵然后直接用 index_select 完成重排。GPU 上一次性做索引取值比逐行循环快很多。4. 实操一个可以直接照抄的训练流程4.1 数据准备与预处理我用的是公开数据集 DRIVE总共有 40 张眼底图像其中 20 张训练、20 张测试分辨率统一为 565×584每张图像都有医生手工标注的血管掩膜。数据预处理有几个关键点取绿色通道作为主输入。眼底图里红色通道过曝蓝色通道噪声大绿色通道的血管对比度最高。做 CLAHE 对比度受限自适应直方图均衡化增强末梢血管和背景的区分度。归一化到 [0, 1]。随机中心裁剪到 512×512没有 padding保持血管结构和标签空间一致。数据增强方面我开了随机旋转 0 到 180 度、水平垂直翻转、轻度随机缩放、弹性形变、亮度对比度扰动和伽马校正。血管分割的训练数据量不大这些增强对防止过拟合非常关键。特别是弹性形变极度适合血管这类细长结构可以模拟不同个体的血管弯曲差异。4.2 模型搭建的关键模块下面是一个最简的 SerpMambaBlock 示意适合搭建初期快速验证效果。import torch import torch.nn as nn from mamba_ssm import Mamba class SerpMambaBlock(nn.Module): def __init__(self, dim, d_state16, d_conv4, expand2): super().__init__() self.norm nn.LayerNorm(dim) self.mamba Mamba( d_modeldim, d_stated_state, d_convd_conv, expandexpand, ) self.local nn.Conv2d(dim, dim, kernel_size3, padding1, groupsdim) self.norm_conv nn.GroupNorm(num_groups8, num_channelsdim) def forward(self, x): B, C, H, W x.shape identity x # 蛇形展开 seq serpentine_flatten(x) # (B, C, H*W) seq seq.permute(0, 2, 1) # (B, H*W, C) seq self.mamba(self.norm(seq)) seq seq.permute(0, 2, 1) # (B, C, H*W) # 折回二维 x serpentine_unflatten(seq, H, W) # 局部卷积补充 x self.local(x) x self.norm_conv(x) return x identity这里的 Mamba 用的是 mamba_ssm 官方实现。d_state 控制状态维度我一般取 16d_conv 是内部的一维卷积卷积核大小取 4expand 控制内部隐藏维度放大倍数取 2。搭建完整模型时我采用结构是一个 stem 卷积把通道数从 3 提到 32然后连续四个阶段下采样通道数依次为 64、128、256、512。每个阶段末尾插入一个 SerpMambaBlock。解码器部分用转置上采样并带上 UNet 风格的跳跃连接。骨架结构和经典的 UNet 很像只是把某些阶段的关键模块换掉了。4.3 损失函数与评价指标怎么选损失函数我用了组合损失实测比单一 Dice 损失更稳def hybrid_loss(pred, target): bce nn.functional.binary_cross_entropy_with_logits(pred, target) prob torch.sigmoid(pred) dice 1 - (2.0 * (prob * target).sum() 1.0) / ( prob.sum() target.sum() 1.0 ) return 0.5 * bce 0.5 * diceBCE 提供稳定的梯度Dice 损失强制关注正样本区域。混合损失在血管类不平衡的场景下比单独用 BCE 收敛快很多。如果追求更强的拓扑保持可以考虑在损失上加一个 clDice 分支。clDice 需要先对预测和真值做拓扑细化提取中心线直接用的话会拖慢训练速度。一个折中方案是训练前期用混合损失训练后期再加 clDice 做微调。评估指标除了 Dice、IoU、AUC还要算 Betti 数误差和 clDice。单纯看 Dice 容易自我欺骗。我的实验里出现过 Dice 涨了 0.01但连通分量误差变大几十个的情况。4.4 训练资源配置与收敛表现我用的配置如下输入512×512 随机裁剪batch size8优化器AdamWlr1e-4调度器余弦退火最小 lr1e-6训练轮数200混合精度开启 AMP硬件单张 RTX 309024GB 显存显存占用大概在 13GB 到 17GB 之间主要取决于中间特征图的通道数和 Mamba 的 expand 参数。实测下来的情况是前 20 轮 Dice 快速涨到 0.80 左右然后进入慢涨阶段。150 轮之后基本收敛Dice 稳定在 0.84 到 0.85 区间。相比同参数量的 UNet 和 ResUNetDice 高出 1 到 2 个点Betti 数误差下降尤其明显。同一张测试图UNet 会把一条细血管断成 5 段Serp-Mamba 输出基本能保持一整条。推理阶段512×512 的单张图像在 3090 上大概 60 到 90 毫秒加上后处理也就 0.1 秒左右完全可以用于批量离线分析。5. 踩坑记录与问题速查5.1 扫描顺序、位置编码与长序列第一次把 Mamba 用在图像上最容易踩的第一个坑是序列长度。512×512 展平后有 26 万个 token如果直接对原始像素做 SSM显存必爆。解决方案是先做一个 4×4 的 patch embedding把序列长度降到 65536或者下采样后每层只处理当前位置的特征图。第二个坑是位置编码。Mamba 本身对序列顺序有一定感知能力但二维结构信息经过蛇形展开后还是会丢失。我的经验是给不同尺度的 token 加上可学习的相对位置偏置或者在 Mamba 层内部引入简化的二维坐标嵌入。加不加位置编码对 Betti 误差影响很大差 3 到 5 个点是常事。第三个和扫描方向有关。只做单向蛇形扫描对近水平血管效果很好但对竖直走向的血管效果就差一些。我建议至少做双向蛇形扫描一条从左到右一条从右到左两条序列并行过 Mamba 再融合。想要更全的话再叠加一个竖直方向的蛇形扫描形成二维蛇形网格。多方向扫描的代价是计算量加倍但分割质量的提升是值得的。5.2 细血管断裂的边界在哪里修如果模型的输出中细血管还是断先别急着换结构先做两件事。第一检查阈值。直接用 sigmoid 输出大于 0.5 作为血管往往会丢掉置信度不高的末梢血管。我一般会在验证集上搜索最佳阈值或者用 Otsu 全局阈值。合理调阈值通常可以让 Dice 涨 0.005 到 0.02。第二做连通域分析。从预测结果里找出所有连通的血管区域删掉面积小于某个阈值的孤立小块。同时可以把断开的端点做一次最近邻匹配如果两个端点距离很近并且走向一致就补一条短线段连上。这个后处理不复杂几分钟就能写完但对肉眼观感提升极大。需要注意的是不要过度后处理。如果错误地把两个本不相连的血管补到一起虽然连通性指标好看了但临床上是另外一种错误。后处理参数要基于验证集确定不要拿测试集反复调。5.3 常见问题速查表现象可能原因解决办法训练不收敛loss 一直震荡Mamba 内部状态维度 d_state 太小或学习率过大把 d_state 调到 16 或 32学习率降到 5e-5加 warmup显存爆炸序列长度过长或 batch 太大加 patch embedding缩小输入分辨率batch_size 降到 4开梯度累积血管断点集中在水平方向扫描方向单一对竖直结构建模不足加反向蛇形扫描或竖直扫描分支Dice 高但连通性差只用了像素级损失缺少拓扑损失训练后期加 clDice或对预测结果做连通域后处理细血管完全丢失下采样次数太多细节磨平减一个下采样层或在不同尺度特征图上都加蛇形 Mamba 模块推理结果带上明显条带感蛇形序列化与反序列化后边缘位置出现伪影在蛇形展开前后加一维卷积过渡或者做重叠分块扫描还有一个容易被忽略的坑是模型初始化。Mamba 的初始化方式对训练稳定性影响很大直接用默认初始化可能在前几十个迭代产生极端梯度。我建议在 SerpMambaBlock 的 LayerNorm 上设置一个较小的初始化权重并保证跳跃连接里的卷积初始化为单位映射或者说初始化为近似恒等这样深层状态可以平稳起步。6. 写在后面医学分割补什么才叫真进步做完这个项目之后我最大的体会是视网膜血管分割的核心不是像素分类而是结构推理。CNN 输在视野太窄所有上下文都是局部堆出来的Transformer 和 Mamba 赢在能建立长程依赖但“怎么组织这个依赖”比“能不能建立依赖”更关键。Serp-Mamba 的蛇形扫描本质上是在告诉模型你顺着这条路往前看血管通常是连续走完的。如果你也想复现我建议按这个顺序来先跑通 UNet确认自己的数据管线没问题然后把中间的卷积模块替换成 SerpMambaBlock加双向蛇形扫描损失函数换成 BCE 加 Dice最后再加 clDice 微调。每一步改动都单独记录指标变化尤其是 Betti 误差。这样即使模型不 work你也能清楚知道是扫描问题、损失问题还是数据问题。最后再分享一个小技巧训练的时候把每一轮的验证结果里断点最多的那张图单独存下来每隔几个 epoch 翻出来看一眼。指标不会骗人但指标也不会告诉你血管到底是怎么断的。图片会。
返回列表