使用自定义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
相关产品推荐
相关产品推荐

