ACDC数据集入门指南:从数据加载到模型训练的全流程解析

1次阅读
没有评论

共计 3716 个字符,预计需要花费 10 分钟才能阅读完成。

image.webp

背景介绍

ACDC(Automatic Cardiac Diagnosis Challenge)数据集是医学图像分析领域常用的心脏 MRI 数据集,包含来自不同患者的短轴心脏 MRI 序列。该数据集主要用于左心室分割和心脏功能评估任务,具有以下特点:

ACDC 数据集入门指南:从数据加载到模型训练的全流程解析

  • 包含 100 例患者的心肌 MRI 数据(70 例训练 /30 例测试)
  • 每例患者提供舒张末期(ED)和收缩末期(ES)时相的图像
  • 标注包含左心室(LV)、右心室(RV)和心肌(Myo)三个区域
  • 图像分辨率为 1.37×1.37mm,切片厚度 5 -8mm

数据加载实战

下面演示如何使用 Python 加载 ACDC 数据集的基本结构:

import os
import nibabel as nib  # 用于读取医学图像格式
import numpy as np

# 假设数据集存放在以下路径
DATA_DIR = './ACDC_dataset/training/'

# 获取患者目录列表
patient_dirs = sorted([d for d in os.listdir(DATA_DIR) if d.startswith('patient')])

# 示例:加载第一个患者的 ED 时相图像和标注
def load_patient_data(patient_id):
    patient_path = os.path.join(DATA_DIR, patient_dirs[patient_id])

    # 加载图像数据(4D numpy 数组:x,y,z,time)img_path = os.path.join(patient_path, f'{patient_dirs[patient_id]}_4d.nii.gz')
    img = nib.load(img_path).get_fdata()

    # 加载标注数据(3D numpy 数组:x,y,z)label_path = os.path.join(patient_path, f'{patient_dirs[patient_id]}_gt.nii.gz')
    label = nib.load(label_path).get_fdata()

    # 提取 ED 和 ES 时相(数据集中第 0 和最后 1 个时相)ed_img = img[..., 0]
    es_img = img[..., -1]

    # 返回图像和对应的标注
    return {'ED': (ed_img, label == 1),  # 左心室
        'ES': (es_img, label == 3),  # 右心室
        'patient_id': patient_dirs[patient_id]
    }

# 示例:加载第一个患者的数据
patient_data = load_patient_data(0)
print(f"患者 ID: {patient_data['patient_id']}")
print(f"ED 图像形状: {patient_data['ED'][0].shape}")
print(f"ES 标注形状: {patient_data['ES'][1].shape}")

数据预处理技巧

医学图像预处理对模型性能至关重要,以下是关键步骤:

  1. 归一化处理
  2. 将像素值缩放到 [0,1] 或标准化(z-score)
  3. 处理不同扫描仪之间的强度差异

  4. 重采样

  5. 将所有样本重采样到相同分辨率
  6. 使用线性插值处理图像,最近邻插值处理标注

  7. ROI 裁剪

  8. 根据标注框裁剪心脏区域
  9. 减少计算量并移除无关区域

  10. 数据增强

  11. 随机旋转(-15°~15°)
  12. 随机缩放(0.9~1.1 倍)
  13. 弹性变形(模拟心脏运动)

示例预处理代码:

import torch
from skimage.transform import resize

def preprocess(data, target_size=(128, 128)):
    img, mask = data

    # 归一化(按病例为单位)img = (img - img.min()) / (img.max() - img.min())

    # 重采样
    img = resize(img, target_size, order=1, preserve_range=True)  # 线性插值
    mask = resize(mask, target_size, order=0, preserve_range=True)  # 最近邻插值

    # 转换为 PyTorch 张量并添加通道维度
    img_tensor = torch.FloatTensor(img).unsqueeze(0)  # [1, H, W]
    mask_tensor = torch.FloatTensor(mask).unsqueeze(0)  # [1, H, W]

    return img_tensor, mask_tensor

PyTorch 数据集集成

创建自定义 Dataset 类将数据集成到训练流程:

from torch.utils.data import Dataset, DataLoader

class ACDCDataset(Dataset):
    def __init__(self, patient_ids, transform=None):
        self.patient_ids = patient_ids
        self.transform = transform

    def __len__(self):
        return len(self.patient_ids) * 2  # 每个患者有 ED 和 ES 两个时相

    def __getitem__(self, idx):
        patient_idx = idx // 2
        phase = 'ED' if idx % 2 == 0 else 'ES'

        # 加载原始数据
        data = load_patient_data(patient_idx)
        img, mask = data[phase]

        # 预处理
        img, mask = preprocess((img, mask))

        # 数据增强
        if self.transform:
            img, mask = self.transform((img, mask))

        return img, mask

# 使用示例
train_dataset = ACDCDataset(patient_ids=range(60))  # 前 60 例患者
val_dataset = ACDCDataset(patient_ids=range(60, 70))  # 后 10 例患者

train_loader = DataLoader(train_dataset, batch_size=8, shuffle=True)
val_loader = DataLoader(val_dataset, batch_size=4, shuffle=False)

性能优化策略

  1. 多进程数据加载

    DataLoader(..., num_workers=4, pin_memory=True)

  2. 预加载缓存

  3. 将预处理后的数据缓存到内存或 SSD
  4. 使用 torch.utils.data.TensorDataset 加速后续读取

  5. 混合精度训练

    scaler = torch.cuda.amp.GradScaler()
    
    with torch.cuda.amp.autocast():
        outputs = model(inputs)
        loss = criterion(outputs, labels)
    
    scaler.scale(loss).backward()
    scaler.step(optimizer)
    scaler.update()

常见问题与解决方案

  1. 内存不足错误
  2. 降低批量大小
  3. 使用梯度累积(accumulate gradients)
  4. 启用 pin_memory 加速 CPU 到 GPU 传输

  5. 类别不平衡

  6. 使用带权重的损失函数

    criterion = torch.nn.BCEWithLogitsLoss(pos_weight=torch.tensor([5.0])  # 正样本权重
    )

  7. 过拟合

  8. 增加 Dropout 层
  9. 使用更激进的数据增强
  10. 添加 L2 正则化

可视化示例

使用 matplotlib 进行简单可视化:

import matplotlib.pyplot as plt

# 显示一个批次的数据
def show_batch(batch, n_display=4):
    images, masks = batch
    fig, axes = plt.subplots(n_display, 2, figsize=(10, n_display*3))

    for i in range(n_display):
        axes[i,0].imshow(images[i][0], cmap='gray')
        axes[i,0].set_title('MRI Slice')
        axes[i,1].imshow(masks[i][0], cmap='jet')
        axes[i,1].set_title('Segmentation')

    plt.tight_layout()
    plt.show()

# 可视化第一个批次
for batch in train_loader:
    show_batch(batch)
    break

思考题

  1. 如何扩展本方法处理 3D 体积数据而非单一切片?
  2. 当训练数据非常有限时(如只有几十例),哪些策略可以提高模型泛化能力?
  3. 医学图像分割常用的评价指标有哪些?如何实现 Dice 系数的 PyTorch 版本?

总结

本文详细介绍了 ACDC 数据集从加载到训练的全流程。关键在于:理解医学图像的特殊性、合理设计预处理流程、高效的数据加载实现。建议读者先在小数据集上验证流程正确性,再扩展到完整训练。后续可尝试 3D 卷积网络、注意力机制等先进方法提升分割精度。

正文完
 0
评论(没有评论)