TT Lab
はじめる
学ぶ 学習パス コース

Transformer — アテンションを手で計算する

アテンションは加重平均だ

TT Labで続きを見る

一言でいうと

アテンションは、「いまこの位置で、どの位置をどれだけ見るか」を決めて、値の加重平均をとる演算です。それがすべてです。

3つの名前

同じ入力から、3つのベクトルを作ります。

図書館のたとえがぴったりです。質問(Q)を持って行き、背表紙のタイトル(K)と照らし合わせ、よく合う本の内容(V)を、その度合いに応じて持ち帰ります。

score  = Q · Kᵀ / √d       ← 얼마나 맞는가
weight = softmax(score)     ← 합이 1이 되게
out    = weight · V         ← 가중 평균

たった3行です。残りはすべて、この3行を何度も、いろいろな方向に行うことです。

なぜ√dで割るのか、数字で確かめる

「そういうものだ」で済まされることが最も多い部分ですが、理由ははっきりしています。

成分が独立で分散が1である2つのd次元ベクトルの内積は、分散がdになります。d=64なら、スコアはおよそ±8の範囲に散らばります。

ソフトマックスに±8のスコアが入ると、どうなるでしょうか。比はe⁸ / e⁻⁸ ≈ 9백만(韓国語で「900万」を意味する語です)になります。最大の1つがほぼ1をすべて持っていき、残りは0になります。加重平均ではなく、単なる「1つ選び」になってしまいます。

そうなると、学習ができません。ソフトマックスが飽和すると、勾配が0に近づくからです。

√dで割ると、分散が再び1に戻ります。このラボのステップ2で、これを実際に測ります。割らなかったときに、最大の確率がどこまで上がるかを見ます。

最大値を引いて計算するソフトマックス

def softmax(xs):
    m = max(xs)                       # 이 한 줄이 없으면
    e = [exp(x - m) for x in xs]      # exp(1000) 에서 터진다
    s = sum(e)
    return [v / s for v in e]

exp(x - max)は、数学的には同じ値ですが、オーバーフローが起きません。実務のコードでこの1行を抜かすと、大きな値が入った瞬間にinfとnanが出ます。

マスク、未来を見せないために

言語モデルは、次のトークンを当てるように学習します。ところがアテンションは、デフォルトではすべての位置を見ます。答えを見て答えを当てるようなものです。

そこで、スコア行列の上三角を-infにします。

      k0    k1    k2
q0   0.3  -inf  -inf
q1   0.1   0.5  -inf
q2   0.2   0.1   0.4

-infはexpを通ると、ちょうど0になります。0を掛けるのではなく、ソフトマックスの前に-infを足すことが核心です。ソフトマックスのあとに0を掛けると、残りの重みの合計が1にならなくなります。

これが「causal」または「decoder」マスクです。BERTのようなエンコーダーモデルにはありません。そのためBERTは文全体を見て理解することに強く、GPTは続きを書くことに強いのです。

マルチヘッド、なぜ分けるのか

d=64を一度に使う代わりに、8個に分けてd=8のアテンションを8回行い、そのあと連結します。計算量はほぼ同じです。

なぜでしょうか。1つのソフトマックスは、1種類の関係しか表現できません。重みの合計が1なので、複数の場所を同時に強く見ることはできません。ヘッドを分けると、あるヘッドは直前の単語を、あるヘッドは文頭の主語を、あるヘッドは引用符の対応を見ます。

学習済みのモデルのヘッドを開いてみると、実際にそのように分かれています。

位置エンコーディング、順序を知らないアテンション

これは、初めて学ぶときに最も驚く部分です。

アテンションには、順序という概念がまったくありません。入力トークンを並べ替えると、出力も同じように並べ替えられて出てくるだけで、値そのものは変わりません(permutation equivariant)。「私はあなたが好きです」と「あなたを私は好きです」が区別されません。

そこで、位置の情報を入力に足してあげます。

ステップ5で、これを実際に確認します。位置エンコーディングなしで入力を並べ替えると、出力がそのまま並べ替えられ、足すとそうなりません。

LayerNormと残差接続

x = x + Attention(LayerNorm(x))
x = x + FFN(LayerNorm(x))

LayerNormをアテンションの前に置くか後ろに置くか(pre-LN vs post-LN)は、実際に学習の安定性を左右します。原論文はpost-LNでしたが、最近はほとんどがpre-LNです。ウォームアップなしでも学習できるからです。

コストはどこから生まれるのか

スコア行列はn × nです。シーケンス長の2乗でメモリと計算が増えます。コンテキストを4kから8kに延ばすと、4倍です。

これを減らそうとする試みが、最近の研究の大きな流れです。

まとめ

アテンションそのものは3行です。残りは、その3行を、安定して、安く、順序を知りながら行う方法についてのことです。次のラボで、その3行を自分で書き、スケーリングと位置エンコーディングがなければ何が崩れるのかを、数字で見ます。

現場では

この計算を自分で組む機会は、ほとんどありません。フレームワークがすべてやってくれるからです。それでも知っておくべき理由は、問題が起きたときにどこを見るかが、ここで分かれるからです。

学習の損失が突然nanになったら、たいていはソフトマックスの前の値があふれており、マスクを間違ってかけると、エラーなしに、静かに未来のトークンを見ながら学習して、評価スコアだけが妙に良く出ます。バッチサイズを変えたのに結果が変わるなら、正規化の軸を疑う必要があり、コンテキストを2倍にしたらメモリが4倍に跳ね上がるのは、バグではなく、構造上当然のことです。

推論サービスを運用するなら、KVキャッシュがそのままメモリの予算です。同時リクエスト数と最大コンテキスト長を掛けた値がGPUメモリに収まるかをまず計算し、収まらなければ、GQAを使うモデルを選ぶか、コンテキストの上限を下げるのが、実際の選択肢です。