ドキュメント

compile=True と TileLang 高速パス:エンジニアリングノート

compile=True と TileLang 高速パス:エンジニアリングノート

これらのノートは、compile=True と fast=True が README の記述を超えてどう振る舞うかを扱います。 #472、#576、#718 の作業中に、RTX 4070 Ti SUPER、torch 2.11、tilelang 0.1.14 で計測した結果です。 次の人がもう一度測らずに済むように置いています。

compile=True はアテンションマスクを実体化する

Eager SDPA は ModernBERT の (rows, 1, L, L) アテンションマスクをブロードキャストビューとして 扱います。compile=True が使う動的な形状の下では、inductor は最後の次元が整列していることを 証明できません。マスクをすべてのヘッドに展開し、rows x heads x L x L の実バッファに詰め込み ます。12 ヘッドの bf16 では、32 行 x 1024 token で 0.8 GB になります。

  • GPU に余裕がある場合。 このバッファは帯域を消費し、長いバッチあたり数十ミリ秒かかります。
  • GPU がほぼ満杯の場合。 キャッシュアロケータがスラッシングし、同じ呼び出しに数十秒かかる ことがあります。

ビジーな GPU で長いバッチをコンパイルするなら、バッチサイズを制限するか(predict_batch(..., batch_size=) or use fast=True を使ってください。TileLang のアテンションは、パックされた QKV バッファを読み、系列長でマスクするので、そのようなバッファを持ちません。

コールドスタート

  • 初回のコンパイル。 グラフごとに数十秒かかります。compile=True には 2 つのグラフが必要 です。1 つはバッチ用、もう 1 つは 1 行用で、後者は torch が特殊化します。agent.warmup() (#718)は、トラフィックが来る前に両方を構築します。
  • 再起動をまたぐ場合。 Inductor の FX グラフキャッシュは、コンパイル済みのグラフを TORCHINDUCTOR_CACHE_DIR の下に保ちます。既定は /tmp の下で、再起動やコンテナの再起動では 残りません。永続ディレクトリ、またはコンテナのボリュームに設定すれば、2 つ目のプロセスは コンパイルせずにグラフを読み込みます。#472 の計測では、これでウォームアップが約 120 秒から 約 50 秒になりました。

AOTInductor:まだ

チェックポイントと GPU アーキテクチャごとに事前コンパイル済みの成果物を配布する (torch._inductor.aoti_compile_and_package)と、コンパイルを完全になくせます。torch 2.11 では、パッケージングで止まります:

  • エクスポートは成功する。 DecisionModel の torch.export は約 5 秒で成功し、動的な rows、 markers、tokens を扱えます。tokens は 16 の倍数として宣言する必要があります(16 * Dim(...))。 通常の range は、エクスポータ自身の L % 8 整列ガードを通りません。これは上と同じマスク整列 です。
  • パッケージングは失敗する。 失敗の仕方は、プログラムをどうエクスポートしたかで決まります:
    • autocast の下では、プログラムが dtype アサートを持ち、AOTI が autocast の外でそれに 引っかかります:Tensor dtype mismatch! Expected: torch.bfloat16, Got: torch.float32。
    • autocast なしの bf16 コピーからでは、トレースがフォワード内で失敗します: mat1 and mat2 must have the same dtype。DecisionModel.forward は action head の前に、 プールされた state と信頼度特徴を fp32 にアップキャストし、autocast が通常それを調整します。

したがって成果物の道には、dtype を明示した action head が必要です。その入力を head の dtype に キャストするか、head を fp32 で動かすかのどちらかです。

TileLang の可搬性:カーネルは CUDA 専用

tilelang は CUDA、HIP、Metal、WebGPU、および C バックエンド用のターゲットを登録します。AMD や Apple のハードウェアがない場合に答えられる問いは、laya/tl_kernels.py がそもそも CPU 向けに lower できるかどうかでした。Linux x86-64 で tilelang.compile(kernel.prim_func, target=...) を使って調べました:

target 結果
"cpu" 最初に拒否される:Target cpu is not supported。tilelang の CPU バックエンドは "c"。
"llvm" Cannot find global function target.build.llvm。この wheel は LLVM バックエンドを積んでいない。
"c" C に lower され CPU テンソルで動くが、言語の一部分に限られる。

すべての Laya カーネルが "c" で失敗し、原因は 3 つのうちのいずれかです:

カーネル target="c" での失敗 使う構文
gemm_kernel, gemm_geglu_kernel CPU fill only supports local and global buffers, but got dst scope local.fragment T.alloc_fragment アキュムレータ
add_ln_kernel CPU reduce only supports local src and local/local.var dst buffers fragment 上の T.reduce_sum / T.reduce_max
rope_kernel Cannot convert type bfloat16 to C type bf16 テンソル
attn_kernel T.alloc_fragment で失敗 fragment

C バックエンドは次を受け付けます:

  • fp32 の要素ごとのループ(T.Parallel);
  • T.Pipelined。これは通常のループに lower されます;
  • T.alloc_local アキュムレータを持つ T.gemm。これはスカラーの三重ループに lower されます。

したがって CPU 版は、ターゲットフラグではなく第 2 のカーネル群になります。その GEMM はブロック 化されないスカラーループになり、標準のフォワードが CPU ですでに使っている MKL/oneDNN の経路とは 競合しません。同じ 3 つの構文は、HIP と Metal で最初に確認すべきものでもあります:fragment、 GemmWarpPolicy を持つ T.gemm、bf16/fp16 対応です。

2 番目の表の 1 行目を再現するには:tilelang.compile(K.gemm_kernel(768, 768).prim_func, target="c")。