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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.29 16:28:03