Triton-TLE 实战:三层扩展写出更快的 GPU 算子

Triton 把 GPU 算子开发拉回了「写 Python 就行」的时代:面向 Tile 编程,线程映射、寄存器分配、流水线全交给编译器。但新硬件(Thread Block Cluster、DSMEM)和新算子(MoE 路由、百万 token 稀疏注意力)一上来,Triton 的抽象就开始漏风——要么手写 CUDA,要么跟 DSL 搏斗。

智源 FlagOS 的开源编译器 FlagTree(Triton 多后端 fork)用 TLE(Triton Language Extensions) 补上这个缺口:在 Triton 语法之上做三层渐进扩展,不替代 Triton,存量 kernel 小改即可上手。

三层设计:按需选择优化深度

  • TLE-Lite:轻量语义提示,面向算法工程师。子 Tile 切片、异步加载、生产者-消费者流水线、Device Mesh,一次编写、多端运行。
  • TLE-Struct:架构感知控制(GPGPU / DSA)。你表达结构化的内存意图,编译器把「本地 Buffer」映射到 GPU 的 Shared Memory 或 DSA 的 Scratchpad。
  • TLE-Raw:原生透传。性能专家可以直接内联 CUDA/汇编或厂商 Intrinsic,走厂商编译管线。

三层最终都经 LLVM IR 汇入同一个 kernel,一个程序里可以混用。

核心原语长什么样

1. 子 Tile 切片:告别手算偏移

经典痛点:从大 Tensor 里取逻辑子块,要自己算偏移、构造 Mask、处理边界。TLE-Lite 的 extract_tile / insert_tile 直接表达意图:

# x 是 [4, 4];按 2x2 子块切分,取 [0, 0] 处的子块
z = x.extract_tile(index=[0, 0], shape=[2, 2])

# y 是 [2, 2];写回 x 的 [0, 0] 子块位置
z = x.insert_tile(y, index=[0, 0])

编译器识别出这是规则 Tile 访问后,能针对数据布局、向量化、Bank Conflict 做优化——稀疏注意力、局部归一化、分块统计这类算子直接受益。

2. tle.pipe:显式的生产者-消费者流水线

tl.range(num_stages=...) 的自动流水是黑盒。tle.pipe 把数据流边显式化:生产者 acquire 槽位→填数据→commit;消费者 wait→读→release

stage_buf = tle.gpu.alloc([2, BLOCK], dtype=tl.float32, scope=tle.gpu.smem)
pipe = tle.pipe(capacity=2, scope="cta", name="x_pipe", x=stage_buf)
writer, reader = pipe.writer(), pipe.reader()
offs = tl.arange(0, BLOCK)

for k in tl.range(0, n_tiles):          # 生产者
    slot = writer.acquire(k)
    tl.store(tle.gpu.local_ptr(slot.x), tl.load(x_ptr + k * BLOCK + offs))
    writer.commit(k)

for k in tl.range(0, n_tiles):          # 消费者
    wait = reader.wait(k)
    x = tl.load(tle.gpu.local_ptr(wait.slot.x))
    acc += x
    reader.release(k)

Barrier、缓冲区复用、同步交给编译器,但结构可分析——这正是编译器能做负载/计算重叠的前提。

3. 集群协作:长序列 TopK 选择器

最有说服力的例子是 DeepSeek Sparse Attention 的 TopK 选择器。batch=1 的长序列单个 block 内几乎没有并行度,TLE 的做法是把一行拆给一个 block cluster 协作处理:

topology = {
    "node": [("node_x", 2), ("node_y", 2)],
    "device": 4,
    "block_cluster": [("cluster_x", 2), ("cluster_y", 2)],
    "block": 4,
}
mesh = tle.device_mesh(topology=topology)

# 每个 block 在共享内存里维护本地直方图
s_histogram = tle.gpu.alloc([4096], dtype=tl.int32, scope=tle.gpu.smem)

# 通过 remote 把各 block 的直方图汇总到 rank 0,再按集群做作用域同步
tle.remote(...)               # 读其他 block 的片上内存
tle.distributed_barrier(mesh) # 只同步集群内,不同步全世界

直方图更新、候选写入、最终排序全部在片上闭环完成,不落全局内存;单行扫描被分散到多个 block。完整 kernel 在仓库 python/tutorials/tle/deepseek_v32/01-topk_selector.py

性能数据

  • TopK 选择器(H800,batch=1,131K token):TLE 集群版 0.030 ms,FlashInfer 0.045 ms、TRT-LLM 0.049 ms,相对 TRT-LLM 最高约 2.5×;512K 长度依然稳定,而 TileLang 此时候选集已溢出。
  • Radix Select:TLE 复刻 TRT-LLM 算法,多组 Shape 下达到原生约 85%–97% 性能。
  • SparseMLA(128K 上下文):用 Pipeline 原语表达协作,性能约为 FlashMLA 基线的 90%。
  • FlagOSTune 自动调优:H20 上 MM 算子搜索空间从 62 万余配置压到 4070 个,调优效率加速 120 倍且性能几乎无损;在 NVIDIA、摩尔线程、沐曦多款算力上实测提速 1.21–7.35×。
  • 编译优化:布局转换消除 68%–79%(最高提速 71%);指令重排提前发射独立 Load,平均提速 1.19–1.61×、峰值 2×。

上手路径

  1. git clone https://github.com/flagos-ai/flagtree.git,切 main 分支(NVIDIA 后端,Triton 3.6;昇腾在 triton_v3.5.x,寒武纪在 triton_v3.2.x)。
  2. NVIDIA 用户手册 编译安装。
  3. python/tutorials/tle/ 下的教程,先从 TopK 选择器开始。
  4. 迁移策略:先在存量 kernel 上加 TLE-Lite 提示;只有 profile 出瓶颈后再上 TLE-Struct / Raw。

小结

TLE 不是新语言,而是 Triton 抽象天花板的出口。写稀疏注意力、MoE 或多芯片算子的同学值得花一个周末试试——TopK 选择器教程是理解集群级 Triton 最快的方式。

资源

滚动至顶部