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

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

ライブラリなしでアテンション

TT Labで続きを見る

目標

アテンションをライブラリなしで自分で実装します。numpyもtorchも使いません。標準ライブラリ(math)だけを使います。

遅いです。そのため、次元を小さく(d=8–64、n=4–8)します。目的は速度ではなく、中が見えることです。

作るもの

/root/work/tf/model.pyに、次のものを定義します。

関数 契約
softmax(xs) 合計が1になります。大きな入力でも壊れません
attention(Q, K, V, mask=None) (out, weights)を返します。√d_kで割ります
causal_mask(n) mask[i][j]が真なら、iはjを見ることができます
multi_head(Q, K, V, h, mask=None) 最後の次元をh等分し、それぞれにアテンションを適用してから連結します
pos_encoding(n, d) 位置ごとに異なる、[-1,1]の範囲のベクトルです
layer_norm(row, eps=1e-5) 平均0、分散1にします
block(X, h, mask=None) X + multi_head(LN(X), ...)を計算します

行列はすべてPythonのリストのリスト([[float]])です。位置 × 次元です。

確認

cd /root/work/tf
python3 -c "import model; print(model.softmax([1,2,3]))"

ステップ

  1. softmax: オーバーフローしない
  2. attention + √dの測定 → 02-scale.txt
  3. causal_mask
  4. multi_head(h=1ならattentionと同じ)
  5. pos_encoding + 並べ替えの実験 → 05-perm.txt
  6. layer_norm
  7. block(pre-LN + 残差)
  8. まとめ → 08-notes.md

参考

採点ツールは、参照実装と1e-6以内で比較します。アテンションは、実装が違っても同じ数字が出る演算なので、値がずれるなら、たいていはスケーリングかマスクの順序が間違っています。

オーバーフローしないソフトマックスを作る

/root/work/tf/model.pyにsoftmax(xs)を作成してください。合計が1になり、[1000, 1001, 1002]のような大きな入力でも、inf・nanを出さずに動作する必要があります。

mkdir -p /root/work/tf。標準ライブラリだけを使います(import math)。最大値を引いてからexpしてください。exp(x - max)は数学的には同じ値ですが、オーバーフローが起きません。この1行がないと、大きな値が入った瞬間に壊れます。

アテンションの3行と√dを実装する

attention(Q, K, V, mask=None)を作成し、(출력, 가중치)(プレースホルダーは出力と重みです)を返してください。スコアは必ず√d_kで割ります。そして、割ったときと割らなかったときで、最大の重みがどれだけ違うかを測り、02-scale.txtに残してください。

Q・K・Vは[[float]](位置 × 次元)です。score[i][j] = dot(Q[i], K[j]) / sqrt(len(K[j]))、w[i] = softmax(score[i])、out[i] = Σ w[i][j] * V[j]です。測定は、d=64の乱数ベクトルで行ってください。割らないと、最大の重みが1.0に張り付きます。それが、「加重平均ではなく1つ選び」になってしまった状態です。

未来を隠す

causal_mask(n)を作成してください。mask[i][j]が真なら、i番目のクエリがjを見られるという意味です。attentionがこのマスクを受け取ったら、隠された位置の重みがちょうど0になっている必要があります。

i >= jなら見られます。実装するときは、ソフトマックスの前に、スコアを-inf(またはごく小さな値)にしてください。ソフトマックスのあとに0を掛けると、残りの重みの合計が1にならなくなります。これがよくあるバグです。

分割して再び結合する

multi_head(Q, K, V, h, mask=None)を作成してください。最後の次元をh等分し、それぞれでアテンションを実行して、再び連結します。出力の形は、入力と同じになる必要があります。

各ヘッドの次元はd // hで、スケーリングも、その小さな次元を基準にする必要があります。h=1なら、attentionとまったく同じ結果が出る必要があります。それが最もよい自己検証です。

アテンションが順序を知らないことを証明する

pos_encoding(n, d)を作成してください。そして、入力の順序を並べ替えたときに、位置エンコーディングがなければ出力がそのまま並べ替えられ、足せばそうならないことを確認して、05-perm.txtに残してください。

サイン/コサインでも、別の方式でもかまいません。値は[-1,1]の範囲で、位置ごとに異なる必要があります。証明の方法は、Xを並べ替えたX'でアテンションを実行した結果が、元の結果を同じ順序で並べ替えたものと等しいかを比較することです。これがpermutation equivarianceです。

LayerNormを特徴軸で計算する

layer_norm(row, eps=1e-5)を作成してください。1つのトークンのベクトルを、平均0、分散1に揃えます。

バッチではなく、そのベクトルの中で平均と分散を求めます。そのため、バッチサイズやシーケンス長の影響を受けません。BatchNormとの決定的な違いです。定数ベクトルが入ってきても0で割らないように、epsを使ってください。

ブロックを組み立てる

block(X, h, mask=None)を作成してください。pre-LN + 残差の構造です: X + multi_head(LN(X), ...)。出力の形は、入力と同じになる必要があります。

rows = [layer_norm(r) for r in X]を行ったあと、マルチヘッドを実行し、元のXを足します。残差がないと、深く積み重ねたときに勾配が入力まで届きません。最近のモデルがpost-LNではなくpre-LNである理由は、ウォームアップなしでも学習できるからです。

3つのことを数字でまとめる

08-notes.mdに、3行以上を書いてください。ステップ2で測った2つの最大の重み、ステップ5が示した性質、そして、シーケンス長を2倍にするとスコア行列が何倍になるか、です。

本文に、스케일링、순서、제곱(韓国語の語で、それぞれ「スケーリング」「順序」「2乗」を意味します)が含まれている必要があります。3つ目は、実務で最もよくぶつかることです。コンテキストを4kから8kに延ばすと、4倍になります。