Keras自定义Sequence类DataGenerator批次乱序、调用次数超量问题求助
问题根因
- 多worker预取机制导致调用次数超标、顺序混乱:你配置了
workers=6,即使use_multiprocessing=False,Keras也会启动多线程预取批次,会提前调用__getitem__加载后续批次,导致调用次数高于预期,且多线程调度下批次返回顺序不固定。 - Sequence实现存在多处错误:
__init__中调用on_epoch_end()但未初始化self.indexes属性,存在隐患__init__已经将整个CSV读入self.lines,但__getitem__中又重复调用pandas.read_csv读文件,冗余IO且容易出现行偏移错误__getitem__调用__data_generation时传入的参数和方法定义不匹配(未定义的self.folder_name、缺少config_file参数)- 未用到
self.shuffle对应的索引逻辑,即使后续开启shuffle也不会生效
- 调用了废弃API:
model.fit_generator在TensorFlow 2.x已经被废弃,内部适配逻辑可能存在兼容性问题。
解决步骤
- 调整训练配置,关闭多worker预取:将
workers设为1,新增max_queue_size=1,禁用预取机制,保证批次按顺序串行调用。 - 修正Sequence实现错误,删除冗余IO,补全缺失属性,统一用内存中已加载的
self.lines做批次切分,避免重复读文件的偏移错误。 - 替换
fit_generator为标准model.fit,适配TensorFlow 2.x的官方推荐用法。
修正后的示例代码
import csv import numpy as np import pandas as pd from tensorflow import keras class DataGenerator(keras.utils.Sequence): def __init__(self, file_name, rows_per_batch=50, shuffle=False, config_file=None): self.rows_per_batch = rows_per_batch self.shuffle = shuffle self.file_name = file_name self.config_file = config_file # 补全缺失的配置参数 # 一次性读入所有行(如果文件过大不适合全量读入,可以替换为记录行偏移量的方案,避免全量加载内存) with open(self.file_name, 'r') as f: reader = csv.reader(f) self.lines = list(reader) # 初始化索引 self.indexes = np.arange(len(self.lines)) self.on_epoch_end() def __len__(self): # 保证所有样本都被用到可以改成 int(np.ceil(len(self.lines) / self.rows_per_batch)) return len(self.lines) // self.rows_per_batch def __getitem__(self, index): # 取当前批次对应的索引范围 batch_indexes = self.indexes[index*self.rows_per_batch : (index+1)*self.rows_per_batch] # 直接从内存的lines里取对应行,不需要重复读文件 batch_lines = [self.lines[k] for k in batch_indexes] X, y = self.__data_generation(batch_lines) return X, y def on_epoch_end(self): if self.shuffle == True: np.random.shuffle(self.indexes) def __data_generation(self, batch_lines): # 直接处理传入的批次行,不需要再读CSV df = pd.DataFrame(batch_lines) # <这里保留你原有数据处理逻辑> X = df.values[:, :-1] y = df.values[:, -1] return X.astype(np.float32), y.astype(np.float32) # 主程序 if __name__ == "__main__": file_name = "your_data.csv" rows_per_batch = 50 config_file = "your_config.yaml" # 按你的实际需求传入 train_gen = DataGenerator(file_name, rows_per_batch, shuffle=False, config_file=config_file) val_gen = DataGenerator(file_name, rows_per_batch, shuffle=False, config_file=config_file) # 替换fit_generator为fit,调整worker和队列配置保证顺序执行 model.fit( x=train_gen, validation_data=val_gen, use_multiprocessing=False, shuffle=False, epochs=10, workers=1, max_queue_size=1 # 关闭预取,保证调用次数和顺序符合预期 )
补充说明
如果你的CSV文件过大无法全量加载到内存,不要在__init__中全量读入所有行,可以替换为提前统计文件行总数、记录每行的字节偏移量的方案,__getitem__时直接根据偏移量读取对应行,既避免全量内存占用,也不会出现重复读文件的行偏移错误。
内容的提问来源于stack exchange,提问作者python_beginner
相关产品推荐
相关产品推荐

