KV キャッシュが節約する乗算を数える
目標
自己回帰生成で、KVキャッシュが何を節約するのかを、乗算の回数を自分で数えて確認します。乗算を数える乗算関数を作り、キャッシュなしで動く版と、キャッシュで動く版をそれぞれ作って、2つの版の出力が同じ値であるかを許容誤差で確認したあと、コンテキスト長ごとに乗算の回数を数えて、表にします。最後に、キャッシュを低い精度で保持すると、その同一性が崩れることを測り、自分で決めた設定で、キャッシュが食う要素の数とバイト数を計算して、記録に残します。
なぜ重要なのか
アテンション1回を手で計算することと、そのアテンションを生成のループの中で呼ぶことは、別の話です。長さ100の答えを作るには、モデルを100回呼び、キャッシュがないと、呼ぶたびに前のコンテキスト全体のk・vを最初から作り直します。最初のトークンのk・vを100回作って、99回捨てるということです。
この無駄は、コードを読んでも見えません。アテンションの関数自体には、間違っているところが1つもないからです。見えるようにするには、数える必要があります。そして、時間を測ってはいけません。時間は、マシンと負荷によって変わりますが、乗算の回数は、同じ入力ならいつも同じで、長さに対してどう増えていくかをそのまま示します。
このラボは、実際のモデルを呼び出しません。このPodのシステムのPythonにはnumpyがなく(/opt/onnx-lab/bin/pythonの中にしかありません)、torchもtransformersもありません。標準ライブラリだけで同じ構造を作り、ここで測った数字だけを使います。そのため、「実際のモデルは何GB食う」や「何倍速くなる」といった話は、ここではしません。
採点ツールは、書いておいた説明を信用しません。自分で作ったモジュールを実際に呼び出し、毎回異なるサイズの入力で関数を直接叩いて、値と乗算の回数を別に計算して照合します。サイズが実行のたびに変わるので、値を暗記して入れることはできません。
ステップ
- /root/work/tf-kv/kv.pyに
CONFIGと、reset_muls()・muls()・mul(a, b)・dot(u, v)を作成してください。mulは、掛け算をしながら数えた回数を1つ上げ、dotは、そのmulだけで掛けます。 project(x, W)を追加して、トークン1つを重み1式に通すようにしてください。Wは、行ごとの長さがlen(x)の行列で、返すリストの長さはlen(W)です。attend(q, K, V)を追加して、クエリ1つが、積み重なったK・Vの全体を見るようにしてください。スコアを√len(q)で割ってソフトマックスをかけ、Vの加重平均をとります。step_nocache(xs, Wq, Wk, Wv)を追加して、キャッシュなしで出力を1つ出すようにしてください。前のすべてのトークンについて、q・k・vを作り直して、最後のクエリでアテンションします。new_cache()・store(cache, k, v, digits=None)・step_cached(x, Wq, Wk, Wv, cache, digits=None)を追加して、キャッシュで同じ出力を出すようにしてください。新しいトークン1つだけを計算して、K・Vに1行追加します。ATOL・RTOL・close(got, want)・compare(xs, Wq, Wk, Wv, digits=None)を追加して、2つの方式の出力が同じ値であるかを比べるようにしてください。等号では比べません。mul_table(xs, Wq, Wk, Wv)を追加して、コンテキスト長1からlen(xs)まで、2つの方式の乗算の回数を数えるようにしてください。返す値は、(문맥길이, 캐시없음, 캐시있음)(プレースホルダーはコンテキスト長、キャッシュなし、キャッシュありです)のペアのリストです。- トークン12個で表を作り、キャッシュが食うメモリを計算して、/root/work/tf-kv/kv_report.jsonと、/root/work/tf-kv/kv_report.mdに記録してください。
参考
- 実行の契約: 採点ツールは、
/root/work/tf-kv/kv.pyをPythonモジュールとして読み込み、CONFIG・reset_muls・muls・mul・dot・project・attend・step_nocache・new_cache・store・step_cached・ATOL・RTOL・close・compare・mul_tableを直接使います。スクリプトとして実行しないので、if __name__ == "__main__"はなくてかまいません。 CONFIGは、{"layers": ..., "heads": ..., "head_dim": ..., "dtype_bytes": ...}の4つのキーを持つ辞書です。値は自分で決めます。範囲は、層が2以上8以下、ヘッドが2以上8以下、ヘッド次元が4以上32以下の偶数、データ型のバイト数が1・2・4のいずれかです。実際のモデルの数を真似る必要はありません。ステップ8のメモリの計算は、自分で決めたこの値だけで行います。mul(a, b)は、a * bを返しながら、モジュール内のカウンターを1上げます。reset_muls()はカウンターを0に、muls()は現在の値を返します。乗算が起きる箇所をすべてmulに通さなければ、数が合いません。dot(u, v)は、内積です。乗算はlen(u)回起きます。足し算・割り算・指数は数えません。project(x, W)の乗算は、len(W)掛けるlen(x)回です。Wの行が、出力の1つの要素を作ります。行と列を取り違えると、値も回数もずれます。attend(q, K, V)の乗算は、スコア側がlen(K)掛けるlen(q)回、加重和側がlen(V)掛けるlen(V[0])回です。割り算と指数は乗算ではないので、数えません。ソフトマックスは、最大値を引いてから指数をとってください。step_nocache(xs, Wq, Wk, Wv)は、xsのすべてのトークンについてq・k・vを作り、最後のクエリで全体のK・Vを見ます。使われるクエリが1つだけなのに、すべてを作るのが、キャッシュがない状態の姿です。new_cache()は、{"K": [], "V": []}を返します。store(cache, k, v, digits=None)は、KとVにそれぞれ1行追加します。上書きしません。digitsが与えられたら、各成分をその小数の桁数に丸めて入れます。step_cached(x, Wq, Wk, Wv, cache, digits=None)は、新しいトークン1つだけを射影して、storeでキャッシュに入れたあとで、そのクエリでキャッシュ全体を見ます。入れる前にアテンションすると、新しいトークンが自分自身を見られません。ATOL = 1e-9、RTOL = 1e-6とし、close(got, want)は、abs(got - want) <= ATOL + RTOL * abs(want)です。なぜこの幅なのかをコメントで書いてください。Pythonのfloatは、IEEE 754の倍精度なので、有効数字が約15桁で、この規模の内積・ソフトマックスは、足す順序を変えただけでは、相対誤差が1e-12程度にとどまります。compare(xs, Wq, Wk, Wv, digits=None)は、長さ1からlen(xs)まで1つずつ増やしながら、2つの方式を動かして、{"lengths": [...], "max_gap": 실수, "all_close": 참거짓, "last_out": [...]}(プレースホルダーは実数と真偽です)を返します。max_gapは、成分同士の差の絶対値のうち最も大きい値で、all_closeは、すべての成分がcloseを通ったかどうかです。mul_table(xs, Wq, Wk, Wv)は、長さごとに測る直前にカウンターを戻します。キャッシュ側は、1式のキャッシュを使い続ける必要があります。行ごとに新しく作るなら、それはキャッシュではありません。- ステップ8は、
d = CONFIG["head_dim"]として、トークン12個、d掛けるdの重み3つで測ります。round_digitsは2を使います。 - ステップ8のトークンのベクトルと重みは、
((i * 7 + j * 5 + salt * 3) % 13 - 6) / 7.0で作ります。iは行の番号、jは列の番号で、saltは、トークンのベクトルが1、Wqが2、Wkが4、Wvが6です。乗算の回数とメモリは、値ではなくサイズからだけ出てくるので、どんな値を使っても表は同じですが、丸めの実験は値に依存します。7で割るのは、値が小数第2位できっちり割り切れないようにするためです。きっちり割り切れる値だけを使うと、小数第2位に丸めても値がそのままなので、丸めの影響が見えません。 - メモリは、
원소 수 = 2 × layers × heads × 길이 × head_dimと바이트 = 원소 수 × dtype_bytesです(韓国語の語は順に、要素数、長さ、バイト数という意味です)。長さは12です。 - このPodにはインターネットがありません。
pip installはできず、torchもtransformersもありません。numpyは/opt/onnx-lab/bin/pythonの中にしかないので、システムのPythonではimport numpyができません。mathだけで十分です。 - 時間を測らないでください。
timeで測った数は、マシンと負荷によって変わるので、判定に使えません。このラボが測るのは、乗算の回数です。 - 公式ドキュメント: Attention Is All You Need・Hugging Face — Cache strategies・Hugging Face — Text generation・Python — math
- よくある間違い:
mulを通さずに直接掛けて、数が0になる、projectで行と列を取り違える、スコアを√dで割らない、キャッシュに上書きする、キャッシュに入れる前にアテンションする、mul_tableでカウンターを戻さない、キャッシュを行ごとに新しく作る。
乗算を数える乗算関数を作る
/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位できっちり割り切れないようにするためです。きっちり割り切れる値だけを使うと、丸めても値がそのままなので、この実験そのものが成り立ちません。時間は測らないでください。