ホームへ戻る

推論系統

InferactがTPU推論用メガカーネルをオープンソース化、Kimi K3は低同時実行数で毎秒700トークン超のデコード

単一のPallas呼び出しで92層をつなぎ、層をまたぐ重みのプリフェッチでメモリ帯域幅の利用率を向上。公式の比較は特定のトポロジーと小規模バッチに限られ、投機的デコーディングの性能はドラフトトークンの受理率にも左右される。

Internet Archive Book Images · No restrictions · Image source
zh-Hant

Inferactは9月23日、TPU v7向けの推論用メガカーネルを公開した。PallasでKimi K3のデコードを実装している。公式の測定では、16個のTPUで単一ユーザー当たり毎秒700トークン超の出力を達成した。比較対象はvLLMのレシピを用いた16基のGB200で、毎秒452トークンだった。グラフでは投機的デコーディングを使用し、受理長を6としているため、結果はこの条件と併せて解釈する必要がある。[公式技術記事](https://inferact.ai/blog/tpu-megakernels)

この実装は92個のMoE層を1回のPallas呼び出しに融合し、現在の層でエキスパートの計算を行う間に、次の層のアテンション投影の重みをプリフェッチする。計算、通信、重みの転送を重ね合わせ、デコード中に高帯域幅メモリのインターフェースが遊休状態になる時間を減らす狙いだ。これは小規模バッチでのサービングで特に意味を持つ。毎回生成するトークンが少数でも、大量のモデル重みを繰り返し読み込む必要があるためだ。[カーネル設計の説明](https://inferact.ai/blog/tpu-megakernels)

この最適化を理解する鍵は、Pallasの実行制御の仕組みにある。JAXのドキュメントによると、TPUは主に命令を逐次実行する一方、DMA転送と行列演算はバックグラウンドで進められる。カーネルはオンチップのVMEMなどのバッファを使い、HBMへのアクセス待ちの時間を短縮する。これにより、エンジニアはデータの生存期間やプリフェッチの順序を制御できる。ただし、バッファの割り当てが容量を超えるとコンパイルに失敗し、融合そのものが高速化を保証するわけではない。[Pallasドキュメント](https://docs.jax.dev/en/latest/pallas/tpu/details.html)

これはコンパイラ研究にも具体的な事例を提供している。層をまたぐデータ転送の順序が、単一演算子の局所的な最適化よりも結果を大きく左右する可能性がある。この設計から推察すると、別のモデルに変更した場合、重みのサイズ、アテンションの状態、一時記憶領域の要件が変わるため、元のスケジュールをそのまま流用できるとは限らない。再現検証では数値的な正しさとメモリのスケジューリングを併せて確認し、どの部分が性能向上に寄与したのかを見極める必要がある。

再現時には、物理チップとフレームワーク上のデバイスも区別しなければならない。Googleのドキュメントでは、各IronwoodチップはJAXに2つのデバイスとして公開され、16チップ構成は4台のホストに分散すると明記されている。したがって、リポジトリに記載された32個のTPUデバイスは、記事の16個のチップと整合する。デプロイする際は、デバイス数だけを比較せず、トポロジーを照合する必要がある。[Googleのアーキテクチャドキュメント](https://docs.cloud.google.com/tpu/docs/tpu7x)

公開コードには対話型デモ、HTTPサービス、正しさを検証するテストが含まれる。ただしREADMEは、CPUテストではTPU向けコンパイル、VMEM容量、DMAスケジューリング、性能を検証できないと明記している。実装は2×2×4トポロジーを対象としており、同時実行数が多い場合には通信とベクトル演算が新たなボトルネックになり得る。サービス運用チームの次の検証では、モデルのバージョンとドラフトトークンの受理率を固定し、長いコンテキストでの性能、最初のトークンが出るまでの遅延、複数リクエスト負荷での性能を測定すべきだ。現時点の数値だけでは、サービス全体のコストは推計できない。[コードとテストの範囲](https://github.com/Inferact/tpu-megakernels)、[適用上の制約](https://inferact.ai/blog/tpu-megakernels)

出典

  1. 700 TPS on Kimi K3: A Case for TPU Megakernels
  2. Inferact/tpu-megakernels
  3. Writing TPU kernels with Pallas
  4. TPU7x (Ironwood)