这就是大名鼎鼎的循环神经网络,你可以去理解它是一种专门处理序列数据的神经网络,比如文本、语音、时间序列。和普通前馈神经网络的最大区别在于 :RNN 有“记忆”,它在处理当前输入的时候,会参考前面已经处理过的信息。

为什么会有 RNN

普通神经网络有一个明显的局限:它通常只能处理“当前输入”,不擅长保留历史信息。但现实中的很多数据并不是彼此独立的,而是天然具有顺序关系,例如:

  • 文字有上下文
  • 语音有前后音节
  • 股票有历史走势
  • 天气有时间变化

也就是说,在这些任务中,当前结果往往不仅取决于当前输入,还取决于之前出现过的信息。因此,就需要一种能够在处理当前数据时,同时保留和利用历史信息的模型,RNN 就是在这样的背景下提出的。

RNN 怎么理解

当前输入 + 历史记忆 = 当前状态

每次处理新数据的时候,不只是看现在这一步,还会把前面记住的东西一起带上。

你可以把它想成一个人看句子:

  • 先看到“我”
  • 再看到“今天”
  • 再看到“很”
  • 最后猜下一个词

因为前面都记住了,所以后面更容易判断。

RNN 的基本结构

它的基本过程可以写成:

ht=tanh⁡(Wxxt+Whht−1+b)h_t = \tanh(W_x x_t + W_h h_{t-1} + b)ht=tanh(Wxxt+Whht1+b)

其中:

  • xtx_txt表示当前输入
  • ht−1h_{t-1}ht1表示上一时刻的记忆
  • hth_tht表示当前时刻的隐藏状态
  • WxW_xWx表示输入到隐藏层的权重
  • WhW_hWh表示上一时刻隐藏状态到当前隐藏状态的权重
  • bbb表示偏置项

在这里插入图片描述

如图所示,RNN 会把上一时刻的隐藏状态传递到下一时刻,从而实现信息在时间维度上的流动。这种“循环”结构让模型具备了处理序列数据的能力。

RNN 的时间展开

从结构上看,RNN 可以理解为同一个网络单元在不同时间步上的重复使用。
如果把它沿时间轴展开,就会变成下面这样的形式:

  • 第 1 步处理 x1x_1x1,得到h1h_1h1
  • 第 2 步处理 x2x_2x2,结合h1h_1h1 得到h2h_2h2
  • 第 3 步处理x3x_3x3,结合h2h_2h2得到h3h_3h3

也就是说,RNN 会沿着时间顺序不断更新自己的隐藏状态,把前面的信息逐步传递到后面。

这里你可以补一句很关键的话:

虽然展开后看起来像多个网络单元,但它们实际上共享同一套参数,这也是 RNN 能够处理变长序列的重要原因。

RNN 的特点

1. 能处理序列数据

它适合处理文本、语音、时间序列等有先后顺序的数据。

2. 具有记忆能力

它能够通过隐藏状态保留一部分历史信息,从而在当前时刻利用上下文。

这两个特点决定了 RNN 在早期自然语言处理和时间序列任务中非常常见。


RNN 的优点和缺点

RNN 的主要优点包括:

  • 能够建模前后顺序关系
  • 能处理长度不固定的序列输入
  • 参数在各时间步共享,结构统一
  • 相比普通前馈网络,更适合上下文相关任务

RNN 的缺点:

虽然 RNN 能处理序列,但它也有比较明显的问题。

  1. 梯度消失和梯度爆炸

当序列很长时,误差在时间维度上反向传播,容易出现梯度越来越小或者越来越大的问题,导致训练困难。

  1. 长期依赖问题

普通 RNN 虽然能记住前面的信息,但“记忆能力”有限。
如果前后间隔太远,模型往往很难保留早期的重要信息。

  1. 并行能力差

RNN 的计算依赖前一个时刻的结果,因此必须按顺序一步一步计算,不像 Transformer 那样容易并行训练。

RNN的简单实现

import torch
import torch.nn as nn

class ManualRNN(nn.Module):
    def __init__(self, input_size, hidden_size):
        super(ManualRNN, self).__init__()
        self.hidden_size = hidden_size
        
        self.Wx = nn.Linear(input_size, hidden_size)
        self.Wh = nn.Linear(hidden_size, hidden_size)

    def forward(self, x):
        # x.shape = (batch_size, seq_len, input_size)
        batch_size, seq_len, _ = x.shape
        
        h = torch.zeros(batch_size, self.hidden_size)
        
        outputs = []
        for t in range(seq_len):
            x_t = x[:, t, :]
            h = torch.tanh(self.Wx(x_t) + self.Wh(h))
            outputs.append(h.unsqueeze(1))
        
        outputs = torch.cat(outputs, dim=1)
        return outputs, h


# 测试
if __name__ == "__main__":
    x = torch.randn(2, 4, 3)  # batch=2, seq_len=4, input_size=3
    model = ManualRNN(input_size=3, hidden_size=5)
    
    outputs, h_last = model(x)
    
    print("每个时间步输出:", outputs.shape)   # (2, 4, 5)
    print("最后隐藏状态:", h_last.shape)     # (2, 5)
Logo

AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。

更多推荐