RNN と LSTM 入門

2017 年より前、系列モデルは状態を運ぶものだった。ここではそれを分解する:状態を計算するループ、数十ステップ戻ったところでそれを殺す導関数の積、その積を生かすために作られたゲート、そしてこの一族を終わらせた量である。

01

トークンを 1 個ずつ

再帰型ネットにはループ 1 つと記憶 1 つしかない。以降のすべてはそこから出てくる。

左から右へ読み、ベクトルを 1 本だけ持ち続け、各ステップでそのベクトルと新しいトークンを混ぜて押し込む:h_t = tanh(W_h·h_{t−1} + W_x·x_t + b)。ここでは 1 次元で走らせるので、どのフレームも手で検算できる。

8 個のトークンは与えられているので最初から灰色だ。すでに計算された状態がその上に並び、ループがいま計算している状態があなたの手の中にある。スライダーで t を進めよう:

t = 0 · h_0 = 0.565

スライダーの右側には何も存在しないことに注目してほしい。あのセルが破線なのは、h_4 をネットが出し惜しみしているからではなく、誰も計算していないからだ。それがこの不変条件で、断言として書く価値がある:各ステップの先頭でh_t は x_0 … x_t だけの関数であり、その先には依存しない。

押し込みは飾りではない。どんな事前活性も −1 から 1 の開区間に折り畳み、いま立っている点でのその傾きこそ、次の節が丸ごと扱うものだ。点を曲線に沿ってドラッグしよう:

z = 0.00 · h = 0.000 · tanh' = 1.000. 左右にドラッグ。矢印キーで 1 段ずつ、Home で最初の状態に戻る
z = 0.00 · h = 0.000 · tanh' = 1.000

接線を見てほしい。z = 0 で最も急になり、読み出しは tanh' = 1.000 を返す。z = 2 では 0.071 まで平らになる。飽和したユニットとは、出力が入力に応えなくなったユニットであり、勾配が戻らなくなったユニットでもある。

W_h は 1 つしかない。同じ数がすべての状態を次の状態へ掛けていく。だから 1 セットの重みでどんな長さの文も読める。動かせば 8 個が一斉に動く —— を試してほしい:

W_h = 0.90 · h_7 = 0.412

重みの共有は、パラメータ数を系列長から切り離す:再帰に d×d 行列 1 つ、入力に 1 つ、バイアス 1 つ。 1.40 では最後の状態が 0.841 を示し、 では 0.310 —— ちょうど tanh(0.8 × 0.4)、自分のトークンだけになる。

答えはそれらの状態から、もう 1 つの行列を通って出てくる。分類は最後の状態から出力 1 個を読み、系列ラベリングは位置ごとに 1 個ずつ読む。切り替えよう:

1 個の答え、最後の状態から

どちらでも、読み出しヘッドが見るのはその真下の状態だけだ。最初のトークンが最後の答えに寄与するものはすべて、W_h を 7 回と押し込みを 7 回くぐって届かねばならない。その道中で何が起きるかが次節だ。

02

勾配がループの中で死ぬ理由

学習は鎖を逆にたどり、ホップごとに掛け算する。1 未満の数を大量に掛けた積は小さい数ではなく、もはや数ではない。

最後のトークンから学ぶには、損失が最初のトークンまで届かねばならない。同じ鎖を逆向きに進み、各ホップで ∂h_t/∂h_{t−1} = tanh'(z_t) · W_h が掛かる。灰色の曲線がその前半の傾きで、バーがホップ 1 回ぶん —— 傾きに重みを掛けたものだ:

z = 0.00 · tanh' = 1.000 · × W_h = 0.900

傾きが 1 になるのは z = 0 だけで、両側で落ちていく。だからバーはどこでも W_h より小さく、ユニットが働いている場所でははるかに小さい。これが 1 ホップだ。勾配を 8 ホップ戻そう:

0 ホップ · 勾配はまだ 1

7 ホップで 1 から 0.057 になる。何も壊れてはいない —— ここの係数はどれも正常なネットの正常な導関数だ。問題は、それが 7 個あることと、文が 8 語より長いことだけである。

では係数だけを、40 ステップの鎖の上で見よう。破線が × 1、縮むか伸びるかの境界だ。W_h をドラッグし、活性化関数を切り替えよう:

W_h = 0.90 · 最大の係数 0.884

tanh' は 1 を超えないので係数も W_h を超えない。既定の 0.90 では最大が 0.884 だ。W_h を まで上げてもバーは破線に届かない —— 飽和がブレーキだ。 ReLU に切り替えるとブレーキは消える。

その係数を掛け合わせたもの —— 今度は 200 ステップの鎖で —— がh_0 へ届く数だ。縦軸は対数で、罫線 1 本ごとに 24 桁上がる。だから一定の係数は直線として描かれる:

W_h = 0.90 · 200 ホップ後 9.8e-48

出荷時の設定では、積は 161 ホップ目で下側の破線(float32 の最小正規数)を割り込み、9.8e-48 で終わる。ReLU に切り替えて にすると逆に上がり続け、189 ホップ目で上側の破線を越える。それを越えた float32 は無限大であり、無限大に触れた損失は NaN だ。

1 次元ではブレーキが必ず勝つ。多次元では、W_h の最大特異値が誰も飽和していない方向に効き、積は本当に暴走する。標準的な歯止めは Pascanu・Mikolov・Bengio の 2013 年のスケーリングで、このステップのクリップ前のノルムは 84.0、しきい値は 5 だ —— しきい値の下までドラッグして戻そう:

ノルム 84.0 → 5.00

クリップは天井であって床ではないので、もう半分には何もしない。消失は静かに失敗する:例外も NaN も出ず、損失は普通に下がり、モデルは直近数トークンだけから予測することを黙って覚える —— そこだけが勾配の届いた証拠だからだ。

03

この積を生かしておくゲート

LSTM は掛け算を小さくしたのではない。ほとんど何も乗っていない第 2 の経路を作ったのだ。

Hochreiter と Schmidhuber の 1997 年の答えは、第 2 の記憶 c を足すことだった。その更新は c_t = f · c_{t−1} + i · ĉ_t —— 数 1 個の掛け算と足し算だけ。c_{t−1} と c_t の間には重み行列も押し込みもない。

絵は §01 のものと部品 1 つだけが違う。トークンの行も、セルも同じで、違うのはつなぎ方だけだ:セル状態は素の線を進み、このステップが書く分が下から入ってくる。前へ歩かせてほしい:

t = 0 · c_0 = 0.415

このつなぎ方を §01 の矢印と比べてほしい。あちらのホップは tanh'(z)·W_h という、ネットが制御しづらい 2 つの積だった。こちらは忘却ゲートという、ネットが各ステップで意図して計算する数 1 個だ。この置き換えがアイデアのすべてである。

どちらのゲートも割合なので、1 回の更新は長さ 2 本を継いだものになる:古い状態のうち生き残る分、候補のうち書き込まれる分、そして捨てられる残り。忘却ゲート、次に入力ゲートを動かしてほしい:

f = 0.75 · c = 0.698

これを にすると捨てられる残りが消え、セル状態は純粋な累算器になる。 0 にすれば古い状態は 1 ステップで消える —— 文が終わり、次の文がそれと無関係なとき、ネットが欲しいのはまさにこれだ。

これが §02 を解く理由は、c_0 から c_k への経路が忘却ゲートの積だけ —— f^k だけ —— だからだ。対比として、前節の素の RNN の積を同じ線形軸にローズで重ねてある。f をドラッグしよう:

f = 0.950 · 100 ステップ後に 0.6% 残る

f = 0.95 なら 100 ステップ後にも 0.006 が残る。小さいが最適化器が使える数だ。RNN の積は 9 ステップ目で 0.01 を下回り、あとはずっと横軸に貼り付く。しかも忘却ゲートは定数ではなく、ネットが次元ごとに学ぶ。

古典的なバグはまさにそこにある。初期化時の重みはほぼ 0 なので f はバイアスの言い値になり、バイアス 0 はセルが毎ステップ自分を半分にすることを意味する。曲線は 20 ステップ後に c_0 がどれだけ残るかで、横軸はそのバイアスだ。縦軸は対数で、8 桁を収めてある:

b_f = 0.0 · f = 0.500 · 20 ステップ後 9.5e-07

既定のフレームがそのバグそのものだ。bias = 0 は f = 0.500 を与え、20 ステップ後に 9.5e-07 —— 自分が直すはずだった勾配消失を生まれつき抱えた LSTM だ。 までドラッグすれば、同じ 20 ステップで 0.002 が残る。b_f = 0 と書けば理由の見えないまま学習が悪くなり、b_f = 1 と書けば —— Jozefowicz・Zaremba・Sutskever、2015 —— そうはならない。

04

ゲートが直さなかったもの

ゲートは勾配を解いた。再帰の時代を終わらせた 2 つの性質には手を付けていない。

1 つ目は容量だ。セル状態がどれだけよく生き延びても、それは固定幅のベクトル 1 本であり、ゲートはそれを大きくしない。

下の状態は幅 512 個の数で、変わることはない。その上にあるのが要約すべきトークンで、最初のトークンの取り分が左端のブロックだ。右へドラッグして系列を伸ばそう:

n = 4 トークン · 1 トークンあたり 128.00 個の数。左右にドラッグ。矢印キーで 1 段ずつ、Home で最初の状態に戻る
n = 4 トークン · 1 トークンあたり 128.00 個の数

4 トークンなら取り分は 128 個の数だ。 では 2.00 になり、上の目盛りは溶けて模様になる。これは調整で消せるバグではない ——「固定長ベクトルで前文を要約する」とはこういう意味であり、300 トークン前の具体的な事実を LSTM に尋ねると、正しくではなくもっともらしく答える理由でもある。

2 つ目はもっと悪く、モデルではなく機械の話だ。行が位置、列が実時間ステップで、その位置の仕事が起きたときにセルが点く。ステップを進めて、この格子のどれだけが実際に働いているか見よう:

8 実時間ステップ中の 1 番目 · 同時に 1 位置

点くのは対角線だけだ。位置 5 は位置 4 が終わるまで始められないので、 8 位置に8 個の逐次ステップがかかり、最後まで進めても 64 マス中 8 マスしか働かない。GPU には数万のレーンがある。別々の系列をバッチにすればいくらかは埋まるが、1 つの系列の内側の仕事は鎖であり、鎖には幅がない。

GRU(Cho ら、2014)は人気の節約版だ。ゲートは 3 つでなく 2 つで、しかも入力ゲートは自由ではない —— 1 − f に固定され、残すものと書くものが 1 つの判断になる。ドラッグしよう:

f = 0.60 · c = 0.686

2 本の長さの合計がトラック全体でなければならないので、GRU は同じステップで古い状態を保ちつつ強く書くことができない。LSTM にはそれができ、余分なゲートが買うのはその 1 点だけだ。節約されるのはゲート 1 つぶんの行列 —— 同じ予算の上で、1 本を残り 2 本と比べてほしい:

d = 1,024 · LSTM は 1 層 8.39 M

どのセルも 2d² + d ブロックの整数倍だ。RNN は 1 つ、GRU は 3 つ、LSTM は 4 つ —— d = 1024 なら 2.10 M、6.29 M、8.39 M。上の 2 点はそれで変わらない。状態は固定長ベクトル 1 本のままで、ループはループのままだ。

05

ループを捨てて得たもの

3 つの量が変わった。そのどれもが Vaswani ら 2017 年の Table 1 の 1 列だ。

自己アテンションは各位置を全位置から一度に計算する。全ペアにスコアを付け、 softmax を取り、重み付き和を作る。運ぶ状態がないので、どの位置も他を待たずに始められる。

§04 の格子にスイッチを付けたものだ。再帰層は 1 列につき1 マスしか点かないが、アテンションは1 列を一度に点ける。系列を伸ばしてから、層を切り替えよう:

n = 8 · 逐次ステップ 8 個

逐次の深さとは中身のある列の数で、左では n、右では1 だ。 GPU が幅の数パーセントしか使えない状態と使い切る状態の差であり、 2017 年の base Transformer が P100 8 枚で 12 時間で学習できた一方、それが破った LSTM 系が数日を要した理由でもある。

2 つ目の量は距離だ。段 1 つが 2 つのトークン間の掛け算 1 回なので、経路の高さがそのまま両者の間に挟まる掛け算の数になる。目標トークンをドラッグしてから、層を切り替えよう:

token 0 → token 12 · 12 ホップ。左右にドラッグ。矢印キーで 1 段ずつ、Home で最初の状態に戻る
token 0 → token 12 · 12 ホップ

層を切り替えて、高さが潰れるのを見てほしい。再帰は token 0 と token k の間に k 回の掛け算 —— §02 の主題そのものの積 —— を置くが、アテンションはそこに辺 1 本を置く。k がいくつでもだ。最大経路長は O(n) から O(1) になり、減衰もそれとともに消える。

3 つ目はデコーダが何を見てよいかだ。RNN エンコーダは最後の状態だけを渡すので、ソース文全体がベクトル 1 本を通り抜けねばならない。アテンションは各出力位置が全入力位置を読むことを許す。つなぎ方を切り替えよう:

n = 12 トークン · デコーダは 512 個の数を読む(d = 512)

ソース 12 トークンでは、ボトルネックはソース長に関わらず 512 個しか渡さない。直結は 6,144 個、24 トークンなら 12,288 個だ。 Bahdanau・Cho・Bengio は 2014 年にこれを見て RNN にアテンションを継ぎ足し、 2017 年はアテンションを残して RNN を消した。

ただではない。再帰層は n·d² 回の積和 —— d×d 行列を n ステップ —— を行い、アテンション層はペアごとに 1 スコアで n²·d 回を行う。下の軸は両方とも対数なので両方とも直線になり、アテンションのほうが急だ。d をドラッグしよう:

d = 512 · n = 512 で交差

急なほうは n = d で緩やかなほうと交わる。両者の代価が等しくなる点だ。 では 512 トークン未満で安く、それより上では高い。

06

リファレンスと、戻り道

セルの全文、ループの費用、そして実際にはどこまで戻して学習するか。

# keep · write · what · show
f = sigmoid(W_f[h,x] + b_f)
i = sigmoid(W_i[h,x] + b_i)
g = tanh(W_g[h,x] + b_g)
o = sigmoid(W_o[h,x] + b_o)

# the only path back to c_0
c = f * c + i * g
h = o * tanh(c)

このループが先頭まで学習されることはない。時間方向の逆伝播は鎖上の活性をすべて保持するので、学習は窓に切り詰める。窓の中のホップだけが勾配を受け取り、切断より先は何も受け取らない。窓をドラッグしよう:

窓 3 · 勾配は h_4 まで、h_0 には届かない

切断が奪うのは情報ではない —— 順方向は状態をそのまま運ぶ —— 奪われるのは帰責だ。窓より前のものは末尾の誤差の責任を問われない。

時間
1 層あたり O(n · d²)、しかもその n ステップは順番に
ステップは重ねられず、GPU では行列が小さすぎて機械を埋められない。支配するのは依存関係そのものだ。
空間
推論は O(d)、学習は O(n · d)
推論は状態 1 本で済む。学習は鎖上の活性をすべて保持し、上の窓が打ち切っているのはこれである。

この計算量が下がらない理由

決め手はどちらでもなく、逐次の深さと最大経路長 —— どちらも n 対 1 だ。アテンションは 1 層 n²·d でそれを買う。

提供側では、この取引がそのまま金額になる。デコーダは見てきた全位置ぶんのKV キャッシュを抱えるが、再帰状態はどんな長さでも同じ大きさだ。文脈を 3 桁ぶんドラッグしよう:

n = 131,072 · KV キャッシュ 64 GB · 再帰状態 256 KB. 左右にドラッグ。矢印キーで 1 段ずつ、Home で最初の状態に戻る
n = 131,072 · KV キャッシュ 64 GB · 再帰状態 256 KB

では、32 層・幅 4096 の fp16 モデルが64 GB を抱えるのに対し、256 KB —— 262,144 倍の差で、重みを載せる前に 80 GB のカード 1 枚の大半を食う。グループ化クエリアテンションは 4 分の 1 に削る。

だからループは交点が示す場所に戻ってきている。 Mamba(Gu と Dao、2023)や RWKV は再帰状態を、学習時に並列計算できる形に組み替えて走らせる。今日の窓では依然としてアテンションの 3 つの勝利が決め手だ。だが持ち帰る価値があるのはそれより古い:掛け算 1 回の経路は減衰しない。