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

如何在Keras有状态LSTM的fit_generator中设置batch_size>1

How to Use Batch Size > 1 with Stateful LSTM and fit_generator in Keras

Hey there! Let's fix that glacial training speed— using batch size 1 is a total bottleneck, and we can absolutely scale this up while keeping the stateful LSTM behavior intact. The key here is understanding how stateful LSTMs handle batches: each position in the batch corresponds to a continuous, independent sequence, so we need to structure our data and generator to match that rule perfectly.

Core Rules for Stateful LSTM Batches

First, let's recap the non-negotiables you need to follow:

  • Your input layer's batch_shape must explicitly define your target batch size (e.g., (B, n_history, n_cols) where B is your desired batch size).
  • Each batch’s i-th sample must be the direct continuation of the i-th sample from the previous batch. This is how the model maintains its state across batches correctly.
  • After every epoch, you must reset the model’s states— otherwise, the next epoch will start with leftover state from the previous sequence, leading to garbage results.

Step 1: Update the Input Layer

First, set your desired batch size in the input layer (let's use B=32 as an example):

from keras.layers import Input

# Replace B with your target batch size
B = 32
input_layer = Input(
    shape=(n_history, n_cols),
    batch_shape=(B, n_history, n_cols),
    dtype='float32',
    name='daily_input'
)

Step 2: Rewrite the Data Generator

The biggest change is how we generate batches. Instead of yielding single samples, we need to split our time series into B continuous sub-sequences, then pull matching windows from each sub-sequence to form a batch. Here's a working generator:

import numpy as np

def training_data(batch_size, n_history, n_cols):
    total_possible_samples = pdf_daily_data.shape[0] - n_history
    # Ensure we can split data evenly into `batch_size` sequences
    samples_per_sequence = total_possible_samples // batch_size
    
    # Split raw data into `batch_size` continuous sub-sequences
    sequences = []
    for seq_idx in range(batch_size):
        # Calculate start/end indices for each sub-sequence
        start = seq_idx * samples_per_sequence
        # Add n_history to the end so we can extract full windows
        end = start + samples_per_sequence + n_history
        sequences.append(pdf_daily_data[start:end].values)  # Convert to numpy if needed
    
    while True:
        # Yield one batch per time step across all sequences
        for time_step in range(samples_per_sequence):
            x_batch = []
            y_batch = []
            for seq in sequences:
                # Extract window for this time step in the sub-sequence
                x = seq[time_step : time_step + n_history].reshape(1, n_history, n_cols)
                x_batch.append(x)
                # Replace this with your actual y(i) logic— e.g., next time step's value
                y = seq[time_step + n_history]
                y_batch.append(y)
            
            # Combine into batch shape (batch_size, n_history, n_cols)
            x_batch = np.concatenate(x_batch, axis=0)
            y_batch = np.array(y_batch)
            yield (x_batch, y_batch)

Step 3: Adjust Training Code with State Reset

We need to set the correct steps_per_epoch and add a callback to reset states after each epoch:

from keras.callbacks import Callback

class ResetStatesCallback(Callback):
    def on_epoch_end(self, epoch, logs=None):
        self.model.reset_states()

# Initialize your generator with target batch size
train_gen = training_data(batch_size=B, n_history=n_history, n_cols=n_cols)

# Calculate steps per epoch (matches samples_per_sequence from the generator)
steps_per_epoch = (pdf_daily_data.shape[0] - n_history) // B

# Start training
model.fit_generator(
    generator=train_gen,
    steps_per_epoch=steps_per_epoch,
    epochs=20,
    callbacks=[ResetStatesCallback()]
)

Critical Notes to Avoid Bugs

  • Data Alignment: Make sure total_possible_samples is divisible by your batch size. If not, trim your dataset slightly or adjust the batch size to avoid leftover data.
  • Consistent Batch Size: Once you set B for training, you must use the same batch size for prediction with this stateful model.
  • Y Logic: Update the y = seq[time_step + n_history] line to match your actual prediction task (e.g., multi-step output, classification labels).

内容的提问来源于stack exchange,提问作者Jose Antonio Martin H

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.15 03:56:59