RNN 的诞生是为了让计算机拥有“记忆”。在它出现之前,传统的神经网络只能处理各自独立的输入数据。但真实世界充满了连续的序列,比如一段语音、一句连贯的话或者每天的股票走势。RNN 通过在网络中引入“循环”机制,让前一刻的信息能够传递到下一刻,巧妙地解决了让机器理解“上下文”的难题。

我们先来看看 RNN 是为了解决什么痛点而诞生的。

背景

在 RNN 出现之前,主流的是前馈神经网络 (Feedforward Neural Network)。这种网络有一个很大的局限性:它假设所有的输入数据都是相互独立的。

想象一下,如果你让传统网络去阅读句子“我出生在法国,所以我流利地讲___”。当它读到最后的空格时,它并没有机制去“记住”前面出现的词汇(比如“法国”),因此很难准确预测出答案是“法语”。

科学家们意识到,现实世界中大量的数据其实是序列数据 (Sequential Data) 📈,比如:

  • 语言和文字:单词的先后顺序直接决定了整句话的含义。
  • 语音片段:声音信号是随时间连续变化的。
  • 时间序列:今天的股票走势或气温通常与前几天的数据息息相关。

为了让机器拥有“记忆”来处理这些具有上下文关联的数据,在 20 世纪 80 年代(特别是通过 John Hopfield 和 Jeffrey Elman 等人的工作),研究人员在神经网络中引入了 “循环 (Recurrence)” 结构。通过这种设计,网络在某个时刻进行计算时,不仅会接收当前的新输入,还会接收自己上一时刻留下的“隐藏状态”(即记忆)。

这就好比我们在阅读一篇文章时,对当前词语的理解是建立在对前面词语记忆的基础之上的。

比如“狗咬人”和“人咬狗”,字完全一样,顺序一变,意思就天差地别,甚至能变成大新闻。
传统的神经网络就像是一个没有短期记忆的人,它看到“狗”、“咬”、“人”时,是把它们当成三个毫不相干的孤立画面去处理,完全不知道谁先谁后。
为了解决这个问题,RNN 引入了一个极其核心的概念:隐藏状态 (Hidden State),你可以把它想象成 RNN 的 “动态记忆本” 📓。

RNN 在处理当前时刻的输入(比如“咬”字)时,会结合上一时刻留下来的隐藏状态(包含了“狗”的信息)一起进行计算。这就好比我们在大脑中不断更新对当前语境的理解。

既然背景我们已经理清了,接下来我们直接深入核心原理实现流程,看看底层的数据流和数学逻辑是如何运转的。

核心原理:隐藏状态的更新逻辑

为了让这个过程更直观,我们可以把 RNN 沿着时间线“展开”。

在展开的图中,你可以看到信息是如何一步步流动的。假设我们有一个长度为 T T T 的序列,在任意一个时间步 t t t

  1. 接收输入:网络接收当前时刻的数据 x t x_t xt(比如某个词的词向量)。
  2. 读取记忆:网络读取上一时刻传过来的隐藏状态 h t − 1 h_{t-1} ht1
  3. 融合与更新:网络内部的权重矩阵会将 x t x_t xt h t − 1 h_{t-1} ht1 进行线性组合,然后再通过一个激活函数(通常是 tanh ⁡ \tanh tanh)来压缩数值范围,生成当前时刻的新隐藏状态 h t h_t ht

这个核心逻辑用数学公式表达非常简洁优美:

h t = tanh ⁡ ( W h x x t + W h h h t − 1 + b h ) h_t = \tanh(W_{hx} x_t + W_{hh} h_{t-1} + b_h) ht=tanh(Whxxt+Whhht1+bh)

这里的 W h x W_{hx} Whx W h h W_{hh} Whh 是权重矩阵,它们在整个时间序列的循环中是完全共享的。这种参数共享机制非常关键,它不仅大大减少了模型的参数量,还让模型能够处理任意长度的序列。

有了当前时刻的隐藏状态 h t h_t ht,如果我们在这步需要输出结果(比如预测下一个词),就可以直接用它计算输出 y t y_t yt

y t = W y h h t + b y y_t = W_{yh} h_t + b_y yt=Wyhht+by


实现流程:从前向传播到 BPTT

理解了上面的公式,整个 RNN 的实现流程(类似于我们在深度学习框架中编写的前向传播逻辑)就非常清晰了:

1. 初始化阶段
在序列开始输入之前,我们需要初始化一个初始的隐藏状态 h 0 h_0 h0。通常情况下,我们会把它初始化为一个全零矩阵,代表“白板”状态,没有任何先验记忆。

2. 前向传播 (Forward Pass)
这是一个按时间步推进的循环过程 (Loop):

  • t = 1 t=1 t=1:输入 x 1 x_1 x1 h 0 h_0 h0,计算得出 h 1 h_1 h1 和(可选的) y 1 y_1 y1
  • t = 2 t=2 t=2:输入 x 2 x_2 x2 h 1 h_1 h1,计算得出 h 2 h_2 h2 和(可选的) y 2 y_2 y2
  • 依此类推,直到序列的最后一个时间步 t = T t=T t=T,计算出最终的 h T h_T hT y T y_T yT

3. 计算损失 (Loss Computation)
将网络在各个时间步的预测输出 y t y_t yt 与真实的标签数据进行对比,计算出总的损失函数 (Loss)。

4. 沿时间反向传播 (BPTT, Backpropagation Through Time)
这是 RNN 训练中最关键也是最难的一步。普通的神经网络反向传播是一层一层往回传,而 RNN 因为存在时间步的循环,它的误差不仅要从输出层传到隐藏层,还要沿着时间线一步步往回溯(从 t = T t=T t=T 一直传回 t = 1 t=1 t=1),以此来计算所有时间步上梯度之和,最终更新那些共享的权重矩阵 ( W h x , W h h , W y h W_{hx}, W_{hh}, W_{yh} Whx,Whh,Wyh)。


梯度消失和梯度爆炸问题

这是一个非常核心且触及 RNN 本质的问题!要真正理解“梯度消失”和“梯度爆炸”,我们需要稍微借助一点微积分中的链式法则 (Chain Rule)

其实,沿时间反向传播 (BPTT) 在数学本质上和普通的反向传播没有任何区别,它只是把链式法则应用在了“按时间展开”的网络结构上。问题的根源,就出在这个**“连乘效应”**上。

我们来一步步拆解这个过程:

1. 罪魁祸首:链式法则的“无限连乘”

回顾一下 RNN 更新隐藏状态的公式:

h t = tanh ⁡ ( W h h h t − 1 + W h x x t + b h ) h_t = \tanh(W_{hh} h_{t-1} + W_{hx} x_t + b_h) ht=tanh(Whhht1+Whxxt+bh)

假设我们现在在时间步 t t t 计算出了一个损失 L L L,我们需要把误差反向传播回很久以前的时间步 k k k(比如 k k k 是第 1 个字, t t t 是第 100 个字)。

根据微积分的链式法则,我们要计算损失 L L L 对早期状态 h k h_k hk 的偏导数(也就是梯度),必须把中间所有时间步的偏导数乘起来:

∂ L ∂ h k = ∂ L ∂ h t ∏ i = k + 1 t ∂ h i ∂ h i − 1 \frac{\partial L}{\partial h_k} = \frac{\partial L}{\partial h_t} \prod_{i=k+1}^{t} \frac{\partial h_i}{\partial h_{i-1}} hkL=htLi=k+1thi1hi

请紧紧盯住这个连乘符号 ∏ \prod 。这里的关键在于 ∂ h i ∂ h i − 1 \frac{\partial h_i}{\partial h_{i-1}} hi1hi,即当前隐藏状态对上一个隐藏状态的导数。

2. 导数里的“套娃”

我们对隐藏状态公式求导,会得到什么呢?

∂ h i ∂ h i − 1 = tanh ⁡ ′ ⋅ W h h \frac{\partial h_i}{\partial h_{i-1}} = \tanh' \cdot W_{hh} hi1hi=tanhWhh

  • W h h W_{hh} Whh: 是隐藏层之间共享的权重矩阵。
  • tanh ⁡ ′ \tanh' tanh: 是激活函数 tanh ⁡ \tanh tanh 的导数。

把这个结果代回前面的连乘公式里,这就意味着,如果要跨越 t − k t-k tk 个时间步往回传导梯度,我们实际上是在把 W h h W_{hh} Whh tanh ⁡ ′ \tanh' tanh 连续相乘 t − k t-k tk

3. 为什么会“消失” (Vanishing Gradient)?

  • 首先, tanh ⁡ \tanh tanh 函数的导数值域是 ( 0 , 1 ] (0, 1] (0,1]。这意味着 tanh ⁡ ′ \tanh' tanh 的值绝大多数情况下都是小于 1 的小数
  • 其次,如果我们在初始化权重时,矩阵 W h h W_{hh} Whh 的特征值域(你可以粗略理解为矩阵里的数值大小)也小于 1

这就好比你在计算 0.9 100 0.9^{100} 0.9100。一个小于 1 的数被不断地连乘,它会以指数级的速度迅速衰减,逼近于 0。

后果: 当梯度传回到早期的网络节点时,数值已经变成了 0.00000…1。早期的权重得不到更新,网络也就无法学习到长距离之前的依赖关系(它患上了“健忘症”)。

4. 为什么会“爆炸” (Exploding Gradient)?

反过来想,如果在训练过程中,权重矩阵 W h h W_{hh} Whh 的值变得比较大,它的特征值域大于 1(比如 1.5),并且由于网络结构原因 tanh ⁡ ′ \tanh' tanh 没有把它压制住。

这就好比你在计算 1.5 100 1.5^{100} 1.5100。一个大于 1 的数被不断连乘,它会以指数级的速度疯狂增长,变成一个天文数字。

后果: 梯度数值大到超出了计算机浮点数的表示范围,导致权重更新出现 NaN (Not a Number),整个模型直接崩溃,无法继续训练


总结来说: RNN 在时间维度上的参数共享(每次都乘同一个 W h h W_{hh} Whh),导致 BPTT 变成了一个疯狂的连乘游戏。由于底数很难做到刚好等于 1,经过几十上百次连乘,结果要么归零(消失),要么上天(爆炸)。

由此引出了长短期记忆网络 (Long Short-Term Memory, LSTM) 来解决这个问题。

Logo

开源鸿蒙跨平台开发社区汇聚开发者与厂商,共建“一次开发,多端部署”的开源生态,致力于降低跨端开发门槛,推动万物智联创新。

更多推荐