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

使用自定义SequenceGenerator训练Keras模型时遇shape属性缺失错误

问题

我尝试用自定义的SequenceGenerator类训练一个简单的tf.keras.Sequential()模型,按官方文档方式将生成器传入model.fit():

model.fit(x=generator, epochs=5)

但持续报错:

AttributeError                            Traceback (most recent call last)
Cell In[35], line 1
----> 1 model.fit(x=generator, epochs=5)

File ~\venv\lib\site-packages\keras\utils\traceback_utils.py:70, in filter_traceback.<locals>.error_handler(*args, **kwargs)
     67     filtered_tb = _process_traceback_frames(e.__traceback__)
     68     # To get the full stack trace, call:
     69     # `tf.debugging.disable_traceback_filtering()`
---> 70     raise e.with_traceback(filtered_tb) from None
     71 finally:
     72     del filtered_tb

File ~venv\lib\site-packages\keras\engine\data_adapter.py:885, in GeneratorDataAdapter.__init__(self, x, y, sample_weights, workers, use_multiprocessing, max_queue_size, model, **kwargs)
    878     except NotImplementedError:
    879         # The above call may fail if the model is a container-like class
    880         # that does not implement its own forward pass (e.g. a GAN or
    881         # VAE where the forward pass is handled by subcomponents).  Such
    882         # a model does not need to be built.
    883         pass
---> 885 self._first_batch_size = int(tf.nest.flatten(peek)[0].shape[0])
    887 def _get_tensor_spec(t):
    888     # TODO(b/226395276): Remove _with_tensor_ranks_only usage.
    889     return type_spec.type_spec_from_value(t)._with_tensor_ranks_only()

AttributeError: 'SequenceGenerator' object has no attribute 'shape'

模型代码如下:

# Create model
model = tf.keras.Sequential()
model.add(tf.keras.layers.LSTM(32,
                               input_shape=(gen_config["num_steps"], 40,),
                               return_sequences=False,
                               stateful=False,
                               ))

model.add(tf.keras.layers.Dense(3, activation='softmax'))

#compile model
model.compile(optimizer='adam',
              loss='sparse_categorical_crossentropy',
              metrics=['accuracy'],
              )

自定义生成器代码:

class SequenceGenerator():

    def __iter__(self): 

        # set seed everytime that generator is restarted
        if self.seed:
            np.random.seed(self.seed)

        # set chunk_index, batch_index
        self.chunk_index = 0
        self.batch_index = 0

        # set chunk_indexer
        self._chunk_indexer_setup()  # self._chunk_indexer_reset()
        self.no_total_chunks = len(self._chunk_indexer)
        self._no_features = self.pipeline_function.no_features

        # Start iteration
        if len(self._chunk_indexer) > 0:

            # Infinite iteration
            while True:

                # Fill pool of chunks if
                # (1) x_ts is not yet initialized
                # (2) self.x_ts is not enough to feed two batches
                if self.x_ts is None or len(self.x_ts) < self.batch_size * 2:
                    chunk_pool = self.get_chunk_pool()

                    if len(chunk_pool) > 0:
                        x, y, flag = [np.vstack(array) for array in zip(*chunk_pool)]
                        x_ts, y_ts = self._get_timeseries(x, y, flag)

                        # 1st iteration of generator
                        if self.x_ts is None:
                            self.x_ts, self.y_ts = x_ts, y_ts
                        # Iterations after 1st, i.e. observations exists: just append
                        else:
                            self.x_ts = np.concatenate((self.x_ts, x_ts))
                            self.y_ts = np.concatenate((self.y_ts, y_ts))

                # if x_ts contains at least one batch, yield batch
                # if generator is exhausted, yield empty batch
                if len(self.x_ts) >= self.batch_size or self.exhausted:
                    yield self.x_ts[:self.batch_size], self.y_ts[:self.batch_size]

生成器每次输出包含两个numpy数组的元组,x形状为(32, 100, 40),y形状为(32, 1),不清楚model.fit()要求生成器具备何种shape属性。

解决方案

报错核心原因是自定义SequenceGenerator未遵循Keras对数据输入的规范。Keras的fit方法仅能正确识别两种自定义输入:生成器函数,或是继承自tf.keras.utils.Sequence的类(需实现__len__和__getitem__方法)。你的类只是普通可迭代类,导致Keras无法推断批次信息,从而抛出错误。

方案一:继承tf.keras.utils.Sequence类(推荐)

这是Keras官方推荐的自定义数据生成方式,线程安全且能被框架正确识别:

import tensorflow as tf
import numpy as np

class SequenceGenerator(tf.keras.utils.Sequence):
    def __init__(self, batch_size, gen_config, pipeline_function, seed=None):
        self.batch_size = batch_size
        self.gen_config = gen_config
        self.pipeline_function = pipeline_function
        self.seed = seed
        self.x_ts = None
        self.y_ts = None
        self.exhausted = False
        self._chunk_indexer_setup()
        self.no_total_chunks = len(self._chunk_indexer)
        self._no_features = self.pipeline_function.no_features
        self._fill_data_pool()

    def _chunk_indexer_setup(self):
        # 替换为你的_chunk_indexer初始化逻辑
        self._chunk_indexer = []

    def get_chunk_pool(self):
        # 替换为你的获取chunk池逻辑
        return []

    def _get_timeseries(self, x, y, flag):
        # 替换为你的时序数据转换逻辑
        return np.zeros((self.batch_size, self.gen_config["num_steps"], 40)), np.zeros((self.batch_size, 1))

    def _fill_data_pool(self):
        while self.x_ts is None or len(self.x_ts) < self.batch_size * 2:
            chunk_pool = self.get_chunk_pool()
            if not chunk_pool:
                self.exhausted = True
                break
            x, y, flag = [np.vstack(array) for array in zip(*chunk_pool)]
            x_ts, y_ts = self._get_timeseries(x, y, flag)
            if self.x_ts is None:
                self.x_ts, self.y_ts = x_ts, y_ts
            else:
                self.x_ts = np.concatenate((self.x_ts, x_ts))
                self.y_ts = np.concatenate((self.y_ts, y_ts))

    def __len__(self):
        # 返回每个epoch的批次总数
        if self.exhausted:
            return len(self.x_ts) // self.batch_size
        # 若为无限生成,需定义固定批次数量,或根据数据总量计算
        return 100

    def __getitem__(self, idx):
        # 返回指定索引的批次数据
        start = idx * self.batch_size
        end = start + self.batch_size
        if end > len(self.x_ts) and not self.exhausted:
            self._fill_data_pool()
        batch_x = self.x_ts[start:end]
        batch_y = self.y_ts[start:end]
        if len(batch_x) < self.batch_size:
            self.exhausted = True
        return batch_x, batch_y

    def on_epoch_end(self):
        # 每个epoch结束时重置种子或打乱数据
        if self.seed:
            np.random.seed(self.seed)
        self._fill_data_pool()

使用方式:

generator = SequenceGenerator(batch_size=32, gen_config=gen_config, pipeline_function=pipeline_function)
model.fit(generator, epochs=5)

方案二:改为生成器函数

若不想继承Sequence类,可将类改写为生成器函数:

import numpy as np

def sequence_generator(batch_size, gen_config, pipeline_function, seed=None):
    if seed:
        np.random.seed(seed)
    chunk_index = 0
    batch_index = 0
    _chunk_indexer = []
    x_ts = None
    y_ts = None
    exhausted = False

    def _chunk_indexer_setup():
        nonlocal _chunk_indexer
        # 替换为你的_chunk_indexer初始化逻辑
        _chunk_indexer = []

    def get_chunk_pool():
        # 替换为你的获取chunk池逻辑
        return []

    def _get_timeseries(x, y, flag):
        # 替换为你的时序数据转换逻辑
        return np.zeros((batch_size, gen_config["num_steps"], 40)), np.zeros((batch_size, 1))

    _chunk_indexer_setup()
    no_total_chunks = len(_chunk_indexer)
    _no_features = pipeline_function.no_features

    if len(_chunk_indexer) > 0:
        while True:
            if x_ts is None or len(x_ts) < batch_size * 2:
                chunk_pool = get_chunk_pool()
                if len(chunk_pool) > 0:
                    x, y, flag = [np.vstack(array) for array in zip(*chunk_pool)]
                    x_ts_batch, y_ts_batch = _get_timeseries(x, y, flag)
                    if x_ts is None:
                        x_ts, y_ts = x_ts_batch, y_ts_batch
                    else:
                        x_ts = np.concatenate((x_ts, x_ts_batch))
                        y_ts = np.concatenate((y_ts, y_ts_batch))
                else:
                    exhausted = True
            if len(x_ts) >= batch_size or exhausted:
                yield x_ts[:batch_size], y_ts[:batch_size]
                x_ts = x_ts[batch_size:]
                y_ts = y_ts[batch_size:]
                if exhausted and len(x_ts) == 0:
                    break

使用方式:

generator = sequence_generator(batch_size=32, gen_config=gen_config, pipeline_function=pipeline_function)
model.fit(generator, epochs=5, steps_per_epoch=100)  # 必须指定steps_per_epoch

关键说明

  • 继承Sequence类时,__len__方法告知Keras每个epoch的批次数量,__getitem__返回对应批次的数据,框架可自动处理epoch结束逻辑。
  • 使用生成器函数时,必须指定steps_per_epoch参数,否则Keras无法判断一个epoch何时结束。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.01 19:40:34