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 が通常それを調整します。
- autocast の下では、プログラムが dtype アサートを持ち、AOTI が 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")。