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:

\[h_t = \tanh(W_h \cdot h_{t-1} + W_x \cdot x_t + b)\]

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

\[ \begin{aligned} \text{Input:} & x_t \text{ (current time step input)}, h_{t-1} \text{ (previous hidden state)}, c_{t-1} \text{ (previous cell state)} \\ \text{Forget gate:} & f_t = \sigma(W_f \cdot [h_{t-1}, x_t] + b_f) \\ \text{Input gate:} & i_t = \sigma(W_i \cdot [h_{t-1}, x_t] + b_i) \\ \text{Candidate value:} & \tilde{g}_t = \tanh(W_g \cdot [h_{t-1}, x_t] + b_g) \\ \text{Output gate:} & o_t = \sigma(W_o \cdot [h_{t-1}, x_t] + b_o) \\ \text{Update cell state:} & c_t = f_t \odot c_{t-1} + i_t \odot \tilde{g}_t \\ \text{Update hidden state:} & h_t = o_t \odot \tanh(c_t) \end{aligned} \]

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

\[ \begin{aligned} \text{Input:} & x_t \text{ (current time step input)}, h_{t-1} \text{ (previous hidden state)} \\ \text{Reset gate:} & r_t = \sigma(W_r \cdot [h_{t-1}, x_t] + b_r) \\ \text{Update gate:} & z_t = \sigma(W_z \cdot [h_{t-1}, x_t] + b_z) \\ \text{Candidate value:} & \tilde{h}_t = \tanh(W_h \cdot [r_t \odot h_{t-1}, x_t] + b_h) \\ \text{Update hidden state:} & h_t = (1 - z_t) \odot h_{t-1} + z_t \odot \tilde{h}_t \end{aligned} \]

Interpretation of the computation logic:

  • Reset gater_tWhen close to 0, the candidate stateh̃_tbarely depends on history, equivalent to starting over
  • Update gatez_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
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):

\[ \begin{aligned} \text{Parameters per layer} & = 4 \times (hidden\_size \times input\_size + hidden\_size \times hidden\_size + hidden\_size) \\ & = 4 \times hidden\_size \times (input\_size + hidden\_size + 1) \end{aligned} \] where 4 corresponds to four groups of weight matrices: forget gate, input gate, candidate value, and output gate.

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
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
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

import torch.nn as nn

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
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
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
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
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
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

lstm = nn.LSTM(input_size=32, hidden_size=64, batch_first=True)

# 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
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
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
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
Other Extensions