使用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
额外的优化建议
- 替换
fit_generator为fit:Keras在较新版本中已经弃用了fit_generator,直接使用model.fit()就可以接收Sequence类型的生成器,代码更简洁:
H = model.fit( training_generator, validation_data=validation_generator, epochs=EPOCHS )
- 处理样本数量不能被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
相关产品推荐
相关产品推荐

