M4-001M4: Sequences & TransformersRecurrent Models (RNN/LSTM/GRU)Easy
Mastery:
Recurrent Models (RNN/LSTM/GRU): 解释 RNN 的前向与 BPTT,并说明它的主要缺陷。
📐 Mathematical Definition
⚡ Executive Summary
Core Concept: 隐状态递归 h_t=f(W_h h_{t-1}+W_x x_t);BPTT 沿时间反传,梯度含 W_h 的 T 次连乘,导致梯度消失/爆炸与串行不可并行。
📌 Key Takeaways
- •BPTT 的梯度含同一权重矩阵的 T 次幂(谱半径≠1 即指数衰减/爆炸)
- •时间维不可并行(必须等 h_{t-1}),训练吞吐低
- •长依赖在远距离上梯度趋 0,实际有效记忆远短于理论
📐 Mathematical Derivations
数学机理:<strong>前向</strong>——RNN 在每个时间步用同一组权重处理输入与上一步隐状态:h_t=tanh(W_h h_{t-1}+W_x x_t),输出 y_t=g(W_y h_t)。<strong>BPTT(Backpropagation Through Time)</strong> 把时间维展开成 T 层的前馈网络再做反向传播,梯度为 ∂L/∂h_t=∂L/∂h_T·∏_{k=t+1}^{T}∂h_k/∂h_{k−1},其中每项 ∂h_k/∂h_{k−1}=diag(1−h_k²)·W_h。<strong>关键差异</strong>:与深层前馈网络不同,RNN 的每一层用的是<strong>同一个 W_h</strong>(权重共享),故连乘变成 <strong>W_h 的 T 次幂</strong>(近似):梯度尺度 ~ ρ(W_h)^T·(tanh 导数的乘积)。<strong>后果</strong>:(a) ρ(W_h)<1 → 梯度指数衰减到 0(<strong>梯度消失</strong>,无法学长依赖);(b) ρ(W_h)>1 → 梯度指数爆炸(训练发散);(c) tanh 导数 ≤1 进一步加剧衰减。<strong>其他缺陷</strong>:(a) <strong>时间维串行</strong>——h_t 依赖 h_{t−1},无法在时间维并行,训练吞吐受限于序列长度的串行链;(b) <strong>固定维隐状态</strong>——所有历史被压进固定维向量,形成信息瓶颈;(c) <strong>长距离依赖</strong>——梯度随距离指数衰减,实际有效记忆远短于理论。<strong>Truncated BPTT</strong> 只反传固定窗口(如 50 步)内的梯度,把 O(T) 的反向成本与内存降到 O(window),代价是无法直接学习超过窗口长度的依赖。
🏭 Production Trade-offs
深度剖析与工程权衡:① <strong>谱半径与稳定性的实用判据</strong>——正交初始化(W_h 正交矩阵,谱半径=1)能让梯度尺度近似保持,是训练长序列 RNN 的经典技巧;这与后来 Transformer 用残差 + LN 稳定深层梯度是同一思想谱系。② <strong>梯度裁剪的角色</strong>——RNN 中爆炸比消失更常见且更危险(一步即 NaN),故 gradient clipping(全局范数裁剪)是 RNN 训练的标配;裁剪对消失无效。③ <strong>LSTM 的解法</strong>——用<strong>加性</strong>的细胞状态更新(c_t=f⊙c_{t−1}+i⊙g)替代乘性递归,使梯度沿 c_t 有一条'导数近似为 f(接近 1)'的直通路径,把连乘变成累加,从而大幅缓解消失。这是'用加法对抗指数衰减'的经典设计。④ <strong>并行性的根本局限</strong>——RNN 的串行性使其无法利用 GPU 的大规模并行;这直接催生了 CNN(可并行但感受野有限)与 Transformer(完全并行 + 全局交互)两条路线。⑤ <strong>现代残留价值</strong>——RNN 在流式/在线场景(需 O(1) 状态、逐 token 低延迟)仍有优势,且 SSM(Mamba)可视为'用并行扫描实现可训练的长程 RNN'的复兴。⑥ <strong>面试要点</strong>——被问'RNN 为什么不行',必须点出'<strong>同一矩阵的 T 次连乘</strong>'与'<strong>时间维串行</strong>'这两条根因;只说'梯度消失'会漏掉并行性这一更致命的问题。
⚠️ Common Interview Pitfalls
- ✕把 RNN 的梯度问题等同于普通深层网络的梯度消失(根源是同一矩阵连乘)
- ✕忽略时间维串行导致的训练不可并行
🎯 Interviewer Follow-ups
- ?为什么 RNN 的梯度问题是'同一矩阵连乘'而非'不同矩阵连乘'?
- ?truncated BPTT 解决了什么、代价是什么?
📚
Associated Knowledge Base Guides & Mindmaps
Explore the comprehensive technical article, exam cards, and global architecture tree.