如何使用MONAI在超声数据上训练肿瘤分割模型
大多数分割教程都是从选择一个模型开始,将图像输入该模型中,然后调整超参数直到相关指标得到改善。但这种方法忽略了通常最为关键的一步:理解数据本身。 在本教程中,我们首先会对数据集进行详细分析,随后会根据这些分析结果来决定MONAI分割流程中的每一个设计细节。 我们将涵盖以下内容: 本教程适合谁? 关于数据集 什么是MONAI,为什么使用它? 什么是Dice评分? 第1部分——建模前的数据分析 类别平衡对分割结果的影响 患者数量对数据划分的影响 第2部分——构建分割流程 单一配置对象 按患者分组的数据划分方式 由快照自动选择的转换操作 模型、损失函数与评估指标 结果解读 预测结果可视化 失败模式比
大多数分割教程都是从选择一个模型开始,将图像输入该模型中,然后调整超参数直到相关指标得到改善。但这种方法忽略了通常最为关键的一步:理解数据本身。
在本教程中,我们首先会对数据集进行详细分析,随后会根据这些分析结果来决定MONAI分割流程中的每一个设计细节。
我们将涵盖以下内容:
本教程适合谁?
本教程假设读者已经具备一定的Python编程基础,也了解神经网络训练的基本原理。它会解释与MONAI相关的特定组件(如字典转换函数、DiceCELoss和DiceMetric)以及医学影像领域的专业术语(如BI-RADS分级、低回声区域、按患者分组的数据划分方式)。不过,学习本教程并不需要事先具备超声检查方面的经验。
关于数据集
该数据集名为BUS-BRA,它是一个公开发布的乳腺超声图像集合,其中包含了经过活检验证的标签信息以及肿瘤分割结果。
每张图像都会标注良性或恶性类别、BI-RADS分级(放射科医生的怀疑程度,分为2到5级)、组织学类型,还会附带二进制的肿瘤掩膜。随数据集提供的CSV文件中还包含了预先定义好的交叉验证划分方案。
这项任务具有二元性:需要将肿瘤与背景区分开来。BUS-BRA数据集包含了来自巴西某癌症研究所的1,064名患者的1,875张B模式乳腺超声图像,这些图像是使用四台扫描设备拍摄获得的。
什么是MONAI,为什么使用它?
MONAI(医疗人工智能开放网络)是一个专为医学影像分析开发的开源PyTorch框架。它是在PyTorch基础上构建的、针对特定领域设计的组件:你仍然可以使用标准的PyTorch训练流程,但MONAI提供了专门用于医学影像处理的模块,因此你无需自行开发这些功能。
MONAI为你提供了以下资源:
专为医疗数据设计的数据处理工具,能够支持DICOM、NIfTI等格式的文件加载,还能对图像进行强度标准化、尺寸调整以及数据增强操作,所有这些操作都通过基于字典的流程来完成,从而确保图像及其分割结果始终保持同步。
多种常用于医学影像分割的网络架构(如U-Net、UNETR、SegResNet等),可以直接使用而无需重新开发。
专为图像分割任务设计的损失函数和评估指标,包括基于Dice系数的评估方法。
作为评估指标,Dice系数用于衡量模型的分割效果。例如,如果验证集上的Dice值为0.876,那就意味着模型预测的肿瘤区域与真实标签对应的肿瘤区域有大约88%的重叠部分。
作为损失函数(具体为
DiceCELoss),模型会通过优化这个损失函数来提升分割精度。对于存在类别不平衡的问题来说,这种损失函数尤为重要——因为Dice系数关注的是区域之间的整体重叠情况,而不是每个像素的匹配程度,所以如果模型将所有区域都标记为背景,其得分就会很低。因此,模型必须努力找到真正的肿瘤区域才能获得较高的评分。
使用MONAI可以显著减少代码冗余,并有效避免图像及其分割结果出现对齐错误的情况。
什么是Dice系数?
Dice系数用于衡量两个区域之间的重叠程度。在图像分割任务中,该系数会将模型预测的分割结果与真实标签进行比较,并返回0到1之间的分数:0表示完全不重叠,1表示完全匹配。
Dice系数的计算公式为:
Dice = 2 × (overlap) / (predicted area + true area)
分子中的“2×”这一因子确保了得分始终处于0到1的范围内,因为分母计算的是两侧重叠的像素数量。
在本教程中,Dice系数扮演着两个重要角色:
第1部分——建模前的数据分析
第一步就是进行数据分析。这个过程会读取每一张图像及其对应的分割结果,然后回答一些简单的问题,这些问题的答案将决定后续数据处理流程的具体构建方式。执行这些检查只需要几秒钟的时间,但却能大大减少后续开发中的猜测工作。
下方的快照总结了那些直接影响管道设计的因素。我们将让这些考量来决定工作流程的每一个步骤。
| 快照所测量的内容 | 对应的数量 | 这会带来什么要求 |
|---|---|---|
| 不同的图像分辨率 | 数百种不同的(宽度、高度)组合 | 在批量处理之前,必须将图像调整为固定大小 |
| 类别平衡性 | 背景与前景的比例约为10.6:1 | 如果使用简单的逐像素损失函数进行训练,系统很可能会主要预测背景部分,因为在这种不平衡的数据集上,这样做确实能够获得较高的像素准确率 |
| 每张图像的亮度分布 | 在整个数据集中,亮度值差异较大 | 因此需要对图像的亮度进行归一化处理,这一操作属于数据处理流程的一部分 |
| 患者数量与图像数量的关系 | 共有1,064名患者,对应的图像数量为1,875张(每名患者的左右侧影像各一张) | 数据分割时必须按照患者来进行分组,否则同一患者的影像会同时出现在训练集和验证集中 |
| 掩码的构成规则 | 每个掩码都应表示一个连续的区域 | 如果预测结果中出现多个不连续的区域,那么该预测肯定是错误的 |
| 像素格式 | 图像为8位灰度图像,掩码为1位二进制图像 | 加载时将掩码视为单通道数据,在加载完成后将其转换为二值格式 |
其中有两个因素尤其值得我们重点关注,因为它们决定了两个最为重要的决策方向。
类别平衡性是决定最终处理方式的关键因素
肿瘤的大小非常小。在整个数据集中,背景像素的数量远远超过肿瘤像素的数量,其比例约为10:1。
如果使用普通的二元交叉熵损失函数进行训练,无论将所有图像都标记为背景,模型的像素准确率也能达到约91%。不过这个高准确率实际上反映的只是数据集的不平衡性,并不能说明模型具备识别肿瘤的能力。
为了解决这个问题,应该采用一种能够奖励那些与实际肿瘤区域有重叠部分的预测方法,而Dice损失函数正好符合这一要求。
患者数量决定了数据的分割方式
由于许多患者的左右侧影像都被纳入了数据集,因此实际的患者数量少于图像的数量。如果随机进行数据分割,导致某位患者的左侧影像被放入训练集,而右侧影像被放入验证集,那么验证集的准确率就会因为这种“泄漏”现象而被抬高。
数据集的作者已经解决了这个问题:他们在CSV文件中添加了一列K5P,该列用于表示5折分割方案,其中P代表按患者进行分组,也就是说,同一位患者的所有影像都会被分到同一个组中。直接使用这个分组方式要比手动重新构建分组规则更加安全。
第2部分 — 构建数据处理流程
以下所有内容在处理分割相关任务时都使用了MONAI:包括数据转换、数据集封装、网络模型构建、损失函数的计算以及评估指标的确定。
单一配置对象
整个数据处理流程所需的所有配置参数都来自同一个数据类。后续处理过程中没有任何地方会硬编码常量值,因此如果想要使用不同的折叠方式或调整图像尺寸来重新运行实验,只需要进行一次修改即可。
from dataclasses import dataclass
from typing import Tuple, Optional
from pathlib import Path
@dataclass
class TrainConfig:
data_root: Optional[Path] = None
fold_column: str = "K5P" # 按患者分组进行的5折分割(用于开发集)
val_fold: int = 1 # 哪个K5P折叠部分用于验证
test_column: str = "HOP" # 按患者分组设置的保留集
test_group: int = 1 # HOP值被指定为测试集
image_size: Tuple[int, int] = (256, 256)
batch_size: int = 16
lr: float = 1e-3
epochs: int = 30
use_amp: bool = True # 是否使用混合精度计算
ckpt_path: str = "best_model.pt"
cfg = TrainConfig()
上述代码定义了一个TrainConfig数据类,其中包含了数据处理流程所需的所有配置参数:患者分组方式、验证所使用的折叠部分、目标图像尺寸、批量大小、学习率、训练周期数、是否使用混合精度计算,以及最佳模型的保存路径。只需创建一次cfg对象,后续的所有步骤都可以从中获取所需的配置信息。
按患者分组的数据分割方式
这种数据分割方式使用了两个预定义的列。HOP列用于标记保留集,这部分数据在整个数据处理过程中都不会被修改。在剩余的开发集中,其中一个K5P折叠部分被用作验证集,其余四个折叠部分则用于训练。通过简单的检查可以确认,没有任何患者会同时出现在多个分割组中。
dev_df = manifest[manifest[cfg.test_column] != cfg.test_group]
test_df = manifest[manifest[cfg.test_column] == cfg.test_group]
train_df = dev_df[dev_df[cfg.fold_column] != cfg.val_fold]
val_df = dev_df[dev_df[cfg.fold_column] == cfg.valFold]
# 确保没有患者会同时出现在多个分割组中
for a, b in [(train_df, val_df), (train_df, test_df), (val_df, test_df)]:
assert not (set(a["Case"]) & set(b["Case"])), "患者信息泄漏"
上述代码首先将HOP保留集分离出来,然后将剩余的开发数据分为验证集(选定的K5P折叠部分)和训练集。接下来会检查每一对分割组,确保它们不包含重复的患者信息。如果发现有任何重复的情况,程序会立即抛出异常。
由快照决定的数据转换操作
MONAI提供的字典式数据转换功能可以针对以特定名称为键的记录进行操作(这些键包括"image"和"label"),并且会对这些记录应用相应的处理操作。这里的每一步设计都是基于第1部分的数据分析结果来确定的。
from monaitransforms import (
Compose, LoadImaged, EnsureChannelFirstd, ScaleIntensityd,
AsDiscreted, Resized, RandFlipd, EnsureTyped,
)
import torch
base = [
LoadImaged(keys=["image", "label"], reader="PILReader", image_only=True),
EnsureChannelFirstd(keys=["image", "label"]),
ScaleIntensityd(keys="image"), # 调整图像的亮度
AsDiscreted(keys="label", threshold=0.5), # 将标签二值化为{0, 1}
Resized.keys=["image", "label"], # 调整图像的大小
spatial_size=cfg.image_size,
mode=("bilinear", "nearest")),
]
train_transforms = Compose(base + [
RandFlipd(keys=["image", "label"], prob=0.5, spatial_axis=1), # 水平翻转
EnsureTyped(keys=["image", "label"], dtype=torch.float32),
])
val_transforms = Compose(base + [
EnsureTypedkeys=["image", "label"], dtype=torch.float32),
])
上述代码首先定义了一组基础处理步骤:加载PNG图像,将图像的通道顺序调整为先通道后颜色,调整图像的亮度范围为[0, 1],将标签二值化为{0, 1},最后将图像和标签的大小都调整为256×256。接下来,这些基础处理步骤被封装到了两个处理流程中:训练流程会增加水平翻转操作,而验证流程则不会执行这一操作,因此评估时看到的图像始终是原始状态。
水平翻转是一种简单的数据增强方法,在这个数据集中使用这种方法可以保持图像的解剖结构完整性。然而,一些更为激进的增强方法,比如大幅旋转或弹性变形,需要谨慎使用,因为这些方法可能会扭曲具有临床意义的结构。
对于图像来说,使用双线性插值方法可以保留图像的亮度梯度;而对于标签来说,则使用最近邻插值方法,这样标签的值就只能保持为0或1。如果对标签使用双线性插值,那么在物体边界处就会产生人为生成的标签值。
模型、损失函数与评估指标
该网络采用MONAI的UNet结构,具有一个输入通道(灰度图像)和一个输出通道(肿瘤的概率分布)。所使用的损失函数正是那些用于实现类别平衡的方法所推荐的。
U-Net由编码器和解码器组成:编码器以逐渐降低的分辨率捕获图像的上下文信息,而解码器则负责重建细节丰富的图像结构。跳跃连接机制使得高分辨率的特征可以直接从编码器传递到解码器,因此U-Net在需要精确区分边界的医学分割任务中表现得尤为有效。
from monai.networks.nets import UNet
from monai.losses import DiceCELoss
from monai.metrics import DiceMetric
from monaitransforms import Activations, AsDiscrete
model = UNet(
spatial_dims=2, in_channels=1, out_channels=1,
channels=(16, 32, 64, 128, 256), strides=(2, 2, 2, 2),
num_res_units=2,
).to(device)
loss_fn = DiceCELoss(sigmoid=True) # Dice损失函数用于处理类别不平衡问题;CE损失函数用于平滑梯度
metric = DiceMetric(include_background=True, reduction="mean")
post_pred = Compose([Activations(sigmoid=True), AsDiscrete(threshold=0.5)])
上述代码创建了一个U-Net模型(包含五个不同的分辨率层次,一个输入通道和一个输出通道),并将其传输到GPU上。随后,代码定义了与该模型相关的三个核心组件:损失函数、验证指标,以及post_pred处理步骤——这个步骤通过应用Sigmoid函数并设置阈值0.5,将模型的原始输出转换为格式规范的0/1二值掩码。
DiceCELoss损失函数由两部分组成。其中“Dice”部分对于前景区域具有尺度不变性,因此无论是小型肿瘤还是大型肿瘤,模型在计算损失时都会给予相同的权重;而“交叉熵”部分则能够在“Dice”分数较低的区域使损失值的变化更加平缓。sigmoid=True这个选项指示损失函数直接使用Sigmoid激活函数进行计算,因此模型会输出原始的logits值,而post_pred步骤会在评估时负责执行Sigmoid变换和阈值处理。最终,这个U-Net模型的参数数量约为160万。
训练流程本身主要遵循PyTorch的标准框架。MONAI框架并未涉及优化逻辑的实现,其中与分割任务相关的部分仅包括损失函数、数据转换操作以及评估指标而已。
for epoch in range(1, cfg.epochs + 1):
model.train()
for batch in train_loader:
img, lab = batch["image"].to(device), batch["label"].to(device)
optimizer.zero_grad(set_to_none=True)
with torch.amp.autocast("cuda", enabled=cfg.use_amp):
loss = loss_fn(model(img), lab)
scaler.scale(loss).backward()
scaler.step(optimizer); scaler.update()
model.eval(); metric.reset()
with torch.no_grad():
for batch in val_loader:
img, lab = batch["image"].to(device), batch["label"].to(device)
pred = post_pred(model(img))
metric(y_pred=pred, y=lab)
val_dice = metric.aggregate().item()
if val_dice > best_dice:
best_dice = val_dice
torch.save(model.state_dict(), cfg.ckpt_path)
在上述代码中,每个训练周期会包含两次处理流程。在训练阶段,每批数据都会被传输到GPU上,模型会在混合精度模式下计算损失值,并通过梯度调整机制更新权重;而在验证阶段,则会关闭梯度计算功能,使用post_pred步骤对预测结果进行处理,并统计所有验证数据的Dice分数。每当某个训练周期的Dice分数超过迄今为止的最佳记录时,模型的权重就会被保存到磁盘上。
结果分析
两条曲线可以概括整个训练过程:训练损失值持续下降,在接近0.12时趋于稳定;而验证数据的Dice分数则从约0.57开始上升,最终在第28个训练周期达到了0.866的峰值。
从这些曲线中我们可以得出一些有意义的信息:损失值呈单调递减趋势,说明模型正在学习中,梯度信号是真实的。
损失值在零以上趋于平稳而非降至零是预期中的现象。
DiceCELoss函数存在一个下限,因为在边界模糊的像素处,交叉熵项永远不会完全消失。如果损失值真的降为零,那反而是一个警示信号,而不是模型成功的标志。在训练过程中,验证集的Dice分数稳定在0.85左右,而训练损失却在持续下降,这通常是模型轻微过拟合的表现。增加训练周期通常会降低训练损失,但验证集的Dice分数却不会随之变化。在这种情况下,30个训练周期已经足够了,不过采用基于耐心值的提前停止策略会更为稳妥。
对于这个数据集而言,0.866的验证Dice分数属于正常范围。不过需要注意的是,验证Dice分数是用同一组数据进行评估得出的结果,因此其评估结果可能会显得较为乐观。
最后需要测试的是HOP测试集。在完成所有训练和模型选择工作之后,会仅对这一测试集进行一次评估,其得分为0.864,与0.866的验证分数基本一致。这说明该模型能够泛化到训练或选择过程中从未见过的新病例上,且验证结果也没有反映出任何模型误差。
预测结果可视化
各项指标虽然可以概括模型的整体性能,但它们并不能展示模型是如何对各个肿瘤进行分割的。
下图展示了一个典型的验证案例。从左到右依次为输入的超声图像、真实标注结果、模型预测的结果,以及将预测结果叠加在原始图像上的效果。
预测结果与真实标注之间的高度吻合,说明了模型能够准确定位病变的位置及其边界。
失败模式的重要性远超平均水平
平均Dice分数为0.866这一数值,实际上可能掩盖了模型截然不同的行为特征。它可能意味着所有案例的表现都一般,或者大部分案例表现优异,而只有少数案例出现了严重错误。
为了区分这些情况,可以按照每张图像的Dice分数对验证集进行排序,并仔细检查得分最低的那些预测结果。
在这个数据集中,299个验证案例中仅有4个案例的Dice分数低于0.5,占比约为1%。分析这4个案例后发现了一个明显的规律:其中3个案例的预测结果都是碎片化的——模型输出了多个相互分离的斑点,而实际上真实情况应该是一个连续的区域;第4个案例则将一个常见的超声伪影(即暗色声学阴影)误认为是肿瘤组织。
这种数据碎片化的现象实际上与一项实验结果密切相关:质量检测结果显示,BUS-BRA数据集中的每一个真实标记区域都构成了一个独立的连通组件。因此,如果预测结果出现多个分割区域,那肯定是由于数据集本身的特性所致。为此,在后处理阶段应该只保留最大的连通组件:
from monaitransforms import KeepLargestConnectedComponent
post_pred = Compose([
Activations(sigmoid=True),
AsDiscrete(threshold=0.5),
KeepLargestConnectedComponent(applied_labels=[1]),
])
这段代码在原有的处理流程中增加了一个步骤。经过sigmoid函数和阈值处理后,生成的二值掩码中只会保留最大的连通区域,其他所有预测出的区域都会被舍弃。这样一来,原本被分割成多个部分的图像又会重新合并为最大的那一个部分,这样就符合数据集“每个标记区域只对应一个连通组件”的特性了。
我在验证集上测试了这个方法,结果发现其效果并不像“简单提高准确率”那么直观。虽然这种方法能够恢复一些被分割的数据点,但整体来看,平均Dice值的提升幅度非常微小,甚至有时还会出现负值。例如,当一个真实的病变被预测为两个相邻的区域时,舍弃较小的那个区域就会导致真正属于正类的区域也被遗漏。因此,这种方法只能针对特定的错误类型进行优化,并不能盲目地应用于所有情况。
对于那些阴影与病变难以区分的情况,这种方法更是无能为力。有时候,要区分一个低回声的肿瘤和一个暗色的阴影区域,需要更多的上下文信息,而简单的灰度裁剪是无法提供这些信息的。这说明,在未来的实验中,提高分辨率或扩大模型的感知范围可能是有必要的方向。
下一步该怎么做
一旦你获得了可靠的基准结果,接下来的实验就会更有意义。不要随意尝试更大的模型,而应该先针对你所观察到的问题来进行优化:
将2D U-Net替换为Attention U-Net或DynUNet。
在更高分辨率下进行训练,以便更好地检测出较小的病变。
在推理过程中有选择地应用连通组件分析算法。
探索在测试阶段对数据进行处理的方法。
对于那些类别分布极不均衡的病变,比较DiceCE损失函数和Focal Tversky损失函数的效果。
总结
整个处理流程的核心就是先分析数据,然后根据分析结果来决定后续的处理方式。
调整图像尺寸是为了满足分辨率要求;损失函数的选取则是为了保证类别平衡;数据分割的依据是患者数量;而最有用的后处理方法,则是在训练开始之前通过对标记区域进行分析得出的。这些决策都不是凭猜测做出的,也不需要通过大规模的实验才能发现。
构建一个模型很容易,但一个每一个决策都有合理依据的模型,才会更值得信赖、更容易调试,也更容易向后续的使用者解释其工作原理。
参考资料
本教程中使用的完整代码可以作为MONAI笔记本获取:busbra_segmentation_monai.ipynb。你可以在Kaggle或Colab平台上直接运行这段代码,如果数据集还没有被加载的话,系统会自动下载它。
相关文章
如何利用Gemini构建人工智能功能:面向开发者的提示工程实用指南
大多数关于提示工程的教学教程都遵循相同的流程:安装SDK,输入API密钥,调用 generateContent 函数,然后打印输出结果。模型会生成一些看似合理的内容,之后教学教程也就结束了。 但当你真正尝试将这个系统投入实际使用时,才会发现其实真正的准备工作根本还没有开始。 “API返回的文本”与“让用户感到可信的实际功能”之间的差距,正是需要耗费大量精力去解决的地方。 这个差距中充满了各种棘手的问题:模型生成的内容听起来和其他聊天机器人没什么两样;它会编造用户从未说过的话;它返回的数据会被用Markdown格式包裹起来;系统会在凌晨2点出现故障;而对于那些只是想得到答案的用户来说,系统展示的
阅读全文
如何使用LangSmith来追踪和监控人工智能代理的行为
在本教程中,我将向您展示如何使用LangSmith来追踪和监控本地的AI代理。我们会构建一个简单的本地AI代理,然后为其启用LangSmith追踪功能,这样我们就能通过Web界面查看模型调用情况、工具使用情况以及请求处理延迟等信息。 我们将使用LangChain v1、Ollama、Qwen以及Python这些工具。除了用于实现观测功能的组件外,所有操作都在您的本地机器上完成,因此代理本身不会产生任何与模型API相关的费用。 目录 背景知识 什么是可观测性与监控? 什么是LangSmith? 开发动机与架构设计 步骤1:安装Ollama并下载模型 步骤2:安装Python相关依赖库 步骤3:启
阅读全文
什么是HyDE?如何利用假设性文档来提升RAG的质量?
检索增强生成技术,通常被称为RAG,已成为利用大型语言模型构建应用程序时最常用的方法之一。 与让大型语言模型完全依据其训练数据来回答问题不同,RAG系统会从外部知识库中检索相关信息,并将这些信息作为上下文提供给模型。 其基本原理非常简单: 将用户的问题转化为嵌入向量。 在向量数据库中搜索语义上相似的文档片段。 将检索到的这些片段传递给大型语言模型。 基于这些片段生成答案。 然而,这个看似简单的过程存在一个重大缺陷:用户提出的问题与包含答案的文档在表达方式上可能存在很大差异。 例如,用户可能会提出这样的问题: 为什么我的AWS Glue作业在处理了几百万条记录后速度会显著下降? 而知识库中相关的
阅读全文
从RPC到gRPC:了解远程过程调用、Protocol Buffers以及现代分布式系统的通信机制
任何应用程序在某些时候都需要与其他系统进行交互。移动应用会与后端服务进行通信;后端服务又会与支付网关对接;认证服务需要与用户服务进行交互;数据传输流程则要与存储系统相连。 问题不在于系统是否需要相互沟通,而在于应该如何实现这种沟通。 多年来,基于HTTP和JSON的REST架构一直是人们的首选方案。它运行稳定、使用简单,而且相关的开发工具也随处可见。然而,随着系统规模的扩大——无论是参与交互的服务数量增加、数据交换量增大,还是对实时通信的需求提高——REST架构逐渐暴露出了其局限性。 这时,远程过程调用、Protocol Buffers以及gRPC这些技术应运而生了。 通过这本手册,你将了解什
阅读全文