如何在Keras有状态LSTM的fit_generator中设置batch_size>1
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_shapemust explicitly define your target batch size (e.g.,(B, n_history, n_cols)whereBis your desired batch size). - Each batch’s
i-thsample must be the direct continuation of thei-thsample 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_samplesis 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
Bfor 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

