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

TensorFlow用Keras自定义DataGenerator训练CNN维度不匹配报错如何解决

问题修复方案

核心错误原因

你遇到的维度报错由多个逻辑错误共同导致,和模型输入要求的(None, 224, 224, 3)格式不匹配:

  1. 通道数不符:模型要求3通道RGB输入,但你将图像转为了单通道灰度图
  2. 维度处理错误:维度扩展的轴位错误,且最终返回时将输入数组套入了额外列表
  3. 样本读取逻辑错误:取数据集行时用了循环序号而非样本id,导致数据读取错位
  4. 输入数组初始化维度缺失:初始化X时未包含通道维度

修正后代码

def __getitem__(self, index):
    index = self.index[index * self.batch_size:(index + 1) * self.batch_size]
    batch = [self.indices[k] for k in index]
    X, y = self.__data_generation(batch)
    return X, y

def __data_generation(self, batch):
    df = self.df
    # 初始化X时明确包含3通道维度,确保shape为(batch_size,224,224,3)
    X = np.empty((self.batch_size, 224, 224, 3), dtype='float32')
    y = np.empty((self.batch_size), dtype=int)
    for i, sample_id in enumerate(batch):
        # 修正:用样本id取对应行,而非循环序号i
        doc_row = df.loc[sample_id]
        path = str(doc_row['PATH'])
        path = os.path.join(dataset_path,path)
        typologie = str(doc_row['TYPOLOGIE'])
        img_i = convert_from_path(path)[0]

        # 图像转换逻辑修正
        img_i = img_i.resize((224,224), Image.ANTIALIAS)
        # 去掉灰度转换,直接转RGB格式确保3通道
        img_i = img_i.convert('RGB')
        array_image_i = np.array(img_i, dtype='float32')
        # 无需额外扩展维度,array_image_i本身shape为(224,224,3)
        X[i,] = array_image_i
        y[i] = self.map_classes[typologie]
    # 修正:去掉额外的列表包裹,直接返回numpy数组
    Y = tf.keras.utils.to_categorical(y, num_classes = self.num_classes)
    return X, Y

额外说明

如果你的业务场景必须使用灰度图,可以保留灰度转换逻辑,将单通道复制3次适配模型输入即可:

gray_img_i = ImageOps.grayscale(img_i)
array_image_i = np.array(gray_img_i,dtype='float32')
# 单通道复制3次,shape从(224,224)转为(224,224,3)
array_image_i = np.repeat(array_image_i[..., np.newaxis], 3, axis=-1)

修改完成后可以在__data_generation返回前打印X.shape,确认输出为(batch_size, 224, 224, 3)即可正常传入模型训练。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 12:15:04