KVキャッシュのメモリ見積もり計算機を作る
目標
KVキャッシュのメモリの公式をコードで実装し、GQAと量子化の効果を計算で確認し、与えられたGPUで可能な最大の同時リクエスト数を算出します。
なぜ重要なのか
「A100 80GB 1枚に、このモデルは入りますか」という質問は、デプロイの前に必ず出ます。重みだけを計算すると、答えが間違います。7Bモデルの重みは、fp16で14GBですが、128Kコンテキストでバッチ8なら、KVキャッシュが512GiBです。重みの36倍です。この計算をせずに始めると、デプロイ当日にOOMに出会い、OOMは負荷が集中するときに起きるので、最も悪い時点で起きます。そして、この計算機は、一度作っておけば、使い続けられます。モデルを変えるとき、コンテキスト長を延ばすとき、量子化を検討するとき、max_num_seqsを決めるとき、毎回必要です。ステップ6で求める最大の同時リクエスト数が、キャパシティプランニングの出発点になる数字です。
ステップ
/root/kv/kv.pyに、kv_bytes(layers, hidden, seq, batch, dtype_bytes, kv_groups=1)を作成してください。基本の公式は、2 * layers * hidden * seq * batch * dtype_bytesです。layers=32, hidden=4096, seq=4096, batch=1, dtype_bytes=2で計算した値を、/root/kv/base.txtに、bytes=<정수> gib=<소수 둘째자리>(プレースホルダーは整数と、小数第2位までの小数です)として書いてください。gibは2.00である必要があります。seq=131072に変えた値を、/root/kv/long.txtに同じ形式で書いてください。gibは64.00である必要があります。kv_groups=8を適用して、/root/kv/gqa.txtに書いてください。gibは8.00である必要があります(128K基準)。dtype_bytes=1を適用して、/root/kv/int8.txtに書いてください。gibは4.00である必要があります(128K、GQA-8基準)。/root/kv/fit.txtに、gpu_gib=80 weights_gib=14 overhead_gib=6 kv_budget_gib=60 per_req_gib=<소수> max_concurrent=<정수>(プレースホルダーは小数と整数です)を書いてください。リクエストあたりの基準は、32Kコンテキスト、GQA-8、fp16です。/root/kv/kv_table.csvに、context,batch,gibのヘッダーと6行を書いてください。コンテキストは4096、32768、131072で、バッチは1と8です。値は、GQAなしのfp16基準です。
参考
- 公式の先頭の2は、キーと値の2式という意味です。
- GQAは、キー・値ヘッドをグループで共有するので、グループ数で割ります。
- 重みの量子化(GPTQ、AWQ)と、KVキャッシュの量子化は、別の設定です。
- よくあるミス1: 重みだけを計算して、KVキャッシュを抜かしてしまうことです。ロングコンテキストでは、KVが支配的です。
- よくあるミス2: GiBとGBを混ぜてしまうことです。このラボは、2^30基準です。
公式を実装する
/root/kv/kv.pyに、kv_bytes(layers, hidden, seq, batch, dtype_bytes, kv_groups=1)を作成してください。基本の公式は、2 * layers * hidden * seq * batch * dtype_bytesです。
先頭の2が何かを理解して入れてください。引数は、レイヤー数、隠れ次元、シーケンス長、バッチ、dtypeのバイト数です。
7Bの4Kの基準値で検証する
layers=32, hidden=4096, seq=4096, batch=1, dtype_bytes=2で計算した値を、/root/kv/base.txtに、bytes=<정수> gib=<소수 둘째자리>(プレースホルダーは整数と、小数第2位までの小数です)として書いてください。gibは2.00である必要があります。
広く引用される基準値があります。ここで合えば、公式が合っています。
ロングコンテキストに拡張する
seq=131072に変えた値を、/root/kv/long.txtに同じ形式で書いてください。gibは64.00である必要があります。
シーケンス長だけを32倍にすれば済みます。結果が、なぜロングコンテキストを難しくするのかを見てください。
GQAを適用して減らす
kv_groups=8を適用して、/root/kv/gqa.txtに書いてください。gibは8.00である必要があります(128K基準)。
キーと値のヘッドを、グループで共有します。グループ数で割られる箇所が、公式のどこなのかを考えてください。
KVキャッシュの量子化を適用する
dtype_bytes=1を適用して、/root/kv/int8.txtに書いてください。gibは4.00である必要があります(128K、GQA-8基準)。
dtypeのバイト数を減らすことです。重みの量子化とは別の設定だという点が重要です。
GPUに入る最大の同時リクエスト数を求める
/root/kv/fit.txtに、gpu_gib=80 weights_gib=14 overhead_gib=6 kv_budget_gib=60 per_req_gib=<소수> max_concurrent=<정수>(プレースホルダーは小数と整数です)を書いてください。リクエストあたりの基準は、32Kコンテキスト、GQA-8、fp16です。
総メモリから、重みとオーバーヘッドを引いて、リクエストあたりのサイズで割ります。この値が、同時実行数の上限です。
シナリオの表を生成する
/root/kv/kv_table.csvに、context,batch,gibのヘッダーと6行を書いてください。コンテキストは4096、32768、131072で、バッチは1と8です。値は、GQAなしのfp16基準です。
コンテキストとバッチを変えながら、表を作ります。あとでキャパシティプランニングのときに、そのまま使える形で作ってください。