PyTorch LSTM / GRU
Recurrent Neural Networks (RNNs) face the vanishing gradient problem when processing sequential data, making it difficult to learn long-range dependencies.
Long Short-Term Memory (LSTM)andGated Recurrent Unit (GRU)By introducing gating mechanisms, they solve this problem and are core models for processing sequence tasks such as time series, natural language, and speech.
1. Limitations of RNN and Gating Mechanisms
A standard RNN combines the current input with the previous hidden state at each time step to compute a new hidden state:
This structure has two core problems:
Vanishing gradient: During backpropagation, the gradient is multiplied step by step by the weight matrix. When the sequence is long, the gradient decays exponentially, causing the parameters at early time steps to be barely updated, and the model cannot learn long-range dependencies.
Exploding gradient: When the largest eigenvalue of the weight matrix is greater than 1, the gradient increases exponentially during backpropagation, making training unstable (usually mitigated by gradient clipping).
The core idea of gating mechanisms is to introduce learnable "switches" that let the network autonomously decide: at the current time step, which information should be remembered, which should be forgotten, and which new information should be written into memory.
LSTM uses three gates (forget gate, input gate, output gate) plus a separate cell state; GRU simplifies the structure to two gates (reset gate, update gate), with fewer parameters and faster training.
2. LSTM Principle
2.1 Core Structure and Three Gates
LSTM maintains two state vectors passed between time steps:
- Cell State\(c_t\): the carrier of long-term memory, in which information can flow almost losslessly
- Hidden State\(h_t\): short-term memory, also the output of the current time step
All three gates are linear transformations activated by Sigmoid, with output values between 0 and 1, acting as "valves":
遗忘门(Forget Gate):决定从细胞状态中丢弃哪些信息 输入门(Input Gate):决定将哪些新信息写入细胞状态 输出门(Output Gate):决定基于细胞状态输出什么
2.2 Forward Computation Formulas
where \(\odot\) denotes element-wise multiplication (Hadamard product), and \(\sigma\) denotes the Sigmoid function.
f_t ⊙ c_{t-1}Interpretation of the computation logic:i_t ⊙ g_t: The forget gate decides how much historical memory to retain; when close to 0 it forgets, when close to 1 it retainsg_t: The input gate decides how much new information to write in,o_t ⊙ tanh(c_t)is the candidate new content
3. GRU Principle
3.1 Core Structure and Two Gates
3.1 Core Structure and Two GatesGRU merges LSTM's forget gate and input gate into theupdate gate
重置门(Reset Gate):决定忽略多少历史状态来计算候选隐藏状态 更新门(Update Gate):决定保留多少历史状态,写入多少新状态
3.2 Forward Computation Formulas
Interpretation of the computation logic:
- Reset gate
r_tWhen close to 0, the candidate stateh̃_tbarely depends on history, equivalent to starting over - Update gate
z_tWhen close to 1, the new state adopts more of the candidate value; when close to 0, it retains more of the historical state - GRU has no separate cell state, and its parameter count is about 75% of that of LSTM
4. LSTM in PyTorch
This section details the parameters, input/output shapes, and hidden state initialization methods of nn.LSTM.
4.1 Detailed Explanation of nn.LSTM Parameters
Example
import torch.nn as nn
lstm = nn.LSTM(
input_size=64, # Dimension of the input vector at each time step
hidden_size=128, # Dimension of the hidden state (and cell state)
num_layers=2, # Number of stacked layers, default is 1
bias=True, # Whether to use bias terms, default True
batch_first=False, # Whether batch is in the first dimension of input/output shape, default False
dropout=0.0, # Dropout probability between layers (only effective when num_layers > 1)
bidirectional=False, # Whether to use bidirectional LSTM, default False
proj_size=0, # Projection layer dimension (LSTM with projection), default 0 means not used
)
# View parameter count
total_params = sum(p.numel() for p in lstm.parameters())
print(f"LSTM parameter count: {total_params:,}")
# When input_size=64, hidden_size=128, num_layers=2, it is approximately 197,632
Parameter count estimation formula (single-layer unidirectional):
4.2 Shapes of Input and Output
This is the most error-prone part of using LSTM, so special attention is required.batch_firstThe impact of parameters.
Example
import torch.nn as nn
# ── batch_first=False (default) ────────────────────────
lstm = nn.LSTM(input_size=32, hidden_size=64, batch_first=False)
# Input shape: (seq_len, batch_size, input_size)
seq_len, batch_size, input_size = 10, 4, 32
x = torch.randn(seq_len, batch_size, input_size)
output, (h_n, c_n) = lstm(x)
print(f"output shape: {output.shape}")
# torch.Size([10, 4, 64]) → (seq_len, batch_size, hidden_size)
# Hidden state output at each time step
print(f"h_n shape: {h_n.shape}")
# torch.Size([1, 4, 64]) → (num_layers * num_directions, batch_size, hidden_size)
# Hidden state at the last time step
print(f"c_n shape: {c_n.shape}")
# torch.Size([1, 4, 64]) → same as h_n, cell state at the last time step
# ── batch_first=True (recommended, more intuitive) ────────
lstm_bf = nn.LSTM(input_size=32, hidden_size=64, batch_first=True)
# Input shape: (batch_size, seq_len, input_size)
x = torch.randn(batch_size, seq_len, input_size)
output, (h_n, c_n) = lstm_bf(x)
print(f"output shape: {output.shape}")
# torch.Size([4, 10, 64]) → (batch_size, seq_len, hidden_size)
print(f"h_n shape: {h_n.shape}")
# torch.Size([1, 4, 64]) → (num_layers, batch_size, hidden_size)
# Note: the shape of h_n is not affected by batch_first
# ── Output shape of multi-layer bidirectional LSTM ──────────────────────
lstm_bd = nn.LSTM(input_size=32, hidden_size=64,
num_layers=3, bidirectional=True, batch_first=True)
x = torch.randn(batch_size, seq_len, input_size)
output, (h_n, c_n) = lstm_bd(x)
print(f"output shape: {output.shape}")
# torch.Size([4, 10, 128])
# hidden_size × 2 = 128, because bidirectional concatenation
print(f"h_n shape: {h_n.shape}")
# torch.Size([6, 4, 64])
# num_layers × num_directions = 3 × 2 = 6
Summary of output shapes:
| Variables | batch_first=False | batch_first=True |
|---|---|---|
output |
(seq_len, N, H * D) |
(N, seq_len, H * D) |
h_n |
(L * D, N, H) |
(L * D, N, H) |
c_n |
(L * D, N, H) |
(L * D, N, H) |
\(N\) = batch_size, \(H\) = hidden_size, \(L\) = num_layers, \(D\) = 2 (bidirectional) or 1 (unidirectional)
4.3 Initialization of Hidden State
Example
import torch.nn as nn
lstm = nn.LSTM(input_size=32, hidden_size=64, num_layers=2, batch_first=True)
batch_size = 8
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
lstm = lstm.to(device)
# Method 1: do not pass an initial state; PyTorch automatically uses zero initialization
x = torch.randn(batch_size, 10, 32).to(device)
output, (h_n, c_n) = lstm(x)
# Method 2: manually initialize to zeros (equivalent to Method 1, but explicitly shows the intent)
num_layers, num_directions = 2, 1
h_0 = torch.zeros(num_layers * num_directions, batch_size, 64).to(device)
c_0 = torch.zeros(num_layers * num_directions, batch_size, 64).to(device)
output, (h_n, c_n) = lstm(x, (h_0, c_0))
# Method 3: stateful mode — pass state across batches
# Suitable for slicing ultra-long sequences (language models, long-text generation, etc.)
# Need to call detach() to cut off the computation graph from the previous batch to prevent GPU memory leakage
h, c = h_0, c_0
for batch_x in data_loader:
batch_x = batch_x.to(device)
output, (h, c) = lstm(batch_x, (h, c))
h = h.detach() # Cut off gradients, keep only the values
c = c.detach()
# Method 4: use Xavier or normal distribution initialization (converges faster in some scenarios)
def init_hidden(lstm_module, batch_size, device):
num_layers = lstm_module.num_layers
hidden_size = lstm_module.hidden_size
directions = 2 if lstm_module.bidirectional else 1
h = torch.zeros(num_layers * directions, batch_size, hidden_size, device=device)
c = torch.zeros(num_layers * directions, batch_size, hidden_size, device=device)
nn.init.orthogonal_(h) # Orthogonal initialization helps stabilize training
return h, c
5. GRU in PyTorch
The GRU interface is almost identical to LSTM; the main difference is that there is no cell state.c。
5.1 Detailed Explanation of nn.GRU Parameters
Example
gru = nn.GRU(
input_size=64,
hidden_size=128,
num_layers=2,
bias=True,
batch_first=True, # Recommended to set to True
dropout=0.3, # Dropout between layers
bidirectional=False,
)
5.2 Basic Usage Example
Example
import torch.nn as nn
gru = nn.GRU(input_size=32, hidden_size=64, batch_first=True)
batch_size, seq_len = 8, 10
x = torch.randn(batch_size, seq_len, 32)
# GRU only returns output and h_n, no c_n
output, h_n = gru(x)
print(f"output shape: {output.shape}")
# torch.Size([8, 10, 64]) → (batch_size, seq_len, hidden_size)
print(f"h_n shape: {h_n.shape}")
# torch.Size([1, 8, 64]) → (num_layers, batch_size, hidden_size)
# Take the output at the last time step (for tasks such as classification)
last_hidden = output[:, -1, :] # (batch_size, hidden_size)
# Or equivalently:
last_hidden = h_n.squeeze(0) # (batch_size, hidden_size)
6. Variants of LSTM and GRU
This section introduces bidirectional, multi-layer stacked, and multi-layer structures with Dropout.
6.1 Bidirectional LSTM / GRU
Bidirectional models process information from both the forward and backward directions of the sequence simultaneously. The output at each time step contains past and future context, making them suitable for tasks that require global context, such as text classification and named entity recognition.
Example
import torch.nn as nn
# Bidirectional LSTM
bilstm = nn.LSTM(
input_size=32,
hidden_size=64,
num_layers=2,
batch_first=True,
bidirectional=True, # Enable bidirectional
)
x = torch.randn(8, 10, 32)
output, (h_n, c_n) = bilstm(x)
print(f"output shape: {output.shape}")
# torch.Size([8, 10, 128])
# Forward 64-dimensional + backward 64-dimensional = 128-dimensional
print(f"h_n shape: {h_n.shape}")
# torch.Size([4, 8, 64])
# num_layers(2) × num_directions(2) = 4
# Separate the final hidden states of the forward and backward directions
# h_n ordering: [forward layer0, backward layer0, forward layer1, backward layer1]
h_forward = h_n[-2, :, :] # (batch_size, hidden_size) last forward layer
h_backward = h_n[-1, :, :] # (batch_size, hidden_size) last backward layer
h_combined = torch.cat([h_forward, h_backward], dim=-1) # (batch_size, 128)
# The usage of bidirectional GRU is exactly the same
bigru = nn.GRU(input_size=32, hidden_size=64,
num_layers=2, batch_first=True, bidirectional=True)
output, h_n = bigru(x)
6.2 Multi-layer Stacking
Example
import torch.nn as nn
# 3-layer stacked LSTM
deep_lstm = nn.LSTM(
input_size=32,
hidden_size=128,
num_layers=3, # Stack 3 layers
batch_first=True,
)
x = torch.randn(8, 20, 32)
output, (h_n, c_n) = deep_lstm(x)
print(f"output shape: {output.shape}")
# torch.Size([8, 20, 128]) only contains the output of the top layer
print(f"h_n shape: {h_n.shape}")
# torch.Size([3, 8, 128]) contains the final hidden state of each layer
# Get the final hidden state of each layer
h_layer1 = h_n[0] # (batch_size, hidden_size) first layer
h_layer2 = h_n[1] # (batch_size, hidden_size) second layer
h_layer3 = h_n[2] # (batch_size, hidden_size) third layer (top layer)
6.3 Multi-layer Structure with Dropout
nn.LSTMThe built-indropoutparameter only applies tobetween layers, not to the output of the last layer. If you need to add Dropout after the last layer as well, you must add it manually:
Example
import torch.nn as nn
class StackedLSTM(nn.Module):
"""
Multi-layer LSTM + inter-layer Dropout + output Dropout
"""
def __init__(self, input_size, hidden_size, num_layers,
num_classes, dropout=0.3):
super().__init__()
self.lstm = nn.LSTM(
input_size=input_size,
hidden_size=hidden_size,
num_layers=num_layers,
batch_first=True,
dropout=dropout if num_layers > 1 else 0.0,
# When num_layers=1, dropout is invalid; set it to 0 to avoid warnings
)
self.dropout = nn.Dropout(dropout) # Dropout after the last layer
self.fc = nn.Linear(hidden_size, num_classes)
def forward(self, x):
output, (h_n, c_n) = self.lstm(x)
# Take the hidden state at the last time step
last_output = output[:, -1, :] # (batch_size, hidden_size)
last_output = self.dropout(last_output)
return self.fc(last_output)
model = StackedLSTM(input_size=64, hidden_size=128,
num_layers=3, num_classes=5, dropout=0.3)
x = torch.randn(16, 20, 64)
print(model(x).shape) # torch.Size([16, 5])
7. Handling Variable-Length Sequences
In real tasks, sequences within the same batch usually have different lengths. PyTorch providespack_padded_sequenceandpad_packed_sequenceto handle this problem, avoiding LSTM from performing invalid computations on padding positions.
7.1 pack_padded_sequence
Example
import torch.nn as nn
from torch.nn.utils.rnn import pack_padded_sequence, pad_packed_sequence, pad_sequence
# Simulate variable-length sequences in a batch (already sorted from longest to shortest)
seq1 = torch.randn(5, 32) # Sequence length 5
seq2 = torch.randn(3, 32) # Sequence length 3
seq3 = torch.randn(2, 32) # Sequence length 2
# pad_sequence automatically pads with zeros to align to the longest sequence
# When batch_first=True, the output shape is (batch_size, max_seq_len, input_size)
padded = pad_sequence([seq1, seq2, seq3], batch_first=True, padding_value=0.0)
lengths = torch.tensor([5, 3, 2])
print(f"Shape after padding: {padded.shape}") # torch.Size([3, 5, 32])
# pack_padded_sequence: compress padding, tell LSTM the real lengths
packed = pack_padded_sequence(
padded,
lengths=lengths,
batch_first=True,
enforce_sorted=True, # Sequences must be sorted in descending order of length
# enforce_sorted=False # Allows any order (internally sorted automatically), recommended to set to False
)
print(type(packed)) # <class 'torch.nn.utils.rnn.PackedSequence'>
7.2 pad_packed_sequence
Example
# Pass the PackedSequence into the LSTM
packed_output, (h_n, c_n) = lstm(packed)
# pad_packed_sequence: restore to a padded tensor
output, output_lengths = pad_packed_sequence(packed_output, batch_first=True)
print(f"Restored output shape: {output.shape}")
# torch.Size([3, 5, 64]) → (batch_size, max_seq_len, hidden_size)
# The output at padding positions is 0
print(f"Actual lengths of each sequence: {output_lengths}")
# tensor([5, 3, 2])
7.3 Complete Variable-Length Sequence Processing Pipeline
Example
import torch.nn as nn
from torch.nn.utils.rnn import pack_padded_sequence, pad_packed_sequence
class LSTMClassifier(nn.Module):
def __init__(self, vocab_size, embed_dim, hidden_size, num_classes, padding_idx=0):
super().__init__()
self.embedding = nn.Embedding(vocab_size, embed_dim, padding_idx=padding_idx)
self.lstm = nn.LSTM(embed_dim, hidden_size, batch_first=True)
self.fc = nn.Linear(hidden_size, num_classes)
def forward(self, x, lengths):
# x: (batch_size, max_seq_len) — word indices
embedded = self.embedding(x) # (batch_size, max_seq_len, embed_dim)
# Pack
packed = pack_padded_sequence(
embedded, lengths.cpu(), batch_first=True, enforce_sorted=False
)
# LSTM forward pass (skip padding positions)
packed_output, (h_n, c_n) = self.lstm(packed)
# Method 1: use the h_n of the last layer as the sequence representation
last_hidden = h_n.squeeze(0) # (batch_size, hidden_size)
# Method 2: After unpacking, take the output at the true last position of each sequence (equivalent to Method 1)
# output, _ = pad_packed_sequence(packed_output, batch_first=True)
# last_hidden = output[range(len(lengths)), lengths - 1, :]
return self.fc(last_hidden)
model = LSTMClassifier(vocab_size=10000, embed_dim=128,
hidden_size=256, num_classes=5)
# Simulate a batch
batch_tokens = torch.randint(1, 10000, (8, 30)) # (batch_size=8, max_len=30)
batch_lengths = torch.randint(5, 31, (8,)) # True length of each sequence
output = model(batch_tokens, batch_lengths)
print(output.shape) # torch.Size([8, 5])
8. Complete Practical: Text Sentiment Classification
Taking IMDB movie review sentiment binary classification as an example, demonstrate the complete pipeline of bidirectional LSTM processing text:
Example
import torch.nn as nn
import torch.optim as optim
from torch.utils.data import Dataset, DataLoader
from torch.nn.utils.rnn import pack_padded_sequence, pad_sequence
from torch.optim.lr_scheduler import ReduceLROnPlateau
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# ── Model Definition ──────────────────────────────────────
class BiLSTMClassifier(nn.Module):
"""
Bidirectional LSTM Text Classification Model
Architecture: Embedding -> BiLSTM -> Dropout -> FC
"""
def __init__(self, vocab_size, embed_dim, hidden_size,
num_layers, num_classes, dropout=0.5, padding_idx=0):
super().__init__()
self.embedding = nn.Embedding(
vocab_size, embed_dim, padding_idx=padding_idx
)
# When using pretrained word vectors:
# self.embedding = nn.Embedding.from_pretrained(pretrained_vectors)
self.lstm = nn.LSTM(
input_size=embed_dim,
hidden_size=hidden_size,
num_layers=num_layers,
batch_first=True,
dropout=dropout if num_layers > 1 else 0.0,
bidirectional=True,
)
self.dropout = nn.Dropout(dropout)
# Bidirectional LSTM: Concatenate forward and backward last hidden states
self.fc = nn.Linear(hidden_size * 2, num_classes)
def forward(self, x, lengths):
# x: (batch_size, max_seq_len)
embedded = self.dropout(self.embedding(x))
packed = pack_padded_sequence(
embedded, lengths.cpu(), batch_first=True, enforce_sorted=False
)
packed_output, (h_n, c_n) = self.lstm(packed)
# h_n: (num_layers * 2, batch_size, hidden_size)
# Take the forward and backward hidden states of the last layer and concatenate them
h_forward = h_n[-2, :, :] # (batch_size, hidden_size)
h_backward = h_n[-1, :, :] # (batch_size, hidden_size)
h_combined = torch.cat([h_forward, h_backward], dim=-1)
# (batch_size, hidden_size * 2)
out = self.dropout(h_combined)
return self.fc(out)
# ── Custom Dataset ────────────────────────────────
class TextDataset(Dataset):
def __init__(self, texts, labels, vocab, max_len=200):
self.data = texts
self.labels = labels
self.vocab = vocab
self.max_len = max_len
def __len__(self):
return len(self.data)
def __getitem__(self, idx):
tokens = self.data[idx][:self.max_len]
ids = [self.vocab.get(t, 1) for t in tokens] # 1 = <UNK>
return torch.tensor(ids, dtype=torch.long), torch.tensor(self.labels[idx])
def collate_fn(batch):
"""Custom collate: Pad variable-length sequences with zeros, record true lengths"""
sequences, labels = zip(*batch)
lengths = torch.tensor([len(s) for s in sequences])
padded = pad_sequence(sequences, batch_first=True, padding_value=0)
labels = torch.stack(labels)
return padded, lengths, labels
# ── Training and Evaluation Functions ────────────────────────────────
def train_epoch(model, loader, optimizer, criterion):
model.train()
total_loss, correct = 0.0, 0
for texts, lengths, labels in loader:
texts, labels = texts.to(device), labels.to(device)
optimizer.zero_grad()
outputs = model(texts, lengths)
loss = criterion(outputs, labels)
loss.backward()
# Gradient clipping: prevent gradient explosion in RNN
nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
total_loss += loss.item() * len(labels)
correct += (outputs.argmax(1) == labels).sum().item()
n = len(loader.dataset)
return total_loss / n, correct / n
def eval_epoch(model, loader, criterion):
model.eval()
total_loss, correct = 0.0, 0
with torch.no_grad():
for texts, lengths, labels in loader:
texts, labels = texts.to(device), labels.to(device)
outputs = model(texts, lengths)
loss = criterion(outputs, labels)
total_loss += loss.item() * len(labels)
correct += (outputs.argmax(1) == labels).sum().item()
n = len(loader.dataset)
return total_loss / n, correct / n
# ── Initialization and Training ──────────────────────────────────
VOCAB_SIZE = 50000
EMBED_DIM = 128
HIDDEN_SIZE = 256
NUM_LAYERS = 2
NUM_CLASSES = 2
DROPOUT = 0.5
EPOCHS = 15
model = BiLSTMClassifier(VOCAB_SIZE, EMBED_DIM, HIDDEN_SIZE,
NUM_LAYERS, NUM_CLASSES, DROPOUT).to(device)
criterion = nn.CrossEntropyLoss()
optimizer = optim.Adam(model.parameters(), lr=1e-3)
scheduler = ReduceLROnPlateau(optimizer, mode="max",
patience=3, factor=0.5)
best_acc = 0.0
for epoch in range(1, EPOCHS + 1):
train_loss, train_acc = train_epoch(model, train_loader, optimizer, criterion)
val_loss, val_acc = eval_epoch(model, val_loader, criterion)
scheduler.step(val_acc)
print(f"Epoch {epoch:2d}/{EPOCHS} | "
f"Train Loss: {train_loss:.4f}, Acc: {train_acc:.4f} | "
f"Val Loss: {val_loss:.4f}, Acc: {val_acc:.4f}")
if val_acc > best_acc:
best_acc = val_acc
torch.save(model.state_dict(), "best_bilstm.pth")
print(f" -> Saving best model, Val Acc: {best_acc:.4f}")
9. Complete Practical: Time Series Prediction
Taking multi-step time series prediction as an example, use LSTM to predict values for the next N steps:
Example
import torch.nn as nn
import torch.optim as optim
import numpy as np
from torch.utils.data import Dataset, DataLoader
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
# ── Sliding Window Dataset ──────────────────────────────
class TimeSeriesDataset(Dataset):
"""
Sliding window segmentation of time series
Input: past input_len steps
Target: future output_len steps
"""
def __init__(self, series, input_len, output_len):
self.data = torch.tensor(series, dtype=torch.float32)
self.input_len = input_len
self.output_len = output_len
def __len__(self):
return len(self.data) - self.input_len - self.output_len + 1
def __getitem__(self, idx):
x = self.data[idx : idx + self.input_len]
y = self.data[idx + self.input_len : idx + self.input_len + self.output_len]
return x.unsqueeze(-1), y # x: (input_len, 1), y: (output_len,)
# ── Model Definition ──────────────────────────────────────
class LSTMForecaster(nn.Module):
"""
Multi-step time series prediction model
Architecture: LSTM -> Dropout -> FC
"""
def __init__(self, input_size, hidden_size, num_layers,
output_len, dropout=0.2):
super().__init__()
self.lstm = nn.LSTM(
input_size=input_size,
hidden_size=hidden_size,
num_layers=num_layers,
batch_first=True,
dropout=dropout if num_layers > 1 else 0.0,
)
self.dropout = nn.Dropout(dropout)
self.fc = nn.Linear(hidden_size, output_len)
def forward(self, x):
# x: (batch_size, input_len, input_size)
output, (h_n, c_n) = self.lstm(x)
# Take the output of the last time step
last = output[:, -1, :] # (batch_size, hidden_size)
last = self.dropout(last)
return self.fc(last) # (batch_size, output_len)
# ── Data Preparation (using sine wave as example)──────────────────────
t = np.linspace(0, 200, 10000)
series = np.sin(t) + 0.1 * np.random.randn(len(t))
INPUT_LEN = 60 # Use past 60 steps
OUTPUT_LEN = 10 # Predict next 10 steps
split = int(len(series) * 0.8)
train_data = series[:split]
val_data = series[split:]
train_dataset = TimeSeriesDataset(train_data, INPUT_LEN, OUTPUT_LEN)
val_dataset = TimeSeriesDataset(val_data, INPUT_LEN, OUTPUT_LEN)
train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True)
val_loader = DataLoader(val_dataset, batch_size=64, shuffle=False)
# ── Training ──────────────────────────────────────────
model = LSTMForecaster(input_size=1, hidden_size=128,
num_layers=2, output_len=OUTPUT_LEN).to(device)
criterion = nn.MSELoss()
optimizer = optim.Adam(model.parameters(), lr=1e-3)
def train_epoch_ts(model, loader, optimizer, criterion):
model.train()
total_loss = 0.0
for x, y in loader:
x, y = x.to(device), y.to(device)
optimizer.zero_grad()
pred = model(x)
loss = criterion(pred, y)
loss.backward()
nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()
total_loss += loss.item() * x.size(0)
return total_loss / len(loader.dataset)
def eval_epoch_ts(model, loader, criterion):
model.eval()
total_loss = 0.0
with torch.no_grad():
for x, y in loader:
x, y = x.to(device), y.to(device)
pred = model(x)
total_loss += criterion(pred, y).item() * x.size(0)
return total_loss / len(loader.dataset)
for epoch in range(1, 31):
train_loss = train_epoch_ts(model, train_loader, optimizer, criterion)
val_loss = eval_epoch_ts(model, val_loader, criterion)
print(f"Epoch {epoch:2d}/30 | Train MSE: {train_loss:.6f} | Val MSE: {val_loss:.6f}")
# ── Multi-step recursive prediction (another strategy)────────────────────
def recursive_forecast(model, init_sequence, steps, device):
"""
Recursive prediction: Predict one step at a time, append the predicted value to the sequence, then predict the next step
Suitable for single-step prediction models with output_len=1
"""
model.eval()
sequence = list(init_sequence)
predictions = []
with torch.inference_mode():
for _ in range(steps):
x = torch.tensor(sequence[-INPUT_LEN:], dtype=torch.float32)
x = x.unsqueeze(0).unsqueeze(-1).to(device) # (1, input_len, 1)
pred = model(x).item()
predictions.append(pred)
sequence.append(pred)
return predictions
10. Comparison and Selection of LSTM and GRU
This section compares the structural differences between LSTM and GRU, and provides selection recommendations.
Structural Comparison
| Comparison Item | LSTM | GRU |
|---|---|---|
| Number of Gates | 3 (Forget, Input, Output) | 2 (Reset, Update) |
| State Vectors | Cell state + Hidden state | Hidden state only |
| Parameter count (same hidden_size) | Baseline | About 75% |
| Training Speed | Slower | Faster |
| Long Sequence Performance | Usually better | Comparable |
| Short Sequence Performance | Similar | Similar |
| Implementation Complexity | Higher | Lower |
Selection Recommendations
数据量小、训练资源有限
-> 优先选 GRU,参数少,不容易过拟合,训练快
序列较长(> 100 步)、长距离依赖重要
-> 优先选 LSTM,细胞状态更擅长保留远期信息
需要快速实验和基线对比
-> 先用 GRU,效果差再换 LSTM
任务对准确率要求极高,有充足数据
-> 两者都试,结合交叉验证选择
2020 年后的新项目
-> 考虑 Transformer 架构(BERT、GPT),在大数据量下通常优于 LSTM/GRU
LSTM/GRU 仍在边缘设备、低延迟推理、数据量较小的场景中有优势
Quick Reference for Common Questions
| Problem | Cause | Solution |
|---|---|---|
| Training Loss does not decrease | Learning rate too high or gradient explosion | Reduce learning rate; add gradient clippingclip_grad_norm_ |
| Validation Loss much higher than training loss | Overfitting | Increase Dropout; reduce number of layers or hidden_size |
| Poor prediction at sequence tail | Vanishing gradient | Increase number of layers; use bidirectional structure; consider attention mechanism |
| Excessive GPU memory usage | Sequence too long or batch too large | Reduce seq_len or batch_size; use gradient checkpointing |
| Abnormal Loss in stateful training | States passed across batches not detached | Call after each batchh.detach_() |
| Bidirectional LSTM concatenation dimension error | h_n index misunderstanding | useh_n[-2](forward) andh_n[-1](backward) |
| PackedSequence error | Sequence lengths not sorted in descending order | Setenforce_sorted=False |
| batch_first confusion | Forgot to set consistently | Recommended to use throughoutbatch_first=True |