Keras自定义Sequence类传入fit_generator报找不到数据适配器错误
报错产生原因
- 核心原因1:
Sequence基类导入来源不匹配。Keras的数据适配器会严格校验传入的生成器是否属于当前运行环境下的keras.utils.Sequence子类,如果你混用了独立安装的Keras和TensorFlow内置的tf.keras,比如从独立keras导入Sequence,却用tf.keras的模型训练,适配器就无法识别你的自定义类,直接抛出适配失败错误。 - 核心原因2:自定义
__getitem__方法存在硬编码逻辑缺陷。你在reshape操作时强制指定第一维为self.batch_size,但__len__方法是按向上取整计算batch总数,最后一个batch的样本量大概率小于设定的batch_size,此时reshape操作会触发维度不匹配错误,导致方法返回异常值(甚至None类型的标签),对应报错信息里的<class 'NoneType'>。 - 附加原因:你使用的
fit_generator在TensorFlow 2.1及以上版本已经被废弃,旧API的适配逻辑存在兼容问题,会进一步提升报错概率。
修复方案
按以下步骤调整即可解决问题:
- 统一所有Keras相关模块的导入来源,全部使用TensorFlow内置的tf.keras模块,不要混用独立安装的keras包。
Sequence基类必须从tensorflow.keras.utils导入。 - 修正
__getitem__中的硬编码reshape逻辑,用当前batch的实际样本量替代固定的self.batch_size,兼容最后一个样本不足量的batch。 - 废弃
fit_generator调用,直接使用model.fit传入Sequence实例即可,新版fit原生支持Sequence生成器输入。
修复后的完整可运行代码如下:
import numpy as np import pandas as pd import tensorflow as tf # 必须从tensorflow.keras导入Sequence,保证类型匹配 from tensorflow.keras.utils import Sequence class Sequence_train(Sequence): def __init__(self, x_data, y_data, batch_size=10, n_variables=8): self.x_data = x_data self.y_data = y_data self.batch_size = batch_size self.n_variables = n_variables def __len__(self): return int(np.ceil(len(self.x_data) / self.batch_size)) def __getitem__(self, item): batch_x = self.x_data.iloc[item*self.batch_size : (item+1)*self.batch_size, :] batch_y = self.y_data[item*self.batch_size : (item+1)*self.batch_size] k = [] for idx in batch_x.index: # 过滤NaN值 k.append(list(filter(lambda v: v == v, batch_x.loc[idx]))) j = tf.keras.preprocessing.sequence.pad_sequences(k, padding='post', value=0) # 用当前batch实际样本数做reshape,避免最后一批量不足时报错 actual_batch_len = len(batch_x) j = j.reshape([actual_batch_len, int(j.shape[-1] / self.n_variables), self.n_variables]) return np.asarray(j), np.asarray(batch_y)
训练调用代码调整为:
# 实例化训练数据生成器,注意确认y_data传入正确,不要传None train_generator = Sequence_train(x_data=train_x_df, y_data=train_y, batch_size=10) # 直接用fit传入生成器,无需使用废弃的fit_generator model.fit(train_generator, epochs=50, verbose=2)
校验提示:实例化Sequence_train时需要确认传入的y_data是有效的标签数组/矩阵,不要传None,否则也会触发NoneType相关的适配错误。
内容的提问来源于stack exchange,提问作者Ana Maria
相关产品推荐
相关产品推荐

