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

使用Keras自定义DataGenerator时出现AttributeError: 'DataGenerator' object has no attribute 'shape'错误的解决方法

解决AttributeError: 'DataGenerator' object has no attribute 'shape'的问题

这个错误的根源在于你的自定义DataGenerator类里存在两个关键问题,导致生成的数据不符合Keras模型的预期格式,我们一步步来修复:

问题1:未初始化self.list_IDs变量

在你的__init__方法里,你只定义了self.file_list,但在__getitem__里却使用了self.list_IDs——这个变量根本没被初始化,这会直接导致索引获取逻辑出错。

问题2:__data_generation没有生成批量数据

当前的__data_generation函数里,你每次循环都会用单个样本覆盖X和y,最后返回的只是最后一个样本的数组,而不是包含batch_size个样本的批量数据。Keras的模型期望输入是(batch_size, ...)形状的张量,但你返回的是单个样本的形状,这就触发了shape相关的错误。


修正后的完整DataGenerator代码

import numpy as np
import os
import keras

class DataGenerator(keras.utils.Sequence):
    def __init__(self, file_list, batch_size):
        """Constructor can be expanded, with batch size, dimentation etc. """
        self.file_list = file_list
        self.batch_size = batch_size
        # 初始化list_IDs,对应传入的file_list
        self.list_IDs = self.file_list
        self.on_epoch_end()

    def __len__(self):
        'Take all batches in each iteration'
        return int(np.floor(len(self.file_list) / self.batch_size))

    def __getitem__(self, index):
        'Generate one batch of data'
        # Generate indexes of the batch
        indexes = self.indexes[index*self.batch_size:(index+1)*self.batch_size]
        # Find list of IDs
        list_IDs_temp = [self.list_IDs[k] for k in indexes]
        # Generate data
        X, y = self.__data_generation(list_IDs_temp)
        return X, y

    def on_epoch_end(self):
        'Updates indexes after each epoch'
        self.indexes = np.arange(len(self.file_list))
        # 可选:训练时打乱索引,提升模型泛化能力
        # np.random.shuffle(self.indexes)

    def __data_generation(self, file_list_temp):
        'Generates data containing batch_size samples'
        loc = os.path.abspath('.')
        # 先获取单个样本的形状,用来初始化批量数组
        sample_x = np.load(os.path.join(loc, "target", file_list_temp[0]))
        sample_y = np.load(os.path.join(loc, "newlabel", file_list_temp[0]))
        
        # 初始化批量数据数组:(batch_size, ...样本形状...)
        X = np.empty((self.batch_size, *sample_x.shape))
        y = np.empty((self.batch_size, *sample_y.shape))

        # 遍历批量内的每个样本,填充到数组中
        for i, ID in enumerate(file_list_temp):
            x_file_path = os.path.join(loc, "target", ID)
            y_file_path = os.path.join(loc, "newlabel", ID)
            # Store sample
            X[i] = np.load(x_file_path)
            # Store class
            y[i] = np.load(y_file_path)

        return X, y

额外的优化建议

  1. 替换fit_generator为fit:Keras在较新版本中已经弃用了fit_generator,直接使用model.fit()就可以接收Sequence类型的生成器,代码更简洁:
H = model.fit(
    training_generator,
    validation_data=validation_generator,
    epochs=EPOCHS
)
  1. 处理样本数量不能被batch_size整除的情况:如果你的数据集长度不是batch_size的整数倍,当前__len__方法会丢弃最后不足一个batch的样本。如果不想丢弃,可以修改为:
def __len__(self):
    return int(np.ceil(len(self.file_list) / self.batch_size))

同时在__data_generation里处理实际的样本数量(比如len(file_list_temp)可能小于batch_size),避免数组越界。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.28 16:42:53