文档导航

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 缓冲区。在 bf16、12 个头的配置下,32 行 x 1024 token 时就是 0.8 GB。

  • GPU 有余量。 这个缓冲区消耗带宽,每个长批次几十毫秒。
  • GPU 接近占满。 缓存分配器会颠簸,同一次调用可能要几十秒。

如果在繁忙的 GPU 上用长批次编译,就限制批大小(predict_batch(..., batch_size=) or use fast=True。TileLang 注意力读取打包的 QKV 缓冲区,并按序列长度做掩码, 所以它没有这样的缓冲区。

冷启动

  • 首次编译。 每个图要几十秒。compile=True 需要两个图:一个用于批,一个用于单行,后者由 torch 特化。agent.warmup()(#718)在流量到来之前把两者都构建好。
  • 跨重启。 Inductor 的 FX-graph 缓存把编译好的图放在 TORCHINDUCTOR_CACHE_DIR 下。默认 在 /tmp 下,重启机器或重启容器就没了。把它设成一个持久目录,或者容器里的一个卷,第二个进程 就会加载这些图,而不是重新编译。在 #472 的实测里,这把预热从约 120 秒降到约 50 秒。

AOTInductor:还不行

为每个 checkpoint 和 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 副本出发,trace 会在前向里失败: 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 后端注册了 target。没有 AMD 或 Apple 硬件时, 还能回答的问题是:laya/tl_kernels.py 到底能不能为 CPU 降低。在 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,能跑 CPU 张量,但只覆盖这门语言的一个子集。

每个 Laya 内核在 "c" 上都失败,原因属于这三类之一:

内核 在 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,它会降低成一个普通循环;
  • 带 T.alloc_local 累加器的 T.gemm,它会降低成一个标量三重循环。

所以一个 CPU 版本会是第二套内核,而不是一个 target 开关。它的 GEMM 会是不分块的标量循环, 也就没法跟原装前向在 CPU 上已经在用的 MKL/oneDNN 路径竞争。这三样构造也是在 HIP 和 Metal 上 首先要检查的:fragment、带 GemmWarpPolicy 的 T.gemm,以及 bf16/fp16 支持。

要复现第二张表的第一行:tilelang.compile(K.gemm_kernel(768, 768).prim_func, target="c")。