我从2022年底开始关注扩散模型在医学图像分割里的应用,到现在也算踩了不少坑、跑了不少实验,发现很多人对“Diffusion分割”的理解还停留在“用stable diffusion生成图像”这个层面,这是个挺大的误解。
这篇内容我想系统聊聊:去噪扩散技术到底怎么用于医学图像分割,它和传统分割模型(像UNet、Transformer)的本质区别在哪里,以及如果你想自己训练一个扩散分割模型,核心要关注哪些环节、会踩哪些坑。内容会尽量兼顾原理和实操,既适合刚入门的研究生,也适合想把这个技术落到具体项目里的工程师参考。
1. 内容整体设计与思路拆解
1.1 扩散分割到底是什么,和图像生成有什么区别
先纠正一个常见的认知偏差。很多人一听到“Diffusion”就想到stable diffusion画图,想到CLIP文本引导,想到各种风格化生成。但在医学图像分割语境下,扩散模型并不是用来“画”分割图的,而是用来“逐步细化”分割结果的。
传统分割模型本质上是判别式的,输入一张图像,直接输出一个概率图或分割掩码,一步到位。而扩散分割走的是另一条路:它把分割问题建模成一个条件生成问题。具体来说,我们不再直接预测分割掩码,而是从一个随机噪声图出发,在给定的医学图像条件下,一步步去噪,最终得到一个清晰、精确的分割结果。
这个思路的转变非常关键,它带来两个直观好处:第一,扩散模型的生成天然具备“多模态”能力,也就是说同一个输入图像,模型可以生成多个略有差异但都合理的分割结果,这对于医学标注中存在边界模糊、医师之间存在主观差异的情况特别有价值;第二,扩散模型的逐步去噪过程对空间细节的保持能力很强,不容易出现传统模型那种“分割边界过于平滑”“小目标直接丢失”的问题。
1.2 为什么医学图像分割需要扩散模型
医学图像分割和自然图像分割有个很大的不同点,就是标注成本极其昂贵。一张高质量的分割标注图,往往需要专业医生花几十分钟甚至几个小时去勾画。而且医学图像本身存在很高的不确定性,同一个病灶在不同医生眼里的边界可能都有差异,甚至同一个医生不同时间的标注也不一样。
这时候传统判别式模型就暴露出一个问题:它强行把一个本来存在多种合理答案的问题,压缩成单一的确定性输出。模型只能学习所有医生的一个平均结果,遇到边界模糊区域时,输出就会变得模糊或者犹豫不决。
扩散模型天然适合这种场景。因为它学到的是分割结果的分布,而不是简单地学到输入到输出的映射函数。在推理时可以多次采样,得到多个可能的分割结果,这样就能同时告诉临床医生:“这里大概率是病灶,但这个边界区域存在一定不确定性。”这种不确定性估计在辅助诊断场景里非常宝贵,因为医生在做最终决策时,不仅想知道“哪里有问题”,还想知道“这个结论有多可靠”。
1.3 整体方案选型背后的考量
做扩散分割,市面上主要有两大类架构思路:
第一类是直接把分割当成条件生成任务,代表工作是MedSegDiff和nnU-Net与扩散模型的结合。这类方案通常以UNet作为扩散模型的骨干网络,把原始医学图像作为条件输入,通过交叉注意或concat的方式注入网络,让模型在去噪过程中不断参考原图结构。
第二类思路是做“分割结果精修”,先用传统分割模型产出一个粗分割结果,再用扩散模型对这个粗糙结果进行迭代优化。这个思路比较务实,运算量更小,因为扩散模型只需要学习“如何让粗糙结果变得更精细”这个相对简单的映射。
从我自己的实验经验来看,如果你的目标是发论文或者做前沿探索,第一类方案更有研究价值;但如果你面对的是一项实际落地任务,比如某个具体器官的分割,第二类方案往往更加实用,训练更稳、收敛更快、显存消耗也更友好。
提示:在实际动手之前,先想清楚你要解决的核心问题,是“边界不确定区域的建模”还是“自动化分割精度提升”。这个问题直接决定了你后续的模型设计方向。
2. 核心原理与关键环节解析
2.1 前向扩散过程:给分割标注加上噪声
扩散过程的核心思想很简单,就是把数据逐步破坏成纯噪声。在医学图像分割任务里,我们破坏的对象是分割掩码(mask),而不是原始的医学图像。
具体来说,给定一个干净的分割标注 $x_0$,我们通过一个固定的马尔可夫链,逐步向 $x_0$ 中添加高斯噪声,经过 $T$ 步之后,原始标注会变成一个接近纯噪声的图。这个过程也叫前向扩散过程,它的数学形式可以写成:
$$q(x_t | x_{t-1}) = \mathcal{N}(x_t; \sqrt{1-\beta_t} x_{t-1}, \beta_t I)$$
这里的 $\beta_t$ 是一个预先定义好的噪声调度表,控制了每一步添加噪声的强度。$\beta_t$ 越大,代表这一步对原始信息的破坏越严重。实际实现中,我们不会一步一步地去加噪声,而是利用重参数技巧,直接一步到位算出任意时刻 $t$ 的加噪结果。这个技巧在训练时能大幅提升效率。
2.2 反向去噪过程:从噪声中恢复分割结构
前向过程是把标注变成噪声,反向过程就是反过来,从噪声出发,在医学图像的条件下一步步还原出分割结果。
反向过程的核心是一个神经网络,通常用UNet作为骨干。这个网络的任务是预测噪声,或者说预测“如果我们要从 $x_t$ 恢复出 $x_{t-1}$,需要移除多少噪声”。训练时我们把网络预测的噪声和实际添加的噪声做对比,计算均方误差损失,然后更新网络参数。
在推理阶段,我们从一个完全随机的噪声图开始,利用网络预测的噪声,按调度表一步步去噪。每一轮去噪都会引入一定的随机性,这也是为什么同一个输入会得到多个稍有差异的分割结果。
这里有一个特别容易混淆的地方,值得单独拎出来说清楚。很多初学者以为扩散模型是直接输出一个干净图,实际上它输出的主目标是对噪声的估计,然后通过噪声估计推算出上一轮的去噪结果。如果你刚开始自己实现代码,这个逻辑一定要理顺,不然debug的时候会非常痛苦。
2.3 条件信息的注入方式:怎么把医学图像作为条件约束
在扩散分割中,原始医学图像相当于“指路人”,它必须在整个去噪过程中持续给模型提供结构引导。如果没有这个条件,模型就只是在生成一张随机的分割图,完全没有任何临床意义。
目前实现条件注入主要有三种方式:
第一种是直接在通道维度拼接。把当前时刻的带噪标注和原始医学图像拼接在一起,作为UNet的输入。这种方式最简单直接,也最稳定,缺点是原始图像信息只在输入端注入了一次,深层特征里对原图的引用能力弱一些。
第二种是交叉注意力机制。在UNet的每个尺度上,用可学习的注意力层,让去噪特征主动去查询原图特征中与自己相关的空间区域。这种方式信息流动更充分,尤其适合处理多模态输入,但实现复杂度明显更高,训练时需要更多显存。
第三种是条件归一化层。把原图特征经过全局池化后作为缩放和平移参数,注入去噪网络的特征图里。这种方式计算开销最小,但对空间细节的引导能力也最弱,适合对精细度要求没那么高的任务。
我在实际项目中通常的做法是,主干特征提取用UNet,条件注入采用“通道拼接+深层交叉注意力”的组合,前几轮去噪时原图信息主要通过拼接提供粗粒度引导,后几轮去噪时交叉注意力开始精细调控细节。这个组合在多个脏器和病灶分割数据集上表现都比较稳定。
注意:条件注入不是越复杂越好。在数据量有限的情况下,复杂的注入结构反而容易导致模型过拟合训练集中的原图纹理,而不是学到真正的结构先验。建议先从拼接方案开始,跑通整个pipeline后再逐步升级注入方式。
2.4 训练目标与损失函数设计
扩散分割模型的训练目标,往简单了说就是预测噪声,但直接采用标准DDPM的噪声预测目标做分割,往往效果不是最优的,因为分割任务对边界精度非常敏感,而标准噪声预测损失对高频细节的关注不够。实际使用中更推荐对噪声预测损失和分割质量指标做联合优化。
我推荐一个组合方案:以噪声预测的MSE损失作为主损失,附加一个Dice损失项。具体做法是,在训练过程中抽样几个去噪中间步,把当前预测的 $x_0$ 解码出来,和真实标注计算Dice损失,它的梯度再回传到网络。这样模型在训练时不仅要学会“噪声长什么样”,还要学会“我最终输出的分割图质量够不够好”。
这个联合优化方案有一个细节要注意,就是中间步的选取。如果中间步过于靠近纯噪声阶段,解码出的 $x_0$ 非常粗糙,计算出来的Dice对训练几乎没用;如果太靠近干净阶段,又失去了指导意义。我常用的做法是随机采样 $t \in [0.2T, 0.8T]$ 区间,让模型在信息量适中但不是完全确定的阶段学习分割质量。
2.5 推理策略:单次采样还是多次采样取平均
训练完成之后,推理阶段的策略选择也会显著影响最终分割精度。最简单的做法是只采样一次,从纯噪声出发,跑到最后一步得到一个结果。这种方式速度快,但随机性会导致单次结果的质量波动。
更推荐的做法是多次采样取平均。以我自己常用的参数为例,对同一个测试样本采样10次,得到10个分割概率图,然后逐像素取平均作为最终输出。这个策略能明显提升分割的鲁棒性和边界平滑度。代价是推理时间成倍增加,一张图可能要一到三秒,这在研究场景可以接受,但如果要部署到实时辅助诊断系统里就是个不小的挑战。
如果你需要在推理速度和结果质量之间找一个折中点,可以试试概率图平均和空洞采样提速结合的方式。关于空洞采样,核心思路是训练时用完整的时间步,推理时只取其中一部分步数来跑,比如训练时 $T=1000$,推理时只跑50步。因为扩散模型学到的去噪能力是连续的,跳着去噪也能取得不错效果,速度可以提升很多。
3. 实操过程与训练调优全记录
3.1 数据准备与格式处理
医学图像分割中最常见的数据格式是NIfTI文件,通常包含一个3D体数据和一个对应的标注文件。直接用3D体数据训练扩散模型非常吃显存,大多数实验室的显卡都扛不住,所以主流的做法是切成2D切片来训练。
切片切得好不好,直接影响模型效果。我看到很多人直接沿着轴向切一刀就开训,结果大量切片里只有极小一部分包含病灶,模型很快就偏向预测背景。更好的做法是先统计一下数据集中每张切片中目标区域的占比,只保留目标区域占比在合理范围内的切片,比如5%到95%之间,过滤掉纯背景和纯前景的无效切片。
如果目标区域较小,还有一个数据增强的锦囊,就是裁剪目标区域的bounding box,在局部区域上训练。这样模型能充分学习目标区域内部的结构特征,而不是把大量容量浪费在学习无关背景上。病灶分割里,我做过对比实验,采用这种局部裁剪训练的方式,病灶区域的Dice指标比直接全图切片训练能提升三到五个百分点。
3.2 扩散模型骨干网络的选择
扩散模型的骨干网络直接决定了模型的特征表达能力和训练资源消耗,这是整个pipeline中最值得花心思的部分。
以2D切片训练为例,我推荐从传统的2D UNet开始。它的结构足够简单,每个尺度都有清晰的skip connection,能让梯度顺畅回流。在2D UNet跑通流程后,再考虑往骨干网络里加入注意力机制,特别是空间注意力模块,它对病灶边缘的建模帮助很大。注意力的加入可以用在UNet瓶颈层以及下采样最深的两个尺度上,不需要所有层都加,否则训练负担会明显变大。
从UNet往Swin UNet或者UTNet这类Transformer-UNet混合结构迁移时需要注意,Transformer结构对训练数据量更敏感。如果你的训练数据只有几十例,Transformer骨干的表现很可能反而不如一个设计良好的卷积UNet。我做过的对比是:600例训练数据下,Swin UNet的效果能超出卷积UNet约1.5个Dice点;但在100例以内时,Swin UNet的优势几乎消失,训练速度却明显更慢。
3D数据方面,如果你有足够显存并且切的是3D patch,可以考虑3D UNet作为扩散模型的骨干。3D模型能利用更多空间上下文信息,在器官分割上的表现通常会比2D模型好一个档次。但显存开销极其惊人,通常需要集群训练,单卡只能做推理。我自己的建议是,除非你的任务确实对三维连续性要求很高,否则先从2D方案起步会更稳。
3.3 训练参数与调度策略
扩散模型的训练有几个核心超参数,噪声调度表、总步数、批次大小和学习率,它们对收敛速度和最终效果都有直接影响,我把我常用的初始化配置整理成了表格,方便你直接参考:
| 超参数 | 推荐值 | 说明 |
|---|---|---|
| 总扩散步数 T | 1000 | 太少则噪声粒度不够细,太多则训练耗时 |
| 噪声调度 β 范围 | 0.0001 ~ 0.02 | 可采用线性调度或余弦调度 |
| 训练批次大小 | 8 ~ 16 | 由显存决定,BatchNorm层需大于4 |
| 基础学习率 | 1e-4 ~ 2e-4 | AdamW优化器配合余弦退火效果最好 |
| 训练轮数 | 200 ~ 500 | 以验证集Dice不再上升为准 |
| 推理采样步数 | 50 ~ 100 | DDIM方式可极大压缩采样步数 |
噪声调度表的选择有个值得注意的细节。线性调度在自然图像上效果不错,但在医学图像这种背景占比高的数据上,线性调度的早期步长对标注掩码的破坏会显得过快。你可以尝试余弦调度,它在中间步数保持更温和的加噪速度,能让模型更充分地学习中等噪声水平下分割结构的恢复,我实测下来余弦调度在细小血管和狭窄组织的分割任务上会比线性调度稳定一些。
优化器方面,扩散模型训练首选AdamW而不是SGD,因为扩散模型的损失面相对崎岖,AdamW的逐参数自适应学习率能更好地稳定训练过程。在医学小数据集上,权重衰减不要设太大,我一般设0.01到0.05之间,权重衰减过大容易让模型在训练后期出现欠拟合。学习率预热也很重要,前500到1000步用很小的学习率热身,之后再进入主训练阶段,能有效避免前期剧烈波动。
3.4 训练过程监控与评价指标
训练过程中的监控指标不能只看噪声预测的loss,那样很容易被表面的收敛欺骗。建议在训练过程中定期跑一个快速推理,挑几个验证集样本做单步预测和完整去噪推理,把生成的分割结果可视化出来直接观察。
医学图像分割的评价指标上,Dice系数是核心指标,但不能只看Dice。边界模糊的情况下,Dice可能相近,但实际分割结果在临床参考意义上差别很大。所以我会额外跟踪两个指标:Hausdorff距离(HD95)和平均表面距离(ASD)。HD95衡量的是分割边界和真实边界之间的最大偏差,95%分位可以排除极端离群点的干扰,对评价边界的准确性更有说服力。ASD则衡量整体边界的贴合程度。理想的分割结果是Dice高、HD95低、ASD也低,三个指标都好看才说明模型真正学到了精细的边界。
推理时的不确定性评估也是扩散分割独有的一个优势。多次采样得到的结果,逐像素计算标准差,就能生成一张不确定性图。我习惯在推演报告里把不确定性图也一并输出,因为它的高值区域往往对应病灶的浸润边界,对我们判断哪些区域需要人工复核非常有参考价值。
3.5 实际训练记录与效果对比
我在一个公开的肝脏肿瘤分割数据集上完整跑过一次扩散分割实验,这里把关键的记录分享出来。数据集一共280例,切成2D切片后保留了约18000张有效切片,选其中80%做训练,10%做验证,10%做测试。
骨干网络采用2D UNet加瓶颈空间注意力,扩散总步数1000,训练批次大小16,在单张24G显存的显卡上跑完300个epoch大约花了两天。训练到第30个epoch时,验证集Dice就超过了0.75,进展挺快;第80个epoch时Dice到了0.84附近,但后续提升明显变慢。第150个epoch时,Dice到了0.87,之后曲线基本走平。最终测试集上的表现是Dice 0.873,HD95 6.3mm,ASD 1.8mm。
同一数据集上我跑了标准的UNet作为对照,Dice是0.842,HD95 是9.1mm。扩散分割在Dice上的提升看起来只有约3个百分点,但在HD95上的优势特别明显,边界精度提升了接近3毫米。这个对比也再次印证了一点:扩散分割的价值更多体现在对边界细节的刻画能力上,而不是在简单的区域重叠率上更胜一筹。
4. 常见问题与排查技巧实录
4.1 训练时loss爆掉或产生NaN
扩散模型的训练中,NaN问题非常常见。多数情况下出现在训练刚开始没几步,loss直接变成NaN,这通常和学习率过大有关。扩散模型初始阶段的梯度量级比较大,如果学习率设置过高,很容易一步到位把权重推到数值溢出区间。解决方法是把学习率先降到1e-5甚至更低,确认训练能稳定跑起来后,再逐步调回正常范围。
如果NaN出现在训练中途,那就不是学习率的问题了,更可能出在噪声调度表上。当 $\beta_t$ 接近上限1时,加噪过程中的方差容易变得数值不稳定。你可以检查一下自己用的 $\beta$ 调度上限是否超过了0.02,如果超过了,可以用余弦调度替代,或者强行限制 $\beta$ 的取值上限。
还有一个经常被忽略的细节是输入数据的归一化范围。UNet的激活函数对输入尺度比较敏感,建议把原始医学图像的像素值都标准化到0到1之间,分割标注则保持0和1的二值格式。如果原图的像素值范围是0到255甚至更大,第一层卷积就有可能被某些离群像素值推爆,训练同样会出现NaN。遇到这种情况,在数据加载阶段先做数值统计,然后做z-score标准化,能解决大部分不稳问题。
4.2 推理结果出现大面积噪声残留
训练一切正常,但推理出来的分割图覆盖了一层明显的椒盐噪声或斑点,这种情况我遇到过不止一次。原因通常是训练和推理时使用的噪声调度表不一致。
有些实现里,训练时会加入一些随机性微小抖动,比如动态调整 $\beta_t$ 或者对条件输入做随机mask增强,但推理时却直接用原始调度。这种不一致会累积误差,噪声在几百步的迭代里被逐步放大,最终导致输出充斥着残差。排查方法很简单:检查训练pipeline里所有涉及随机扰动的代码路径,确保推理时的采样流程和训练流程中的平稳过程保持完全一致。
另一个常见原因是采样步数太少。如果推理时用的去噪步数不足以让模型充分收敛到干净状态,输出就会半生不熟。我一般推荐至少用50步。如果50步还是一堆噪声,往上调到100步看变化。如果100步还不行,那基本可以确定问题出在调度表的一致性上,而不是步数不够。
4.3 多次采样结果差异过大,说明什么
扩散分割的多次采样结果本身就是一项有价值的信息,但如果差异过大,比如同一个样本采10次,得到的Dice从0.75跨度到0.91,那就说明模型对当前输入不够自信。
造成这种情况的最常见原因,是输入图像的特征和训练数据分布有明显偏移。举个例子,训练集里多是增强CT的动脉期图像,而测试时来了一张平扫或者静脉期的图像,扩散模型的自适应能力就会大幅减弱,采样方差也随之变大。这个问题的根源不在模型本身,而在数据域适配。我常用的解决方案是,对目标域的少量数据做domain adaptation,即用小学习率在目标域数据上fine-tune原始模型几十轮。
还有一种情况是目标区域本身就高度模糊。比如一些小病灶和周围组织的灰度差异极小,人眼都未必能稳定勾画边界,这时候模型多次采样结果差异大其实是对数据固有模糊度的真实反映。这种情况不需要追求更稳定的输出,反而应该把这种不确定性作为临床提示信息输出给医生,帮助他们决策是否需要进行进一步检查。
4.4 显存不够怎么办
显存不够是扩散分割落地中最常见的工程瓶颈。一个标准的2D UNet扩散模型加上完整的条件注入分支,batch size 8的情况下可能就需要16G以上的显存。
最直接的手段是缩小batch size,但要注意,如果UNet中使用了BatchNorm层,batch size降到4以下时,BatchNorm的统计量会变得非常不稳定,影响训练效果。这时候可以改用GroupNorm或InstanceNorm替代BatchNorm,它们不依赖batch内的统计信息,即使batch size为1也能稳定训练。
梯度累积是另一种有效的方案。它做的事情是,先让模型在多个小batch上分别计算梯度,把这些梯度加起来,再更新一次参数。这样可以模拟出更大的batch size效果,同时显存开销保持不变。我实际训练时经常用batch size 2配合16步梯度累积来模拟batch size 32的效果,在整个训练周期里稳定性非常好。
推理阶段的显存问题也有技巧。如果你不需要反向传播梯度,可以把模型切换成半精度推理,加上torch.inference_mode(),这会大幅降低显存使用量。如果还要进一步压缩,可以考虑分块推理,把切片切成若干个重叠块,分别推理后拼回去,重叠区域做平均。这个方法能轻松解决超大尺寸切片的显存问题,只是会稍微增加推理耗时。
5. 扩散分割的局限性与后续扩展方向
扩散分割虽然优势明显,但在实际落地时也要客观认识到它的短板。第一是推理速度慢,一次分割需要进行多次迭代去噪,即便DDIM压缩到50步,依然远慢于传统的一步推理CNN模型。第二是训练资源需求高,扩散模型对训练数据的充分性和多样性要求明显更高,在小样本医学数据集上很容易出现细节幻觉。第三是解释性相对较弱,医生对“逐步去噪得到结果”这个过程缺乏直观理解,这会在临床应用推广时带来一定阻力。
后续的扩展方向我个人最看好三个。一个是基于扩散模型的分割不确定性引导主动学习,让模型自己判断哪些样本最值得补充标注,用更少的标注成本持续提升模型性能。另一个是跨模态条件扩散分割,比如用CT图像作为条件,引导分割出MRI模态下的对应结构,这类跨模态泛化能力是传统模型很难实现的技术路线。还有一个是轻量化扩散分割,通过蒸馏或减少推理步数把耗时压缩到百毫秒级别,这可能会让扩散分割在临床实时辅助系统中真正站稳脚跟。
在一个真实项目里,我用扩散分割产出的不确定性图去辅助医生标注,让医生可以优先处理模型认为最不确定的区域。经过两轮迭代,整体标注效率提升了不少。这个方向我还会继续做下去,后面的实践经验也会持续整理出来。