你知道RNN吗?简述RNN
这就是大名鼎鼎的循环神经网络,你可以去理解它是一种专门处理序列数据的神经网络,比如文本、语音、时间序列。和普通前馈神经网络的最大区别在于 :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+Whht−1+b)
其中:
- xtx_txt表示当前输入
- ht−1h_{t-1}ht−1表示上一时刻的记忆
- 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 能处理序列,但它也有比较明显的问题。
- 梯度消失和梯度爆炸
当序列很长时,误差在时间维度上反向传播,容易出现梯度越来越小或者越来越大的问题,导致训练困难。
- 长期依赖问题
普通 RNN 虽然能记住前面的信息,但“记忆能力”有限。
如果前后间隔太远,模型往往很难保留早期的重要信息。
- 并行能力差
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)
AtomGit 是由开放原子开源基金会联合 CSDN 等生态伙伴共同推出的新一代开源与人工智能协作平台。平台坚持“开放、中立、公益”的理念,把代码托管、模型共享、数据集托管、智能体开发体验和算力服务整合在一起,为开发者提供从开发、训练到部署的一站式体验。
更多推荐



所有评论(0)