共计 2188 个字符,预计需要花费 6 分钟才能阅读完成。
传统 RNN 的困境与 LSTM 的诞生
1997 年提出的 LSTM 网络解决了传统 RNN 在处理长序列时的根本缺陷。让我们先通过数学视角理解这个问题:

在标准 RNN 中,隐藏状态的更新公式为:
$$h_t = \sigma(W_h h_{t-1} + W_x x_t + b)$$
其中梯度通过链式法则反向传播时会出现连乘项:
$$\frac{\partial h_t}{\partial h_k} = \prod_{i=k+1}^t \frac{\partial h_i}{\partial h_{i-1}}$$
当 $\frac{\partial h_i}{\partial h_{i-1}} < 1$ 时,经过多次连乘会导致梯度指数级衰减,这就是著名的梯度消失问题。
LSTM 的核心架构
LSTM 通过三个门控机制解决上述问题:
-
遗忘门:决定丢弃哪些历史信息
$$f_t = \sigma(W_f \cdot [h_{t-1}, x_t] + b_f)$$ -
输入门:控制新信息的存储
$$
i_t = \sigma(W_i \cdot [h_{t-1}, x_t] + b_i) \
\tilde{C}t = \tanh(W_C \cdot [h, x_t] + b_C)
$$ -
输出门:决定当前输出的内容
$$
o_t = \sigma(W_o \cdot [h_{t-1}, x_t] + b_o) \
h_t = o_t * \tanh(C_t)
$$
完整细胞状态更新公式:
$$C_t = f_t * C_{t-1} + i_t * \tilde{C}_t$$
PyTorch 实战:字符级文本生成
以下是完整的实现流程(测试环境:RTX 3080-10GB):
数据准备
from torch.nn.utils.rnn import pad_sequence
class CharDataset(Dataset):
def __init__(self, text, seq_length=100):
self.chars = sorted(list(set(text)))
self.char_to_idx = {c:i for i,c in enumerate(self.chars)}
self.data = [self.char_to_idx[c] for c in text]
self.seq_length = seq_length
def __getitem__(self, index):
# 滑动窗口截取序列
inputs = self.data[index:index+self.seq_length]
targets = self.data[index+1:index+self.seq_length+1]
return torch.LongTensor(inputs), torch.LongTensor(targets)
模型定义
class CharLSTM(nn.Module):
def __init__(self, vocab_size, embedding_dim=128, hidden_dim=512):
super().__init__()
self.embedding = nn.Embedding(vocab_size, embedding_dim)
self.lstm = nn.LSTM(embedding_dim, hidden_dim, batch_first=True)
self.fc = nn.Linear(hidden_dim, vocab_size)
def forward(self, x, hidden=None):
x = self.embedding(x)
out, hidden = self.lstm(x, hidden)
return self.fc(out), hidden
训练关键点
# 温度采样函数
def sample_with_temp(logits, temperature=1.0):
probs = F.softmax(logits / temperature, dim=-1)
return torch.multinomial(probs, num_samples=1)
# 初始化隐藏状态
hidden = (torch.zeros(1, batch_size, hidden_dim).to(device),
torch.zeros(1, batch_size, hidden_dim).to(device))
生产环境优化建议
-
变长序列处理:
# 正确使用 pack_padded_sequence lengths = [len(seq) for seq in batch] packed = pack_padded_sequence(embeddings, lengths, batch_first=True, enforce_sorted=False) -
双向 LSTM 优化:
self.lstm = nn.LSTM(embedding_dim, hidden_dim//2, bidirectional=True, layer_norm=True) -
TPU 显存优化:
- 使用
torch.xla提供的优化器 - 梯度累积代替大 batch
- 混合精度训练
延伸实验方向
- 尝试将 LSTM 应用于语法纠错任务,考虑如何设计错误检测机制?
- 对比 LSTM 与 Transformer 在长文本生成中的性能差异(指标建议:困惑度 / 生成连贯性)
- 探索 LSTM 在语音识别中的时序建模能力,如何处理不同长度的音频帧?
通过本教程,我们不仅理解了 LSTM 的理论基础,还掌握了从实验到生产的完整技术路径。建议读者动手实现每个代码片段,并尝试回答最后的思考问题来深化理解。
