Production Ready Tilelang Ops Framework
在生产环境下使用 Tilelang,需要考虑几个问题: 算子缓存 与 torch.compile 兼容 算子缓存 算子定义 在了解算子缓存之前,我们需要了解一个 GPU 算子应该如何定义,通常由三部分组成: 参数类型 RMSNorm 中的实例 作用 问题规格(Workload / specialization parameters) M, N, eps, dtype 确定一类具体计算问题 调度参数(Schedule / config / meta-parameters) block_m, threads 确定该 workload 的一种实现方案 运行时参数(Runtime arguments / operands) x, weight 每次调用时传入的实际数据 在代码里对应得非常直接: _rms_norm_kernel(M, N, eps, dtype) # workload parameters def _func(block_m, threads): # schedule/meta-parameters def main(x, weight, y): # runtime arguments 这三类参数不是一次性传给同一个函数,而是在 kernel 从定义到执行的过程中逐层绑定: jit_impl = _rms_norm_kernel(M, N, eps, dtype) jit_kernel = jit_impl(block_m=4, threads=128) y = jit_kernel(x, weight) 第一步绑定 workload 参数,得到 TileLang JITImpl 描述一个确定的计算问题,例如处理形状为 (M, N)、数据类型为 dtype 的 RMSNorm,但尚未确定具体的调度方案。 第二步绑定 schedule 参数,得到可以执行的 JITKernel block_m/threads 可以来自默认配置、手动配置或 autotune 不同 schedule 对应同一个 workload 的不同实现版本 最后传入 x/weight 等 runtime arguments,启动已经选定的 JITKernel 运行数据只参与本次调用,不改变前面已经确定的 workload 和 schedule 算子调用流程 一次真实的算子调用从业务 Op 开始,依次确定 workload、选择 schedule、取得编译结果,最后才把 Tensor 作为 runtime arguments 启动 kernel ...