资讯详情

SAM-Med 2D医学图像分割复现与微调实战:从环境配置到Dice提升

📅 2026/9/27 23:10:42 | 华诺云谱 👁 阅读
SAM-Med 2D医学图像分割复现与微调实战:从环境配置到Dice提升
简介面向医学图像分割与大模型复现需求这份资源包提供SAM-Med 2D在脊椎分割任务上的完整数据集与训练方案适用对象为具备一定深度学习基础的研究者或开发者。压缩包含2000个文件以1949张PNG图像为主体配套Python训练脚本、JSON标签映射、Markdown说明文档及Word格式的实验报告总大小约243.76MB目录结构清晰便于按流程取用。目前已有1918人学习下载。数据集存放于RawData目录运行process脚本即可自动生成data_demo训练数据无需修改参数即可直接执行train脚本启动训练大幅降低复现门槛docx文档附有脊椎SAM模型的分割预测结果和训练数据制作详解另有Jupyter Notebook示例支持逐步运行调试可帮助理解完整流程。尤其适合需要快速上手医学图像分割大模型、或在自定义数据集上微调SAM-Med 2D的读者可用于科研试验或项目落地。1. SAM-Med 2D视觉大模型复现先别急着跑训练把 SAM-Med 2D 视觉大模型复现这条路走通是我最近拆过的医学图像项目里最值得记录的一个。它的定位很明确——通用 SAM 在自然图像上很强但一碰到 CT、MRI 这类灰度医学图像就明显水土不服椎体边界经常糊成一团而 SAM-Med 2D 是在大规模医学图像上做过适配的视觉大模型拿来分割脊椎、器官这类结构比直接套 SAM 靠谱得多。这篇笔记解决三件事怎么把官方流程复现出来并跑通推理怎么把你手头标好的脊椎数据转成它能吃的训练格式以及微调时那些让人挠头的坑。适合手里有标注数据或标注工具、准备做医学影像分割但不想从零搭网络的人。2. 复现首步环境、权重与官方推理链路2.1 环境版本对照一张表看清 PyTorch 和 CUDA 怎么配先讲环境因为这步最容易被忽略但翻车率也最高。SAM-Med 2D 本质上是基于 SAM 的 ViT 骨架做的微调模型所以它对 PyTorch 版本的依赖基本跟随 SAM 走。我复现时用的组合是 Python 3.10 PyTorch 2.0.1 CUDA 11.8这套组合最稳。原因很直接官方仓库里大量代码依赖 transformer 和 timm 的特定接口版本太新容易遇到 API 改名太旧又跑不动 ViT-B 所需的算子。用 conda 建环境常见做法是这样conda create -n sammed2d python3.10 -y conda activate sammed2d pip install torch2.0.1 torchvision0.15.2 --index-url https://download.pytorch.org/whl/cu118 pip install opencv-python pillow numpy scipy tqdm这里说明一下为什么这么装。PyTorch 单独指定--index-url到 cu118 的 wheel 源是为了避免 conda 默认源装了 CPU 版——CPU 版能加载权重也能推理但速度慢到没法做训练。detectron2 我没写进命令因为它在官方仓库里主要服务于可视化工具训练和推理主链路用不到如果后面跑官方 demo 脚本报 ImportError再补装也不迟。提示装完先跑一句python -c import torch; print(torch.cuda.is_available())输出 True 再进行下一步这一步能过滤掉八成环境问题。我把关键依赖的版本对照放在表里方便你按自己的显卡调整组件推荐版本踩坑说明Python3.103.11 也能跑但 opencv 等老依赖容易出兼容告警PyTorch2.0.11.13 也能训练但 2.x 的 AMP 混合精度更稳CUDA11.8显卡驱动需不低于 520用 nvidia-smi 确认timm0.6.13别装 0.9 版本build_model_with_cfg 接口不兼容numpy1.24.x2.x 改了随机数接口某些数据增强代码会报错2.2 权重下载与目录组织避免最常犯的路径错误代码 clone 下来后我习惯先把目录和数据骨架建好而不是直接跑脚本。这一步看起来多余但后面训练时所有路径都依赖这个结构git clone https://github.com/your-cloned-repo/SAM-Med2D cd SAM-Med2D mkdir -p checkpoints data/images data/labels data/json_datas权重我建议从官方指定的渠道下载文件一般是sam_med2d_b.pth放到 checkpoints 目录。这里有个常见的坑权重文件命名和代码里写死的名字不一致时加载会静默失败不报错但模型权重是随机的或者直接抛 KeyError。所以下载后先核对文件名如果实际叫sam-med2d_b.pth之类的要改回脚本预期的名字。我在这上面浪费过半小时血泪经验。目录结构里data/images放原始医学图像data/labels放掩码 PNGdata/json_datas放训练提示 JSON。这三个子目录是官方训练脚本默认找数据的路径如果你把数据散放在别处后面 finetune.py 的参数就要逐个改没必要。2.3 拆官方推理脚本从加载权重到出掩码环境配好、权重就位之后先跑通推理再谈训练。我一般拿一张自己的脊椎图像试不用官方示例图因为自己的数据能直观看出模型对医学图像的适配程度。import torch from sam_med2d import sam_model_registry from utils import get_mask_from_points sam sam_model_registry[vit_b](checkpoint./checkpoints/sam_med2d_b.pth) sam.cuda().eval() # image 是读入的 CT/MRI 图像h,w,c 的 uint8 数组 box [x1, y1, x2, y2] # 框提示坐标原图坐标系 mask get_mask_from_points(sam, image, [], [box], devicecuda) print(mask.shape) # (1, 256, 256)这段代码是官方推理链路的浓缩。sam_model_registry[vit_b]按 backbone 名字注册模型结构checkpoint参数加载预训练权重get_mask_from_points把框提示编码进 prompt encoder再走 mask decoder 出掩码。关键点是 box 坐标必须和image对应同一个坐标系函数内部会按输入尺寸做 resize 映射所以传入的 image 是原始尺寸box 也用原始坐标即可。跑通这一步意味着环境、权重、代码链路都没问题。接下来要做的是把训练数据准备好这也是自定义数据集复现最花时间的部分。3. 准备脊椎分割训练集把标注转成 SAM-Med 2D 能吃的 JSON3.1 医学标注格式与训练格式的差异SAM-Med 2D 的微调数据格式和 COCO 类似但更轻量每组数据对应一个 JSON包含图片路径、掩码标签路径以及 prompt 信息——框和点。标签图是 PNG 单通道像素值 0 为背景、255 或 1 为前景。这和 ITK-SNAP 直接导出的 NIfTI 格式不一样也和 labelme 保存的多边形 JSON 不一样所以必须做一次转换。以脊椎分割为例典型的输入是 CT 或 MRI 的某一个断面。你用 ITK-SNAP 分割完保存的是 .nii.gz或者用 labelme 手动描完椎体轮廓保存的是多边形坐标。这两种格式都不能直接喂给训练脚本需要统一转成「原图 PNG 掩码 PNG 提示 JSON」三件套。3.2 转换脚本从掩码 PNG 到训练 JSON下面这个脚本是我实际用来做转换的输入是一张掩码 PNG输出是对应的训练 JSON。它做的事情很简单读掩码、找连通域、把每个连通域的外接矩形写到 boxes 字段里。import json import numpy as np from PIL import Image from skimage import measure from pathlib import Path image_dir Path(data/images) label_dir Path(data/labels) json_dir Path(data/json_datas) for label_path in label_dir.glob(*.png): mask (np.array(Image.open(label_path)) 0).astype(np.uint8) props measure.regionprops(mask) boxes [] for p in props: y1, x1, y2, x2 p.bbox boxes.append([int(x1), int(y1), int(x2), int(y2)]) sample { image: str(image_dir / (label_path.stem .png)), label: str(label_path), boxes: boxes, points: [] } (json_dir / (label_path.stem .json)).write_text( json.dumps(sample, indent2))代码逻辑分三段。regionprops遍历掩码里每个连通域返回其属性bbox给出外接矩形的边界但注意它返回的顺序是[y1, x1, y2, x2]我写入时把 x 和 y 调换回图像坐标系这个细节不处理训练时框提示就会整体偏移。最后生成的 JSON 里boxes是框提示列表points留空后面再决定要不要加。注意掩码读入后做了 0的二值化。如果你的标签图里像素值是 255这一句能自动归一化成 1但如果你的标签图里除了 0 和 1 之外还有别的灰度值比如多类别标注regionprops会把每个灰度级当成一个连通域导致框数量爆炸。多类别场景建议先按类别拆成多张单类掩码再逐类转换。3.3 点提示怎么生成对脊椎分割尤其重要只靠框提示微调模型容易把整个框内区域都当成前景对细长的椎体边界不够敏感。椎体在 CT 里是边缘清晰但不规则的结构框提示给的信息太粗所以微调时我建议混合使用点提示。常见做法是对前景掩码做距离变换取中心点作为正提示再在背景区域随机采样几个负提示点from scipy.ndimage import distance_transform_edt dist distance_transform_edt(mask) pos_y, pos_x np.unravel_index(np.argmax(dist), mask.shape) points [[int(pos_x), int(pos_y), 1]] neg_y, neg_x np.where(mask 0) if len(neg_y) 0: idx np.random.choice(len(neg_y), min(3, len(neg_y))) for i in idx: points.append([int(neg_x[i]), int(neg_y[i]), 0])distance_transform_edt计算前景区域每个像素到背景的最近距离距离最大的点就是椎体内部最中心的点作为正提示最有代表性。负提示点在背景区域随机采 3 个作用是告诉模型「这些位置不是前景」。第三位数字是 label——1 表示前景0 表示背景。训练时 prompt encoder 会把点和框一起编码模型学习「在提示指定的位置附近分割」这比单纯给框要精准得多。生成的点要更新到上一小节的 JSON 里吗我的做法是直接合到转换脚本里每张图同时生成框和点一步到位。这样训练时 80% 样本走点提示 框提示混合20% 只走框提示模型对两种模式的输入都能适应。4. 微调实操加载预训练权重让 Dice 从 0.5 爬到 0.94.1 训练脚本参数解读一张表看清每个 flag数据准备好后进入微调阶段。官方 finetune.py 脚本参数比较多我第一次跑的时候逐个查文档查了半天这里直接给参数表参数我的取值参数含义与理由--data_root./data数据根目录脚本会去它下面找 images/labels/json_datas--sam_checkpoint./checkpoints/sam_med2d_b.pth预训练权重必须和 model_type 匹配--model_typevit_bbackbone 类型对应 SAM-Med 2D 的 B 版--batch_size4显存够大可以到 8不够就减到 2 并配合梯度累积--lr1e-4微调学习率医学图像数据量小低于 5e-5 收敛太慢--num_epochs50单组数据量不大时 50 轮足够多了容易过拟合--image_size256输入分辨率显存吃紧可以降到 192--json_data_dir./data/json_datas提示标注 JSON 目录每个参数都不是孤立的。image_size直接影响显存占用和训练速度256 是直观感受中精度和速度的平衡点lr太低会导致 Dice 曲线爬得很慢太高又容易在预训练权重附近震荡。我第一次跑的时候用了 1e-3结果 loss 曲线像心电图一样上下跳后来降到 1e-4 才正常。4.2 启动训练命令、日志与断点恢复参数确认后启动命令如下python finetune.py \ --data_root ./data \ --sam_checkpoint ./checkpoints/sam_med2d_b.pth \ --model_type vit_b \ --batch_size 4 \ --lr 1e-4 \ --num_epochs 50 \ --image_size 256 \ --json_data_dir ./data/json_datas \ --output_dir ./runs/spine训练过程中日志会打印每个 epoch 的 loss 和验证集指标。关键要盯的是val_dice这个值它是模型在验证集上算出的 Dice 系数。前 10 个 epoch 你会看到 dice 涨得很快从 0.1 冲到 0.7 都正常但从 0.7 往 0.9 爬的阶段会明显变慢甚至连续十几个 epoch 看起来在原地踏步。这时候别急于加轮数先确认两点一是训练集和验证集的数据分布是否一致二是点提示有没有正确加载。我遇到过验证集里混了一张没有标注的空白切片结果 dice 被拉到 0.3 以下排查了半小时才找到。断点恢复是个容易被忽略但很实用的功能。训练中断时 finetune.py 会在 output_dir 下保留最近的 checkpoint重新启动时指定--resume ./runs/spine/latest.pth就能接着训。我建议每 5 个 epoch 手动备份一下权重因为训练到后期一次崩溃可能损失十几个 epoch 的进度。4.3 监控指标Dice 曲线和 loss 曲线怎么配合看训练时同时看两个曲线别只看 loss。loss 是多个损失函数的加权和SAM 官方用的是 focal loss dice loss L1 的组合它对微小边界的惩罚不够直观而 Dice 指标直接衡量分割区域和真实区域的重叠程度更贴近你的最终目标。epoch 10: loss 0.42, val_dice 0.61 epoch 20: loss 0.31, val_dice 0.74 epoch 30: loss 0.27, val_dice 0.81 epoch 40: loss 0.24, val_dice 0.85 epoch 50: loss 0.22, val_dice 0.87上面这个趋势是典型的健康收敛曲线。如果出现 loss 降到 0.2 以下但 val_dice 只有 0.6说明模型过拟合了训练集中的背景区域对椎体边界的细节没学会这时候要回查数据和提示点而不是继续训下去。反过来如果 loss 还在 0.5 以上但 val_dice 已经 0.8说明分割目标区域相对集中模型学得比 loss 显示的要好这种情况可以提前收工。5. 避坑排查复现和训练里最常见的五个翻车现场5.1 加载权重后推理全黑像素范围问题现象权重加载成功推理也不报错但输出的掩码全黑或者全是噪声点。原因医学图像本质是 16bit 灰度图直方图范围在 0-4096 甚至 0-65535而 SAM-Med 2D 的输入归一化假设图像是 8bit0-255。图像不缩放直接送入模型ViT 的 patch embedding 会把高像素值当成极端特征导致 mask decoder 输出完全偏离。解决读图后先做归一化常见做法是取图像的最大最小值做 min-max 缩放或者直接np.clip(image, 0, 255).astype(np.uint8)。我在推理脚本里统一用 min-max 归一化效果最稳定因为不同 CT 设备的像素范围差异很大固定阈值容易失效。5.2 显存 OOM连 batch_size2 都跑不动现象启动训练后几秒内报 CUDA out of memorybatch_size 降到 2 也没用。原因ViT-B 本身参数就有 90M 左右输入 256x256 时中间特征图占显存很可观再加上训练时同时要计算图像编码器、prompt encoder、mask decoder 三部分的梯度显存占用比推理高一个量级。解决优先把--image_size从 256 降到 192这一步能省接近一半显存再配合梯度累积--accumulation_steps 2相当于用 2 个 step 模拟一次有效更新。如果还是不够检查一下机器上是不是有别的进程占了显存nvidia-smi看一眼经常是之前跑过的训练进程没杀掉。5.3 训练时图像和掩码尺寸错位现象训练刚开始报 shape mismatch或者 loss 直接变成 nan。原因转换脚本里图像做了 resize但掩码和 JSON 里的框坐标没跟着变导致掩码是原尺寸图像是 256x256两者在计算损失时维度对不上。更隐蔽的情况是掩码和图像都 resize 了但插值方式不同——图像用双线性掩码用了相同的双线性导致掩码边缘出现 0.7 这样的中间值。解决所有和掩码相关的转换强制使用 nearest 插值。我的统一规则是图像可以随意做几何变换但掩码只能用Image.fromarray(mask).resize(size, Image.NEAREST)。同时JSON 里的框坐标也必须在 resize 后同步更新不能只缩放图像。5.4 loss 一直掉但 Dice 卡在 0.5 附近现象训练 30 个 epochloss 从 0.6 降到 0.25但 val_dice 始终在 0.5-0.55 晃动上不去。原因这是分割任务里最典型的翻车——模型学会了「把整个框内区域都预测成前景」。由于椎体在框内占比高这个策略算出来 loss 很低但分割出来的区域把背景也包进去了边界精度很差。本质是框提示太粗模型没有学到精细边界的表征。解决给训练数据增加点提示让正提示点落在椎体内部负提示点落在椎体边缘外的背景上模型才能学会区分边界。我的做法是把 3.3 的距离变换逻辑做成离线脚本直接改掉 JSON 数据再重新启动训练。改完后 Dice 通常在 10 个 epoch 内突破 0.7。5.5 提示点坐标错位导致分割结果偏移现象推理时给了一个很准的框提示但分割区域整体往左上或右下偏移了几个像素。原因训练时图像是 resize 后的尺寸但推理时你输入的 box 坐标是原始图像坐标系。如果推理脚本内部没有把 box 映射到 resize 后的坐标prompt encoder 拿到的提示点就和实际位置对不上输出就跟着偏。解决推理前做坐标映射resized_x original_x * (resized_width / original_width)box 的四个坐标都要换算。我把这段逻辑写成了固定函数每次推理前强制调用从那以后再也没有因为坐标映射翻过车。6. 推理验证与进阶把模型接到你的批量分割流程6.1 带坐标映射的推理脚本训练完成后把模型接到自己的批量分割流程里核心是写一个带坐标映射和后处理的推理函数def infer_trained(sam, image, boxes, scale_x, scale_y): masks [] for box in boxes: bx1, by1, bx2, by2 box bx1, by1 int(bx1 * scale_x), int(by1 * scale_y) bx2, by2 int(bx2 * scale_x), int(by2 * scale_y) mask get_mask_from_points(sam, image, [], [[bx1, by1, bx2, by2]], cuda) mask (mask 0.5).astype(np.uint8) masks.append((bx1, by1, mask)) return masks推理后处理我一般做两步先按阈值 0.5 二值化再用形态学闭运算把椎体内部的小空洞填掉。如果单张图里有多个椎体建议把每个连通域单独存一张掩码方便后续按椎体编号统计。6.2 进阶把 2D 模型接到 3D 序列最后一个进阶技巧针对 CT/MRI 这种连续多断面的场景。SAM-Med 2D 是 2D 模型但实际使用时我们面对的是一个三维 volume。之前的做法是逐切片推理再把结果堆叠成 3D 掩码但这样切片间会出现不连续——某一层椎体被分割下一层突然缺了一块。我一般用的折中方案是先抽中间层和一个靠近上下边界的层做推理确认该病例的灰度分布和椎体形态然后针对每一层做以下操作——用上一层的分割结果做形态学扩张作为当前层的框提示来源再配合当前层图像输入模型。这样做的好处是充分利用了层间连续性而且避免了大量重复标注框的时间。如果你面对的是整个 spine volume 的分割还可以考虑在得到 2D 分割结果后用简单的连通域追踪把相邻切片粘连起来这一步不需要额外训练但能把 2D 模型的输出平滑成更像 3D 的结果。这套流程跑下来我最大的教训是数据集转换阶段多花一小时训练阶段就能少踩三天的坑。从那以后我每次做医学图像微调都强制先走一遍「掩码二值化检查 坐标映射自测 单样本 JSON 可视化」再启动训练三个检查做完才敢把数据交给训练脚本。希望帮到你。本文还有配套的精品资源点击获取
📝

华诺云谱内容团队

资深建站顾问 · 行业研究员

10年+企业数字化服务经验,专注智能建站、SEO优化与品牌营销,持续输出建站技巧、行业洞察与营销干货,已帮助5000+企业实现数字化增长。

你可能需要的服务

订阅华诺云谱资讯周报

每周一封,精选建站技巧、SEO与营销干货,直达邮箱。已有 8,000+ 企业主订阅,助你少走弯路。

↑