クエリヘッドはそのまま、キー・バリューヘッドだけ減らす
一言でいうと
MHA・MQA・GQAを分けるつまみは、1つだけです。キー・値ヘッドをいくつ置くかです。クエリヘッドの数には手を付けません。そのため、減るのはKVキャッシュが占めるメモリで、アテンションの乗算の回数は、ほぼそのままです。
なぜ必要なのか
生成は、トークンを1つずつ進みます。新しいトークンを1つ作るたびに、モデルは、これまでに出たすべての位置のキーと値を見直します。そのため、そのキーと値を、消しておいて作り直すのではなく、保持しておきます。それが、KVキャッシュです。
問題は、このキャッシュが、コンテキストが長くなるほど、層が多いほど、同時に処理するリクエストが多いほど、そのまま大きくなるという点です。モデルの重みは、コンテキスト長と無関係に固定ですが、キャッシュはそうではありません。重みをすべて載せたあとに残った場所をキャッシュが食うので、ある地点からは、キャッシュが同時処理数を決めます。
しかも、トークン1つを作るときに行う計算は、わずかです。その少ない計算を行うために、キャッシュ全体をメモリから読み込む必要があります。MQAの原論文が、インクリメンタルデコーディングのボトルネックとして挙げたのが、まさにこの読み直しです。タイトルがそのまま答えです。使うヘッドは1つで十分なのです。
ここから、MQAの決定が出てきます。キー・値ヘッドを1つだけ置きます。クエリヘッドは、8個でも64個でもそのままにして、そのすべてが同じK・Vを見ます。キャッシュは、クエリヘッドの数分の1に減ります。
その代わりに、失うものがあります。ヘッドごとに別の場所を見るように作ってあった構造で、「何を差し出すか」を決める側を1つにまとめてしまったのです。GQAの原論文は、その間に場所を作ります。クエリヘッドをいくつかのグループに分け、グループごとにキー・値ヘッドを1つずつ置きます。
3つの名前は、1本の軸の上の3つの位置
クエリヘッドがh個、グループがg個だとします。このとき、キー・値ヘッドの数が、そのままgです。
| グループ数g | 名前 | キー・値ヘッド | 1つのK・Vを見るクエリヘッド |
|---|---|---|---|
| g = h | MHA | h個 | 1個 |
| 1 < g < h | GQA | g個 | h/g個 |
| g = 1 | MQA | 1個 | h個 |
3行を別々に覚えるものではありません。1つのつまみをどこまで回したかの違いです。条件は1つだけです。hがgで割り切れる必要があります。そうでないと、グループの大きさがふぞろいになり、あるグループは3つのヘッドを、あるグループは2つのヘッドをまとめることになります。実際の設定ファイルでこの条件を破ると、たいていモデルを読み込む場面で止まります。
何が減り、何が減らないのか
ここがこのテーマの核心で、最も間違えやすい箇所です。式を立ててみれば、すぐに見えます。
キャッシュに入る実数の個数は、次のとおりです。
KV 캐시 원소 수 = 2 × 층 수 × (키·값 헤드 수) × 길이 × 헤드 차원
このコードブロックの韓国語の式は、KVキャッシュの要素数が、2×層の数×(キー・値ヘッドの数)×長さ×ヘッド次元だ、という意味です。
最初の2は、KとVの2式です。クエリヘッドの数が、この式にありません。クエリは、キャッシュに残さないからです。そのため、キー・値ヘッドを8個から1個に減らせば、キャッシュはちょうど8分の1になります。
今度は、乗算の回数を数えてみましょう。長さn、クエリヘッドh、ヘッド次元dの1つの層を、1回通過するときです。
점수 Q·Kᵀ h × n × n × d ← g 가 없다
가중합 × V h × n × n × d ← g 가 없다
K 투영 n × d_model × (g × d)
V 투영 n × d_model × (g × d)
このコードブロックの韓国語は、行頭の語が順にスコア、加重和、K射影、V射影を指し、行末の2つの韓国語コメントは、どちらもgが含まれていないという意味です。
アテンション本体に、gがありません。理由は、実装を見ればはっきりします。畳んでおいたキー・値ヘッドを使うときは、クエリヘッドの数だけ繰り返して広げて、いつもどおりに動かします。異なる値の種類が減っただけで、アテンションが動く回数は、クエリヘッドの数が決めます。減る項は、K・V射影の2つだけで、それは全体の中で大きな取り分ではありません。
そのため、こう整理されます。GQAが減らすのは、メモリです。速くなるのは、そのメモリを読む量が減るからであって、乗算を減らしたからではありません。この区別ができないと、「GQAに変えたのに、なぜプリフィルがそのままなのか」という問いに答えられません。
グループをどう統合するのか
すでに学習済みのマルチヘッドのモデルを、GQAに移すとします。グループの中のキー・値ヘッド4つを1つにしなければなりませんが、どう作ればよいのでしょうか。
1つだけ残して捨てることも、平均をとることも、最初から減らしたサイズで学習し直すこともできます。GQAの原論文が選んだ道は、グループの中のキー・値ヘッドを平均して初期値とし、そこからもう少し学習することです。平均をとるだけで終わらないという点が重要です。平均は出発点であって、答えではありません。
次のラボも、平均を使います。ただし、学習はしません。そのため、このラボで出てくる「出力がこれだけ変わった」という数字は、品質がそれだけ悪くなるという意味ではなく、同じ重みのままキー・値だけをまとめると、答えが変わるという事実を示す値です。この区別を、レポートにも書くことになります。
現場での姿
第1に、設定ファイルに、ヘッド数が2行で書かれています。最近のモデルの設定には、クエリヘッドの数とキー・値ヘッドの数が別々にあります。2つの値が同じならMHA、キー・値の側が1ならMQA、その間ならGQAです。3つの名前を探すまでもなく、2つの数字の比を見れば済みます。
第2に、「GQAなのに、なぜ速くならないのか」が出てきます。プリフィルのように、一度に多くの位置を処理するときは、計算がボトルネックで、そこにはgがほとんど影響しません。得になるのは、長いコンテキストを持ってトークンを1つずつ取り出す区間と、同時リクエストの数を上げる場面です。
第3に、同時処理数の上限が、キャッシュで決まります。残りのメモリを、キャッシュ1式のサイズで割った値が、おおよその上限です。キー・値ヘッドを4つに減らすと、その1式が小さくなるので、上限が上がります。このとき増えるのはスループットであって、1つのリクエストの速度ではありません。
第4に、テンソル並列で、キー・値ヘッドの数が障害になります。ヘッドを複数のデバイスに分けて載せるとき、キー・値ヘッドがデバイスの数より少ないと、分けるものが足りなくなり、同じキー・値を複数のデバイスが持つことになります。キャッシュを減らすためにgを下げたのに、デバイスごとに複製されると、期待したほどは減りません。
第5に、ヘッドの数を変えると、別のモデルになります。グループを平均して作った重みを、そのままサービスに載せると、出力が変わります。原論文も、移したあとに追加の学習を行っています。「設定だけ変えればよい」のではありません。
実務で本当に大切なこと
- 減るのは、メモリです。乗算の回数は、ほぼそのままです。得がどこから来るのかを知っていれば、どこに使うかを決められます。
- hをgで割り切れる必要があります。グループの大きさがそろわないと、実装が成り立ちません。
- 2つの端は、同じ軸の上にあります。g=hならMHA、g=1ならMQAです。コード1式で3つとも動かせるように作っておくと、比較が簡単になります。
- 移すことと学習することは、別の作業です。平均は初期値にすぎません。
次のラボですること
/root/work/tf-gqa/gqa.pyを、1ステップずつ育てていきます。クエリヘッドを連続した塊としてグループにまとめることから始めて、グループの中のキー・値ヘッドを平均で畳み、それをクエリヘッドの数だけ繰り返して広げたあと、1ヘッドのアテンションを経て、3つの方式を同じ重みで動かします。gをhにすれば、MHAと1か所も違わない必要があり、gを減らせば、出力が変わる必要があります。その差を、数字で測ります。
そのあとが、このラボの要点です。KVキャッシュの要素数と乗算の回数を、それぞれ数えて並べます。キャッシュは、gにそのまま比例して減るのに、アテンション本体の乗算は一度も減らないことを、他の人の説明ではなく、自分が数えた数字で確認します。最後に、その2つの表と、出力の差を、レポートに残します。
時間は測りません。このPodにはGPUがなく、CPUも他の作業と共用しているので、「速くなった」は、ここで測れるものではありません。numpyも、システムのPythonにはありません(/opt/onnx-labの中にしかありません)。標準ライブラリだけで十分で、判定はすべて、要素数・乗算の回数のような整数の計数と、許容誤差を設けた数値の照合です。採点ツールは、自分で作ったモジュールを実際に呼び出し、毎回異なるヘッド数とグループ数で関数を直接叩いて、別に計算した値と照合します。