手學(xué)深度學(xué)習(xí)》循環(huán)神經(jīng)網(wǎng)絡(luò)(RNN)實(shí)戰(zhàn)指南:序列模型、語(yǔ)言建模與 BPTT 全解析)
人工智能深度學(xué)習(xí)機(jī)器學(xué)習(xí)教程【免費(fèi)下載鏈接】d2l-zh《動(dòng)手學(xué)深度學(xué)習(xí)》面向中文讀者、能運(yùn)行、可討論。中英文版被70多個(gè)國(guó)家的500多所大學(xué)用于教學(xué)。項(xiàng)目地址https://gitcode.com/GitHub_Trending/d2/d2l-zh點(diǎn)擊查看免費(fèi)下載導(dǎo)讀本文以《動(dòng)手學(xué)深度學(xué)習(xí)》d2l-zh倉(cāng)庫(kù)中 chapter_recurrent-neural-networks/index.md 一章為骨架系統(tǒng)梳理從為什么需要序列模型到RNN 訓(xùn)練中的梯度問(wèn)題的完整技術(shù)脈絡(luò)。文章覆蓋序列模型的統(tǒng)計(jì)基礎(chǔ)、文本預(yù)處理流水線、語(yǔ)言模型與長(zhǎng)序列采樣、RNN 的前向計(jì)算原理、從零實(shí)現(xiàn)與高級(jí) API 簡(jiǎn)潔實(shí)現(xiàn)以及通過(guò)時(shí)間反向傳播BPTT的數(shù)學(xué)細(xì)節(jié)同時(shí)結(jié)合倉(cāng)庫(kù) d2l 包中各框架實(shí)現(xiàn)源碼進(jìn)行佐證。讀完本文你將理解 RNN 為何能處理序列信息、如何構(gòu)建字符級(jí)語(yǔ)言模型并用困惑度評(píng)估以及梯度爆炸/消失的成因與工程化解方案。為什么需要循環(huán)神經(jīng)網(wǎng)絡(luò)序列數(shù)據(jù)打破獨(dú)立同分布假設(shè)在之前的章節(jié)中表格數(shù)據(jù)與圖像數(shù)據(jù)都默認(rèn)樣本來(lái)自某種分布且獨(dú)立同分布i.i.d.。然而真實(shí)世界的大多數(shù)數(shù)據(jù)并非如此文章中的單詞按順序書(shū)寫(xiě)一旦順序被隨機(jī)重排原意便難以理解視頻中的圖像幀、對(duì)話中的音頻信號(hào)、網(wǎng)站的瀏覽行為天然有序我們不僅能接收一個(gè)序列作為輸入還期望繼續(xù)猜測(cè)序列的后續(xù)例如預(yù)測(cè)2, 4, 6, 8, 10, ...這在股市波動(dòng)、患者體溫曲線、賽車(chē)加速度等時(shí)間序列分析中非常常見(jiàn)。簡(jiǎn)言之如果說(shuō)卷積神經(jīng)網(wǎng)絡(luò)擅長(zhǎng)處理空間信息那么本章的循環(huán)神經(jīng)網(wǎng)絡(luò)recurrent neural networkRNN則擅長(zhǎng)處理序列信息——它通過(guò)引入狀態(tài)變量存儲(chǔ)過(guò)去的信息與當(dāng)前輸入從而確定當(dāng)前輸出。本章共包含七個(gè)小節(jié)構(gòu)成一條完整的學(xué)習(xí)鏈路小節(jié)相對(duì)路徑核心議題序列模型sequence.md序列數(shù)據(jù)的統(tǒng)計(jì)工具與預(yù)測(cè)挑戰(zhàn)文本預(yù)處理text-preprocessing.md詞元化與詞表構(gòu)建語(yǔ)言模型和數(shù)據(jù)集language-models-and-dataset.md概率建模、$n$ 元語(yǔ)法與長(zhǎng)序列采樣循環(huán)神經(jīng)網(wǎng)絡(luò)rnn.mdRNN 前向計(jì)算原理與困惑度循環(huán)神經(jīng)網(wǎng)絡(luò)的從零開(kāi)始實(shí)現(xiàn)rnn-scratch.md手寫(xiě) RNN 字符級(jí)語(yǔ)言模型循環(huán)神經(jīng)網(wǎng)絡(luò)的簡(jiǎn)潔實(shí)現(xiàn)rnn-concise.md用框架高級(jí) API 實(shí)現(xiàn)同一模型通過(guò)時(shí)間反向傳播bptt.md梯度計(jì)算、截?cái)嗯c穩(wěn)定性序列模型自回歸、馬爾可夫與 $k$ 步預(yù)測(cè)從統(tǒng)計(jì)工具看序列預(yù)測(cè)處理序列數(shù)據(jù)需要新的統(tǒng)計(jì)工具。以股票價(jià)格為例用 $x_t$ 表示時(shí)間步 $t$ 觀察到的價(jià)格交易員的目標(biāo)是估計(jì) $P(x_t \mid x_{t-1}, \ldots, x_1)$。直接回歸面臨核心矛盾輸入數(shù)量隨 $t$ 增長(zhǎng)而增長(zhǎng)。為此有兩種經(jīng)典策略自回歸模型autoregressive models假設(shè)足夠長(zhǎng)的歷史 $x_{t-1}, \ldots, x_1$ 并不必要只取長(zhǎng)度為 $\tau$ 的時(shí)間跨度 $x_{t-1}, \ldots, x_{t-\tau}$參數(shù)數(shù)量在 $t \tau$ 時(shí)保持恒定從而可以用普通深度網(wǎng)絡(luò)訓(xùn)練。隱變量自回歸模型latent autoregressive models保留對(duì)過(guò)去觀測(cè)的總結(jié) $h_t$同時(shí)更新預(yù)測(cè) $\hat{x}t$ 與總結(jié) $h_t$即 $\hat{x}t P(x_t \mid h_t)$ 且 $h_t g(h{t-1}, x{t-1})$。由于 $h_t$ 從未被觀測(cè)到故稱(chēng)為隱變量。這正是后續(xù) RNN 的雛形。馬爾可夫條件與一階模型若用 $x_{t-1}, \ldots, x_{t-\tau}$ 代替 $x_{t-1}, \ldots, x_1$ 來(lái)估計(jì) $x_t$ 足夠精確稱(chēng)序列滿足馬爾可夫條件Markov condition。當(dāng) $\tau1$ 時(shí)即一階馬爾可夫模型$$ P(x_1, \ldots, x_T) \prod_{t1}^T P(x_t \mid x_{t-1}) \text{ 當(dāng) } P(x_1 \mid x_0) P(x_1). $$當(dāng) $x_t$ 為離散值時(shí)可用動(dòng)態(tài)規(guī)劃沿馬爾可夫鏈精確計(jì)算例如 $P(x_{t1} \mid x_{t-1}) \sum_{x_t} P(x_{t1} \mid x_t) P(x_t \mid x_{t-1})$只需考慮很短的過(guò)去歷史。因果關(guān)系與靜止性在時(shí)間上存在自然的前進(jìn)方向未來(lái)事件不能影響過(guò)去。因此解釋 $P(x_{t1} \mid x_t)$ 比解釋 $P(x_t \mid x_{t1})$ 更容易正向估計(jì)通常也更可行。同時(shí)統(tǒng)計(jì)上假設(shè)序列的動(dòng)力學(xué)不變靜止stationary整個(gè)序列的聯(lián)合概率可分解為條件概率之積$$ P(x_1, \ldots, x_T) \prod_{t1}^T P(x_t \mid x_{t-1}, \ldots, x_1). $$若處理離散對(duì)象如單詞則用分類(lèi)器而非回歸模型來(lái)估計(jì)條件概率。動(dòng)手實(shí)驗(yàn)正弦序列的預(yù)測(cè)倉(cāng)庫(kù)文檔 sequence.md 用一個(gè)可復(fù)現(xiàn)的實(shí)驗(yàn)演示序列預(yù)測(cè)的難度生成 $T1000$ 個(gè)時(shí)間步的正弦加噪數(shù)據(jù)x d2l.sin(0.01 * time) d2l.normal(0, 0.2, (T,))取嵌入維度tau 4構(gòu)造特征標(biāo)簽對(duì)$y_t x_t$$\mathbf{x}_t [x_{t-\tau}, \ldots, x_{t-1}]$前 600 個(gè)樣本用于訓(xùn)練batch_size 16。訓(xùn)練模型用兩層全連接10 個(gè)隱藏單元ReLU 激活加平方損失的 MLPAdam 優(yōu)化器學(xué)習(xí)率 0.01訓(xùn)練 5 輪net nn.Sequential(nn.Linear(4, 10), nn.ReLU(), nn.Linear(10, 1)) net.apply(init_weights) loss nn.MSELoss(reductionnone) trainer torch.optim.Adam(net.parameters(), lr)評(píng)估兩種預(yù)測(cè)單步預(yù)測(cè)one-step-ahead prediction直接輸入真實(shí)觀測(cè)即使超出訓(xùn)練范圍n_train tau 604結(jié)果仍可信多步預(yù)測(cè)$k$-step-ahead-prediction一旦觀測(cè)止于 $x_{604}$后續(xù)預(yù)測(cè)必須使用自己上一步的預(yù)測(cè)結(jié)果遞歸外推$$ \hat{x}{605} f(x{601}, x_{602}, x_{603}, x_{604}),\quad \hat{x}{606} f(x{602}, x_{603}, x_{604}, \hat{x}_{605}),\ldots $$實(shí)驗(yàn)結(jié)果非常直觀多步預(yù)測(cè)經(jīng)過(guò)若干步后迅速衰減為常數(shù)。原因是誤差累積——步驟 1 引入誤差 $\epsilon_1$ 后步驟 2 的輸入被擾動(dòng)誤差按 $\epsilon_2 \bar\epsilon c\epsilon_1$ 遞歸放大。對(duì)比 $k1,4,16,64$ 步預(yù)測(cè)可以發(fā)現(xiàn)超過(guò) 4 步的預(yù)測(cè)幾乎無(wú)價(jià)值。這也呼應(yīng)了天氣預(yù)報(bào)24 小時(shí)內(nèi)較準(zhǔn)、再遠(yuǎn)精度驟降的現(xiàn)象并預(yù)告了本章后續(xù)RNN 及更復(fù)雜模型要解決的核心問(wèn)題。文本預(yù)處理從原始字符串到詞元索引文本是最常見(jiàn)的序列數(shù)據(jù)。預(yù)處理流水線通常包含四步加載文本、拆分為詞元、構(gòu)建詞表、轉(zhuǎn)換為數(shù)字索引序列。該流程的完整實(shí)現(xiàn)沉淀在 text-preprocessing.md 與 d2l/torch.py及 mxnet/tensorflow/paddle 各版本中。讀取數(shù)據(jù)集以 H. G. Wells 的《時(shí)光機(jī)器》The Time Machine為例僅 3 萬(wàn)多個(gè)單詞足夠小規(guī)模實(shí)驗(yàn)。read_time_machine將文本讀成文本行列表并忽略標(biāo)點(diǎn)與字母大小寫(xiě)d2l/torch.pydef read_time_machine(): #save 將時(shí)間機(jī)器數(shù)據(jù)集加載到文本行的列表中 with open(d2l.download(time_machine), r) as f: lines f.readlines() return [re.sub([^A-Za-z], , line).strip().lower() for line in lines]詞元化與詞表tokenize(lines, tokenword)將每條文本行拆分為詞元列表支持word按空白切分與char按字符切分兩種粒度d2l/torch.py。詞元是字符串而模型需要數(shù)字因此構(gòu)建**詞表vocabulary**將詞元映射到從 0 開(kāi)始的索引。Vocab類(lèi)的核心設(shè)計(jì)見(jiàn) text-preprocessing.md按出現(xiàn)頻率對(duì)詞元降序排序min_freq過(guò)濾低頻詞元以降低復(fù)雜度索引 0 固定為未知詞元unk語(yǔ)料中不存在或已被刪除的詞元統(tǒng)一映射到它可通過(guò)reserved_tokens預(yù)留填充詞元pad、序列開(kāi)始詞元bos、序列結(jié)束詞元eos等特殊詞元提供__getitem__詞元/詞元列表 → 索引/索引列表與to_tokens索引 → 詞元雙向映射。整合為load_corpus_time_machine為簡(jiǎn)化后續(xù)訓(xùn)練文檔將所有功能打包進(jìn)load_corpus_time_machine(max_tokens-1)d2l/torch.py并做了兩點(diǎn)調(diào)整改用字符級(jí)詞元化而非單詞因文本行不一定是完整句子將corpus展平為單個(gè)詞元索引列表。max_tokens 0時(shí)可截?cái)嗾Z(yǔ)料規(guī)模默認(rèn)參數(shù)下返回(corpus, vocab)。語(yǔ)言模型與數(shù)據(jù)集概率建模與長(zhǎng)序列采樣語(yǔ)言模型的目標(biāo)給定長(zhǎng)度 $T$ 的詞元序列 $x_1, x_2, \ldots, x_T$**語(yǔ)言模型language model**的目標(biāo)是估計(jì)聯(lián)合概率 $P(x_1, x_2, \ldots, x_T)$。它按鏈?zhǔn)椒▌t展開(kāi)為$$ P(x_1, x_2, \ldots, x_T) \prod_{t1}^T P(x_t \mid x_{t-1}, \ldots, x_1). $$一個(gè)理想的語(yǔ)言模型可基于自身生成自然文本即使達(dá)不到理解它也能消除語(yǔ)音識(shí)別中的歧義to recognize speech 與 to wreck a nice beach或判斷我想吃奶奶與我想吃奶奶的正常程度。$n$ 元語(yǔ)法與拉普拉斯平滑直接統(tǒng)計(jì)計(jì)數(shù)估計(jì)條件概率會(huì)遇到稀疏性問(wèn)題二元組 deep learning 的出現(xiàn)頻率遠(yuǎn)低于單詞 deep三元及以上組合更是大量存在于合理語(yǔ)言中卻在數(shù)據(jù)集中缺席。經(jīng)典補(bǔ)救是拉普拉斯平滑Laplace smoothing在計(jì)數(shù)中加入小常量例如$$ \hat{P}(x) \frac{n(x) \epsilon_1/m}{n \epsilon_1}, $$其中 $n$ 為訓(xùn)練集單詞總數(shù)$m$ 為唯一單詞數(shù)$\epsilon_1 0$ 表示不平滑$\epsilon_1 \to \infty$ 時(shí) $\hat{P}(x)$ 趨向均勻分布 $1/m$。但這類(lèi)模型存在根本缺陷需存儲(chǔ)所有計(jì)數(shù)、完全忽略單詞語(yǔ)義、長(zhǎng)序列幾乎從未出現(xiàn)而表現(xiàn)不佳。馬爾可夫假設(shè)將依賴截?cái)酁楣潭A數(shù)一元語(yǔ)法unigram、二元語(yǔ)法bigram、三元語(yǔ)法trigram分別只依賴 0、1、2 個(gè)前驅(qū)詞元。齊普夫定律為什么統(tǒng)計(jì)平滑不可行在時(shí)光機(jī)器語(yǔ)料上統(tǒng)計(jì)發(fā)現(xiàn)最常用詞基本是停用詞stop words且詞頻衰減極快——第 10 高頻詞的頻率不足第 1 高頻詞的 1/5。雙對(duì)數(shù)坐標(biāo)下詞頻近似直線即滿足齊普夫定律Zipfs law$$ n_i \propto \frac{1}{i^\alpha}, \quad \log n_i -\alpha \log i c. $$一元、二元、三元語(yǔ)法均服從該規(guī)律只是指數(shù) $\alpha$ 不同。這帶來(lái)三個(gè)結(jié)論計(jì)數(shù)平滑建模會(huì)高估尾部低頻詞的頻率$n$ 元組總量并不巨大說(shuō)明語(yǔ)言存在大量可利用的結(jié)構(gòu)大量 $n$ 元組極少出現(xiàn)使拉普拉斯平滑不適合語(yǔ)言建?!虼诵枰谏疃葘W(xué)習(xí)的模型如 RNN。讀取長(zhǎng)序列隨機(jī)采樣與順序分區(qū)序列本質(zhì)連續(xù)需將任意長(zhǎng)的語(yǔ)料切分為固定長(zhǎng)度如num_steps個(gè)時(shí)間步的小批量子序列。language-models-and-dataset.md 給出兩種策略兩者都從隨機(jī)偏移量開(kāi)始切分以兼顧覆蓋性與隨機(jī)性隨機(jī)采樣random samplingseq_data_iter_random(corpus, batch_size, num_steps)d2l/torch.py先從random.randint(0, num_steps - 1)偏移處開(kāi)始把序列切為長(zhǎng)度num_steps的子序列并random.shuffle起始索引使相鄰小批量的子序列在原始序列上不一定相鄰標(biāo)簽Y是特征X移位一個(gè)詞元的結(jié)果。以序列range(35)、batch_size2、num_steps5為例可生成 $\lfloor (35-1)/5 \rfloor 6$ 個(gè)特征標(biāo)簽對(duì)即 3 個(gè)小批量。順序分區(qū)sequential partitioningseq_data_iter_sequentiald2l/torch.py保證相鄰兩個(gè)小批量中的子序列在原始序列上相鄰適合訓(xùn)練時(shí)需要跨批量延續(xù)隱狀態(tài)的場(chǎng)景。兩者再被包裝進(jìn)SeqDataLoader迭代器并由load_data_time_machine(batch_size, num_steps, use_random_iterFalse, max_tokens10000)d2l/torch.py統(tǒng)一返回?cái)?shù)據(jù)迭代器與詞表默認(rèn)max_tokens10000限制語(yǔ)料規(guī)模、默認(rèn)采用順序分區(qū)。循環(huán)神經(jīng)網(wǎng)絡(luò)原理隱狀態(tài)與參數(shù)共享從無(wú)隱狀態(tài)到有隱狀態(tài)回顧單隱藏層 MLP給定小批量 $\mathbf{X} \in \mathbb{R}^{n \times d}$隱藏層輸出$$ \mathbf{H} \phi(\mathbf{X} \mathbf{W}_{xh} \mathbfh), \qquad \mathbf{O} \mathbf{H} \mathbf{W}{hq} \mathbf_q. $$RNN 的關(guān)鍵區(qū)別在于引入隱狀態(tài)時(shí)間步 $t$ 的隱變量不僅依賴當(dāng)前輸入還依賴前一時(shí)間步的隱變量新增權(quán)重 $\mathbf{W}_{hh} \in \mathbb{R}^{h \times h}$$$ \mathbf{H}t \phi(\mathbf{X}t \mathbf{W}{xh} \mathbf{H}{t-1} \mathbf{W}_{hh} \mathbf_h), \qquad \mathbf{O}_t \mathbf{H}t \mathbf{W}{hq} \mathbf_q. $$需要特別區(qū)分兩個(gè)易混淆概念隱藏層是從輸入到輸出路徑上以觀測(cè)角度理解的層隱狀態(tài)是給定步驟所做任何事情的輸入只能通過(guò)先前時(shí)間步的數(shù)據(jù)計(jì)算。由于當(dāng)前步隱狀態(tài)的定義與前一步相同計(jì)算是循環(huán)的recurrent執(zhí)行該計(jì)算的層稱(chēng)為循環(huán)層。兩個(gè)關(guān)鍵性質(zhì)參數(shù)共享循環(huán)神經(jīng)網(wǎng)絡(luò)在不同時(shí)間步始終復(fù)用同一組參數(shù)$\mathbf{W}{xh}, \mathbf{W}{hh}, \mathbfh, \mathbf{W}{hq}, \mathbf_q$因此參數(shù)開(kāi)銷(xiāo)不隨時(shí)間步增加而增加——這正是它相對(duì) $n$ 元語(yǔ)法需存儲(chǔ) $|\mathcal{V}|^n$ 個(gè)數(shù)字的核心優(yōu)勢(shì)。拼接等價(jià)性$\mathbf{X}t \mathbf{W}{xh} \mathbf{H}{t-1} \mathbf{W}{hh}$ 等價(jià)于把輸入與隱狀態(tài)沿列拼接、把兩個(gè)權(quán)重沿行拼接后再做矩陣乘法。文檔用隨機(jī)矩陣X(3,1)、W_xh(1,4)、H(3,4)、W_hh(4,4)驗(yàn)證兩種寫(xiě)法得到相同形狀 $(3,4)$ 的輸出。字符級(jí)語(yǔ)言模型與困惑度RNN 的典型應(yīng)用是字符級(jí)語(yǔ)言模型輸入序列 machin 預(yù)測(cè)標(biāo)簽 achine逐時(shí)間步輸出經(jīng) softmax 后與標(biāo)簽計(jì)算交叉熵?fù)p失第 3 個(gè)時(shí)間步的輸出由 m、a、c 共同決定。實(shí)踐中批量大小為 $n1$、每個(gè)詞元用 $d$ 維向量表示時(shí)間步 $t$ 的輸入 $\mathbf{X}_t$ 為 $n \times d$ 矩陣。評(píng)估語(yǔ)言模型質(zhì)量用困惑度perplexity——平均交叉熵的指數(shù)$$ \exp\left(-\frac{1}{n} \sum_{t1}^n \log P(x_t \mid x_{t-1}, \ldots, x_1)\right). $$它可理解為下一個(gè)詞元實(shí)際選擇數(shù)的調(diào)和平均數(shù)完美預(yù)測(cè)時(shí)困惑度為 1均勻基線時(shí)困惑度等于詞表唯一詞元數(shù)這是無(wú)壓縮存儲(chǔ)的理論上限任何實(shí)際模型必須超越預(yù)測(cè)概率為 0 時(shí)困惑度趨于無(wú)窮。較短的序列更可能、絕對(duì)似然難以跨文檔比較而困惑度使不同長(zhǎng)度文檔可比。從零實(shí)現(xiàn) RNN 字符級(jí)語(yǔ)言模型rnn-scratch.md 從零手寫(xiě)完整訓(xùn)練流程訓(xùn)練配置為batch_size32, num_steps35數(shù)據(jù)經(jīng)load_data_time_machine讀取。獨(dú)熱編碼詞元索引 $i$ 映射為長(zhǎng)度len(vocab)的全 0 向量、第 $i$ 位為 1。小批量批量大小時(shí)間步數(shù)經(jīng)one_hot轉(zhuǎn)成三維張量并轉(zhuǎn)置為**時(shí)間步數(shù)批量大小詞表大小**以便沿最外層維度逐步更新隱狀態(tài)。參數(shù)初始化get_params(vocab_size, num_hiddens, device)返回五組參數(shù)——隱藏層 $\mathbf{W}{xh} \in \mathbb{R}^{\text{vocab} \times h}$、$\mathbf{W}{hh} \in \mathbb{R}^{h \times h}$、$\mathbfh$輸出層 $\mathbf{W}{hq} \in \mathbb{R}^{h \times \text{vocab}}$、$\mathbf_q$輸入與輸出共享同一詞表維度用標(biāo)準(zhǔn)差 0.01 的正態(tài)分布初始化。前向計(jì)算rnn函數(shù)逐時(shí)間步執(zhí)行 $\mathbf{H}t \tanh(\mathbf{X}t \mathbf{W}{xh} \mathbf{H}{t-1} \mathbf{W}_{hh} \mathbf_h)$輸出層 $\mathbf{O}_t \mathbf{H}t \mathbf{W}{hq} \mathbf_q$init_rnn_state返回形狀為批量大小隱藏單元數(shù)的全零初始隱狀態(tài)。梯度裁剪訓(xùn)練中每批更新前調(diào)用d2l.grad_clippingd2l/torch.py將梯度范數(shù)裁剪到閾值 $\theta$默認(rèn) 1防止梯度爆炸。代碼中可見(jiàn)param.grad.data被統(tǒng)一縮放這正是后面 bptt.md 要解釋的工程化處理。預(yù)測(cè)與困惑度predict_ch8在給定前綴如 time traveller 后逐字符生成文本先獨(dú)熱編碼再迭代隱狀態(tài)取概率最高的索引作為下一個(gè)字符evaluate_metric用困惑度評(píng)估訓(xùn)練打印的指標(biāo)即困惑度值。實(shí)驗(yàn)結(jié)果在《時(shí)光機(jī)器》上訓(xùn)練后模型能生成類(lèi)似 the time machine by h. g. wells 的文本片段當(dāng)use_random_iterTrue隨機(jī)采樣時(shí)因無(wú)法跨批量延續(xù)隱狀態(tài)訓(xùn)練稍慢但困惑度仍持續(xù)下降。簡(jiǎn)潔實(shí)現(xiàn)高級(jí) API 一行構(gòu)造 RNNrnn-concise.md 用框架高級(jí) API 復(fù)現(xiàn)同一模型代碼量大幅縮減num_hiddens 256 rnn_layer nn.RNN(len(vocab), num_hiddens) # PyTorch輸入維度為詞表大小 state torch.zeros((1, batch_size, num_hiddens)) # (層數(shù), 批量大小, 隱藏單元數(shù)) Y, state_new rnn_layer(X, state)要點(diǎn)各框架對(duì)應(yīng)構(gòu)造分別為 MXNetgluon.rnn.RNN(num_hiddens)、TensorFlowSimpleRNNCellkeras.layers.RNN(time_majorTrue, return_sequencesTrue, return_stateTrue)、Paddlenn.SimpleRNN(len(vocab), num_hiddens, time_majorTrue)隱狀態(tài)形狀統(tǒng)一為隱藏層數(shù)批量大小隱藏單元數(shù)rnn_layer的輸出Y并不含輸出層計(jì)算——它返回每個(gè)時(shí)間步的隱狀態(tài)供后續(xù)nn.Dense(vocab_size)輸出層使用state_new是最后時(shí)間步的隱狀態(tài)可用于順序分區(qū)下一個(gè)小批量的隱狀態(tài)初始化剩余部分嵌入層、輸出層、損失、預(yù)測(cè)、訓(xùn)練循環(huán)與從零實(shí)現(xiàn)高度一致最終困惑度同樣穩(wěn)定下降驗(yàn)證了手寫(xiě)實(shí)現(xiàn)的正確性。通過(guò)時(shí)間反向傳播BPTT梯度如何流動(dòng)與截?cái)郻ptt.md 深入序列模型的梯度計(jì)算解釋此前反復(fù)出現(xiàn)的梯度爆炸梯度消失分離梯度等概念。核心思想BPTT 就是反向傳播在 RNN 上的特定應(yīng)用——把計(jì)算圖按時(shí)間步展開(kāi)一次基于鏈?zhǔn)椒▌t反向計(jì)算并存儲(chǔ)梯度。簡(jiǎn)化模型下 $h_t f(x_t, h_{t-1}, w_h)$目標(biāo) $L \frac{1}{T}\sum_{t1}^T l(y_t, o_t)$。求 $\partial L / \partial w_h$ 時(shí)遇到遞歸項(xiàng)$$ \frac{\partial h_t}{\partial w_h} \frac{\partial f(x_t, h_{t-1}, w_h)}{\partial w_h} \sum_{i1}^{t-1}\left(\prod_{ji1}^{t} \frac{\partial f(x_j, h_{j-1}, w_h)}{\partial h_{j-1}}\right) \frac{\partial f(x_i, h_{i-1}, w_h)}{\partial w_h}. $$當(dāng) $t$ 很大時(shí)鏈極長(zhǎng)且矩陣高次冪$\mathbf{W}_{hh}^\top$ 的 $T-i$ 次冪導(dǎo)致小于 1 的特征值使梯度消失大于 1 的特征值使梯度發(fā)散爆炸。四種梯度計(jì)算策略完全計(jì)算計(jì)算上式全部總和。理論上精確但速度極慢且易梯度爆炸初始條件微小變化引發(fā)蝴蝶效應(yīng)實(shí)踐中幾乎不用截?cái)鄷r(shí)間步常規(guī)截?cái)嘣?$\tau$ 步后終止求和即把 $\partial h_{t-\tau}/\partial w_h$ 當(dāng)作零——這正是從零實(shí)現(xiàn)中detach分離梯度的本質(zhì)。它近似真實(shí)梯度、計(jì)算可行且將估計(jì)偏向更簡(jiǎn)單更穩(wěn)定的模型側(cè)重短期影響是實(shí)踐主流隨機(jī)截?cái)嘤闷谕_的隨機(jī)變量 $\xi_t$$P(\xi_t 0) 1-\pi_t$$P(\xi_t \pi_t^{-1}) \pi_t$$E[\xi_t] 1$替換遞歸項(xiàng)實(shí)現(xiàn)不同長(zhǎng)度序列的加權(quán)和。理論上吸引人但實(shí)踐中常不比常規(guī)截?cái)喔枚谭秶聪騻鞑ヒ炎阋圆东@依賴、方差增加抵消長(zhǎng)梯度精確性、短范圍交互恰是理想模型性質(zhì)實(shí)際訓(xùn)練中通常以給定數(shù)量時(shí)間步后分離梯度的方式實(shí)現(xiàn)截?cái)唷T敿?xì)的梯度推導(dǎo)在忽略偏置、恒等激活的線性 RNN 中隱狀態(tài)與輸出為 $\mathbf{h}t \mathbf{W}{hx}\mathbf{x}t \mathbf{W}{hh}\mathbf{h}_{t-1}$、$\mathbf{o}t \mathbf{W}{qh}\mathbf{h}_t$。沿計(jì)算圖反向遍歷可得$$ \frac{\partial L}{\partial \mathbf{W}{qh}} \sum{t1}^T \frac{\partial L}{\partial \mathbf{o}_t} \mathbf{h}_t^\top,\quad \frac{\partial L}{\partial \mathbf{h}T} \mathbf{W}{qh}^\top \frac{\partial L}{\partial \mathbf{o}_T}, $$任意 $t T$ 時(shí)隱狀態(tài)梯度遞歸為$$ \frac{\partial L}{\partial \mathbf{h}t} \mathbf{W}{hh}^\top \frac{\partial L}{\partial \mathbf{h}{t1}} \mathbf{W}{qh}^\top \frac{\partial L}{\partial \mathbf{o}_t}. $$展開(kāi)后梯度陷入 $\left(\mathbf{W}{hh}^\top\right)^{T-i}$ 的高次冪——特征值小于 1 消失、大于 1 發(fā)散數(shù)值上即梯度消失/爆炸。BPTT 交替進(jìn)行前向傳播與反向傳播并緩存中間值如 $\partial L/\partial \mathbf{h}t$供 $\partial L/\partial \mathbf{W}{hx}$ 與 $\partial L/\partial \mathbf{W}{hh}$ 復(fù)用避免重復(fù)計(jì)算。更復(fù)雜的序列模型如 LSTM正是從架構(gòu)層面進(jìn)一步緩解這一問(wèn)題——這將在后續(xù)章節(jié)chapter_recurrent-modern 下的 lstm、gru 等展開(kāi)。小結(jié)與閱讀路徑本章在《動(dòng)手學(xué)深度學(xué)習(xí)》中處于承上啟下的位置承上復(fù)用此前 MLP、反向傳播、softmax 與信息論知識(shí)將建模對(duì)象從靜態(tài)數(shù)據(jù)擴(kuò)展到序列數(shù)據(jù)啟下RNN 的隱狀態(tài)與 BPTT 梯度問(wèn)題是后續(xù) chapter_recurrent-modern 中 GRU、LSTM、深度 RNN、雙向 RNN、機(jī)器翻譯等更復(fù)雜模型的理論基石。核心結(jié)論可歸納為序列數(shù)據(jù)打破 i.i.d. 假設(shè)RNN 通過(guò)循環(huán)計(jì)算的隱狀態(tài)捕獲直到當(dāng)前時(shí)間步的歷史信息且參數(shù)數(shù)量不隨時(shí)間步增長(zhǎng)文本預(yù)處理 加載 → 詞元化 → 詞表映射 → 數(shù)字索引序列倉(cāng)庫(kù)提供read_time_machine、tokenize、Vocab、load_corpus_time_machine等可復(fù)用 API語(yǔ)言模型估計(jì)聯(lián)合概率$n$ 元語(yǔ)法受齊普夫定律與稀疏性制約長(zhǎng)序列讀取以隨機(jī)采樣與順序分區(qū)為主RNN 可構(gòu)建字符級(jí)語(yǔ)言模型用困惑度評(píng)估質(zhì)量BPTT 是鏈?zhǔn)椒▌t在時(shí)間維度上的應(yīng)用矩陣高次冪導(dǎo)致梯度爆炸/消失工程上以梯度裁剪與時(shí)間步截?cái)嗷獠⒕彺嬷虚g值保證效率。相關(guān)實(shí)現(xiàn)證據(jù)可繼續(xù)在倉(cāng)庫(kù)中查閱d2l/torch.pyPyTorch 實(shí)現(xiàn)與 d2l/mxnet.py、d2l/tensorflow.py、d2l/paddle.py 對(duì)應(yīng)、本章各小節(jié)文檔以及各框架下#save標(biāo)記的復(fù)用函數(shù)——它們就是本文全部代碼示例的權(quán)威出處。贊分享人工智能深度學(xué)習(xí)機(jī)器學(xué)習(xí)教程【免費(fèi)下載鏈接】d2l-zh《動(dòng)手學(xué)深度學(xué)習(xí)》面向中文讀者、能運(yùn)行、可討論。中英文版被70多個(gè)國(guó)家的500多所大學(xué)用于教學(xué)。項(xiàng)目地址https://gitcode.com/GitHub_Trending/d2/d2l-zh點(diǎn)擊查看免費(fèi)下載相關(guān)推薦基于 PyTorch 的 DDP 分布式訓(xùn)練實(shí)戰(zhàn)用 torchrun 編排多節(jié)點(diǎn)多卡應(yīng)用基于 PyTorch 的 DDP 分布式訓(xùn)練實(shí)戰(zhàn)用 torchrun 編排多節(jié)點(diǎn)多卡應(yīng)用 本篇技術(shù)指南圍繞當(dāng)前倉(cāng)庫(kù) distributed/ddp 示例系統(tǒng)人工智能深度學(xué)習(xí)機(jī)器學(xué)習(xí)教程循環(huán)神經(jīng)網(wǎng)絡(luò)終極指南斯坦福CS229深度學(xué)習(xí)手冊(cè)中的序列模型實(shí)戰(zhàn)解析循環(huán)神經(jīng)網(wǎng)絡(luò)終極指南斯坦福CS229深度學(xué)習(xí)手冊(cè)中的序列模型實(shí)戰(zhàn)解析 斯坦福CS229機(jī)器學(xué)習(xí)課程的VIP手冊(cè)為深度學(xué)習(xí)愛(ài)好者提供了全面的理論與實(shí)踐指導(dǎo)其中文檔教程機(jī)器學(xué)習(xí)深度學(xué)習(xí)500問(wèn)第六章精讀循環(huán)神經(jīng)網(wǎng)絡(luò)RNN結(jié)構(gòu)圖解、BPTT 推導(dǎo)與 LSTM/GRU 變體實(shí)戰(zhàn)指南深度學(xué)習(xí)500問(wèn)第六章精讀循環(huán)神經(jīng)網(wǎng)絡(luò)RNN結(jié)構(gòu)圖解、BPTT 推導(dǎo)與 LSTM/GRU 變體實(shí)戰(zhàn)指南 導(dǎo)讀本文以《深度學(xué)習(xí)500問(wèn)》第六章為骨架系統(tǒng)深度學(xué)習(xí)機(jī)器學(xué)習(xí)計(jì)算機(jī)視覺(jué)NLP教程知識(shí)庫(kù)上一篇k-skill 實(shí)戰(zhàn)韓國(guó)高速公路實(shí)時(shí)路況與 CCTV 查詢技能 highway-traffic-status 深度解析下一篇Terraform Provider AWS 權(quán)限自檢實(shí)戰(zhàn)深入解析 aws_iam_principal_policy_simulation 數(shù)據(jù)源創(chuàng)作聲明:本文部分內(nèi)容由AI輔助生成(AIGC),僅供參考