# How to Scale Your Model — Part 7: Inference [[wiki/entities/How to Scale Your Model|How to Scale Your Model]](Google DeepMind による Transformer スケーリング解説シリーズ)の第 7 部。Transformer の学習と推論がいかに異なるかを、レイテンシという新しい制約軸を中心に体系立てて解説する。学習ではチップあたりのスループットだけを見ればよいが、推論では Time To First Token(TTFT)とトークンあたりレイテンシという 2 つの指標が加わる。 ## Transformer 推論の基礎: Prefill と Generation 素朴なサンプリングは毎ステップでプレフィックス全体を再計算するため $\Theta(n^2)$ の実行時間になる。これを避ける鍵が [[KVキャッシュ]](KV cache)である。各トークンの key/value 射影を保存しておけば、後続トークンは過去のトークンに対して新規に FLOPs を払わずに済む。 これにより推論は性質の異なる 2 段階に分かれる。 - **Prefill**: 長いプロンプトの全トークンを一括処理し、key-value 射影を KV キャッシュに保存する。最後のトークンの logits も保存する。 - **Generation**: KV キャッシュと直前の logits から 1 トークンをサンプリングし、それをモデルに戻して次の logits を作る。新トークンの KV 射影をキャッシュに追記する。`<EOS>` または最大長に達するまで繰り返す。 ![[_attachments/scaling-book-inference/fig01-kv-cache-sampling.webp]] (KV キャッシュを使ったサンプリングの模式図。Prefill はプロンプトを処理してキャッシュを初期化し compute-bound になる。Generation はキャッシュと直前 logits から 1 トークンずつ生成し、memory bandwidth-bound になる。) KV キャッシュ導入後の計算量は FFW で $O(n)$、attention で $O(n^2)$ に下がる(素朴実装は FFW が $O(n^2)$、attention が $O(n^3)$)。Prefill と Generation は「見た目は同じ Transformer 呼び出しだが中身は別物」というのが本稿の中心テーマである。 ## 何を最適化すべきか 学習で見るべき指標はチップあたりスループットのみだが、推論では TTFT とトークンあたりレイテンシが新たに重要になる。 - **オフラインバッチ推論**(eval・データ生成): 個々のサンプルのレイテンシは無視してよく、バルクコストだけが問題。 - **チャット/ストリーミング**: 低い TTFT と、人間の読み速度を超えるトークン生成速度が必要。 - **エッジ推論**(`llama.cpp` など): 同時に 1 ユーザーだけを最小レイテンシで捌けばよいが、ハードウェア制約が厳しい。 ハードウェア利用率(MFU)の最大化はコストと TTFT には効くが、学習と違って個々のユーザー体験に自動的には直結しない。 ## 線形演算のボトルネック: 演算強度と critical batch size MLP の $W_{in}/W_{out}$ や attention の QKV/O 射影といった行列積は、[[演算強度]](arithmetic intensity)がバッチサイズ $B$ に依存する。$\text{bf16}[B,D]$ と $\text{bf16}[D,F]$ の行列積で compute-bound になる条件を解くと、臨界バッチサイズは $B_\text{crit} = \beta \cdot \alpha_\text{hbm} = \frac{\text{bits per param}}{\text{bits per activation}} \cdot \frac{C}{W_\text{hbm}}$ bf16 の TPU v5e では $B_\text{crit} \approx 240$ トークン、H100 では約 280 トークンとなる。int8 パラメータ + bf16 FLOPs なら 120 に下がり、int8 FLOPs + int8 パラメータなら再び 240 に戻る。 - **Prefill** はプロンプトが数百〜数千トークンあるため、基本的に常に compute-bound になる。MFU を最大化するだけでコストとレイテンシ(TTFT)の両方が最大化される。 - **Generation** はトークンを逐次生成するため、この臨界バッチサイズを超えるには複数リクエストを同時にバッチする必要がある。これは実運用では難しく(240 の同時リクエストと 240 個の独立した KV キャッシュを意味する)、**generation を FLOPs で飽和させるのは prefill よりはるかに難しい**。 ## Attention のボトルネック Flash Attention 融合における 1 attention head の演算強度は $\frac{ST}{S+T}$ ($S$: KV キャッシュの長さ、$T$: クエリの長さ)。Prefill は自己注意なので $S=T$ となり $T/2$、すなわち $\Theta(T)$ で成長するため長いシーケンスなら compute-bound になりやすい。一方 Generation は $T=1$ なので $ST/(S+T) \approx 1$ に潰れ、**attention は generation 時ほぼ常に memory bandwidth-bound** になる。線形演算はパラメータをバッチ全体で使い回せるので compute-bound になりやすいのに対し、KV キャッシュはリクエストごとに個別なので、バッチを増やすほど KV キャッシュ由来のメモリ負荷も比例して増える非対称性がここにある。 ## レイテンシとスループットの理論下限 小バッチの generation では attention・MLP とも memory bandwidth-bound と仮定してよく、 $\text{Theoretical Min Step Time} = \frac{\text{Batch Size} \times \text{KV Cache Size} + \text{Parameter Size}}{\text{Total Memory Bandwidth}}$ バッチが大きくなり FLOPs がパラメータロードを上回ると、より一般的な式 $\text{Step Time} = \underbrace{\frac{B \times \text{KV Cache Size}}{W_\text{hbm}}}_{\text{attention(常に bandwidth-bound)}} + \max\left(\underbrace{\frac{2B \times \text{Params}}{\text{FLOPs/s}}}_{\text{MLP(compute-bound になりうる)}}, \frac{\text{Params Size}}{W_\text{hbm}}\right)$ になる。attention 項は roofline を必要としない(常に bandwidth-bound)。TPU v5e 4x4・30B dense モデル・int8・8192 context・100kB/token の KV キャッシュという設定で、バッチ 4 なら約 2.5ms、バッチ 256 なら約 21ms という具体的な下限が計算できる。バッチサイズはレイテンシとスループットのトレードオフを直接動かすノブであり、[ESTI 論文](https://arxiv.org/pdf/2211.05102)の PaLM 実測でもスループットはバッチ 240 付近で頭打ちになる。 ## KV キャッシュのメモリフットプリント LLaMA 2-13B(L=40, D=5120, F=13824, N=K=40, H=128)を例に取ると、パラメータは合計 13e9(bf16 で 26GB)。一方 KV キャッシュのサイズは $\text{KV cache size} = 2 \cdot \text{bytes per float} \cdot H \cdot K \cdot L \cdot T$ 8192 トークンの単一シーケンスで 6.7GB(bf16)にもなり、**わずか 4 シーケンス分でパラメータサイズを超える**。LLaMA 2 は KV キャッシュサイズ最適化がされていない例だが(LLaMA-3 は $K$ をずっと小さくしている)、この非対称性を無視すると推論のメモリ・レイテンシ見積もりを誤る。 実際、KV ヘッド数を 1:5 に削減する([[Grouped-Query Attention|GMQA]])だけでバッチ 240 での理論スループットは 963 tok/s から 4,529 tok/s へと 4.7 倍に伸び、最大バッチサイズ自体も広がる。 ## 生成スループット/レイテンシ改善のテクニック いずれも KV キャッシュを小さくする方向の工夫である。 - **[[Grouped-Query Attention]](GMQA/GQA)**: KV ヘッド数を減らし複数の Q ヘッドで共有する。極端な場合は全 Q ヘッドで単一の KV ヘッドを共有する(Multi-Query Attention)。品質への影響は比較的小さいとされる。 ![[_attachments/scaling-book-inference/fig02-gmqa-attention-variants.webp]] (Multi-head・Grouped-query・Multi-query attention の比較。Grouped-query は Q ヘッドのグループごとに 1 組の KV ヘッドを共有し、両極端の中間に位置する。) - **ローカル attention の混在**: attention のコンテキストを一定窓に制限する層を混ぜることで、その層の KV キャッシュ上限を抑える。 - **層間での KV 共有**: 複数層で同じ KV キャッシュを共有する。キャッシュサイズは減るが HBM から複数回読み直す必要があり、ステップタイムは必ずしも改善しない。 - **量子化**: パラメータと KV を int8/int4/fp8 などに量子化し、メモリ帯域を節約する。学習後量子化(post-training quantization)も可能。 - **ragged HBM read と [[PagedAttention]]**: 8k 分の枠を確保していても実際にはそこまで使わないリクエストが多いため、パディング部分を読まないカーネルが有効。PagedAttention は KV キャッシュを OS のページテーブル風に管理し、パディングをほぼ排除する。 ![[_attachments/scaling-book-inference/fig03-paged-attention.webp]] (PagedAttention のイメージ。クエリベクトルは非連続なブロックに分散配置された KV キャッシュへアクセスする。) これらを組み合わせると KV キャッシュサイズを標準 MHA 比で一桁以上削減でき、推論コストも一桁改善しうる。 ## 複数アクセラレータへの分散 ### Prefill のシャーディング Prefill は学習とほぼ同じ roofline 構造を持ち、モデル(Megatron)並列・シーケンス並列・パイプライン・FSDP まで学習と同じ手法がそのまま使える。基本方針は「ICI バウンドになるまでモデル並列(だいたい 4〜8 way)、その先はシーケンス並列」。 ### Generation のシャーディング Generation は事情が大きく異なる。 1. **FSDP は不可能**: パラメータと KV キャッシュを HBM から MXU へ運ぶ帯域がボトルネックなので、それらを ICI 経由で動かすのは論外。動かすべきは activation の方である。 2. **データ並列に意味がない**: 単にモデルのコピーを増やすのと同じで、パラメータロードは速くならない。 3. **シーケンス並列も不可**: そもそも生成時のシーケンス次元は 1。 残るのはモデル並列のバリエーションのみだが、generation は memory bandwidth-bound であることが多いため、学習で使う ICI バウンドを超えてモデル並列度を上げ、スループットをわずかに犠牲にしてレイテンシを改善できる。$\beta = W_\text{hbm}/W_\text{ici}$(TPU v5e/v6e でおよそ 8)として、$Y > F/(B \cdot \beta)$ まではモデル並列を増やしてよい。 ### KV キャッシュのシャーディング KV キャッシュはできる限り複製を避けたい。まず head 次元で Megatron シャーディングし($K$ way が上限)、それ以上はバッチ次元でシャーディングする。この構成では、activation をモデルシャーディングからバッチシャーディングへ切り替えるために attention 層ごとに AllToAll が 2 回必要になる。 ![[_attachments/scaling-book-inference/fig06-kv-cache-sharding.webp]] ((a) 純粋なモデルシャーディングの Multi-head attention と (b) KV キャッシュをバッチシャーディングする Multi-query attention の比較。activation を model sharding から batch sharding へ移すために 2 回の AllToAll が追加で必要になる。) ## 推論エンジンの設計 素朴な実装(Prefill バッチをまとめて処理してから Generation バッチを回す)には以下の欠点がある。 1. TTFT が悪い(全 prefill が終わるまでユーザーは何も見えない) 2. 生成の短いリクエストが長いリクエストにブロックされ、バッチスロットが無駄になる 3. Prefill が最長シーケンスにパディングされ計算を浪費する 4. Prefill と Generation が同じシャーディングを共有せざるを得ない Prefill をバッチサイズ 1 で行い Generation だけ複数リクエストをバッチする **interleaved** 構成は TTFT を改善するが、Prefill 実行中は他リクエストの Generation が止まってしまう。これを解決するのが [[Prefill-Decode分離]](**disaggregated** serving)で、Prefill サーバーと Generate サーバーを分離し、KV キャッシュをネットワーク越しに転送する。 ![[_attachments/scaling-book-inference/fig04-disaggregated-serving.webp]] (Disaggregated serving の構成。Prefill スライス群が history cache と未バッチの prefill キューを持ち、生成された KV キャッシュは insert queue を経由して Generate スライスの継続的バッチへ挿入される。) 利点は (1) 他ユーザーの prefill に律速されない低レイテンシ、(2) Prefill/Generation それぞれに最適なシャーディング戦略とハードウェアを独立に選べる specialization。欠点は KV キャッシュのネットワーク転送コスト。 ### 継続的バッチング [[動的バッチングと継続的バッチング|Continuous batching]]は、可変長コンテキストを扱い KV バッファへ結果を挿入する prefill 関数と、現在アクティブな全リクエストに対して 1 ステップ生成する generate 関数を、空いた generate スロットに応じて orchestrator が呼び分ける方式である。 ### プレフィックスキャッシュ 自己回帰性から、`["I","like","dogs"]` と `["I","like","cats"]` は先頭 2 トークンの KV キャッシュが同一になる。この重複を[[プレフィックスキャッシュ|プレフィックスキャッシュ(prefix caching)]]として再利用すれば prefill の計算量を大きく削減できる。チャットボットの往復対話や few-shot プロンプト・システム指示で特に効果が大きい。実運用では HBM の空きだけでなく Host DRAM(8xTPUv5e で約 450GiB、HBM 128GiB よりずっと遅いが read には十分速い)も使われ、キャッシュと検索は LRU のトライ木で自然に表現できる。 ![[_attachments/scaling-book-inference/fig05-prefix-caching-trie.webp]] (LRU トライ木として実装したプレフィックスキャッシュ。プレフィックスを共有することで KV メモリの重複を避けられる。) ### 実装例: JetStream Google がオープンソース化した [[JetStream]] は、prefill engine と generate engine を別 TPU スライス上に置き単一のコントローラで orchestrate する。prefill thread・generate thread・KV キャッシュ転送を担う transfer thread の 3 つで構成され、Engine インターフェースは `prefill`(トークン列から KV キャッシュを生成)・`insert`(KV キャッシュを generate 中のバッチへ挿入)・`generate`(バッチ化された KV キャッシュから 1 トークンずつ生成)の 3 メソッドを持つ。PyTorch 版の [jetstream-pytorch](https://github.com/google/jetstream-pytorch) も存在する。 ## Worked Problems の要点 本文には LLaMA-2 13B ベースの架空モデル(L=64, D=4096, F=16384, N=32, K=8, H=256)を使った演習が 7 問収録されている。主な結論のみ抜粋する。 - パラメータ数は約 18.4B、KV キャッシュは int8 で 1 トークンあたり 262kB。 - TPUv5e 4x4・128k コンテキスト・int8 なら、フルシャーディングしても最大バッチサイズは約 7(K=1 の MQA なら約 56 まで拡大)。 - パラメータロードのみで見た per-step latency 下限は約 1.4ms。 - MoE(E=16, k=2)にすると総パラメータは約 12 倍(212B)だが activated パラメータは 2 倍未満(31.2B)にしかならず、compute-bound に必要なトークン数は $E/k$ 倍(この設定で約 1920 トークン)に増える。KV キャッシュサイズは dense と変わらない。 - 2D weight-stationary sharding(Appendix B、[ESTI 論文](https://arxiv.org/abs/2211.05102))は、$F=4D$ のとき $N > 32 \cdot (F/D) \cdot (3/4)^2$ を満たすチップ数で 1D モデル並列より通信量が小さくなる。 ## Appendix の要点 - **Appendix A**: バッチサイズ 240 が compute-bound の境界という単純ルールは、TPU が通信中にも重みをプリフェッチできるため厳密なステップ関数にはならず、実測では緩やかに立ち上がってから線形になる。 - **Appendix C(レイテンシバウンドな通信)**: ICI 通信はバイト数が小さいとレイテンシ項に支配される。TPU では 1 ホップあたりの転送時間が 1μs を下回るとディスパッチの固定オーバーヘッドが支配的になり、8-way Megatron では 360kB 未満で発生する。推論の小バッチはこの領域に入りやすい。 - **Appendix D**: [[Speculative Decoding|投機的サンプリング(speculative sampling)]]は小さなドラフトモデルで複数トークンを先に生成し、大モデルで 1 回のフォワードパスにまとめて検証する。大モデル側は compute-bound ではないため、検証を「ついで」に行えるのがレイテンシ面の勝因。draft ヘッドを本体モデルに埋め込む手法(パラメータ共有でドラフト自体も高速化)にも触れている。greedy decoding では最長一致プレフィックスを採用し、非 greedy な場合は Metropolis–Hastings 的な棄却サンプリングで分布を保つ。 ## 出典 - ソース: `.raw/articles/scaling-book-inference-2026-09-23.md` - 原文: [How To Scale Your Model — Inference](https://jax-ml.github.io/scaling-book/inference/)([[wiki/entities/How to Scale Your Model|How to Scale Your Model]] Part 7、著者陣は Google DeepMind、現在は一部 MatX 在籍)