海马体MRI分割:从3D nii.gz到2D U-Net的切片数据集与训练实践
2026/9/12 21:41:08 网站建设 项目流程

简介:面向医学图像分割任务的海马体切片数据集,适用于深度学习的语义分割模型训练与算法验证。数据源自左右海马体3D nii.gz文件,已分别沿横截面(x)、冠状面(y)、矢状面(z)切分为2D图像,并自动剔除前景区域占比不足3%的样本;mask标注中1为前部海马体、2为后部海马体、0为背景。包内共2000个文件,含1998张png格式的原图与对应标签图,以及1个可视化脚本show.py和1个json配置文件,整体压缩包约26.18MB;x轴切片为35×51分辨率共4958对,y轴为35×35分辨率共9465对,z轴为51×35分辨率共4603对,可直接划分训练集。show.py可随机抽取图片展示原始图像、GT标签及叠加效果,方便快速检查数据质量。当前已有353人学习,适合医学图像入门及分割方向研究者、开发者使用。

1. 为什么是左右海马体切片分割:从 3D nii.gz 到 2D PNG 的降维设计

拿到一批左右海马体的 nii.gz 三维标注数据,很多人的第一反应是直接上 3D U-Net。但实际操作中,显存占用、预处理成本和标注校验会先卡住你:单张 MRI 体数据动辄几百层,直接把三维 patch 喂进网络,batch size 稍微大一点就会爆显存。这套资源做了一个很实用的降维处理:将 3D 海马体标注按横断面(x)、冠状面(y)、矢状面(z)切成 2D 切片,并自动丢掉前景面积不足 3% 的纯背景帧,最后导出成图像和标签均为 PNG 的医学图像分割数据集。它解决的是从 nii.gz 到可训练样本之间的脏活累活,适合做 2D 分割模型验证、多切面对比实验,或者在上 3D 模型之前先用 2D 网络快速跑通基线。

2. 三个切面的目录结构与标签编码:先看懂数据再写加载器

这个数据集的目录不是常见的单张 image/mask 同级目录,而是按切面轴各自组织。拿到压缩包后先看根目录,通常会看到dataset.jsonx.pngz.png以及hippocampus_164_18.png这类文件,其中dataset.json是元数据入口。无论你是准备用 PyTorch 写 DataLoader,还是直接丢给 nnU-Net 做训练,第一步都要确认三件事:图像和 mask 是否同名、每个切面目录下有多少样本、mask 里到底有几个类别。下面按这三件事来拆。

2.1 文件命名、目录结构与 dataset.json

摘要里写得很清楚:x 轴 35×51 分辨率,images 图片目录加 masks 模板目录,4958 张图片和 4958 个对应的 mask;y 轴 35×35,9465 张;z 轴 51×35,4603 张。也就是说每个轴都有独立的 images 和 masks 目录,且图像文件名和 mask 文件名一一对应,例如hippocampus_164_18.png同时出现在 images 和 masks 里。这种命名在医学影像切片中很常见:前面是受试者或原始 3D 卷的 ID,后面是切片序号,hippocampus_164_18.png可以理解为编号 164 的 3D 卷,第 18 层切片。

先用一小段代码确认目录结构和 mask 形状:

import json from pathlib import Path from PIL import Image import numpy as np root = Path("hippocampus_slices") for axis in ["x", "y", "z"]: images = list((root / f"{axis}_images").glob("*.png")) masks = list((root / f"{axis}_masks").glob("*.png")) print(axis, len(images), len(masks)) sample = sorted(images)[0] mask_sample = sorted(masks)[0] print("image", np.array(Image.open(sample)).shape, "mask", np.array(Image.open(mask_sample)).shape)

这段代码用glob匹配 PNG 文件,sorted保证每次抽样顺序一致。代码输出的 shape 是(height, width),也就是 Pillow 读进来后 numpy 数组的行列顺序,而摘要里说的 35×51 是宽×高,所以 x 轴的 numpy 数组实际 shape 是(51, 35)。这个维度顺序在第 4 章写模型输入时一定要统一,否则会出现张量维度对不上或结果整体转置的问题。

如果dataset.json里放了文件列表或数据划分,直接用json.load读出来看结构:

with open(root / "dataset.json", "r", encoding="utf-8") as f: meta = json.load(f) print(meta.keys())

我一般会先打印keys,再看里面是 train/val/test 划分还是 images/masks 文件列表。这个数据集的作者没有在摘要里说明 JSON 的具体字段,所以不要假设字段名,先打印再决定怎么往下写。

2.2 mask 类别:0 为背景,1 为前部海马体,2 为后部海马体

标签文件不是二值图,而是三类别语义分割。其中 mask 中 1 为前部海马体、2 为后部海马体、0 为背景。这意味着如果你的网络输出通道数为 3,通道 0、1、2 正好对应背景、前部、后部。很多新手在读取 PNG 标签时会直接按 RGB 读成三通道,导致标签维度变成(3, H, W),训练时和单通道图像对不上;正确做法是用convert("L")读成单通道灰度图,或者用cv2.imread(path, cv2.IMREAD_GRAYSCALE)

检查类别分布是否正常的代码:

unique, counts = np.unique(mask, return_counts=True) print(dict(zip(unique.tolist(), counts.tolist())))

正常情况下unique里应该同时出现 0、1、2。如果某些切片只有 0 和 1,说明该层刚好不包含后部海马体,这是可以接受的;但如果某个 axis 的所有 mask 都只有 0,那就要检查数据解压是否完整,或者原始的 nii.gz 标签是否在转换时丢失了前景。

另外要注意,标题里的“左右海马体”和 mask 中的“前部/后部”不是同一个维度的划分。原始 3D 标注可能同时包含左右海马体,而导出的 2D mask 把分割目标定义成前部和后部海马体,左右海马体不再用标签区分。如果你确实需要左/右海马体独立 mask,必须回到原始 nii.gz 重新切,或者用连通域分析把预测结果按左右脑拆开。这一点写论文时要说明,否则评审会认为标签定义不一致。

2.3 前景过滤:为什么要丢掉前景不足 3% 的切片

摘要里提到“自动去除了前景区域不足 3% 的数据”。处理逻辑很简单:计算 mask 中类别 1 和类别 2 的像素总数,除以整张 mask 像素数,如果比例小于 0.03 就删除该切片。好处是训练时不会出现大量纯背景帧,模型收敛更快,也能缓解 Dice Loss 里背景占比过高带来的类别不平衡。

不过这步过滤是有代价的。用这类数据集训练出来的模型,在真实 3D 体积上推理时,遇到海马体很小的切片会倾向于预测为背景,因为训练分布里漏掉了那些低前景比例样本。我的经验是:如果后续要拼回 3D 做完整分割,最好自己在原始 nii.gz 上重新切一遍,保留所有切片;或者至少把前景过滤阈值从 0.03 放宽到 0.01,让模型见过更多边界帧。数据集给的是“已经清洗过的版本”,不是原始切片全集,这一点要牢记。

提示:写论文或做实验记录时,数据集章节需要写明过滤比例以及过滤前后的样本数量差异,否则别人在同一个数据集上复现实验会对不上指标。

3. 从标签文件到可视化:show.py 拆解与两个增强改法

作者在项目里提供了show.py,作用是从 images 里随机挑一张图,把原始图像、GT mask 和 GT 叠加图生成到当前目录。对于医学图像分割数据集,这种可视化脚本的价值不只是“看一眼”,它同时是数据质量检查工具。mask 边界是否贴合图像、类别颜色是否可区分、PNG 是否在解压时损坏,都能一眼看出来。

3.1 读取标签文件时的通道与数值检查

先说读取。图像是灰度切片,直接用convert("L");mask 虽然是 PNG,但本质是索引标签,后续要转成long喂给网络,不能做归一化。读取和合法性检查:

from PIL import Image import numpy as np def load_slice(image_path, mask_path): img = np.array(Image.open(image_path).convert("L")) mask = np.array(Image.open(mask_path)) assert img.shape == mask.shape, f"shape mismatch: {img.shape} vs {mask.shape}" assert mask.dtype == np.uint8, f"mask dtype is {mask.dtype}, expected uint8" unique = np.unique(mask) assert set(unique.tolist()).issubset({0, 1, 2}), f"unexpected label: {unique}" return img, mask

这里的断言值得保留。下载类数据集经常出现尺寸不一致、mask 被误存成 RGB、标签数值漂移等问题,用assert在加载阶段就把问题暴露出来,能省下大量训练结束后才发现指标异常的排查时间。mask.dtype检查也很重要,因为有些 PNG 会以 uint16 保存,后面做one_hotlong()转换时可能出现隐式类型问题。

3.2 三栏可视化与半透明叠加

show.py的核心是三栏图:左边是原始图像,中间是 GT mask,右边是 GT 叠加在原始图上的效果。右侧叠加需要把 mask 转成彩色三通道后再做 alpha 融合。用 matplotlib 实现:

import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np def show_slice(image_path, mask_path, save_path="vis.png"): img, mask = load_slice(image_path, mask_path) color_mask = np.zeros((*mask.shape, 3), dtype=np.uint8) color_mask[mask == 1] = [255, 0, 0] # 前部海马体,红色 color_mask[mask == 2] = [0, 255, 0] # 后部海马体,绿色 fig, axes = plt.subplots(1, 3, figsize=(12, 4)) axes[0].imshow(img, cmap="gray") axes[0].set_title("image") axes[1].imshow(mask, cmap="gray") axes[1].set_title("GT") axes[2].imshow(img, cmap="gray") axes[2].imshow(color_mask, alpha=0.5) axes[2].set_title("overlay") for ax in axes: ax.axis("off") plt.tight_layout() plt.savefig(save_path, dpi=150) plt.close()

由于 Pillow 读取的 mask 是(H, W)color_mask直接沿用同一个 shape 再扩展到 3 通道,所以坐标是对齐的。alpha=0.5是叠加透明度,如果想更清楚地看 mask 边界,可以降到 0.3;如果想突出类别分布,可以升到 0.7。这里有两个容易踩的坑:一是color_mask的 dtype 必须是uint8,不能是 float 类型,否则叠加时 matplotlib 可能显示成空白;二是figsize中宽度要足够,三栏图如果宽度太窄,细节会被压缩。

3.3 批量生成对比图,快速发现标注错位

只看一张不够,我一般会把一个 axis 下前几十张图全部生成缩略图,拼成网格。这个操作对检查海马体是否总出现在图像中央、是否有切片发生左右翻转、mask 是否有整体偏移很有帮助。

import os from math import ceil image_dir = "x_images" mask_dir = "x_masks" files = sorted(os.listdir(image_dir))[:30] cols = 6 rows = ceil(len(files) / cols) fig, axes = plt.subplots(rows, cols, figsize=(cols * 3, rows * 3)) for ax, name in zip(axes.ravel(), files): img, mask = load_slice( os.path.join(image_dir, name), os.path.join(mask_dir, name) ) overlay = img.astype(float) overlay[mask == 1] = 255 ax.imshow(overlay, cmap="gray") ax.axis("off") plt.tight_layout() plt.savefig("grid_check.png", dpi=120)

这里把 mask 为 1 的位置直接置成白色,适合快速浏览整批数据的空间分布。如果 mask 为 2 的位置也需要突出,可以再用红色通道画一层,做法和show_slice里的color_mask一样。批量生成图还有一个好处:发现个别样本的亮度范围与其他切片差异特别大时,说明该样本可能来自不同采集序列,训练时需要加入归一化或直方图匹配。

4. 训练自己的海马体分割模型:2D U-Net 数据管线与 Dice Loss

三套切面数据准备好后,下一步就是接入分割网络。虽然原始数据是 3D nii.gz,但作者已经把三个轴切成了 2D PNG,所以可以先训练 2D U-Net,把基线跑出来,再决定要不要上 3D U-Net 做最终版本。这一章给出一个可复现的 PyTorch 训练管线,并且会对不同切面的分辨率差异做说明。

4.1 为什么先用 2D 模型而不是直接上 3D U-Net

3D U-Net 在医学图像分割里依然是金标准,尤其是海马体这种体积小、解剖结构固定的小器官。但 3D 模型对显存、patch 大小和重采样策略非常敏感,一个标准的 3D patch 通常要 96×96×96 或 128×128×128,batch size 稍微大一点,显存就爆了。这份数据集已经把每个 3D 卷拆成三个正交方向的 2D 切片,天然适合先用 2D 模型验证分割思路。常见做法是三个切面分别训练三个独立的 2D U-Net,推理时把三个方向的概率图平均;也可以只选一个切面训练,比如冠状面 y 轴样本数最多,达到 9465 张,单个模型就能训练得比较充分。

选型上,2D U-Net 的输入通道是 1(灰度 MRI),输出通道是 3。如果只把海马体当成前景,可以输出 2 通道;但本数据集的标签为前部海马体和后部海马体,所以 3 通道更合适。网络编码器部分可以用 ResNet 预训练权重,但第一层卷积要改成输入 1 通道,或者在读取数据时把单通道复制成三通道以适配 ImageNet 预训练权重。

4.2 自定义 Dataset:同时消费三个切面

写一个通用的 PyTorch Dataset,让它既能加载x_images,也能加载y_images,只要传入不同目录即可:

import os import torch from torch.utils.data import Dataset from PIL import Image import numpy as np class SliceDataset(Dataset): def __init__(self, image_dir, mask_dir, transform=None): self.names = sorted(os.listdir(image_dir)) self.image_dir = image_dir self.mask_dir = mask_dir self.transform = transform def __len__(self): return len(self.names) def __getitem__(self, idx): name = self.names[idx] img = np.array(Image.open(os.path.join(self.image_dir, name)).convert("L")) mask = np.array(Image.open(os.path.join(self.mask_dir, name))) if self.transform is not None: aug = self.transform(image=img, mask=mask) img, mask = aug["image"], aug["mask"] img = torch.from_numpy(img).float().unsqueeze(0) / 255.0 mask = torch.from_numpy(mask).long() return img, mask

参数说明:image_dirmask_dir分别指向某个切面的 images、masks 目录;transform使用 albumentations 时会返回字典,image是 H×W 数组,mask是 H×W 的标签数组。最后把图像除以 255 归一化到 [0,1],mask 保持原始数值,用long作为交叉熵的 target。注意不要对 mask 做归一化,也不要将 mask 转成 one-hot 再返回,直接在损失函数里用类索引更省内存。

4.3 Dice Loss 和评估指标的计算细节

海马体在切片里通常只占很小一块,直接用交叉熵会出现严重的类别不平衡。标准做法是把交叉熵和 Dice Loss 混合。Dice Loss 每次计算一个类别的 Dice,然后把类别 1 和类别 2 的 loss 平均,背景类别不参与:

import torch import torch.nn.functional as F def dice_loss(pred, target, num_classes=3, smooth=1.0): probs = torch.softmax(pred, dim=1) # (B, C, H, W) target_onehot = F.one_hot(target, num_classes).permute(0, 3, 1, 2).float() loss = 0.0 for c in range(1, num_classes): inter = (probs[:, c] * target_onehot[:, c]).sum() union = probs[:, c].sum() + target_onehot[:, c].sum() + smooth loss += 1.0 - (2.0 * inter + smooth) / union return loss / (num_classes - 1)

这里的smooth用来防止分子分母同时为 0。当某个切片只有背景时,类别 1 和类别 2 的interunion都为 0,加上smooth之后 loss 接近 1,不会出现 NaN。评估时同样应该报告类别 1 和类别 2 的平均 Dice,不报告背景 Dice,因为背景占比过高会虚高整体指标。如果你在验证集上看到整体 Dice 很高,但前部/后部海马体单独看很差,多半就是背景类别参与了平均。

训练循环核心:

device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model = UNet2D(in_channels=1, out_channels=3).to(device) opt = torch.optim.AdamW(model.parameters(), lr=1e-4) for epoch in range(epochs): model.train() train_loss = 0.0 for img, mask in loader: img, mask = img.to(device), mask.to(device) pred = model(img) loss = dice_loss(pred, mask) + 0.5 * F.cross_entropy(pred, mask) opt.zero_grad() loss.backward() opt.step() train_loss += loss.item() print(epoch, train_loss / len(loader))

这里把dice_losscross_entropy相加,交叉熵权重是 0.5,让 Dice Loss 主导。如果训练初期 loss 波动太大,可以把交叉熵的权重改成 1.0,让模型先学会大致类别分布,再逐步提高 Dice Loss 占比。AdamWlr=1e-4在 2D 分割任务里通常比较稳,不需要一开始就用余弦退火,等验证集 Dice 不再上升时再降学习率。

4.4 超参数与不同切面尺寸的匹配

三个切面的图像尺寸不一样,不能直接共用同一个 padding。整理成表:

切面宽×高样本数推荐输入尺寸
x 轴横断面35×51495864×64(pad)
y 轴冠状面35×35946564×64(pad 或 resize)
z 轴矢状面51×35460364×64(pad)

z 轴宽 51、高 35,接近 64 的一半,pad 之后补零信息不多。y 轴本身就是方形,最省事。输入尺寸统一到 64×64 可以保证三个切面的模型结构完全一致,只在数据加载时做 resize 或 padding。我不会为了省事直接 resize,因为海马体本身尺寸小,resize 会破坏解剖比例;用 pad 到 64×64 更安全。做 online 增强时再随机 crop、旋转或水平翻转。

epoch 可以按 100 设置,batch size 在 32 附近。三个切面分别训练三个模型,保存三套权重。推理时对同一张切片做 test time augmentation,例如水平翻转后取两次预测的平均概率,Dice 通常能涨 1 到 2 个百分点。对于海马体这种左右对称结构,翻转增强基本不会引入错误。

5. 跨切面集成与三维重建:把 2D 预测拼回海马体体积

只在一个切面上训练,2D U-Net 很容易在切片方向产生不连续预测。真正可靠的评估不是只看单个切面的 2D Dice,而是把三个切面的结果拼回三维体,再计算三维指标。做法是:每个切面模型输出该方向上的 2D 概率图,然后根据切片序号映射到三维 volume 的对应索引。

5.1 使用切片序号恢复三维索引

hippocampus_164_18.png中的 18 表示原始 3D 卷的第 18 层切片。如果 x、y、z 三个方向的切片序号来自同一个 volume,那么可以直接用三维数组承接:

import numpy as np # 假设原始 3D 数组 shape 为 (nx, ny, nz) vol_pred = np.zeros((nx, ny, nz), dtype=np.float32) # x 方向:每张切片对应 volume[slice_idx, :, :] for slice_idx, pred_2d in enumerate(pred_x_list): vol_pred[slice_idx] = pred_2d # y 方向:volume[:, slice_idx, :] # z 方向:volume[:, :, slice_idx]

这里的pred_2d要保证是概率图或 argmax 之后的类别图,尺寸与原始切面分辨率一致。由于三个方向的宽高不同,需要先把每个方向的预测 resize 回原始 2D 切片的宽高,再做拼接。实际操作时,要从dataset.json里取原始 3D shape,否则只能按各轴已知尺寸组合。不同切面读入的 numpy 数组 shape 顺序不同,这是最容易出错的地方。

5.2 三维 Dice 与 Hausdorff 距离验证

拼接结束后,如果原始 nii.gz 中的 mask 还在,可以用 SimpleITK 计算三维指标:

import SimpleITK as sitk pred_np = vol_pred.astype(np.uint8) true_sitk = sitk.ReadImage("label.nii.gz") pred_sitk = sitk.GetImageFromArray(pred_np) pred_sitk.CopyInformation(true_sitk) dice_filter = sitk.LabelOverlapMeasuresImageFilter() dice_filter.Execute(pred_sitk, true_sitk) print("Dice label1:", dice_filter.GetDiceCoefficient(1)) print("Dice label2:", dice_filter.GetDiceCoefficient(2))

用 SimpleITK 的好处是可以同时拿到表面距离信息,也可以直接算 Hausdorff 距离,不需要自己写距离变换。必须注意,重建后的pred_sitk要和真实标签的体素间距、方向保持一致,所以pred_sitk要调用CopyInformation(true_sitk)。如果两个图像的 origin 或 spacing 不一致,Dice 可能仍然偏高,但 Hausdorff 距离会明显异常。

做三个方向模型融合时,建议把每个方向的概率图先对齐到同一个体素网格,再按平均概率做 argmax。直接在 2D 层面拼接,等到转 3D 时很容易因为切片索引错位导致边界区域出现层间噪声。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询