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

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

KV キャッシュが節約する乗算を数える

TT Labで続きを見る

目標

自己回帰生成で、KVキャッシュが何を節約するのかを、乗算の回数を自分で数えて確認します。乗算を数える乗算関数を作り、キャッシュなしで動く版と、キャッシュで動く版をそれぞれ作って、2つの版の出力が同じ値であるかを許容誤差で確認したあと、コンテキスト長ごとに乗算の回数を数えて、表にします。最後に、キャッシュを低い精度で保持すると、その同一性が崩れることを測り、自分で決めた設定で、キャッシュが食う要素の数とバイト数を計算して、記録に残します。

なぜ重要なのか

アテンション1回を手で計算することと、そのアテンションを生成のループの中で呼ぶことは、別の話です。長さ100の答えを作るには、モデルを100回呼び、キャッシュがないと、呼ぶたびに前のコンテキスト全体のk・vを最初から作り直します。最初のトークンのk・vを100回作って、99回捨てるということです。 この無駄は、コードを読んでも見えません。アテンションの関数自体には、間違っているところが1つもないからです。見えるようにするには、数える必要があります。そして、時間を測ってはいけません。時間は、マシンと負荷によって変わりますが、乗算の回数は、同じ入力ならいつも同じで、長さに対してどう増えていくかをそのまま示します。 このラボは、実際のモデルを呼び出しません。このPodのシステムのPythonにはnumpyがなく(/opt/onnx-lab/bin/pythonの中にしかありません)、torchもtransformersもありません。標準ライブラリだけで同じ構造を作り、ここで測った数字だけを使います。そのため、「実際のモデルは何GB食う」や「何倍速くなる」といった話は、ここではしません。 採点ツールは、書いておいた説明を信用しません。自分で作ったモジュールを実際に呼び出し、毎回異なるサイズの入力で関数を直接叩いて、値と乗算の回数を別に計算して照合します。サイズが実行のたびに変わるので、値を暗記して入れることはできません。

ステップ

  1. /root/work/tf-kv/kv.pyにCONFIGと、reset_muls()・muls()・mul(a, b)・dot(u, v)を作成してください。mulは、掛け算をしながら数えた回数を1つ上げ、dotは、そのmulだけで掛けます。
  2. project(x, W)を追加して、トークン1つを重み1式に通すようにしてください。Wは、行ごとの長さがlen(x)の行列で、返すリストの長さはlen(W)です。
  3. attend(q, K, V)を追加して、クエリ1つが、積み重なったK・Vの全体を見るようにしてください。スコアを√len(q)で割ってソフトマックスをかけ、Vの加重平均をとります。
  4. step_nocache(xs, Wq, Wk, Wv)を追加して、キャッシュなしで出力を1つ出すようにしてください。前のすべてのトークンについて、q・k・vを作り直して、最後のクエリでアテンションします。
  5. new_cache()・store(cache, k, v, digits=None)・step_cached(x, Wq, Wk, Wv, cache, digits=None)を追加して、キャッシュで同じ出力を出すようにしてください。新しいトークン1つだけを計算して、K・Vに1行追加します。
  6. ATOL・RTOL・close(got, want)・compare(xs, Wq, Wk, Wv, digits=None)を追加して、2つの方式の出力が同じ値であるかを比べるようにしてください。等号では比べません。
  7. mul_table(xs, Wq, Wk, Wv)を追加して、コンテキスト長1からlen(xs)まで、2つの方式の乗算の回数を数えるようにしてください。返す値は、(문맥길이, 캐시없음, 캐시있음)(プレースホルダーはコンテキスト長、キャッシュなし、キャッシュありです)のペアのリストです。
  8. トークン12個で表を作り、キャッシュが食うメモリを計算して、/root/work/tf-kv/kv_report.jsonと、/root/work/tf-kv/kv_report.mdに記録してください。

参考

乗算を数える乗算関数を作る

/root/work/tf-kv/kv.pyに、CONFIGと、reset_muls()・muls()・mul(a, b)・dot(u, v)を作成してください。CONFIGは、layers・heads・head_dim・dtype_bytesの4つのキーを持つ辞書で、値は自分で決めます(層が2以上8以下、ヘッドが2以上8以下、ヘッド次元が4以上32以下の偶数、データ型のバイト数が1・2・4のいずれか)。mulは、掛け算をしながらカウンターを1上げ、dotは、そのmulだけで掛けます。

カウンターは、モジュール内の整数1つで十分です。関数の中で変更するには、globalが必要です。dotがsum(a * b for a, b in zip(u, v))のように直接掛けると、何も数えなくなるので、必ずmulを通してください。足し算は数えません。行列計算の値は、乗算の側から出てきます。時間を測るコードは入れないでください。

トークン1つを射影する

project(x, W)を追加してください。Wは、行ごとの長さがlen(x)の行列で、行1つが出力の1つの要素を作ります。返すリストの長さはlen(W)で、乗算はlen(W)掛けるlen(x)回起きます。

前に作ったdotを、行ごとに1回ずつ呼べば終わりです。1行で書けます。行と列を取り違えると、正方行列では値だけが間違い、長方形では長さまでずれるので、返すリストの長さがlen(W)であるかを、まず確認してください。qもkもvも、この関数1つで作ります。

クエリ1つがキャッシュ全体を見る

attend(q, K, V)を追加してください。qと各kの内積を√len(q)で割ってスコアを出し、最大値を引いてから指数をとってソフトマックスをかけ、その重みでVの加重平均をとります。乗算は、スコア側がlen(K)掛けるlen(q)回、加重和側がlen(V)掛けるlen(V[0])回です。

割り算と指数は乗算ではないので、mulで包まないでください。包むと、回数がずれます。加重和の乗算だけがmulを通ります。Vの1行の長さはqの長さと異なることがあるので、出力のリストはlen(V[0])で作ってください。Kが1行だけなら、重みが1の1つなので、出力はV[0]と同じである必要があります。

キャッシュなしで、前の全体を再計算する

step_nocache(xs, Wq, Wk, Wv)を追加してください。xsのすべてのトークンについて、q・k・vを作り、最後のクエリで全体のK・Vを見ます。使われるクエリは1つだけなのに、すべてを作るのが、キャッシュがない状態の姿です。

リスト内包表記を3つと、attendを1回使えば済みます。最後のクエリだけを使うからといって、最後のトークンだけを射影してはいけません。キーと値は、前のトークンすべてについて必要で、このステップの要点は、それを毎回作り直すという事実です。乗算の回数は、コンテキスト長に比例して増えます。

キャッシュで、1行だけ追加する

new_cache()・store(cache, k, v, digits=None)・step_cached(x, Wq, Wk, Wv, cache, digits=None)を追加してください。new_cache()は{"K": [], "V": []}を返し、storeは、KとVにそれぞれ1行追加し(digitsが与えられたら、その小数の桁数に丸めて)、step_cachedは、新しいトークン1つだけを射影してキャッシュに入れたあとで、そのクエリでキャッシュ全体を見ます。

上書きせず、appendしてください。前の行は、新しいトークンが付いても値が変わりません。因果マスクのために、各位置が自分より前だけを見るからで、それがキャッシュが成り立つ理由です。入れる前にアテンションすると、新しいトークンが自分自身を見られなくなり、キャッシュなしの版と答えが違ってきます。乗算の回数は、射影の側が、コンテキスト長と無関係に固定されます。

2つの方式の出力は同じ値か確認する

ATOL = 1e-9・RTOL = 1e-6・close(got, want)・compare(xs, Wq, Wk, Wv, digits=None)を追加してください。closeはabs(got - want) <= ATOL + RTOL * abs(want)で、なぜこの幅なのかをコメントで書きます。compareは、長さ1から1つずつ増やしながら2つの方式を動かして、{"lengths": [...], "max_gap": 실수, "all_close": 참거짓, "last_out": [...]}(プレースホルダーは実数と真偽です)を返します。

等号で比べないでください。2つの方式がたまたま同じ順序で足せば、ビットまで同じに出ることもありますが、足す順序が少し違うだけで、最後の桁が揺らぎます。そのため、正しいテストは許容誤差です。digitsは、そのままstep_cachedに渡します。丸めて入れると、all_closeが偽になりますが、それがこのステップで見たいことです。max_gapは、すべての長さ・すべての成分を通じて、最も大きい差です。

長さごとに乗算を数える

mul_table(xs, Wq, Wk, Wv)を追加してください。コンテキスト長1からlen(xs)まで、2つの方式の乗算の回数を数えて、(문맥길이, 캐시없음, 캐시있음)(プレースホルダーはコンテキスト長、キャッシュなし、キャッシュありです)のペアのリストを返します。測る直前ごとにカウンターを戻し、キャッシュ側は、1式のキャッシュを使い続けます。

reset_muls()を2回呼びます。キャッシュなしの版を測る前に1回、キャッシュの版を測る前に1回です。戻さないと、後ろの数に前の数が混ざり込みます。キャッシュを行ごとに新しく作ると、キャッシュ側の数が、キャッシュなしの側のように増えてしまいます。それはキャッシュではありません。表を見ると、キャッシュなしの側は長さに比例して増え、キャッシュの側は、アテンションの取り分だけが増えます。

節約できたものと支払ったものを、あわせて書く

d = CONFIG["head_dim"]として、トークン12個、d掛けるdの重み3つで、mul_tableとcompareを動かしてください。値は((i * 7 + j * 5 + salt * 3) % 13 - 6) / 7.0で作り、saltは、トークンのベクトルが1、Wqが2、Wkが4、Wvが6です。compareは、丸めなしで1回、digits=2で1回動かします。そして、/root/work/tf-kv/kv_report.jsonにはlayers・heads・head_dim・dtype_bytes・tokens・table・total_nocache・total_cached・saved_muls・cache_elems・cache_bytes・round_digits・max_gap_exact・all_close_exact・max_gap_rounded・all_close_roundedを、/root/work/tf-kv/kv_report.mdには## 무엇을 쟀나 ## 캐시 없이 하면 무엇을 다시 계산하나 ## 캐시가 먹는 메모리 ## 두 방식의 출력이 같은가の4つの節で書いてください(韓国語の見出しは順に、「何を測ったか」「キャッシュなしだと何を再計算するのか」「キャッシュが食うメモリ」「2つの方式の出力は同じか」という意味です)。

数字は手で書かず、自分のコードを実際に動かして得た値で埋めてください。total_nocache・total_cachedは、表の各欄を足したもので、saved_mulsは、その差です。cache_elemsは2 × layers × heads × 12 × head_dim、cache_bytesは、それにdtype_bytesを掛けた値です。all_close_exactは真、all_close_roundedは偽である必要があります。キャッシュ自体は近似ではありませんが、キャッシュにより不正確に書き込むことは近似です。値を作る式で7で割るのは、値が小数第2位できっちり割り切れないようにするためです。きっちり割り切れる値だけを使うと、丸めても値がそのままなので、この実験そのものが成り立ちません。時間は測らないでください。