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 通常会把这个协调好。
- 在 autocast 下,程序带着 dtype 断言,AOTI 在 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")。