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

Keras自定义数据生成器训练3DCNN时输入张量不匹配报错

解决Keras 3D CNN训练时输入张量不匹配的问题

错误原因分析

你遇到的ValueError核心问题是错误地将生成器输出的单批次数据直接传给model.fit(),而非传入生成器本身。当调用next(iter(train_data_generator))时,你得到的是拆分后的16个独立图像张量,但你的3D CNN模型期望接收一个形状为(batch_size, 208, 150, 50, 1)的单一输入张量,两者维度逻辑不匹配。

同时需要确认你的CustomDataGenerator实现是否合规:__getitem__方法必须返回(batch_images, batch_labels)结构,其中batch_images的形状必须严格对应模型输入的(batch_size, width, height, depth, 1)。

修正后的代码

1. 核心训练代码修正

移除手动获取批次的逻辑,直接将生成器传入fit:

# 初始化自定义数据生成器
train_data_generator = CustomDataGenerator(
    batch_size = 16, 
    dataset_directory = "NIFTI_train_codegenerator"
)

epochs = 100
# 直接传入生成器,无需手动拆分批次
model.fit(
    train_data_generator,
    epochs=epochs,
    shuffle=True,
    verbose=2,
    callbacks=[checkpoint_cb, early_stopping_cb],
)

2. 确保CustomDataGenerator的正确性(关键实现要点)

生成器必须继承keras.utils.Sequence,并在__getitem__中正确构造批次数据,示例参考:

import numpy as np
import nibabel as nib
from tensorflow import keras

class CustomDataGenerator(keras.utils.Sequence):
    def __init__(self, batch_size, dataset_directory):
        self.batch_size = batch_size
        # 替换为你的实际逻辑:读取所有NIfTI文件路径
        self.data_paths = self._load_all_data_paths(dataset_directory)
        # 替换为你的实际逻辑:读取对应标签
        self.labels = self._load_all_labels(dataset_directory)

    def _load_all_data_paths(self, dir_path):
        # 实现读取目录下所有NIfTI文件路径的逻辑
        pass

    def _load_all_labels(self, dir_path):
        # 实现读取对应标签的逻辑
        pass

    def __len__(self):
        # 返回训练的总批次数
        return len(self.data_paths) // self.batch_size

    def __getitem__(self, idx):
        # 获取当前批次的文件路径和标签
        batch_paths = self.data_paths[idx*self.batch_size : (idx+1)*self.batch_size]
        batch_labels = self.labels[idx*self.batch_size : (idx+1)*self.batch_size]
        
        # 初始化批次图像数组,严格匹配模型输入维度
        batch_images = np.zeros((self.batch_size, 208, 150, 50, 1), dtype=np.float32)
        
        for i, path in enumerate(batch_paths):
            # 读取NIfTI文件数据
            img_data = nib.load(path).get_fdata()
            # 确保图像尺寸与模型输入一致,不一致则添加resize/裁剪逻辑
            # 增加通道维度,适配模型的输入格式
            batch_images[i] = np.expand_dims(img_data, axis=-1)
        
        return batch_images, np.array(batch_labels)

额外注意事项

  • 确认NIfTI图像的实际尺寸是否与模型定义的(208,150,50)一致,若存在差异,需在生成器中加入尺寸调整逻辑(如使用skimage.transform.resize)。
  • 继承keras.utils.Sequence而非普通迭代器,能更好地支持Keras的多线程训练、epoch间数据洗牌等功能。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 17:30:30