基于SAM-Med 2D的脊椎CT图像分割:从通用大模型到专业医学影像的微调实战 📅 发布时间:2026/8/28 8:02:58 👁 浏览次数: 简介图像分割是计算机视觉的核心任务之一旨在将图像划分为多个有意义的区域。其原理通常基于深度学习模型学习像素级特征表示实现像素分类。在医学影像领域精准的图像分割技术具有重要价值它是疾病诊断、手术规划和疗效评估的基础。针对数据稀缺、标注成本高的专业场景如何高效利用预训练大模型成为关键。本文聚焦于医学图像分割这一应用场景以SAM-Med 2D这一视觉大模型为基础详细阐述了如何通过微调Fine-tuning和参数高效微调PEFT技术将其适配到脊椎CT图像分割这一具体任务实现从通用能力到专业精度的跨越。1. 项目缘起当通用大模型遇上专业医学图像最近在折腾一个挺有意思的项目用SAM-Med 2D这个视觉大模型来做脊椎CT图像的分割。这事儿听起来可能有点“杀鸡用牛刀”毕竟分割脊椎在传统图像处理里也不算特别新鲜。但真正上手后我发现这背后其实是一个很典型的场景——如何让一个强大的、预训练好的通用基础模型Foundation Model快速适配到一个数据稀缺、标注成本高昂的专业垂直领域。SAMSegment 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 torch1.13.1cu117 torchvision0.14.1cu117 torchaudio0.13.1 --extra-index-url https://download.pytorch.org/whl/cu117OpenCV医学图像读取和预处理必备。注意opencv-python和opencv-python-headless的区别。如果你在无GUI的服务器上跑或者不需要cv2.imshow这类功能用headless版本更轻量避免一些不必要的系统依赖。pip install opencv-python-headless4.8.1SimpleITK 或 NiBabel对于处理3D的CT数据我们最终处理的是2D切片但数据源是3D的需要库来读取DICOM或NIfTI格式。SimpleITK功能强大但稍重NiBabel更轻量。由于我们主要关心像素数据和简单元信息我选择了NiBabel。pip install nibabel5.1.0MONAI这是一个医学影像AI的PyTorch专属框架。SAM-Med 2D的预处理、数据增强流程很可能借鉴或兼容MONAI的风格。即使原代码未直接使用引入MONAI的transforms来进行数据增强也是极好的选择它提供了大量针对医学图像的增强操作如随机弹性形变、Gamma变换等。pip install monai1.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腰椎L12腰椎L2...。我们需要决定是训练一个模型分割所有椎体多类分割还是每个椎体单独训练一个模型二分类分割。对于SAM-Med由于其提示机制更自然的做法是进行二分类分割即模型只学习分割“脊椎骨”这个整体或者更进一步通过不同的提示来区分不同椎体。在初期我建议先做二分类脊椎骨 vs 背景这样问题更简单。过滤无效切片很多CT切片在头部或尾部并不包含脊椎。我们可以通过计算2D掩码中前景像素的比例来过滤掉这些“空”切片节省存储和训练时间。3.2 标注格式适配与数据集类编写SAM-Med 2D的原始数据加载器可能期望某种特定格式。常见的有两种格式A图像文件夹 掩码文件夹。要求文件名一一对应如001.png和001_mask.png。掩码图为单通道PNG前景为255背景为0二分类。格式BCOCO格式的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, transformNone): 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(imageimage, maskmask) 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 (i1) % 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, threshold0.5, min_area50): 对模型输出的概率图进行后处理。 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, connectivity8) # 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.92Hausdorff距离控制在5个像素以内对于后续的脊柱形态测量分析这个精度已经足够作为可靠的输入。这个过程让我深刻体会到用好一个视觉大模型三分在模型七分在数据。数据的质量、预处理的方式、与模型预期的匹配程度往往比调参更能决定项目的成败。SAM-Med 2D提供了一个强大的基础但如何将它“调教”成你专属领域的专家考验的是你对业务脊椎解剖、数据CT影像和模型原理的综合理解。本文还有配套的精品资源点击获取