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

自定义DataGenerator返回BatchDataset适配Keras模型训练求助

时间序列数据生成器适配Keras训练的解决方案

问题背景

我正在构建一个DataGenerator,用于加载存储在文件中的DataFrame数据,通过tf.keras.utils.timeseries_dataset_from_array转换为时间序列以训练模型。由于全量数据量极大,tf.keras.utils.timeseries_dataset_from_array会将数据全部加载到内存,因此我将数据拆分为多个部分存储,想通过生成器配合GPU训练模型。

当前实现的代码如下:

class DataGenerator(tf.keras.utils.Sequence):
    'Generates data for Keras'
    def __init__(self, df_index, train_mean, train_std, seq_len=128, batch_size=32, shuffle=False):
        'Initialization'
        self.df_index   = df_index
        self.data_index = list(df_index.index)
        self.train_mean = train_mean
        self.train_std  = train_std
        self.seq_len    = seq_len
        self.batch_size = batch_size
        self.shuffle    = shuffle

    def __len__(self):
        'Denotes the number of batches per epoch'
        return int((self.df_index.length.cumsum().iloc[-1] - seq_len) // batch_size)

    def __getitem__(self, ds_index):
        'Generate one batch of data'
        # Generate indexes of the batch
        # Generate data
        file_id = self.data_index[ds_index]
        data_obj = joblib.load(f"{data_path}{data_prefix}_{file_id}.pkl")
        X_norm = (data_obj.data_var - self.train_mean) / self.train_std
        data_ds = tf.keras.utils.timeseries_dataset_from_array(X_norm[:-self.seq_len],
                                                       data_obj.targets[self.seq_len:],
                                                       batch_size=self.batch_size,
                                                       sequence_length=self.seq_len,
                                                       shuffle=self.shuffle)
        return data_ds

    def on_epoch_end(self):
        'Updates indexes after each epoch'
        random.shuffle(self.data_index)

遇到的问题

__getitem__方法返回的是BatchDataset对象,但模型的fit方法期望直接接收批次数据(即特征数组和标签数组的元组)。我考虑通过循环遍历Dataset来输出批次,但不知如何实现(需考虑遍历完一个文件后加载下一个文件并生成新的Dataset)。


解决方案

方案1:修改Sequence生成器,返回直接可用的批次数据

调整生成器逻辑,提前建立全局批次到文件及文件内批次的映射,让每个__getitem__调用返回单个批次的特征和标签数组:

import random
import joblib
import tensorflow as tf
import pandas as pd

class DataGenerator(tf.keras.utils.Sequence):
    'Generates data for Keras'
    def __init__(self, df_index, train_mean, train_std, seq_len=128, batch_size=32, shuffle=False):
        'Initialization'
        self.df_index   = df_index
        self.train_mean = train_mean
        self.train_std  = train_std
        self.seq_len    = seq_len
        self.batch_size = batch_size
        self.shuffle    = shuffle

        # 预计算全局批次与文件、文件内批次的映射
        self.batch_mapping = []
        for file_id in df_index.index:
            file_total_len = df_index.loc[file_id, 'length']
            # 计算单个文件可生成的批次数量
            file_batch_num = max(0, (file_total_len - seq_len) // batch_size)
            if file_batch_num > 0:
                self.batch_mapping.extend([(file_id, inner_idx) for inner_idx in range(file_batch_num)])
        
        self.current_mapping = self.batch_mapping.copy()
        if self.shuffle:
            random.shuffle(self.current_mapping)

    def __len__(self):
        'Denotes the number of batches per epoch'
        return len(self.current_mapping)

    def __getitem__(self, batch_idx):
        'Generate one batch of data'
        file_id, inner_batch_idx = self.current_mapping[batch_idx]
        # 加载目标文件数据
        data_obj = joblib.load(f"{data_path}{data_prefix}_{file_id}.pkl")
        X_norm = (data_obj.data_var - self.train_mean) / self.train_std
        
        # 生成当前文件的时间序列Dataset
        data_ds = tf.keras.utils.timeseries_dataset_from_array(
            X_norm[:-self.seq_len],
            data_obj.targets[self.seq_len:],
            batch_size=self.batch_size,
            sequence_length=self.seq_len,
            shuffle=False  # 全局洗牌在on_epoch_end处理
        )
        
        # 遍历到目标批次并返回numpy数组
        for idx, (X_batch, y_batch) in enumerate(data_ds):
            if idx == inner_batch_idx:
                return X_batch.numpy(), y_batch.numpy()
        
        return None, None

    def on_epoch_end(self):
        'Updates indexes after each epoch'
        if self.shuffle:
            random.shuffle(self.current_mapping)

方案2:改用TensorFlow Dataset API构建完整数据管道

直接利用tf.data.Dataset的异步加载和预处理能力,更适配GPU训练场景:

import tensorflow as tf
import joblib

def load_and_process_file(file_id, train_mean, train_std, seq_len, batch_size):
    # 加载文件数据
    data_obj = joblib.load(f"{data_path}{data_prefix}_{file_id}.pkl")
    X_norm = (data_obj.data_var - train_mean) / train_std
    # 生成时间序列批次Dataset
    return tf.keras.utils.timeseries_dataset_from_array(
        X_norm[:-seq_len],
        data_obj.targets[seq_len:],
        batch_size=batch_size,
        sequence_length=seq_len
    )

def create_time_series_dataset(df_index, train_mean, train_std, seq_len=128, batch_size=32, shuffle=True):
    # 创建文件ID的基础Dataset
    file_ids = list(df_index.index)
    ds = tf.data.Dataset.from_tensor_slices(file_ids)
    
    # 并行加载文件并展开批次
    ds = ds.interleave(
        lambda file_id: tf.data.Dataset.from_generator(
            lambda: load_and_process_file(file_id.numpy().decode(), train_mean, train_std, seq_len, batch_size),
            output_signature=(
                tf.TensorSpec(shape=(None, seq_len, X_norm.shape[-1]), dtype=tf.float32),
                tf.TensorSpec(shape=(None,), dtype=tf.float32)  # 根据你的标签类型调整shape和dtype
            )
        ),
        num_parallel_calls=tf.data.AUTOTUNE
    )
    
    # 全局洗牌(可选)
    if shuffle:
        ds = ds.shuffle(buffer_size=100)
    
    # 预取数据,配合GPU异步计算
    ds = ds.prefetch(tf.data.AUTOTUNE)
    return ds

使用时直接将生成的Dataset传入模型训练:

train_ds = create_time_series_dataset(df_index, train_mean, train_std)
model.fit(train_ds, epochs=10)

方案对比

  • 方案1保留了Keras Sequence的使用习惯,适合对生成器模式熟悉的场景,但每次__getitem__需要加载文件并遍历到目标批次,存在一定性能开销。
  • 方案2基于TensorFlow原生Dataset API,支持异步加载、并行预处理和预取,能更高效地利用GPU资源,是大规模数据训练的推荐方案。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.22 11:04:58