简介:图像分割是计算机视觉的核心任务之一,旨在将图像划分为多个有意义的区域。其原理通常基于深度学习模型学习像素级特征表示,实现像素分类。在医学影像领域,精准的图像分割技术具有重要价值,它是疾病诊断、手术规划和疗效评估的基础。针对数据稀缺、标注成本高的专业场景,如何高效利用预训练大模型成为关键。本文聚焦于医学图像分割这一应用场景,以SAM-Med 2D这一视觉大模型为基础,详细阐述了如何通过微调(Fine-tuning)和参数高效微调(PEFT)技术,将其适配到脊椎CT图像分割这一具体任务,实现从通用能力到专业精度的跨越。
1. 项目缘起:当通用大模型遇上专业医学图像
最近在折腾一个挺有意思的项目:用SAM-Med 2D这个视觉大模型,来做脊椎CT图像的分割。这事儿听起来可能有点“杀鸡用牛刀”,毕竟分割脊椎在传统图像处理里也不算特别新鲜。但真正上手后,我发现这背后其实是一个很典型的场景——如何让一个强大的、预训练好的通用基础模型(Foundation Model),快速适配到一个数据稀缺、标注成本高昂的专业垂直领域。
SAM(Segment Anything Model)大家应该不陌生,Meta搞出来的那个“分割一切”的模型,其核心思想是通过提示(point, box, text)来引导模型进行零样本(zero-shot)分割,泛化能力极强。而SAM-Med 2D,顾名思义,是SAM在大量医学图像(主要是2D的X光、CT、MRI切片等)上进一步预训练或微调后的版本。它继承了SAM强大的提示分割和泛化能力,同时对医学图像的纹理、对比度、解剖结构有了更好的先验知识。
那么,为什么还要复现它,并且用自定义的脊椎数据集来训练呢?原因有三:
第一,精度天花板。尽管SAM-Med 2D在通用医学图像上表现不错,但“通用”意味着在特定任务上(比如精确分割每一节椎体及其附件)可能达不到临床或科研所需的精度。椎体边缘的骨皮质、椎间盘、可能存在的病变(如骨折、骨赘),这些细节需要模型有更强的针对性。
第二,提示方式的效率。SAM系列模型依赖提示。在科研或批量处理中,我们可能希望模型能自动识别并分割出所有椎体,而不是每张图都手动去点一下或画个框。这就需要模型具备一定的“自动实例分割”能力,或者我们通过训练让其对“脊椎”这个特定概念产生更强的响应。
第三,数据与流程的闭环。很多团队积累了自己的脊椎影像数据集(可能是特定设备采集的、特定人群的),这些数据有其独特性。将SAM-Med 2D在自己的数据上微调,不仅能提升模型在本中心数据上的性能,更能将整个流程——从数据准备、模型训练到推理部署——内化,形成可控的技术资产。
所以,这个项目的目标很明确:复现SAM-Med 2D的工作环境,利用我们自己的、已标注的脊椎CT切片数据集,对模型进行微调(Fine-tuning),使其成为一个专精于脊椎分割的利器。下面,我就把从环境搭建、数据准备、模型训练到推理测试的全流程,以及中间踩过的坑和总结的经验,毫无保留地分享出来。
2. 环境复现:依赖管理与版本锁定的艺术
复现任何一篇顶会论文或开源项目,第一步永远是最头疼但也最关键的:环境配置。SAM-Med 2D基于PyTorch,但其依赖链可能比想象中复杂,特别是涉及到一些特定的图像处理库和CUDA版本兼容性问题。
2.1 核心依赖解析与选型
原论文或代码仓通常会提供一个requirements.txt。但直接pip install -r requirements.txt常常是噩梦的开始。我们需要理解核心依赖,并做出适合自己的选择。
PyTorch 与 CUDA:这是基石。首先确认你的显卡驱动支持的CUDA最高版本(
nvidia-smi查看)。SAM-Med 2D通常需要PyTorch 1.11+。我个人的选择是PyTorch 1.13.1 + CUDA 11.7。这是一个相对稳定、兼容性广的组合。太旧的版本可能缺少某些API,太新的版本(如PyTorch 2.0+)可能带来意料之外的变动。# 示例安装命令,请根据你的CUDA版本和Python版本调整 pip install torch==1.13.1+cu117 torchvision==0.14.1+cu117 torchaudio==0.13.1 --extra-index-url https://download.pytorch.org/whl/cu117OpenCV:医学图像读取和预处理必备。注意,
opencv-python和opencv-python-headless的区别。如果你在无GUI的服务器上跑,或者不需要cv2.imshow这类功能,用headless版本更轻量,避免一些不必要的系统依赖。pip install opencv-python-headless==4.8.1SimpleITK 或 NiBabel:对于处理3D的CT数据(我们最终处理的是2D切片,但数据源是3D的),需要库来读取DICOM或NIfTI格式。
SimpleITK功能强大,但稍重;NiBabel更轻量。由于我们主要关心像素数据和简单元信息,我选择了NiBabel。pip install nibabel==5.1.0MONAI:这是一个医学影像AI的PyTorch专属框架。SAM-Med 2D的预处理、数据增强流程很可能借鉴或兼容MONAI的风格。即使原代码未直接使用,引入MONAI的
transforms来进行数据增强也是极好的选择,它提供了大量针对医学图像的增强操作(如随机弹性形变、Gamma变换等)。pip install monai==1.2.0
注意:版本锁定的重要性。强烈建议使用
pip freeze > requirements_lock.txt来生成一个你当前成功环境的确切版本列表。这能保证你未来在任何机器上重建环境时的一致性。分享项目时,提供这个_lock文件比原始的requirements.txt更有价值。
2.2 SAM-Med 2D 代码获取与结构梳理
从GitHub上找到官方或高星的复现仓库。关键不是直接git clone完事,而是要花时间看代码结构。
- 模型定义 (
modeling/): 找到sam_med2d.py或类似文件。这里定义了模型的主干网络(通常是ImageEncoderViT)、提示编码器、掩码解码器。你需要确认的是预训练权重加载的接口。权重文件(通常是.pth或.pt)如何被加载到这些模块中。 - 配置文件 (
configs/): 任何严肃的项目都会有配置文件(yaml或json)。这里定义了模型尺寸(如vit_b,vit_l,vit_h)、输入图像大小、训练超参数等。这是你调整实验的入口。 - 数据加载 (
data/): 查看dataset.py和transforms.py。这是适配自定义数据集最关键的部分。你需要弄清楚它期望的数据标注格式是什么?是COCO格式的JSON?还是简单的图像和掩码(mask)文件对?它如何处理医学图像(如窗宽窗位调整)? - 训练脚本 (
train.py):主训练循环。关注优化器(Optimizer)、学习率调度器(Scheduler)、损失函数(Loss)的设置。SAM-Med通常使用组合损失,如交叉熵损失+Dice损失。 - 推理/演示脚本 (
demo.py或predict.py):用于验证模型效果。
我的做法是,先尝试在不修改任何代码的情况下,用项目提供的示例数据或脚本跑通推理,确保基础环境没问题。比如,用一张公开的脊柱X光图和对应的提示点,看模型能否输出一个合理的分割掩码。
3. 数据准备:从3D CT到2D切片与标注转换
这是我们项目的核心输入。假设你有一批脊椎CT的3D数据(DICOM序列或NIfTI文件),以及对应的3D分割标注(可能是用ITK-SNAP、3D Slicer等工具标注的,保存为另一个NIfTI文件)。我们的目标是将它处理成SAM-Med 2D训练所需的2D图像-掩码对。
3.1 3D数据预处理与切片提取
CT数据通常包含多个序列(如平扫、增强)。我们首先需要确认使用的是哪个序列,并统一空间坐标和方向。
- 读取与重采样:使用
NiBabel读取image.nii.gz和label.nii.gz。检查它们的affine矩阵(空间信息)是否一致。如果不一致,需要将标注重采样到图像的空间。可以使用monai.transforms.Spacingd进行各向同性重采样(例如,将所有体素间距统一为1mm x 1mm x 1mm),这能减少后续因分辨率差异带来的问题。 - 窗宽窗位调整:CT值是亨氏单位(HU),范围很广(-1000到+3000)。我们需要将其映射到灰度图范围(如0-255)。这不是简单的线性缩放,而是应用窗宽(Window Width)和窗位(Window Level)。对于脊椎骨骼,常用的窗宽是1500-2000 HU,窗位是300-500 HU。这个操作能极大增强骨骼与软组织的对比度。
import numpy as np def apply_window(image_hu, window_center, window_width): """将CT值(HU)通过窗宽窗位映射到灰度值.""" lower = window_center - window_width / 2 upper = window_center + window_width / 2 image_hu = np.clip(image_hu, lower, upper) # 截断 image_hu = (image_hu - lower) / (upper - lower) * 255.0 # 归一化到0-255 return image_hu.astype(np.uint8) - 轴向切片提取:沿着CT的轴向(通常是Z轴)逐层提取2D切片。同时,从3D标注文件中提取对应层的2D掩码。注意,标注文件可能是一个多标签的整数数组(如0背景,1腰椎L1,2腰椎L2...)。我们需要决定是训练一个模型分割所有椎体(多类分割),还是每个椎体单独训练一个模型(二分类分割)。对于SAM-Med,由于其提示机制,更自然的做法是进行二分类分割:即模型只学习分割“脊椎骨”这个整体,或者更进一步,通过不同的提示来区分不同椎体。在初期,我建议先做二分类(脊椎骨 vs 背景),这样问题更简单。
- 过滤无效切片:很多CT切片在头部或尾部并不包含脊椎。我们可以通过计算2D掩码中前景像素的比例来过滤掉这些“空”切片,节省存储和训练时间。
3.2 标注格式适配与数据集类编写
SAM-Med 2D的原始数据加载器可能期望某种特定格式。常见的有两种:
- 格式A:图像文件夹 + 掩码文件夹。要求文件名一一对应,如
001.png和001_mask.png。掩码图为单通道PNG,前景为255,背景为0(二分类)。 - 格式B:COCO格式的JSON标注。包含图像信息列表和标注信息列表,标注信息中包含
segmentation字段(多边形点集)或bbox字段。
我们的2D切片和掩码天然适合格式A。处理步骤如下:
- 将调整窗宽窗位后的2D图像保存为
.png或.jpg。 - 将对应的2D二值掩码(0和1或0和255)同样保存为单通道的
.png。 - 划分训练集、验证集和测试集(例如70%/15%/15%)。务必按病例(Patient)划分,而不是随机打乱切片!否则,同一个病人的不同切片会同时出现在训练集和测试集,导致数据泄露,评估结果会虚高。
- 编写自定义的
Dataset类。这个类需要继承torch.utils.data.Dataset,在__getitem__方法中返回image和mask两个张量。这里就是加入数据增强(如旋转、翻转、亮度对比度扰动)的好地方。对于医学图像,在应用空间变换(如旋转)时,必须同时对图像和掩码进行相同的变换,这是铁律。
import torch from torch.utils.data import Dataset import cv2 import os from albumentations import Compose, HorizontalFlip, RandomRotate90, ShiftScaleRotate, RandomBrightnessContrast class SpineDataset(Dataset): def __init__(self, image_dir, mask_dir, transform=None): self.image_dir = image_dir self.mask_dir = mask_dir self.image_names = sorted(os.listdir(image_dir)) self.transform = transform # 使用albumentations库的增强管道 def __len__(self): return len(self.image_names) def __getitem__(self, idx): img_name = self.image_names[idx] img_path = os.path.join(self.image_dir, img_name) mask_path = os.path.join(self.mask_dir, img_name.replace('.png', '_mask.png')) # 假设掩码文件名规则 image = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE) # 以灰度图读取 mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE) # 确保mask是二值的 _, mask = cv2.threshold(mask, 127, 255, cv2.THRESH_BINARY) if self.transform: transformed = self.transform(image=image, mask=mask) image = transformed['image'] mask = transformed['mask'] # 添加通道维度 (C, H, W) -> (1, H, W) image = torch.from_numpy(image).unsqueeze(0).float() / 255.0 mask = torch.from_numpy(mask).unsqueeze(0).float() / 255.0 return image, mask4. 模型微调策略:轻量化与针对性优化
拿到了SAM-Med 2D的预训练模型,我们不是从头训练,而是微调。微调的策略选择直接影响效果和效率。
4.1 解冻哪些参数?—— 参数高效微调
SAM模型参数量巨大(ViT-Huge backbone有超过6亿参数)。全参数微调不仅需要海量显存,也容易在小数据集上过拟合。因此,参数高效微调(Parameter-Efficient Fine-Tuning, PEFT)是更明智的选择。
- 仅微调解码器:这是最保守、最常用的策略。冻结Image Encoder和Prompt Encoder的所有参数,只训练Mask Decoder。因为Encoder负责提取通用的图像特征,而Decoder负责根据特征和提示生成掩码。让Decoder去适应“脊椎”这个特定任务,是合理的。这种方法速度快,显存占用小,适合数据量较少(几百到几千张切片)的场景。
- 微调特定层 + 解码器:如果效果不佳,可以考虑解冻Encoder的最后几层(例如,ViT的最后几个Transformer Block)。这些高层特征更偏向于语义信息,针对特定任务调整它们可能有益。
- 引入适配器(Adapter)或LoRA:这是更先进的PEFT方法。不在原始模型权重上直接更新,而是插入一些小的、可训练的模块(Adapter),或者对权重矩阵进行低秩分解更新(LoRA)。这能极大减少可训练参数量(通常只有原模型的1%-10%),同时保持甚至提升效果。对于SAM-Med这类大模型,我强烈推荐尝试LoRA。你需要找到社区中已经实现的SAM-LoRA代码,或者自己实现(主要是在注意力模块的QKV投影层旁添加低秩矩阵)。
在我们的脊椎分割任务中,我采取了“策略1 + 策略3”的混合模式:首先尝试仅微调Mask Decoder。如果验证集Dice系数达到平台期后仍不理想,再尝试在Image Encoder的注意力模块中加入LoRA进行微调。
4.2 损失函数与评估指标的选择
医学图像分割的损失函数通常是组合拳。
损失函数:
- Dice Loss: 直接优化Dice相似系数,对前景背景像素不平衡的数据集非常友好。脊椎切片中,骨骼区域通常只占图像的一小部分,属于典型的不平衡问题。Dice Loss是首选。
- Cross-Entropy Loss: 标准的分类损失。可以结合Dice Loss使用,提供更稳定的梯度。
- Focal Loss: 如果数据中存在大量难以分割的边界像素(如椎体边缘模糊),Focal Loss可以降低易分样本的权重,让模型更关注难例。 我的常用配方是:
Loss = DiceLoss + 0.5 * BCEWithLogitsLoss。这个比例可以根据验证集效果调整。
评估指标:
- Dice Similarity Coefficient (DSC): 核心指标,范围0-1,越接近1越好。它衡量的是预测掩码和真实掩码的重叠面积。
- Hausdorff Distance (HD): 衡量两个轮廓之间的最大距离,对分割边界的准确性非常敏感。对于要求精确轮廓的脊椎手术规划,这个指标很重要。
- Precision & Recall: 从像素分类的角度看模型的查准率和查全率。 在训练过程中,我主要监控验证集上的平均Dice系数。同时,会定期可视化一些验证集样本的预测结果,直观判断模型是在学习正确的特征,还是只是记住了训练集。
4.3 训练超参数设置与技巧
- 批量大小(Batch Size):受限于显存,可能只能设置到4、8或16。可以使用梯度累积(Gradient Accumulation)来模拟更大的批量大小。例如,实际批量大小=4,设置累积步数=4,效果上就等价于批量大小16,但显存占用仅为4。
# 伪代码示例 optimizer.zero_grad() for i, (data, target) in enumerate(train_loader): output = model(data) loss = criterion(output, target) loss = loss / accumulation_steps # 损失按累积步数缩放 loss.backward() if (i+1) % accumulation_steps == 0: optimizer.step() optimizer.zero_grad() - 学习率(Learning Rate):对于微调,学习率要设置得比从头训练小得多。一个常见的起点是
1e-4到5e-5。使用余弦退火(Cosine Annealing)或带热重启的余弦退火(Cosine Annealing with Warm Restarts)调度器,通常比阶梯下降(Step Decay)效果更好。 - 优化器:
AdamW是目前的主流选择,比经典的Adam具有更好的权重衰减(Weight Decay)处理方式,泛化性能更优。 - 早停(Early Stopping):持续监控验证集损失或Dice系数。如果其在连续多个epoch(如20个)内没有提升,则停止训练,并回滚到验证集指标最好的那个模型检查点。这是防止过拟合的必备手段。
5. 实战训练与问题排查
理论说完,进入实战。假设我们的数据集已经准备好,模型代码也适配好了自定义数据集类。
5.1 训练循环中的关键检查点
- 初始损失值检查:开始训练的第一个epoch,观察第一个batch的损失值。如果损失值异常大(如几十上百),可能是数据归一化出了问题(例如,图像像素值没有归一化到[0,1]或[-1,1]),或者损失函数输入格式不对。
- 训练/验证损失曲线:这是最重要的监控图表。理想情况是训练损失平稳下降,验证损失也同步下降。如果出现以下情况:
- 训练损失下降,验证损失上升:典型的过拟合。需要加强数据增强、增加Dropout、减小模型容量(如果解冻了太多参数)、或者收集更多数据。
- 训练和验证损失都几乎不变:模型可能没有在学习。检查学习率是否太小、梯度是否被裁剪(Gradient Clipping)得过小、或者模型的大部分参数是否被意外冻结了。
- 中间结果可视化:每隔几个epoch,从验证集中取几个样本,让模型预测并保存预测的掩码图。与真实标注对比。这能帮你发现一些指标无法反映的问题,比如模型是否总是漏掉某个特定位置的椎体(可能是该位置在训练集中出现少),或者分割边界是否特别粗糙。
5.2 我遇到的两个典型“坑”及解决方案
坑一:数据泄露导致的虚假高精度最初我随机划分了所有2D切片,结果验证集Dice系数轻松达到0.95以上,让我欣喜若狂。但当我用来自新病人的CT数据测试时,效果骤降到0.7左右。这就是典型的数据泄露——同一个病人的相邻切片在空间上高度相似,它们分别进入了训练集和验证集,导致模型实际上是在“回忆”而不是“泛化”。
解决方案:严格按病例ID划分数据集。确保同一个病人的所有切片,只出现在训练、验证、测试三个集合中的一个里。可以按病人ID排序,然后按比例切分。
坑二:二值掩码边界处的“锯齿”和“空洞”训练出的模型,其预测掩码的边缘有时会出现难看的锯齿状,或者椎体内部出现不应该有的小空洞。这可能是多个原因造成的:
- 原始标注质量问题:医生标注时可能用了较粗的笔刷,或者标注工具本身会导致边界不光滑。需要在数据准备阶段进行后处理,比如对标注掩码进行轻微的形态学闭运算(先膨胀后腐蚀)来填充小空洞和平滑边界。
- 模型容量或训练不足:如果只微调了很少的参数,模型可能没有足够的能力学习到光滑的边界特征。可以尝试解冻更多层,或者使用更强大的损失函数(如结合边界损失)。
- 后处理缺失:模型输出的通常是概率图(每个像素是前景的概率)。我们用一个阈值(如0.5)将其二值化。这个简单的阈值化会放大边界的不连续性。可以改用连通组件分析(Connected Component Analysis):先阈值化,然后找出所有的连通区域,只保留面积最大的那个区域(假设一个切片只有一个主要的脊椎结构),最后对这个区域进行形态学平滑处理。
import cv2 import numpy as np def post_process_mask(pred_prob, threshold=0.5, min_area=50): """ 对模型输出的概率图进行后处理。 pred_prob: [H, W] 概率图,范围0-1 """ # 1. 阈值化 binary_mask = (pred_prob > threshold).astype(np.uint8) * 255 # 2. 连通组件分析 num_labels, labels, stats, centroids = cv2.connectedComponentsWithStats(binary_mask, connectivity=8) # stats: [num_labels, 5], 每一行: [x, y, width, height, area] if num_labels > 1: # 至少有1个背景标签+1个前景标签 # 找到面积最大的前景区域(跳过背景,索引0) max_area_idx = np.argmax(stats[1:, 4]) + 1 largest_component = (labels == max_area_idx).astype(np.uint8) * 255 else: largest_component = binary_mask # 3. 形态学平滑(可选) kernel = np.ones((3,3), np.uint8) smoothed_mask = cv2.morphologyEx(largest_component, cv2.MORPH_CLOSE, kernel) # 闭运算填充小洞 smoothed_mask = cv2.morphologyEx(smoothed_mask, cv2.MORPH_OPEN, kernel) # 开运算去除小毛刺 return smoothed_mask6. 推理部署与提示工程探索
模型训练好后,我们要用它来分割新的、未见过的脊椎CT切片。
6.1 基础推理流程
对于一张新图像,流程如下:
- 预处理:应用与训练时完全相同的窗宽窗位调整、归一化(除以255)等操作。
- 模型前向传播:将处理后的图像输入模型。这里有一个关键点:SAM-Med需要提示(Prompt)。在训练时,我们可能采用了“自动”生成提示的方式(例如,用标注掩码的中心点或边界框作为提示)。在推理时,我们需要提供类似的提示。
- 生成提示:
- 自动提示:如果我们希望模型自动分割出整个脊椎,一个简单的方法是使用一个覆盖整个脊柱区域的大边界框作为提示。这个框可以基于图像直方图或简单的启发式规则(如强度较高的区域)来粗略估计,但更可靠的方法是使用一个轻量级的目标检测模型(比如YOLO)先检测出脊椎的大致区域,再用这个检测框作为SAM-Med的提示。
- 交互式提示:在科研或临床辅助场景中,可以由用户在图像上点击一点(点提示)或画一个框(框提示)。SAM-Med对这种稀疏提示的响应非常好。
- 后处理:对模型输出的概率图或低分辨率掩码进行上采样、阈值化和上述提到的后处理(连通组件分析、平滑),得到最终的分割结果。
6.2 超越二分类:多椎体实例分割的思考
我们之前训练的是二分类模型(脊椎/非脊椎)。但临床往往需要区分不同的椎体(如L1, L2, L3...)。如何用SAM-Med实现?
思路一:训练多个二分类模型。分别训练分割L1、L2...的模型。推理时串行或并行运行所有模型。这种方法简单粗暴,但计算成本高,且可能因为椎体间相似性导致误判。
思路二:基于提示的实例区分。这是SAM的核心优势所在。我们可以训练模型学会响应不同的点提示。例如,在训练时,不仅提供图像和整个脊椎的掩码,还提供每个椎体中心点的坐标作为提示。模型需要学习将“靠近L1中心的点提示”映射到“L1椎体掩码”。这需要更精细的标注数据(每个椎体的中心点)和修改训练代码以支持多点提示。这更接近SAM原始论文的设定,潜力更大,但实现也更复杂。
思路三:结合实例分割模型。先用一个二分类的SAM-Med模型分割出整个脊椎区域,然后在这个区域内,使用一个传统的实例分割模型(如Mask R-CNN)或聚类算法(如基于距离变换的分水岭)来区分各个椎体实例。这是一种两阶段(coarse-to-fine)的混合方案。
在实际项目中,我首先实现了思路一(多模型),因为它能最快出结果验证可行性。但对于一个追求优雅和效率的系统,思路二才是最终方向,它真正发挥了提示式分割大模型的威力。
整个项目从环境搭建到训练出第一个可用的模型,大约花了一周时间。其中大部分时间都耗在了数据预处理、标注格式转换和调试训练管道上。模型本身的微调训练,在单张RTX 4090上,对于约3000张切片的数据集,仅微调Mask Decoder的话,50个epoch大概只需要3-4个小时。最终在独立测试集上,Dice系数达到了0.92,Hausdorff距离控制在5个像素以内,对于后续的脊柱形态测量分析,这个精度已经足够作为可靠的输入。
这个过程让我深刻体会到,用好一个视觉大模型,三分在模型,七分在数据。数据的质量、预处理的方式、与模型预期的匹配程度,往往比调参更能决定项目的成败。SAM-Med 2D提供了一个强大的基础,但如何将它“调教”成你专属领域的专家,考验的是你对业务(脊椎解剖)、数据(CT影像)和模型原理的综合理解。
本文还有配套的精品资源,点击获取