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

LLMサービング

KVキャッシュのメモリ見積もり計算機を作る

TT Labで続きを見る

目標

KVキャッシュのメモリの公式をコードで実装し、GQAと量子化の効果を計算で確認し、与えられたGPUで可能な最大の同時リクエスト数を算出します。

なぜ重要なのか

「A100 80GB 1枚に、このモデルは入りますか」という質問は、デプロイの前に必ず出ます。重みだけを計算すると、答えが間違います。7Bモデルの重みは、fp16で14GBですが、128Kコンテキストでバッチ8なら、KVキャッシュが512GiBです。重みの36倍です。この計算をせずに始めると、デプロイ当日にOOMに出会い、OOMは負荷が集中するときに起きるので、最も悪い時点で起きます。そして、この計算機は、一度作っておけば、使い続けられます。モデルを変えるとき、コンテキスト長を延ばすとき、量子化を検討するとき、max_num_seqsを決めるとき、毎回必要です。ステップ6で求める最大の同時リクエスト数が、キャパシティプランニングの出発点になる数字です。

ステップ

  1. /root/kv/kv.pyに、kv_bytes(layers, hidden, seq, batch, dtype_bytes, kv_groups=1)を作成してください。基本の公式は、2 * layers * hidden * seq * batch * dtype_bytesです。
  2. layers=32, hidden=4096, seq=4096, batch=1, dtype_bytes=2で計算した値を、/root/kv/base.txtに、bytes=<정수> gib=<소수 둘째자리>(プレースホルダーは整数と、小数第2位までの小数です)として書いてください。gibは2.00である必要があります。
  3. seq=131072に変えた値を、/root/kv/long.txtに同じ形式で書いてください。gibは64.00である必要があります。
  4. kv_groups=8を適用して、/root/kv/gqa.txtに書いてください。gibは8.00である必要があります(128K基準)。
  5. dtype_bytes=1を適用して、/root/kv/int8.txtに書いてください。gibは4.00である必要があります(128K、GQA-8基準)。
  6. /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です。
  7. /root/kv/kv_table.csvに、context,batch,gibのヘッダーと6行を書いてください。コンテキストは4096、32768、131072で、バッチは1と8です。値は、GQAなしのfp16基準です。

参考

公式を実装する

/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基準です。

コンテキストとバッチを変えながら、表を作ります。あとでキャパシティプランニングのときに、そのまま使える形で作ってください。