3.5 循环神经网络

1.有些数据是有先后顺序的

我们前面聊过卷积网络,它在图像任务上确实表现很好,一张猫的图片喂进去,分类结果往往清清楚楚。可一旦遇上文本、语音、股价这类带有时间顺序的数据,它就显得有些力不从心了。

说起来,一句话的意思常常和词的排列顺序紧紧绑在一起。狗咬人和人咬狗,字一模一样,顺序一换主角就反过来了。再看田径场上的接力赛,四名选手交接棒的先后顺序直接决定成绩是否有效,先跑后跑大有讲究。一只股票今天的价位也离不开昨天的价位,只盯着今天这一根K线去看,就好比不看走势只看当天涨跌的散户,难免要吃亏。

这种有先有后、当前依赖之前的数据,我们称之为序列数据。处理它,光看当下这一刻的输入远远不够,得有一种能记住之前发生了什么、再结合当下做判断的模型,这就是循环神经网络(RNN)。

2.隐藏状态:模型的短期记忆

那么循环网络是怎么记住东西的呢?它给每一个时间点都配了一个叫隐藏状态的东西,不妨把它理解成模型的短期记忆。

每到一个新的时间点,模型都会把新输入和上一刻的记忆融合到一起,更新出一份新记忆。就这么一边读一边更新,把截至当前的历史信息,全都压缩记在这个隐藏状态里。

我们用下标来表示时间。第 tt 个时间点的输入记作 xtx_t(下标 tt 表示第几个时间点,从1开始数),这一刻的隐藏状态记作 hth_thh 取自hidden的首字母,专门用来标记这种记忆)。最基本的循环关系写成:

ht=tanh(Wxxt+Whht1+b)h_t=\tanh(W_xx_t+W_hh_{t-1}+b)

我们把右边每个符号都拆开说说。ht1h_{t-1} 是前一个时间点的旧状态,也就是上一刻的短期记忆。WxW_x 是专门用来处理当前输入 xtx_t 的权重矩阵。WhW_h 是专门用来处理旧状态 ht1h_{t-1} 的权重矩阵。bb 是偏置。tanh\tanh 是一种把输出压到 1-111 之间的激活函数,你回想一下前面章节提过的那个 tanh\tanh,就是同一个东西。

打个比方,这就好比你一边追剧一边记剧情。每看到一个新片段,脑子里就会把这个新片段和刚才记住的剧情融到一起,更新一下自己对这一集的理解。你脑子里当前这一摊内容,就是隐藏状态 hth_t。我之前翻过一本讲认知心理学的书,书里说到人的工作记忆大概也是这般机制,容量有限,只能边接收边压缩,读完这一节你大概也能体会到这种拘束。

3.同一组权重,能吃下任意长度

循环网络还有一个特别实用的特点:每一个时间点用的都是同一组权重 WxW_xWhW_hbb。这跟卷积里的参数共享是一个思路。正因为参数不随时间变化,同一个模型既能处理五个字的短句,也能处理五十个字的长句,句子长短变了,模型本身一个参数都不用改。

这个特点对做题的我们来说其实非常友好。训练的时候不用为每种长度单独学一套参数,数据利用效率很高,短句子和长句子都能直接拿来喂,学一组权重就够用了。

当前状态同时取决于当前的输入 xtx_t 和上一刻的状态 ht1h_{t-1},所以模型能结合上下文来做判断。比如读到打这个字,前面是下雨还是打电话,会决定它该往哪个意思上靠。模型之所以能结合上下文,全靠隐藏状态把前文一路带到了当下。

4.把时间铺平:展开后它其实是个很深的网络

循环网络的循环结构画在纸上就那么一个小圈带个回头箭头,看着挺简洁,但训练的时候我们往往会换一种更直观的画法,把按时间转圈的那个圈沿着时间轴拆开,每一个时间点的计算一字排开画出来。这种画法叫按时间展开(unrolling)。

展开之后你会看到,原来那个反复转圈的小网络,其实等价于一个非常深的前馈网络,每一个时间点对应其中的某一层。输入 x1x_1 算出 h1h_1h1h_1x2x_2 一起算出 h2h_2h2h_2 再跟 x3x_3 算出 h3h_3,一节一节往后传。展开图上看着有好多层,层与层之间用的权重 WxW_xWhW_hbb 全都是同一份,每一步用的都是同一组参数,这就是上一节说的参数共享在展开图上的具体样子。

你不妨把它想成一条流水线,工位一个接一个,每个工位用的都是同一本操作手册,差别只在于传到这个工位的原料和上一个工位留下的半成品不一样。句子有多长,这条流水线就有多长,模型那本手册一个字都不用改。这也是循环网络能吃下任意长度序列的真正原因。

展开图除了直观,还直接告诉了我们梯度是怎么沿着时间一步步往回传的,这就引出了下一节那个最让人头疼的长距离依赖问题。

5.序列一长,记忆就开始漏

但这种结构也埋着一个很要紧的隐患,叫长距离依赖问题。

你还记得吧,反向传播在循环网络里是沿着时间一步一步往回传的。每传一个时间点,就要乘一次那组权重 WhW_h。如果 WhW_h 每乘一次都把梯度越缩越小(它的特征值小于1),那传不了几步,梯度就缩到几乎为零了,这就是梯度消失,开头的内容就这样被一点点遗忘掉。如果 WhW_h 每乘一次都让梯度变大(它的特征值大于1),梯度又会越乘越大,直接爆掉,这就是梯度爆炸。

这套沿着时间往回传梯度的算法,专门有个名字,叫随时间反向传播(Backpropagation Through Time,缩写BPTT)。它干的事情跟普通反向传播一模一样,只不过梯度除了在网络层与层之间传,还要沿着时间轴一步一步往历史方向传。我们拿一个具体场景说清楚。假设序列一共有 TT 个时间点(TT 就是序列的总长度),最终的损失 LL 是从最后一个时间步 hTh_T 一路相关过来的。我们想算损失 LL 对很早某个中间状态 hkh_k 的梯度 Lhk\frac{\partial L}{\partial h_k}Lhk\frac{\partial L}{\partial h_k} 读作 LLhkh_k 的偏导数,衡量的是 hkh_k 抖一下,LL 会跟着变多少),那这个梯度就得沿着 hTh_ThT1h_{T-1}hT2h_{T-2} 一路连乘回到 hkh_k,每倒退一步就要乘一次那个处理旧状态的权重 WhW_h。把这条连乘链子写出来,大致长这样:

Lhk=LhTi=k+1Thihi1\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}}

式子里那个大写的 \prod 表示一连串相乘,跟表示一连串相加的 \sum 是一对(它俩都是大写希腊字母,\prod 读派,\sum 读西格玛)。hihi1\frac{\partial h_i}{\partial h_{i-1}} 是相邻两个时间步之间的梯度,其中 ii 是从 k+1k+1 取到 TT 的那个时间步编号,\partial 这个符号就是偏导数专用的求导记号。每一步的 hihi1\frac{\partial h_i}{\partial h_{i-1}} 里头都藏着一份 WhW_h,连乘起来 WhW_h 就被反复叠在一起了。

这串连乘就是梯度消失和梯度爆炸真正的根源所在。当 WhW_h 的特征值(你可以粗略理解成它每乘一次会把向量拉长还是压短的那个倍数)小于 11 时,连乘个几十步,梯度就越来越小,最后小到几乎为零,模型几乎收不到来自遥远过去的信号,开头的内容就这么被忘干净了,这就是梯度消失。反过来,要是 WhW_h 的特征值大于 11,每往后退一步梯度就被放大一截,连乘下去数值越涨越离谱,直接冲成NaN让训练当场崩溃,这就是梯度爆炸。

梯度爆炸相对好对付,一个叫梯度裁剪(gradient clipping)的办法就能压住:每次算完梯度,去看它的范数(你可以理解成这个梯度向量的整体大小)有没有超过一个手动设好的上限,超了就把它按比例缩回来,办法朴素但管用。梯度消失就要难处理得多,裁剪根本救不回来,这也是后面门控循环网络专门要解决的核心痛点。

这件事换个说法也许更清楚。某部很有名的侦探小说里,侦探要从一长串转述过的口供追溯到最早的证人原话,每多经过一个人,证词就可能被添油改味一道。循环网络沿时间传梯度也是这般境遇,序列一旦变长,模型就很难再记起开头说过什么。比如一个长段落,开头埋了个人名小王,到结尾才用到,普通循环网络经常接不住这种隔了老远的呼应。

说白了就是,短期记忆的容量有了,但记得不太牢,越早的事情忘得越快。

6.双向和多层:把循环网络叠起来用

到这里我们聊的循环网络都是单向的,只能从左往右读序列。可有些活儿光从左往右读是吃亏的。比方说完形填空里出现我住在这个___,到底是村里、城里还是外地,光看前半句根本定不下来,往往得往后看几眼才能填准。也就是说,后文对理解前面那个位置也是有用的。

双向循环网络(Bidirectional RNN)就是为这种需求生的。它同时跑两条循环网络,一条从左往右读,产出正向隐藏状态 ht\overrightarrow{h_t},另一条从右往左读,产出反向隐藏状态 ht\overleftarrow{h_t},最后在每个时间点上把这两份拼到一起,当成这个位置的最终表示。箭头符号 \overrightarrow{}\overleftarrow{} 分别标出了状态是从哪个方向扫过来的。这么一搭配,每个位置都同时拿到了左边的上下文和右边的上下文,理解起来自然更到位。代价也是有的,双向网络必须拿到完整序列才能开算,对那种得边接收边输出的实时场景(比如实时语音转文字)就用不上。

光加深还不够的话,我们还可以把循环网络像盖楼一样一层一层叠上去,这就是多层循环网络(也叫多层RNN或者Stacked RNN)。第一层循环网络的输出当成第二层的输入,第二层的输出再喂给第三层,以此类推。底层的网络负责抓一些比较浅的模式(比如某个词跟它前后一两个词的搭配),越往上层抽象层次越高,能捕捉更长更复杂的结构。层数叠加能换来更强的表达力,当然代价是算得更慢、参数也更多,叠过头同样容易过拟合。

这两种结构在后面的章节里你会反复见到,比如BERT里那条双向的思路,再比如很多翻译模型encoder用的就是多层双向结构。先在这儿混个脸熟,后面真碰到了不至于慌。

7.换个角度看:一台会更新自己的状态机

换个角度理解,循环网络就像一台随着时间不断更新自己的状态机。每来一个新输入,模型就修改一下内部状态,把没用的旧信息挤出去,把有用的新信息融进来。

状态就像一个容量有限的笔记本,写不下所有的东西,所以模型必须学会把历史压缩成对后续任务最有用的那一小撮表示。这就好比运动员赛前整理战术笔记,你不可能把整个赛季所有比赛的过程都抄上去,你得学会挑重点,把教练反复强调的那几条带进赛场。我记得有本讲篮球训练的书里写到,优秀选手的赛前笔记通常只列三五条,记得少反而记得牢,正是同样的道理。

也正是这个容量限制加上长距离依赖的毛病,才有了后面两章的出场机会:门控循环网络想办法让模型自己管理记忆,注意力机制则干脆让模型能直接回头查看任意一个历史位置。循环网络的这些痛点,正好是后面那些更强模型的灵感来源。

练习

Q1. RNN 的隐藏状态 hth_t 是怎么算出来的?它为什么能当成"短期记忆"?

基本循环关系是 ht=tanh(Wxxt+Whht1+b)h_t=\tanh(W_xx_t+W_hh_{t-1}+b):把当前输入 xtx_t 和上一刻的旧状态 ht1h_{t-1} 各自经权重矩阵处理后相加,再过 tanh\tanh 压到 (1,1)(-1,1)。每到一个新时间点就把新输入和旧记忆融合成新记忆,一边读一边更新,把截至当前的历史信息全压缩记在 hth_t 里,所以它就是模型的短期记忆。而且每个时间点用的是同一组权重 Wx,Wh,bW_x,W_h,b(参数共享),所以能吃下任意长度的序列。

Q2. 把 RNN 按时间展开后,它等价于一个什么样的网络?这和 BPTT 有什么关系?

展开后它等价于一个非常深的前馈网络,每个时间点对应其中一层,层与层之间共用同一组权重。BPTT(随时间反向传播)就是沿着这条展开的时间轴一步步把梯度往回传,每倒退一步都要乘一次那个处理旧状态的权重 WhW_h。这条连乘链 i=k+1Thihi1\prod_{i=k+1}^{T}\frac{\partial h_i}{\partial h_{i-1}} 正是梯度消失/爆炸的根源。

Q3. RNN 里梯度消失和梯度爆炸分别是怎么产生的?哪个好对付、哪个难处理?

沿时间回传时每步都乘 WhW_h,若它的特征值小于 1,连乘几十步梯度越来越小、近乎为零,模型收不到来自遥远过去的信号(梯度消失);若特征值大于 1,每步放大一截、连乘越涨越离谱冲成 NaN(梯度爆炸)。梯度爆炸相对好办,用梯度裁剪(范数超上限就按比例缩回)即可;梯度消失难处理得多,裁剪救不回来,这正是后面门控循环网络要解决的核心痛点。

Q4.(面试题) 为什么 RNN 用同一组权重能处理任意长度的序列?请从参数共享和梯度连乘两个角度说明它的好处和隐患。

好处来自参数共享:每个时间点都用同一组 Wx,Wh,bW_x,W_h,b,不随长度变化,所以长短序列共用一套参数,训练时数据利用效率高、不必为每种长度单独学权重。隐患来自按时间展开后的梯度连乘:每回传一步乘一次 WhW_h,序列一长连乘次数就多,WhW_h 特征值小于 1 就梯度消失、开头信息被遗忘,大于 1 就梯度爆炸。所以参数共享给了 RNN 通用性,却也埋下了长距离依赖的病根,这才催生了 LSTM/GRU 的门控机制和注意力机制。

相关标签
深度学习循环神经网络序列