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

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

キー・バリューヘッドだけ減らしてみる

TT Labで続きを見る

目標

クエリヘッドの数はそのままにして、キー・値ヘッドの数だけを減らす構造を、標準ライブラリだけで作ります。クエリヘッドを連続したかたまりでグループにまとめ、グループの中のキー・値ヘッドを平均で畳み、それをクエリヘッドの数だけ繰り返して広げたあと、アテンションを動かします。グループ数gを変えるだけで、MHA(g=h)・GQA(1同じ重みで動かして、出力の差を測り、最後に、KVキャッシュの要素数と乗算の回数をそれぞれ数えて、何が減り、何が減らないのかを確認します。

なぜ重要なのか

最近のモデルの設定には、クエリヘッドの数とキー・値ヘッドの数が、別々に書かれています。2つの値が同じならMHA、キー・値の側が1ならMQA、その間ならGQAです。3つの名前を覚えるよりも、何が減るのかを知ることが重要です。 減るのは、KVキャッシュのメモリです。キャッシュの要素数は、2 x 층 x 키·값 헤드 수 x 길이 x 헤드 차원(式の韓国語の語は順に、層の数、キー・値ヘッドの数、長さ、ヘッド次元という意味です)で、この式にクエリヘッドの数がありません。一方、アテンション本体の乗算の回数は、h x n x n x 헤드차원(韓国語の語はヘッド次元という意味です)なので、グループ数gがまったく入りません。畳んだキー・値を、クエリヘッドの数だけ繰り返して広げて、いつもどおりに動かすからです。この2つの式を自分で数えてみれば、「GQAに変えたのに、なぜプリフィルがそのままなのか」という問いに答えられます。 このラボは、時間を測りません。PodにGPUがなく、CPUも他の作業と共用しているので、ここで測った速度は、何も語ってくれません。判定はすべて、要素数・乗算の回数のような整数の計数と、許容誤差を設けた数値の照合です。 グループの中のキー・値ヘッドを平均で合わせるのは、このラボの仮定です。GQAの原論文が、学習済みのマルチヘッドのチェックポイントを移すときに使う方式と同じですが、原論文は、平均をとったあとに追加の学習を行います。ここで出てくる出力の差は、品質がそれだけ悪くなるという意味ではなく、同じ重みのままキー・値だけをまとめると、答えが変わるという事実を示す値です。 採点ツールは、書いておいた説明を信用しません。自分で作ったモジュールを実際に呼び出し、毎回異なるヘッド数とグループ数で関数を直接叩いて、採点ツールが別に計算した値と照合します。

ステップ

  1. /root/work/tf-gqa/gqa.pyにmake_heads(count, rows, dim, seed)と、group_of(head, h, g)・group_members(h, g)を作成してください。サンプルは、同じ引数ならいつも同じ値が出る必要があり、グループは連続したかたまりに分けます。
  2. fold_kv(heads, g)を追加して、グループの中のキー・値ヘッドを、位置ごとに平均して、g個に畳むようにしてください。
  3. expand_kv(folded, h)を追加して、畳んだヘッドを隣り合わせに繰り返して、h個に広げるようにしてください。
  4. attend(q, k, v)を追加してください。スコアをsqrt(헤드 차원)(プレースホルダーはヘッド次元です)で割り、ソフトマックスは、最大値を引いてから指数をとります。
  5. heads_out(q_heads, k_heads, v_heads, g)とmax_gap(left, right)を追加して、3つの方式を同じ重みで動かし、出力の差を測るようにしてください。
  6. kv_cache_elems(layers, kv_heads, seq_len, head_dim)とkv_table(layers, h, groups, seq_len, head_dim)を追加して、キャッシュの要素数を数えるようにしてください。
  7. mults(n, d_model, h, g, head_dim)を追加して、1つの層を1回通過するときの乗算の回数を、項目ごとに数えるようにしてください。
  8. サンプルとモデルの規模を決めて2つの表を作り、/root/work/tf-gqa/gqa_report.jsonと、/root/work/tf-gqa/gqa_report.mdに結果を記録してください。

参考

クエリヘッドをグループに分ける

/root/work/tf-gqa/gqa.pyにmake_heads(count, rows, dim, seed)と、group_of(head, h, g)・group_members(h, g)を作成してください。サンプルは、同じ引数ならいつも同じ値が出る必要があり(randomモジュールは禁止、絶対値4以下)、グループは連続したかたまりに分けます。h=8・g=2なら、0, 1, 2, 3がグループ0です。均等に分けられなければ、ValueErrorを出してください。

サンプルは、小さな線形合同法1つで十分です。状態を整数で持ち、state = (a * state + c) % mを繰り返しながら、state / m - 0.5を取り出せば、同じseedから同じ数列が出ます。グループの番号は、head // (h // g)です。head % gは交互の割り当てなので、あとで広げるときにずれます。group_membersは、group_ofをh回呼んで埋めれば、2か所の規則が分かれることがありません。

グループ内のキー・値を1つに畳む

fold_kv(heads, g)を追加してください。グループの中のキー・値ヘッドを位置ごとに平均して、g個に畳みます。ヘッドの数がgで割り切れなければ、ValueErrorです。gがヘッドの数と同じなら、値がそのまま出る必要があります。

グループは、ステップ1と同じ規則、つまり連続したかたまりです。heads[start:start + size]で1つのグループを切り出し、位置ごとに足して、グループの大きさで割れば済みます。最初のヘッドだけを残したり、グループを無視してすべてを平均したりすると、g=1のときだけ偶然同じに見えます。平均は、このラボの仮定です。GQAの原論文がチェックポイントを移すときに使う方式と同じですが、原論文は、そのあとに追加の学習を行います。

クエリヘッドの数だけ再び広げる

expand_kv(folded, h)を追加してください。畳んだヘッドを隣り合わせに繰り返して、h個に広げます。[A, B]を4個に広げると、[A, A, B, B]です。均等に広げられなければ、ValueErrorです。

外側の繰り返しは畳んだヘッド、内側の繰り返しはh // len(folded)回です。リスト全体を掛けて連結すると(folded * size)、[A, B, A, B]になって、クエリヘッドiが、別のグループのキー・値を見ることになります。ステップ1のgroup_ofと対応が合っているかを確認してください。クエリヘッドiが見るものは、folded[group_of(i, h, g)]である必要があります。行は、新しいリストにコピーして返せば、あとで1か所を書き換えて、複数のヘッドがいっしょに変わることがありません。

1ヘッドのアテンションを作る

attend(q, k, v)を追加してください。スコアはsqrt(헤드 차원)(プレースホルダーはヘッド次元です)で割り、ソフトマックスは、最大値を引いてから指数をとります。返す値は、qと同じ形の表です。マスクは使いません。

クエリの行ごとに、すべてのキーの行と内積をとり、割って、ソフトマックスを通して、値の行たちの加重平均をとります。sqrtで割らないと、次元が大きくなるほど、ソフトマックスが1か所に偏ります。最大値を引かないと、スコアが大きいサンプルでmath.expがあふれて、OverflowErrorが出ます。採点ツールが、わざと大きな値を入れてみます。

3つの方式を同じ重みで動かす

heads_out(q_heads, k_heads, v_heads, g)とmax_gap(left, right)を追加してください。heads_outは、キー・値をg個に畳んでからh個に広げたあと、クエリヘッドごとにattendを1回ずつ呼びます。g = len(q_heads)なら、ふつうのマルチヘッドと1か所も違わない必要があり、gを減らせば、出力が変わる必要があります。max_gapは、2つの出力の間の、最も大きい絶対値の差を返します。

3行で終わります。キーを畳んで広げ、値を畳んで広げ、クエリヘッドごとにattendです。クエリヘッドには手を付けないという点が、このステップのすべてです。max_gapは、ヘッド・行・列をすべて走査して、最も大きい差を1つだけ残します。gを変えながら呼び出してみると、g=hで差がちょうど0になり、gが小さくなるほど差が大きくなるのが見えます。

キャッシュの要素数を数える

kv_cache_elems(layers, kv_heads, seq_len, head_dim)とkv_table(layers, h, groups, seq_len, head_dim)を追加してください。要素数は2 * layers * kv_heads * seq_len * head_dimで、クエリヘッドの数は入りません。kv_tableは、[(무리 수, 원소 수), ...](プレースホルダーはグループ数と要素数です)を返し、均等に分けられないグループ数があれば、ValueErrorです。

最初の2は、KとVの2式です。クエリヘッドの数を掛けたくなる場面がありますが、クエリはキャッシュに残さないので、式にありません。kv_tableでは、そのグループ数のキー・値ヘッドの数が、そのままgであるという点を、そのまま使えば済みます。バイト数ではなく要素数で数えるのは、データ型によって、1要素が2バイトにも4バイトにもなるからです。

乗算の回数を数える

mults(n, d_model, h, g, head_dim)を追加してください。キーはproj_q・proj_k・proj_v・scores・weighted・proj_out・totalで、値は整数です。scoresとweightedはh * n * n * head_dim、射影はn * d_model * (헤드 수 * head_dim)(プレースホルダーはヘッドの数です)の形で、totalは、残りの6つの合計です。均等に分けられなければ、ValueErrorです。

どの項にgが入り、どの項に入らないかが、このステップのすべてです。キー・値を広げたあとは、アテンションがクエリヘッドの数だけ動くので、scoresとweightedにgはありません。gに結び付いているのは、K射影とV射影の2つだけです。乗算だけを数え、足し算とソフトマックスの指数は数えません。見たいのは総量ではなく、どの項が減るかです。

何が減り、何が減らないのかを記録する

サンプルは、make_heads(8, 6, 4, 101)・make_heads(8, 6, 4, 202)・make_heads(8, 6, 4, 303)をQ・K・Vとし、グループ数は[8, 4, 2, 1]とします。モデルの規模は、layers = 32、seq_len = 4096、head_dim = 128、h = 8、d_model = 1024、n_tokens = 4096です。/root/work/tf-gqa/gqa_report.jsonにはsample_h・sample_n・sample_head_dim・groups・diff_table・layers・seq_len・head_dim・h・d_model・n_tokens・kv_table・kv_ratio・mult_table・mult_ratio・core_mults・core_same_for_all_gを、/root/work/tf-gqa/gqa_report.mdには## 무엇을 쟀나 ## 무리를 줄이면 출력이 얼마나 달라지나 ## 메모리는 줄고 곱셈은 안 준다 ## 어디에 쓸 것인가の4つの節で書いてください(韓国語の見出しは順に、「何を測ったか」「グループを減らすと出力がどれだけ変わるか」「メモリは減るが、乗算は減らない」「どこに使うのか」という意味です)。

数字は手で書かず、自分の関数を実際に動かして得た値で埋めてください。diff_tableは、グループ数ごとにmax_gap(그 무리의 출력, g=h 의 출력)(プレースホルダーはそのグループの出力と、g=hの出力です)なので、最初の欄が0.0です。kv_ratio・mult_ratioは、MHAを1とした倍数なので、MHA 값 / 그 무리의 값(プレースホルダーはMHAの値とそのグループの値です)です。core_multsはscores + weightedで、グループ数が変わっても同じ値なので、core_same_for_all_gが真になります。mdには、キャッシュの倍数と乗算の倍数を並べて書いておいてください。片方は8倍まで減り、もう片方はほとんどそのままだというのが、このラボの結論です。