
简介基于 PyTorch 的医学图像分割基础框架设计源码面向医学影像分析领域的研究者与开发者聚焦器官、肿瘤等感兴趣区域的分割任务提供从数据预处理、网络搭建到模型训练、验证与测试的完整代码流程适合快速开展分割实验与算法验证。压缩包共 118 个文件大小约 3.78MB其中包含 75 张 PNG 样本图、18 个 Python 源文件、12 个 pyc 编译文件、11 个 txt 说明/配置文件以及 gitignore 与 keep 文件PNG 用于展示数据集样本py 文件覆盖模型构建、训练与评估等核心模块pyc 为编译缓存文件txt 便于查阅配置与日志目录结构清晰便于按需修改与扩展。目前已有 509 人学习下载。该框架针对数据准备、网络结构调整与工具复用做了模块化设计如 networks、utils、dataprepare、datasets 等目录可帮助使用者快速适配不同的分割需求亦可作为入门医学图像分割与二次开发的基础工程。1. 医学图像分割的PyTorch基础框架先解决“可复现”再谈“涨点”很多同学第一次接触医学图像分割时都经历过同一个循环从网上找一份能跑的 U-Net 脚本换到自己的 CT/MRI 数据集上报错、改路径、调参数最后跑出来的结果却说不清是哪一步起了作用。PyTorch 生态里现成的分割库不少但基础框架的价值不在“跑通”而在把数据加载、模型、损失、训练、验证这五块边界拆清楚让每次实验改动都能被追踪、被比较。医学图像分割的数据格式和自然图像差别极大窗口化、病人级划分、类别不均衡这些前提没立住后续再好的模型也白搭。这套框架适合正在做毕设或小团队自研的同学也适合需要把论文方法落地的工程师——先用最简结构抓到可靠基准线再谈提点。2. 工程骨架Anaconda装PyTorch、Config设计与源码目录怎么摆2.1 环境先行PyTorch 的 GPU 版本要按驱动选不是装最新做医学图像分割GPU 基本是刚需但 PyTorch 安装并不是“越新越好”。我见过不止一次这样的场景机器上nvidia-smi显示驱动正常torch.cuda.is_available()却返回 False查到最后是 pip 默认装成了 CPU 版或者驱动版本太老带不动新 CUDA 轮子。医学影像圈里大家习惯用 Anaconda理由也简单一台机器可以同时放两套环境一套 PyTorch 1.x 做旧代码复现一套 2.x 跑新实验互不搞坏。我一般这样建一个干净的 medseg 环境conda create -n medseg python3.10 -y conda activate medseg pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu121选 Python 3.10 不是图新鲜而是 SimpleITK、nibabel、monai 这些医学影像常用库在 3.10 上的预编译 wheel 最齐全少碰编译源码的麻烦。--index-url指向 cu121 的轮子仓库意思是让 pip 从 PyTorch 官方频道拉 CUDA 12.1 对应的版本而不是从默认 PyPI 拉一个可能不带 CUDA 的包。装完第一件事不是急着 import 模型而是验证环境nvidia-smi python -c import torch; print(torch.__version__, torch.version.cuda, torch.cuda.is_available())nvidia-smi看的是显卡驱动能支持的最高 CUDA 版本torch.version.cuda看的是当前 PyTorch 自带运行时两个能对上一般就没问题。这里有个常见误区如果用的是 WSL2GPU 驱动安装在 Windows 宿主机侧WSL 里面不要自己再装一遍 NVIDIA 驱动否则容易把内核模块搞乱。在 WSL2 里跑医学分割框架的同学经常遇到的翻车点就是把驱动装进了 WSL 内部最后nvidia-smi直接看不到卡。确认torch.cuda.is_available()返回 True 之后再进下一步否则后面所有训练代码都是在黑匣子里跑。2.2 源码目录把数据、模型、训练、推理拆成四个模块基础框架的源码组织核心目标只有一个换数据集时只改一个文件做实验时只改配置不碰训练循环。我常用的目录结构长这样medseg_framework/ ├── configs/ │ └── brain_tumor.yaml # 每个实验一份配置留档用 ├── data/ │ └── Task01_BrainTumour/ # 原始数据建议软链不拷进项目 ├── src/ │ ├── dataset.py # 数据读取、预处理、增强 │ ├── model.py # 模型定义默认 U-Net │ ├── losses.py # Dice / CE / 组合损失 │ ├── train.py # 训练循环、验证、模型保存 │ ├── infer.py # 推理、滑动窗口、指标计算 │ └── utils.py # 公共工具种子、日志、指标 ├── checkpoints/ # 训练产物按实验名分文件夹 └── README.mddataset.py是换数据集时最常改的文件。医学影像不像自然图像那样一个文件夹全是 jpg它可能是 DICOM 序列、NIfTI 单文件、多模态多文件把这些差异全部装进dataset.py上层训练代码就不用感知。losses.py单独拆出来是因为医学分割的损失函数经常要试不同组合拆开后可以在配置里直接指定。train.py里只放训练循环本身不写数据读取、不写网络结构这样当你想换 loss 或换模型时训练脚本几乎不用动。2.3 Config 设计用 dataclass 把训练参数集中管理别散落在训练脚本里新手写训练代码最常见的习惯是把 batch size、学习率、epoch 数全写成脚本里的变量跑几次实验就不知道当时用的是哪组参数。基础框架里我倾向于先用 dataclass 做一个最小配置from dataclasses import dataclass from pathlib import Path dataclass class TrainConfig: data_dir: Path Path(./data/Task01_BrainTumour) model_dir: Path Path(./checkpoints/baseline_unet) n_classes: int 3 # 背景 肿瘤增强区 肿瘤核心区 patch_size: tuple (256, 256) # 训练切块大小受显存限制 batch_size: int 8 # 显存不够时先降到 4别动 patch lr: float 1e-4 # AdamW 常用初始学习率 epochs: int 150 seed: int 42 # 固定种子保证可复现 num_workers: int 4 # 数据加载进程数Windows 上别超过 CPU 核数dataclass 的好处是所有训练参数集中在一个类里实例化后直接传给训练函数想跑不同实验复制一份配置改成新名字即可。第一次跑通时只需要动patch_size、batch_size、epochs这三个参数。后续如果实验变多再把这套 dataclass 序列化成 YAML 存到configs/目录每跑一次实验就留一份配置档这就是后悔药。注意seed只固定了 Python 和 PyTorch 的随机数数据加载线程的随机性不一定完全可控。严格可复现还需要配合 DataLoader 的generator参数基础框架阶段先固定全局种子就够用。3. 数据端NIfTI 切片转 Dataset预处理与增强的参数别乱抄3.1 体积数据落到 2D 切片读 NIfTI 还是 DICOM医学影像原始格式主要有两种DICOM 是医院设备直接输出的文件集合一个检查可能几十上百个文件还带一堆元信息NIfTI 是科研界常用的单文件格式把三维体积压缩在一个.nii.gz里读取方便。做分割框架时我一般建议统一转成 NIfTI 再喂给 Dataset。如果手上只有 DICOM 序列常见做法是用 dcm2niix 转换dcm2niix -o ./output_dir ./dicom_input_dir这个工具会把一个序列转成一个.nii.gz文件同时生成一个.json存放扫描参数。转到 NIfTI 之后读数据的代码就只有一种路径不用在 Dataset 里同时维护两套读取逻辑。至于为什么基础框架先从 2D 切片做起而不是直接上 3D U-Net原因很实际2D 切片训练显存压力小、可视化直观、调试速度快而且几乎所有 2D 上的预处理和增强经验都能平移到 3D 版本。先把 2D 跑通拿到基准线再根据任务需求升级到 3D是一条稳妥路径。3.2 Dataset 核心实现窗口化、归一化、随机取切片下面这段代码是框架数据端的核心负责把 NIfTI 文件变成模型能吃到的张量。以 CT 脑肿瘤数据为例重点看预处理那几行import numpy as np import torch from torch.utils.data import Dataset import nibabel as nib from pathlib import Path class SegDataset(Dataset): def __init__(self, case_ids, image_dir, label_dir, patch_size(256, 256), n_classes3, window(None, None), augmentFalse): self.case_ids case_ids self.image_dir Path(image_dir) self.label_dir Path(label_dir) self.patch_size patch_size self.n_classes n_classes self.window window # CT 形如 (-200, 400)MRI 用 None self.augment augment def __len__(self): return len(self.case_ids) def __getitem__(self, idx): case_id self.case_ids[idx] img nib.load(str(self.image_dir / f{case_id}.nii.gz)).get_fdata() lab nib.load(str(self.label_dir / f{case_id}.nii.gz)).get_fdata() lab lab.astype(np.int64) # 随机选一个轴向切片如果想固定看中间层把 randint 改成固定索引 slice_idx np.random.randint(0, img.shape[-1]) img2d img[..., slice_idx].astype(np.float32) lab2d lab[..., slice_idx] # CT 窗口化把窗宽外的像素截断再线性映射到 [0, 1] if self.window[0] is not None: low, high self.window img2d np.clip(img2d, low, high) img2d (img2d - low) / (high - low) else: # MRI 没有固定窗宽窗位用百分位截断防止个别亮点拉高整体对比度 lo, hi np.percentile(img2d, (2, 98)) img2d np.clip(img2d, lo, hi) img2d (img2d - lo) / (hi - lo 1e-6) if self.augment: img2d, lab2d self.apply_augment(img2d, lab2d) # 转成 (C, H, W) 格式C1 表示单通道灰度 img_tensor torch.from_numpy(np.ascontiguousarray(img2d)).unsqueeze(0).float() lab_tensor torch.from_numpy(np.ascontiguousarray(lab2d)).long() return {image: img_tensor, label: lab_tensor, case_id: case_id}get_fdata()读出来是浮点数组NIfTI 的轴序一般是 (x, y, z)slice_idx沿最后一维取切片。window参数是最容易抄错的地方CT 的窗宽窗位由检查部位决定肝部常用 -200 到 400肺部常用 -1000 到 400数值来自成像协议不来自哪个开源项目MRI 没有 CT 那样的物理单位用 2%-98% 百分位截断更稳妥。np.ascontiguousarray这步很多新手会漏切片和 transpose 之后数组不连续放进 CUDA 张量时可能报 stride 相关的错。lab_tensor保持整数标签不做 one-hotone-hot 放到损失函数里做这样 dataset 代码保持简单。3.3 数据增强与病人划分先按 case_id 分层再来谈翻转旋转医学图像分割的增强要克制。自然图像可以随便裁剪、旋转、调色但医学影像翻转可能改变左右语义大角度旋转可能让解剖结构失真。基础框架里我默认只做水平翻转、按 90 度旋转和轻微亮度扰动def apply_augment(self, img2d, lab2d): # 图像和标签必须做完全相同的变换 if np.random.rand() 0.5: img2d np.flip(img2d, axis1) lab2d np.flip(lab2d, axis1) k np.random.choice([0, 1, 2, 3]) if k 0: img2d np.rot90(img2d, k, axes(0, 1)) lab2d np.rot90(lab2d, k, axes(0, 1)) return img2d.copy(), lab2d.copy()图像和标签必须用同一组变换旋转角度和翻转轴必须完全一致这是铁律否则模型学的是错位标签。弹性形变是医学分割里很有用的增强但参数难调变形过大会把血管、肿瘤边界扭曲到不可信我一般等基准线跑稳之后再加入。另一个比增强更重要的步骤是病人级划分。如果按切片随机 split同一个病人的相邻切片会被同时分进训练集和验证集模型记住了病人纹理而不是病灶特征验证集 Dice 虚高。正确做法是按 case_id 分组划分from sklearn.model_selection import GroupShuffleSplit gss GroupShuffleSplit(n_splits1, test_size0.2, random_state42) train_idx, val_idx next(gss.split(case_ids, groupscase_ids)) train_cases [case_ids[i] for i in train_idx] val_cases [case_ids[i] for i in val_idx]GroupShuffleSplit的groups参数传 case_id保证同一病人全部落在同一侧。这个细节直接决定后续所有指标是否可信我在第一次搭框架时也在这里翻过车。4. 模型端轻量 U-Net、Dice 损失与 AMP 训练的最小编法4.1 通道与归一化基础 U-Net 不该从复杂结构开始医学分割基础框架的模型首选就是 U-Net它结构简单、对少量数据友好、效果稳定。U-Net 本质上是一个带跳跃连接的编码器解码器常见做法是把基础通道数设为 32 或 64每下采样一次翻倍。下面是最核心的卷积块import torch.nn as nn class DoubleConv(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.conv nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding1, biasFalse), nn.BatchNorm2d(out_ch) if use_bn else nn.InstanceNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, 3, padding1, biasFalse), nn.BatchNorm2d(out_ch) if use_bn else nn.InstanceNorm2d(out_ch), nn.ReLU(inplaceTrue) ) def forward(self, x): return self.conv(x)编码器做三次下采样每次MaxPool2d(2)后接 DoubleConv解码器用双线性上采样加跳跃连接最后接一个Conv2d(64, n_classes, 1)。这里有个容易忽略的点通道数不是越大越好。医学分割数据集通常只有几十到几百个病例模型容量过大直接过拟合基准阶段先压小通道数跑通再逐步翻倍看收益。归一化层选择也要看 batch sizebatch_size 大于等于 8 时 BatchNorm 表现稳定一次只喂 2 到 4 张图时BN 的均值方差估计会抖换 InstanceNorm 反而更稳。基础框架默认用 BN但如果你的显存只允许小 batch提前改成 InstanceNorm 能省很多调试时间。4.2 Dice 与 CE 结合类别不均衡时不能只靠 CrossEntropy医学分割的标签有个显著特点背景像素占比极高病变区域可能只占几个百分点。只算 CrossEntropy 时模型学到的是“全都预测背景也能拿到很低 loss”Dice Loss 按区域重合度计算能直接缓解这个问题。基础框架里我一般把 Dice 和 CE 按 0.5 对 0.5 加权import torch import torch.nn.functional as F def dice_loss(pred, target, smooth1.0): # pred: (B, C, H, W) logitstarget: (B, H, W) 整数标签 n_classes pred.shape[1] pred_soft F.softmax(pred, dim1) # 按通道做概率归一化 target_onehot F.one_hot(target, n_classes) # (B, H, W, C) target_onehot target_onehot.permute(0, 3, 1, 2).float() intersection (pred_soft * target_onehot).sum(dim(2, 3)) denom pred_soft.sum(dim(2, 3)) target_onehot.sum(dim(2, 3)) dice (2.0 * intersection smooth) / (denom smooth) return 1.0 - dice.mean() def combined_loss(pred, target): ce F.cross_entropy(pred, target) dice dice_loss(pred, target) return 0.5 * ce 0.5 * dicesmooth是平滑项防止某个类别在 batch 里完全不出现时除零。F.one_hot要求标签从 0 开始连续编号如果你的掩膜只有 1 和 2 而没有 0 类先把标签减 1 再传进来。基础阶段建议让背景类也参与 Dice 计算跑通之后如果想追求和论文一致的评估口径再改成只统计前景类。4.3 训练循环AMP 加速与验证集保存的骨架训练循环是框架里最需要稳定的部分。基础框架用混合精度训练可以省一半显存代码只多三行。下面是一个最小可用的训练骨架import torch scaler torch.cuda.amp.GradScaler() def train_one_epoch(model, loader, optimizer, scheduler, device): model.train() total_loss 0.0 for batch in loader: img batch[image].to(device) lab batch[label].to(device) optimizer.zero_grad() with torch.amp.autocast(cuda): pred model(img) loss combined_loss(pred, lab) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() total_loss loss.item() scheduler.step() return total_loss / len(loader)autocast让前向传播里能用半精度的地方自动降精度GradScaler防止梯度下溢。注意scaler.step(optimizer)代替了原来的optimizer.step()不要漏。学习率调度我用CosineAnnealingLR配合AdamW默认lr1e-4相比固定学习率后期收敛更稳。验证时记得切回model.eval()并包上torch.no_grad()否则 BN 的统计量会被验证数据污染验证集 Dice 会偏低。验证之后按 Dice 保存最优模型是基础框架必须具备的能力。建议每轮验证结束计算验证集平均 Dice只有比上一轮高才写torch.save(model.state_dict(), save_path)这样实验中断或过拟合时checkpoints 里留的始终是验证集上最好的那版参数。5. 易踩的四个坑从 CUDA 识别到病人级数据泄漏的排查记录5.1 torch.cuda.is_available() 返回 False先看驱动版本再查包来源现象按网上的教程一步步装完torch.cuda.is_available()还是 False模型只能跑 CPU训练慢到怀疑人生。原因大概率是 pip 从默认 PyPI 装了 CPU 版 PyTorch或者显卡驱动太老跑不动新版 CUDA 轮子。还有一部分同学在 WSL2 里重复安装了 NVIDIA 驱动导致 WSL 内的显卡被识别成 unkonwn。解决按顺序做三件事。先跑nvidia-smi看驱动支持的最高 CUDA 版本再跑python -c import torch; print(torch.__version__, torch.version.cuda)看当前包实际带的是哪个 CUDA 运行时最后pip list | findstr torchWindows或pip list | grep torchLinux确认包来源。驱动版本允许的话直接重装对应的 cu121 或 cu118 轮子不要混装 conda 的 cudatoolkit。装完重新验证还不行就把环境删了重建别在同一套环境里反复补丁浪费时间。5.2 输入轴序错了从 (H, W) 到 (1, 1, H, W) 的那一步现象训练报RuntimeError: Expected 4D input或者 loss 计算时 target 和 pred 维度对不上。原因NIfTI 的get_fdata()返回三维数组切片后是二维直接torch.from_numpy(slice_2d)得到的张量没有通道维和 batch 维卷积网络不认。看起来只是少一次unsqueeze但这正是医学影像数据端最容易踩的轴序坑。解决统一在 dataset 里把切片转成 (1, 1, H, W)img_tensor torch.from_numpy(np.ascontiguousarray(img2d)).float().unsqueeze(0).unsqueeze(0)contiguous是因为 flip、rot90 之后底层内存可能不连续不处理时某些算子会报错。养成在 dataset 里就完成全部 reshape 的习惯训练代码只认四维张量排查范围就缩到数据端单侧。5.3 验证集 Dice 0.9测试崩到 0.6病人级数据泄漏现象训练过程验证曲线很漂亮最终 Dice 接近 0.9换一批新病人测试直接掉到 0.6。原因划分训练集和验证集时按切片随机打乱同一个病人的多张切片同时出现在两侧。相邻切片外形极其相似模型记住的是“这个病人的纹理”而不是“肿瘤长什么样”。解决用GroupShuffleSplit按 case_id 分组保证同一病例的所有切片只属于一侧。更严格的做法是把测试集独立出来调参期间完全不碰。这个坑做自然图像的人几乎不会遇到却是医学分割框架里最致命的一个我在这上面吃过亏现在每套数据集都先检查 case_id 分组再开训。5.4 推理和训练预处理不一致window 和 resize 必须共用同一条路径现象离线验证指标不错部署到新数据上效果明显变差甚至出现整片误判。原因训练时做了 CT 窗口化和 patch 裁剪推理脚本里图省事直接加载原始数组就进网络数值分布和训练时完全对不上。解决把窗口化、归一化、重采样这些预处理全部抽成src/utils.py里的同一个函数训练、验证、推理都调用它不搞第二份实现。下面是一个最小示例def preprocess_slice(slice_2d, window): if window is not None: low, high window slice_2d np.clip(slice_2d, low, high) lo, hi np.percentile(slice_2d, (2, 98)) slice_2d (slice_2d - lo) / (hi - lo 1e-6) return slice_2d.astype(np.float32)推理时如果输入是完整大尺寸图像预处理后还要做和训练一致的重采样先重采样再窗口化顺序反了结果也会差一截。记住一条原则训练管线里出现的任何一步操作推理管线都必须原样复刻一遍。6. 进阶验证滑动窗口推理与 DSC/HD95 双指标6.1 大体积数据用滑动窗口推理别直接缩小整图训练时模型见过的是 256×256 的 patch推理时如果为了省事把整张切片强 resize 到 256 再预测小病灶和边界细节会丢得不成样子。标准做法是用滑动窗口对每个 patch 单独推理再把结果拼回原图。相邻窗口之间设置 overlap重叠区域取多次预测的平均值能有效减少接缝伪影def sliding_slice_infer(model, slice_2d, patch256, stride192, devicecuda): h, w slice_2d.shape preds torch.zeros((n_classes, h, w), devicedevice) counts torch.zeros((h, w), devicedevice) model.eval() with torch.no_grad(): for i in range(0, h - patch 1, stride): for j in range(0, w - patch 1, stride): crop torch.from_numpy(slice_2d[i:ipatch, j:jpatch]).float() crop crop.unsqueeze(0).unsqueeze(0).to(device) out torch.softmax(model(crop), dim1)[0].cpu() preds[:, i:ipatch, j:jpatch] out counts[i:ipatch, j:jpatch] 1 preds preds / counts.clamp(min1) return preds.argmax(dim0).numpy()stride192表示相邻窗口有 64 像素重叠重叠区每个像素累积了多次 softmax 输出再求平均比直接取一次预测更稳。边缘位置如果不足 patch用reflect模式补边基础框架先不考虑边缘切块把主体区域做对之后再去补边。推理和训练用的预处理函数保持一致这是唯一靠谱的路径。6.2 用 DSC 与 HD95 两套指标别只盯着一个数字训练日志里如果只记录一个指标我最常用的是 DSCDice 相似系数因为它对体积重合度敏感、数值直观。但 DSC 有一个明显盲区对边界位置不敏感预测区域整体向外扩一圈DSC 可能还有 0.9临床却完全不接受。所以基础框架的验证端会同时输出 HD9595 分位豪斯多夫距离指标反映什么什么时候容易骗人DSC预测与真值的体积重合比例小病灶漏掉时 DSC 下降不明显HD95表面距离的 95 分位单位 mm边界外扩或轻微错位时迅速变大HD95 的计算需要额外的几何工具常见做法是调 medpy 的binary.hd95。对器官分割我会盯 DSC对肿瘤和细小结构我更相信 HD95。现在每跑一轮实验我至少会同时记录这两个数再加一个训练 loss然后把配置文件和 checkpoints 一起归档。三个月后再翻这些实验记录就能说清楚当时每一个改动到底值不值。希望这套框架思路能帮你少走我走过的弯路。本文还有配套的精品资源点击获取