You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于Train/Val/Test子文件夹MRI影像训练CNN实现三类患者分类方案咨询

适配分层MRI数据集的Keras CNN实现方案

1 数据集结构适配

你的数据集是按「拆分集→受试者目录→单张MRI」的三层结构存储,标签绑定到受试者而非单张影像,无法直接用默认的flow_from_directory接口,可按以下步骤适配:

  • 第一步:构建受试者元数据表
    遍历Train、Val、Test三个根目录,提取每个受试者的所属拆分、分类标签、目录下所有MRI路径,存入结构化表,示例代码如下:
    import os
    import pandas as pd
    
    root_dir = "你的数据集根目录"
    splits = ["Train", "Val", "Test"]
    label_map = {"1":0, "2":1, "3":2} # 对应三类的标签映射,可根据你的目录命名规则调整
    
    meta = []
    for split in splits:
        split_path = os.path.join(root_dir, split)
        for label in os.listdir(split_path): # 若标签未放在二级目录,可从受试者其他属性文件读取
            label_path = os.path.join(split_path, label)
            for subject_id in os.listdir(label_path):
                subject_path = os.path.join(label_path, subject_id)
                mri_paths = [os.path.join(subject_path, p) for p in os.listdir(subject_path) if p.endswith((".png",".jpg",".nii.gz"))]
                meta.append({
                    "split": split,
                    "subject_id": subject_id,
                    "label": label_map[label],
                    "mri_paths": mri_paths
                })
    meta_df = pd.DataFrame(meta)
    
  • 第二步:自定义Keras数据生成器
    继承tf.keras.utils.Sequence实现自定义生成器,训练阶段按单张MRI采样,自动继承受试者标签,同时可避免同一受试者的影像出现在同一个batch降低过拟合风险,核心逻辑示例:
    import tensorflow as tf
    import numpy as np
    from tensorflow.keras.preprocessing import image
    
    class MRIDataGenerator(tf.keras.utils.Sequence):
        def __init__(self, meta_df, batch_size=32, img_size=(224,224), is_train=True):
            self.meta_df = meta_df
            self.batch_size = batch_size
            self.img_size = img_size
            self.is_train = is_train
            # 展开所有单张MRI的路径和对应标签
            self.all_samples = []
            for _, row in meta_df.iterrows():
                for path in row["mri_paths"]:
                    self.all_samples.append((path, row["label"]))
            if is_train:
                np.random.shuffle(self.all_samples)
        
        def __len__(self):
            return int(np.ceil(len(self.all_samples) / self.batch_size))
        
        def __getitem__(self, idx):
            batch_samples = self.all_samples[idx*self.batch_size : (idx+1)*self.batch_size]
            X = np.zeros((len(batch_samples), *self.img_size, 1)) # 单通道MRI,三通道输入可改为3
            y = np.zeros(len(batch_samples))
            for i, (path, label) in enumerate(batch_samples):
                img = image.load_img(path, target_size=self.img_size, color_mode="grayscale")
                X[i] = image.img_to_array(img) / 255.0 # 归一化逻辑可根据你的预处理规则调整
                y[i] = label
            return X, tf.keras.utils.to_categorical(y, num_classes=3)
        
        def on_epoch_end(self):
            if self.is_train:
                np.random.shuffle(self.all_samples)
    
    训练前分别实例化训练、验证、测试生成器即可,后续训练逻辑和常规图像分类任务完全一致。

2 受试者级批量预测实现

要一次性调用同一位受试者的所有MRI输出最终分类结果,按以下逻辑实现即可:

  • 先把单个受试者的所有90张MRI做和训练阶段一致的预处理,打包成批量张量输入模型
  • 得到所有单张影像的预测概率后,做聚合得到受试者级的最终结果,常用聚合方法包括均值概率取最高类别、多数投票两种,示例代码:
    def predict_subject(model, subject_mri_paths, img_size=(224,224), agg_method="mean"):
        # 加载所有MRI并完成预处理
        X = np.zeros((len(subject_mri_paths), *img_size, 1))
        for i, path in enumerate(subject_mri_paths):
            img = image.load_img(path, target_size=img_size, color_mode="grayscale")
            X[i] = image.img_to_array(img) / 255.0
        # 批量预测单张MRI结果
        pred_probs = model.predict(X, verbose=0)
        # 聚合得到受试者级分类结果
        if agg_method == "mean":
            mean_prob = pred_probs.mean(axis=0)
            return mean_prob.argmax() + 1 # 转回1/2/3的标签格式
        elif agg_method == "vote":
            pred_classes = pred_probs.argmax(axis=1)
            vote_count = np.bincount(pred_classes)
            return vote_count.argmax() + 1
    
    调用时只需传入对应受试者的MRI路径列表,即可直接得到该受试者的分类结果。

可选优化方案

如果想要进一步提升准确率,可以修改模型结构,在特征提取backbone后增加注意力聚合层,直接让模型学习同一受试者不同MRI切片的权重,自动完成聚合,不需要后续人工规则判断。

内容的提问来源于stack exchange,提问作者Chanda Steven

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.09.27 17:36:07