先问自己一个问题:为什么需要 RNN?
上一道题学的词嵌入,只能处理单个词:
- "我" 有一个向量
- "打" 有一个向量
- "你" 有一个向量
但 "我打你" 和 "你打我" 用的词完全一样,意思却相反!
问题出在哪?词嵌入忽略了顺序。它只知道 "有哪些词",不知道 "词是怎么排列的"。
RNN 就是来解决这个问题的—— 它能记住 "之前看到了什么",所以能理解顺序。
为什么需要 "隐藏状态 h"?(RNN 的核心)
把 RNN 想象成一个"边读边记的人":
你在读这句话:"The cat sat"
读到第1个字'T' → 脑子里记下:"我看到了T"
读到第2个字'h' → 脑子里更新:"我看到了Th"(结合之前的T)
读到第3个字'e' → 脑子里更新:"我看到了The"(结合之前的Th)
读到第4个字' ' → 脑子里更新:"我看到了The "(知道The是一个完整的词)
读到第5个字'c' → 脑子里更新:"我看到了The c"(预测下一个可能是'a',因为The后面常跟cat)
隐藏状态 h 就是这个人的 "脑子"(短期记忆):
- 每读一个字符,就更新一次记忆
- 更新时要参考 "之前的记忆" + "当前看到的字符"
- 这样到后面,模型就知道 "前面出现过什么"
如果没有 h(没有记忆):模型每看到一个字符都是 "全新的",不知道之前出现过 'T'、'h'、'e',也就不可能预测出下一个是空格(因为它不知道 "The" 已经拼完了)。
为什么需要 W2?(记忆怎么传递)
RNN 有三套权重,其中 W2 是最特别的,也是你最可能困惑的:
表格
| 权重 | 作用 | 类比 |
|---|---|---|
| W1 | 处理 "当前看到的字符" | 眼睛 |
| W2 | 处理 "之前的记忆" | 记忆神经 |
| W3 | 根据记忆输出预测 | 嘴巴 |
前向传播公式:
h_t = tanh(W1·x_t + W2·h_{t-1} + b1)
↑ ↑
当前看到的 之前的记忆
字符x_t h_{t-1}
W2 干的事:把上一步的记忆 h_{t-1} "带" 到当前步。
- 如果没有 W2:
h_t = tanh(W1·x_t),每一步的记忆只由当前字符决定,和之前无关 → 这就退化成普通神经网络了,没有记忆! - 有了 W2:
h_t既看当前字符,又看之前的记忆 → 真正实现了 "边读边记"
一句话:W2 就是 RNN 的 "记忆通道",没有 W2 就没有 RNN。
为什么要逐个字符输入?(而不是整句一次性输入)
因为语言是有顺序的序列:
- "The" 的意思,要看到 'T'→'h'→'e' 三个字符按顺序出现才能理解
- 如果一次性把整句扔进去,模型不知道哪个字符在前、哪个在后
- 逐个输入,模型才能 "按顺序读",并在每一步更新记忆
这就像你读书 —— 你是一个字一个字按顺序读,不是一眼把整页同时看进去。
为什么要 "预测下一个字符"?(自监督学习)
这是另一个容易困惑的点:为什么训练目标是 "预测下一个字符",而不是别的?
因为这是一种"自监督学习"—— 不需要人工标注,文本本身就是答案:
输入序列:T h e q u i c k ...
目标序列:h e q u i c k ...
↑ ↑ ↑ ↑ ↑ ↑ ↑ ↑ ↑
每个输入的"下一个字符"就是目标
- 输入
'T',目标是'h' - 输入
'h',目标是'e' - 输入
'e',目标是' '(空格) - ...
模型学会预测下一个字符,就等于学会了语言的模式:
- 哪些字母常跟在哪些字母后面('q' 后面几乎总是 'u')
- 空格什么时候出现(单词拼完后)
- 常见单词怎么拼("the"、"quick"、"brown")
- 句子的语法结构
这和你学英语是一样的:你读了大量英文后,看到 "The quick brown f..." 就能猜到下一个是 "o"(fox),因为你见过太多次了。RNN 也是通过大量 "预测下一个字符" 的练习,学会了语言模式。
现在的 GPT 也是这个思路:给它上文,它预测下一个词,一个词一个词地生成回答。只是 GPT 用的是 Transformer(比 RNN 更高级的记忆机制),但 "预测下一个" 的核心思想是一样的。
为什么用交叉熵损失?(而不是上一题的 MSE)
上一道词嵌入题用的是 MSE(均方误差),这道题用的是 交叉熵,为什么不一样?
因为问题类型不同:
表格
| 词嵌入题 | RNN 题 | |
|---|---|---|
| 任务类型 | 回归(让两个向量相等) | 分类(29 个字符里选 1 个) |
| 输出 | 一个向量(3 维) | 一个概率分布(29 维,每个字符的概率) |
| 损失 | MSE(量两个向量的距离) | 交叉熵(量预测分布和真实分布的差距) |
交叉熵在干嘛:
- 模型输出:29 个字符的概率(比如 'h' 概率 0.8,'e' 概率 0.1,其他 0.1)
- 真实答案:'h'(one-hot:第 10 位为 1,其他为 0)
- 交叉熵 = 衡量 "模型给的概率分布" 和 "真实分布" 差多远
- 如果模型给真实字符 'h' 的概率很高(0.9),损失小
- 如果给得很低(0.01),损失大
一句话:分类问题用交叉熵,回归问题用 MSE。预测下一个字符是分类(29 选 1),所以用交叉熵。
为什么用 BPTT?(梯度怎么算)
RNN 的梯度下降叫 BPTT(沿时间反向传播),比普通反向传播多了 "沿时间" 三个字。
为什么? 因为 W1/W2/W3 在每个时间步都复用(同一套权重):
时间步1:用 W1/W2/W3 算 h1, p1
时间步2:用同一套 W1/W2/W3 算 h2, p2
时间步3:用同一套 W1/W2/W3 算 h3, p3
...
所以总损失 = 每个时间步损失的和,而每个时间步的损失都依赖 W1/W2/W3。
算梯度时:
- 要从最后一步往前算(因为后面的隐藏状态依赖前面的)
- 每个时间步算出的梯度要加起来(因为权重被复用了)
- 这就是 "沿时间反向传播"
你不需要手推公式,记住:BPTT = 普通反向传播 + 沿着时间轴从后往前 + 梯度累加。代码里已经帮你算好了。
一张图总结 RNN 的 "为什么"
为什么需要RNN? → 词嵌入忽略顺序,"我打你"≠"你打我"
↓
为什么需要h? → h是"短期记忆",每读一个字符更新一次
↓
为什么需要W2? → W2是"记忆通道",把上一步的记忆传到当前步
↓
为什么逐个输入? → 语言是有序序列,要按顺序读才能理解
↓
为什么预测下一个?→ 自监督学习,文本本身就是答案,学会预测=学会语言模式
↓
为什么交叉熵? → 预测字符是分类问题(29选1),分类用交叉熵
↓
为什么BPTT? → 权重在每步复用,梯度要沿时间从后往前累加
现在检验你懂没懂(3 个问题)
- 如果把 W2 去掉(设成全 0),RNN 还能记住之前的字符吗?为什么?
- 输入 'q',RNN 大概率会预测下一个字符是什么?为什么?
- 如果把损失函数从交叉熵改成 MSE,会有什么问题?(提示:MSE 适合回归,不适合分类)答不上来的那题,就是你还没通的点,告诉我,我再拆。
数据来自实际运行 RNN 字符级文本生成代码(训练数据 "The quick brown fox jumps",25 个字符,每 50 轮记录一次平均交叉熵损失)。
图 1:hidden_size(记忆容量)对比
- hidden_size=10(蓝线):稳步下降,500 轮后损失 0.011,学会了
- hidden_size=50(绿线):中间有波动(第 100 轮冲到 4.87),但后面快速下降,最终损失 0.006,学得更好
- hidden_size=200(红线):损失直接冲到 10 以上,越训越高,完全没学会
图 2:学习率 lr(步长)对比
- lr=0.005(蓝线):下降非常慢,500 轮后损失还有 0.33,没走完(步子太小)
- lr=0.05(绿线):虽然中间波动,但最终降到 0.006,刚刚好
- lr=0.5(黄线):第 50 轮就冲到 12.5,之后在 10-15 之间乱跳,完全发散
- lr=2.0(紫线):一开始就 6.15,之后一直在 13-16 之间震荡,彻底飞了
一句话总结这两个参数
表格
| 参数 | 太小会怎样 | 太大会怎样 | 怎么选 |
|---|---|---|---|
| hidden_size | 记忆不够,复杂任务学不会 | 参数太多,小数据集训练不稳定 | 任务简单选小(10-50),任务复杂选大(128-512) |
| lr | 下降太慢,训练轮数不够 | 梯度爆炸,损失震荡 / 发散 | 从 0.01 开始试,看损失曲线,稳定下降就对了 |
调参的实用技巧(以后工作也用得上)
- lr 先试 0.01:这是最安全的起点。如果损失下降太慢,调大到 0.05/0.1;如果损失震荡,调小到 0.001。
- 看损失曲线判断 lr 是否合适:
- 稳步下降 → lr 合适
- 下降极慢 → lr 太小
- 剧烈震荡 / 越来越高 → lr 太大
- hidden_size 看任务复杂度:
- 简单任务(短文本、小数据集):16-64 就够
- 复杂任务(长文本、大数据集):128-512
- 不是越大越好,大模型在小数据集上容易训练不稳定
- RNN 特别注意梯度爆炸:RNN 的 BPTT 容易梯度爆炸,所以一定要加梯度裁剪(代码里的
np.clip(dparam, -5, 5)),并且 lr 不要太大。
梯度爆炸 = 梯度突然变得超级超级大,导致一步更新就把权重改飞了,模型直接 "崩了"。
通俗类比:下山遇到悬崖
继续用 "下山找最低点" 的类比:
- 正常情况:你在山坡上,梯度告诉你 "往这个方向走,一步能下降 0.1 米"。你迈一步,稳稳下降 0.1 米,继续走。
- 梯度爆炸:你突然走到一个悬崖边,梯度告诉你 "往这个方向跳,一步能下降 10000 米!"。你信了,奋力一跳 —— 结果直接飞过了山谷,撞到对面的山上,甚至飞出了山区(损失变成 NaN)。
正常: 一步 ↓0.1米 → 一步 ↓0.1米 → 一步 ↓0.1米 → ... 稳步到山脚
爆炸: 一步 ↓10000米 → 飞出去了 → 不知道飞到哪了 → 损失=NaN
为什么 RNN 特别容易梯度爆炸?(核心原因)
这和 RNN 的 BPTT(沿时间反向传播) 有关。
RNN 的权重在每个时间步都复用,所以算梯度时,要从最后一个时间步往回传,每经过一个时间步,梯度就要乘一个数(就是 W2 的某个值):
时间步25的梯度
↓ 乘一个数(比如1.5)
时间步24的梯度
↓ 再乘一个数(1.5)
时间步23的梯度
↓ 再乘一个数(1.5)
...
↓ 乘了25次
时间步1的梯度 = 原始梯度 × 1.5^25 ≈ 原始梯度 × 25000 倍!
如果每次乘的数 > 1,经过很多个时间步后,梯度就会被放大成千上万倍—— 这就是 "爆炸"。
序列越长(时间步越多),爆炸的风险越大。这道题只有 25 个字符,爆炸风险还不算最高;如果是几百字的长文本,更容易爆炸。
对比一下:如果每次乘的数 < 1(比如 0.8),乘 25 次后变成 0.8^25 ≈ 0.0038,梯度就变得几乎为 0—— 这叫梯度消失。RNN 不仅会爆炸,还会消失,所以后来才有了 LSTM/GRU 来解决这两个问题。
梯度爆炸时你会看到什么?
表格
| 现象 | 意思 |
|---|---|
损失突然变成 NaN | Not a Number,不是数字了,权重被改飞了 |
损失突然变成 inf | infinity,无穷大 |
| 损失剧烈震荡(一会儿 10,一会儿 15) | 每次更新都飞过头,在山坡上来回跳 |
| 生成的文本全是乱码或重复字符 | 权重乱了,模型不会预测了 |
你之前跑 lr=0.5 时看到的:
epoch 0: 3.00
epoch 50: 12.58 ← 一下飞上去了
epoch 100: 12.32
epoch 150: 11.78
epoch 200: 12.71 ← 在高位乱跳
这就是轻度梯度爆炸—— 还没到 NaN,但已经飞上去下不来了。如果 lr 再大一点(比如 5),就会直接变成 NaN。
怎么解决梯度爆炸?(3 个办法)
办法 1:梯度裁剪(代码里已经用了)
这是最直接的办法:把太大的梯度砍掉。
代码里这一行就是干这个的:
for dparam in [dW1, dW2, dW3, db1, db2]:
np.clip(dparam, -5, 5, out=dparam) # 把超过5的梯度砍成5,低于-5的砍成-5
类比:悬崖太陡了,你规定 "不管坡度多大,我每步最多只迈 5 米",这样就不会飞出去了。
办法 2:降低学习率 lr
lr 小,即使梯度大,lr × 梯度 的更新量也不会太离谱。
梯度 = 10000,lr = 0.5 → 更新量 = 5000(飞了)
梯度 = 10000,lr = 0.001 → 更新量 = 10(还能接受)
办法 3:用 LSTM / GRU 代替普通 RNN
这是更根本的解决办法。LSTM/GRU 内部有"门控" 机制,可以控制梯度的流动 —— 该传的传,不该传的挡住,从结构上避免梯度爆炸 / 消失。
你课件后面应该会讲到 LSTM,它就是为了解决普通 RNN 的梯度问题发明的。现在大模型用的 Transformer,也是另一种解决方案。
一句话总结
梯度爆炸 = 梯度在 BPTT 沿时间回传时被反复放大,变得超级大,一步更新把权重改飞了。 解决办法:梯度裁剪(砍梯度)、降低 lr(小步走)、换 LSTM/GRU(从结构上控制梯度)。
公式总结
一、第 1 题:前向传播相关公式
1. 独热编码(One-hot)
x_t ∈ {0, 1}^V,只有对应字符的位置为1,其余为0
- 意思:每个字符用一个 V 维向量表示(V = 字符表大小,这道题 V=29)
- 对应代码:
def one_hot(c): v = np.zeros(V) v[char_to_idx[c]] = 1 return v
2. 隐藏状态更新(RNN 核心公式)⭐⭐⭐
h_t = tanh(W1 · x_t + W2 · h_{t-1} + b1)
- 意思:当前隐藏状态 = 当前输入 + 上一步的记忆,过 tanh 激活
W1·x_t:当前字符的信息W2·h_{t-1}:上一步的记忆(RNN 的 "循环" 就在这)b1:偏置项tanh:把值压缩到 -1~1 之间
- 对应代码:
h = np.tanh(W1 @ x + W2 @ h_prev + b1) - 必考程度:⭐⭐⭐ 必须背,这是 RNN 的定义
3. 输出计算
y_t = W3 · h_t + b2
- 意思:把隐藏状态映射回字符表维度,得到每个字符的 "分数"(未归一化)
- 对应代码:
y = W3 @ h + b2
4. Softmax(转成概率分布)⭐⭐
p_t = softmax(y_t) = exp(y_t - max(y_t)) / Σ exp(y_t - max(y_t))
- 意思:把输出分数转成概率分布,所有概率加起来 = 1,每个值在 0~1 之间
- 减去
max(y_t)是为了数值稳定(防止 exp 太大溢出)
- 减去
- 对应代码:
p = np.exp(y - np.max(y)) / np.sum(np.exp(y - np.max(y))) - 必考程度:⭐⭐ 常考
二、第 2 题:梯度计算相关公式
5. 交叉熵损失(单步)⭐⭐⭐
L_t = -log p_t(true_char)
- 意思:真实字符的概率越小,损失越大
- 如果真实字符概率 = 0.9,损失 = -log (0.9) ≈ 0.1(小)
- 如果真实字符概率 = 0.01,损失 = -log (0.01) ≈ 4.6(大)
- 对应代码:
(total_loss += -np.log(p[targets[t], 0] + 1e-8)+1e-8防止 log (0)) - 必考程度:⭐⭐⭐ 必须背
6. 平均交叉熵损失(整个序列)
L = (1/T) · Σ_{t=1}^{T} L_t = (1/T) · Σ_{t=1}^{T} -log p_t(true_char)
- 意思:整个序列的平均损失,T = 序列长度(这道题 T=25)
- 对应代码:
return total_loss / len(X)
7. 输出层梯度(BPTT 第一步)⭐⭐
dy_t = p_t - y_true (y_true 是真实字符的 one-hot)
- 意思:预测概率减去真实分布,就是输出层的梯度
- 如果预测概率 = 真实分布,梯度 = 0(不用更新)
- 如果预测差得远,梯度大(要大更新)
- 对应代码:
dy = p_list[t].copy() dy[targets[t]] -= 1 # 真实字符位置减1,等价于 p - onehot_true
8. 输出权重梯度
dW3 = dy_t · h_t^T
db2 = dy_t
- 意思:输出权重 W3 和偏置 b2 的梯度
- 对应代码:
dW3 += dy @ h_list[t].T db2 += dy
9. 隐藏层梯度(反向传递的核心)⭐⭐
dh_t = W3^T · dy_t + dh_{next}
- 意思:当前隐藏状态的梯度 = 从输出层传回来的 + 从后一个时间步传回来的
W3^T · dy_t:输出层梯度反向传到隐藏层dh_{next}:后一个时间步的隐藏层梯度(BPTT 的 "沿时间回传")
- 对应代码:
dh = W3.T @ dy + dh_next
10. tanh 的导数
dh_raw = (1 - h_t²) · dh_t
- 意思:tanh 函数的导数是
1 - tanh²(x),因为h_t = tanh(...),所以导数就是1 - h_t² - 对应代码:
dh_raw = (1 - h_list[t] ** 2) * dh
11. 输入权重和隐藏层权重梯度
dW1 = dh_raw · x_t^T
dW2 = dh_raw · h_{t-1}^T
db1 = dh_raw
- 意思:输入权重 W1、隐藏层权重 W2、偏置 b1 的梯度
dW2用到h_{t-1}(上一步的隐藏状态),t=0 时用全 0 初始状态
- 对应代码:
dW1 += dh_raw @ x_list[t].T h_prev = h_list[t-1] if t > 0 else np.zeros((hidden_size, 1)) dW2 += dh_raw @ h_prev.T db1 += dh_raw
12. 梯度沿时间传递(BPTT 的关键)
dh_{next} = W2^T · dh_raw
- 意思:把当前步的梯度通过 W2 传到前一个时间步,这就是 "沿时间反向传播"(BPTT)
- 每经过一个时间步,梯度就乘一次 W2^T
- 如果 W2 的特征值 < 1,梯度会越来越小 → 梯度消失
- 如果 W2 的特征值 > 1,梯度会越来越大 → 梯度爆炸
- 对应代码:
dh_next = W2.T @ dh_raw
13. 梯度下降更新(所有参数)⭐⭐⭐
W1 ← W1 - η · dW1
W2 ← W2 - η · dW2
W3 ← W3 - η · dW3
b1 ← b1 - η · db1
b2 ← b2 - η · db2
- 意思:沿着梯度方向更新所有参数,η = 学习率
- 对应代码:
W1 -= lr * dW1 W2 -= lr * dW2 W3 -= lr * dW3 b1 -= lr * db1 b2 -= lr * db2 - 必考程度:⭐⭐⭐ 必须背
14. 梯度裁剪(防止梯度爆炸)⭐
dparam = clip(dparam, -5, 5)
- 意思:把超过 5 的梯度砍成 5,低于 -5 的砍成 -5,防止梯度爆炸导致参数飞掉
- 对应代码:
for dparam in [dW1, dW2, dW3, db1, db2]: np.clip(dparam, -5, 5, out=dparam)
三、辅助公式(文本生成时用)
15. 按概率采样(生成文本)
next_char ~ Categorical(p_t) (按概率分布随机采样)
- 意思:根据预测的概率分布随机选下一个字符(概率大的更容易被选中)
- 也可以用
argmax(p_t)直接选概率最大的(但会重复、不自然)
- 也可以用
- 对应代码:
idx = np.random.choice(V, p=p.flatten()) # 按概率采样 # 或 idx = np.argmax(p) # 选最大概率
四、公式 → 代码 翻译对照表(最实用)
表格
| 数学公式 | 代码写法 | 出现位置 |
|---|---|---|
| 独热向量 x_t | np.zeros(V); v[idx]=1 | 数据准备 |
h_t = tanh(W1·x_t + W2·h_{t-1} + b1) | np.tanh(W1 @ x + W2 @ h_prev + b1) | 前向传播 |
y_t = W3·h_t + b2 | W3 @ h + b2 | 前向传播 |
softmax(y_t) | np.exp(y-np.max(y))/np.sum(np.exp(y-np.max(y))) | 前向传播 |
L = -log p(true) | -np.log(p[true_idx] + 1e-8) | 损失函数 |
dy = p - true_onehot | dy = p.copy(); dy[true_idx] -= 1 | 反向传播 |
dW3 = dy·h^T | dy @ h.T | 反向传播 |
dh = W3^T·dy + dh_next | W3.T @ dy + dh_next | 反向传播 |
tanh导数 = 1 - h² | (1 - h**2) * dh | 反向传播 |
dW1 = dh_raw·x^T | dh_raw @ x.T | 反向传播 |
dW2 = dh_raw·h_prev^T | dh_raw @ h_prev.T | 反向传播 |
dh_next = W2^T·dh_raw | W2.T @ dh_raw | 反向传播 |
W ← W - η·dW | W -= lr * dW | 参数更新 |
| 梯度裁剪 | np.clip(d, -5, 5, out=d) | 防止爆炸 |
五、必背公式清单(考试前过一遍)
第 1 题(前向传播)必背:
h_t = tanh(W1·x_t + W2·h_{t-1} + b1)—— RNN 的定义p_t = softmax(W3·h_t + b2)—— 输出概率
第 2 题(梯度计算)必背: 3. L = -log p(true_char) —— 交叉熵损失 4. dy = p - true_onehot —— 输出层梯度 5. dh = W3^T·dy + dh_next —— 隐藏层梯度(沿时间回传) 6. W ← W - η·dW —— 梯度下降更新
理解即可(不用死背推导): 7. BPTT 的完整梯度推导(dW1/dW2/dW3 的具体公式) 8. 梯度裁剪的具体阈值 9. tanh 导数的推导过程
这就是做那道 RNN 练习题用到的全部公式了。和上一道词嵌入题对比一下,你会发现套路是一样的:
表格
| 词嵌入题 | RNN 题 | |
|---|---|---|
| 输入 | 词的 one-hot | 字符的 one-hot |
| 参数 | 一个嵌入矩阵 E | 三套权重 W1/W2/W3 + 偏置 |
| 前向 | v = onehot @ E | h = tanh(W1·x + W2·h_prev + b1) |
| 损失 | MSE(类比误差) | 交叉熵(预测下一个字符) |
| 更新 | 梯度下降 | 梯度下降 + BPTT + 梯度裁剪 |
核心区别:RNN 多了一个 W2·h_{t-1}(记忆传递),所以反向传播时要沿时间回传(BPTT),也因此容易梯度消失 / 爆炸,需要梯度裁剪。
转载自 CSDN-专业IT技术社区
原文链接:https://blog.csdn.net/2501_93775482/article/details/166142096




