Skip to main content
  1. Posts/

重新理解 RNN:Hidden State、Output 和 Prediction 的区别

·3506 words·7 mins
Table of Contents
Note: This article is available in Chinese only. 本文暂无英文版本。 View original

起因
#

最近在复习 MIT 6.S191 的 RNN 部分,看到 lecture 里画的那张经典展开图时停了一下。

图上标着 hidden state、output、prediction,这些词我都认识。但仔细想了想,发现自己对它们三者之间的关系,其实没有想象中那么清楚。

不是不知道 RNN 怎么用。干了这么多年 ML Infra,RNN、LSTM 这些东西在工程上早就不陌生了。但真要解释"hidden state 和 output 到底是不是同一个东西"“PyTorch 返回的 output 是 prediction 吗"这类问题时,发现自己的回答开始含糊。

更麻烦的是,翻了几份不同的材料——教科书、PyTorch 文档、各种示意图——发现它们对这几个词的用法并不一致。同一个 “output”,在不同地方指的可能是完全不同的东西。

于是决定把这件事理一遍。下面是整理出来的笔记。

RNN 的核心:State Transition
#

先从最本质的一行公式开始。

$$h_t = f_W(h_{t-1}, x_t)$$

一句话说清楚:给定上一步的状态 \(h_{t-1}\) 和当前输入 \(x_t\),算出当前状态 \(h_t\)。

具体到 vanilla RNN,展开就是:

$$h_t = \tanh(W_{hh} h_{t-1} + W_{xh} x_t + b)$$

这个公式定义了一个 state transition function。它接收两个东西——上一步的内部状态和当前时刻的新输入——然后产出一个新的内部状态。

把它沿时间展开:

flowchart LR
    h0["h₀"] -->|"f_W(h₀, x₁)"| h1["h₁"]
    h1 -->|"f_W(h₁, x₂)"| h2["h₂"]
    h2 -->|"f_W(h₂, x₃)"| h3["h₃"]
    h3 -->|"..."| h4["h_T"]
    x1["x₁"] --> h1
    x2["x₂"] --> h2
    x3["x₃"] --> h3

\(h_t\) 编码了从 \(x_1\) 到 \(x_t\) 的所有历史信息(理论上)。它是 RNN 对"到目前为止发生了什么"的压缩表示。

这个 state transition 就是 RNN 的本体。后面所有的讨论都围绕一件事:把 state transition 和建立在它之上的其他操作区分开来。

“Output” 到底指什么
#

“Output” 这个词在 RNN 语境下至少有三种不同的含义,经常被混着用。

数学公式中的 output
#

很多教材在写 RNN 公式时,会写成这样:

$$h_t = \tanh(W_{hh} h_{t-1} + W_{xh} x_t + b_h)$$$$\hat{y}_t = W_{hy} h_t + b_y$$

第一行是 state transition,第二行是从 hidden state 到 output 的映射。这里的 \(\hat{y}_t\) 才是数学意义上的 “output”——它是 hidden state 经过一个 linear transformation 后的结果。

但问题在于:在 vanilla RNN 中,很多讲解直接省掉第二行,把 \(h_t\) 本身当作 output。也就是说,hidden state = output。

这不是错误——vanilla RNN 确实可以这么理解。但它制造了一个隐含假设:hidden state 和 output 是同一个东西。一旦读者带着这个假设去看其他材料,混淆就开始了。

框架 API 中的 “output”
#

看 PyTorch 的 nn.RNN

1output, h_n = rnn(x, h_0)

这里返回了两个东西。自然的理解是:output 是模型的输出,h_n 是最终的 hidden state。

但实际上,output 的 shape 是 [seq_len, batch, hidden_size]——它是 每一个时间步的 hidden state 拼在一起的序列。也就是 \([h_1, h_2, \ldots, h_T]\)。

h_n 是最后一个时间步的 hidden state \(h_T\)。

对于单层、单向的 RNN:

1output[-1] == h_n.squeeze(0)   # True

换句话说,PyTorch 把"所有时间步的 hidden state 序列"叫做 output它不是 prediction,也不是经过任何 task head 之后的结果。它就是 hidden state。

这个命名选择不能说有错——从 RNN cell 的角度看,每一步的 hidden state 确实是这个 cell 对外暴露的输出。但它很容易让人误以为拿到 output 就已经拿到了模型的预测。

示意图中的混用
#

很多教学示意图在 RNN cell 的顶部画一个向上的箭头,标注 “output” 或 “\(\hat{y}_t\)"。有时这个箭头指的是 \(h_t\)(hidden state 本身),有时指的是 \(h_t\) 经过 output layer 之后的 prediction。

更令人困惑的情况是:同一张图里,从 cell 顶部引出两个箭头——一个水平传给下一步(标注 \(h_t\)),一个垂直向上(标注 “output”)——但这两个箭头指的其实是同一个 tensor。

看到 “output” 时,第一反应应该是问:这是 \(h_t\) 还是 \(\hat{y}_t\)?

Prediction 只是 Readout
#

把前面的混淆理清之后,可以给出一个干净的框架。

RNN 有两层操作,它们的职责完全不同:

$$\text{State transition:} \quad (x_t, h_{t-1}) \rightarrow h_t$$$$\text{Readout:} \quad h_t \rightarrow \text{task head} \rightarrow \hat{y}_t$$

第一层是 RNN 的本体——state transition。它不关心具体任务是什么,只负责维护一个随时间演进的内部状态。

第二层是 readout——从 hidden state 中提取当前任务需要的信息。这个 task head 可以是一个 linear layer + softmax(语言模型),可以是一个 classifier(情感分类),可以是任何东西。它是任务特定的,不属于 RNN 本身。

RNN 的核心抽象是 state transition。Prediction 只是对 state 的 readout。

不同的任务共享同一个 state transition 机制,只是 readout 的方式不同:

flowchart LR
    x1["x₁"] --> h1["h₁"]
    x2["x₂"] --> h2["h₂"]
    x3["x₃"] --> h3["h₃"]
    h1 --> h2 --> h3

    subgraph lm ["语言模型:每步 readout"]
        h1 --> y1["ŷ₁"]
        h2 --> y2["ŷ₂"]
        h3 --> y3["ŷ₃"]
    end

    subgraph cls ["序列分类:只在最后一步 readout"]
        h3 --> yc["ŷ"]
    end

语言模型在每个时间步都做 readout(预测下一个 token)。序列分类只在最后一步做 readout(把整个序列映射到一个 label)。序列标注则在每步做 readout,但可能用不同的 task head。

底层的 state transition 是一样的。变化的只是"从 state 里读什么、在哪里读”。

Weight Sharing
#

RNN 还有一个容易被忽略的 design choice:同一组参数 \(W\) 在所有时间步复用。

$$h_t = f_W(h_{t-1}, x_t)$$

不是 \(f_{W_1}, f_{W_2}, \ldots, f_{W_T}\)。是同一个 \(f_W\)。

这意味着:无论序列有多长,RNN 用的 transition function 都是同一个。第 1 步和第 1000 步应用的变换规则完全相同。

这不只是"省参数”。它编码了一个归纳偏置:序列中每一步的 transition rule 是相同的。

类比一下 CNN。CNN 的 kernel 在空间维度上共享——同一个 3x3 filter 滑过整张图片的每个位置。这背后的假设是:局部 pattern(比如边缘、纹理)不依赖于它出现在图片的哪个位置。

RNN 的 weight sharing 是同一个思想在时间维度上的版本:序列中的局部 transition rule 不依赖于它发生在序列的哪个位置。

Weight sharing 还带来一个实际后果:RNN 的参数量不随序列长度增长。一个处理长度为 10 的序列和处理长度为 10000 的序列的 RNN,参数完全一样。这让 RNN 能够处理任意长度的输入。

为什么 Vanilla RNN 特别容易让人混淆
#

回到最开始的问题:为什么这几个概念这么容易搞混?

一个重要原因是 vanilla RNN 的结构太"简单"了——简单到 hidden state 和 output 之间没有任何区分。

Vanilla RNN 的 hidden state \(h_t\) 直接就是暴露给外部的东西。没有额外的 gate 来控制"哪些信息对外可见"。所以在 vanilla RNN 的语境下,说"hidden state 就是 output"确实没什么问题。

但 LSTM 打破了这个等式。

LSTM 引入了 cell state \(c_t\) 作为真正的内部状态,然后用 output gate 控制 \(c_t\) 中哪些信息暴露出去:

$$h_t = o_t \odot \tanh(c_t)$$

这里 \(o_t\) 是 output gate,\(h_t\) 是 \(c_t\) 经过筛选后的结果。

在 LSTM 中,\(c_t\) 才是完整的内部状态,\(h_t\) 更像是"cell state 的一个 filtered view"。但 PyTorch 的 nn.LSTM 仍然把 \(h_t\) 叫做 “hidden state”,把 \(c_t\) 叫做 “cell state”——这又制造了新一层命名上的不对称。

Vanilla RNN 中 hidden state = output 的等式,是一个特例,不是通用规则。一旦引入 gate 机制(LSTM、GRU),hidden state 和对外暴露的 output 就不再是同一个东西了。

重新看 6.S191 的 RNN 示意图
#

回到 6.S191 lecture 的那张 RNN 展开图。现在可以很清楚地把图上的每个元素对应到框架里:

  • 水平方向传递的箭头:state transition,\(h_{t-1} \rightarrow h_t\)
  • 从底部进入的箭头:当前输入 \(x_t\)
  • 从顶部引出的箭头:如果直接标注 \(h_t\),那是 hidden state 本身;如果标注 \(\hat{y}_t\),那是经过 task head 之后的 prediction

整个 RNN 可以用两行伪代码概括:

1# state transition(RNN 本体)
2h_t = f(h_{t-1}, x_t; W)
3
4# readout(任务特定)
5y_t = task_head(h_t)

之前那篇 training loop 的笔记讨论了 forward / backward / optimizer 的职责划分。这里是类似的思路:把 RNN 拆成"状态演进"和"状态读取"两个独立的职责,很多混淆就消失了。

结论
#

  1. RNN 的本体是 state transition:\((x_t, h_{t-1}) \rightarrow h_t\)。Prediction 只是对 state 的 readout。

  2. 框架 API 中的 “output”(如 PyTorch nn.RNN 的返回值)通常是每步的 hidden state 序列,不是 prediction。看到 “output” 时先问:这是 \(h_t\) 还是 \(\hat{y}_t\)?

  3. Weight sharing 不只是参数效率——它让 RNN 能处理任意长度的序列,并且编码了"transition rule 在时间上不变"这个归纳偏置。

参考资料
#

Related