Terra.K头像
关注
自然语言处理2封面图

自然语言处理2

先问自己一个问题:为什么需要 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 个问题)

  1. 如果把 W2 去掉(设成全 0),RNN 还能记住之前的字符吗?为什么?
  2. 输入 'q',RNN 大概率会预测下一个字符是什么?为什么?
  3. 如果把损失函数从交叉熵改成 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 开始试,看损失曲线,稳定下降就对了

调参的实用技巧(以后工作也用得上)

  1. lr 先试 0.01:这是最安全的起点。如果损失下降太慢,调大到 0.05/0.1;如果损失震荡,调小到 0.001。
  2. 看损失曲线判断 lr 是否合适:
    • 稳步下降 → lr 合适
    • 下降极慢 → lr 太小
    • 剧烈震荡 / 越来越高 → lr 太大
  3. hidden_size 看任务复杂度:
    • 简单任务(短文本、小数据集):16-64 就够
    • 复杂任务(长文本、大数据集):128-512
    • 不是越大越好,大模型在小数据集上容易训练不稳定
  4. 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 来解决这两个问题。


梯度爆炸时你会看到什么?

表格

现象意思
损失突然变成 NaNNot a Number,不是数字了,权重被改飞了
损失突然变成 infinfinity,无穷大
损失剧烈震荡(一会儿 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_tnp.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 + b2W3 @ 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_onehotdy = p.copy(); dy[true_idx] -= 1反向传播
dW3 = dy·h^Tdy @ h.T反向传播
dh = W3^T·dy + dh_nextW3.T @ dy + dh_next反向传播
tanh导数 = 1 - h²(1 - h**2) * dh反向传播
dW1 = dh_raw·x^Tdh_raw @ x.T反向传播
dW2 = dh_raw·h_prev^Tdh_raw @ h_prev.T反向传播
dh_next = W2^T·dh_rawW2.T @ dh_raw反向传播
W ← W - η·dWW -= lr * dW参数更新
梯度裁剪np.clip(d, -5, 5, out=d)防止爆炸

五、必背公式清单(考试前过一遍)

第 1 题(前向传播)必背:

  1. h_t = tanh(W1·x_t + W2·h_{t-1} + b1) —— RNN 的定义
  2. 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 @ Eh = 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

文章来源转载

评论

赞0

评论列表

微信小程序
QQ小程序

关于作者

点赞数:0
关注数:0
粉丝:0
文章:0
关注标签:0
加入于:--