PyTorch Helion 增加 TPU 后端:用高层 DSL 写异构硬件内核
导语
PyTorch 生态正在把自定义高性能内核的门槛继续往上层抽象推进。PyTorch Blog 最新介绍的 Helion TPU 后端,核心目标不是替代底层优化专家,而是让更多 PyTorch 用户能够用相对熟悉的高层 DSL 编写 TPU 内核,并由编译器生成 Pallas 代码、自动探索更合适的流水线方案。
这项工作由 PyTorch 团队与 Google 合作完成。按照原文披露,在 Flash Attention 工作负载上,Helion 生成的 TPU 内核在 TPU v7 上达到 838 TFLOPs,约为单个 tensor core 的 79% MFU。对一个仍在向异构硬件扩展的内核 DSL 来说,这个结果说明其性能潜力已经接近需要严肃对待的阶段。
核心要点
- Helion 面向“性能可移植”内核编写:开发者用 PyTorch 风格更强的 Helion 写 kernel,后端再面向不同硬件生成优化代码。此次新增的是将 Helion 编译到 Pallas 的 TPU 路径。
- TPU 编程难点在显式内存与流水线:与 GPU 依靠大规模并行线程和硬件管理缓存不同,TPU 更强调顺序执行、宽向量寄存器、显式内存空间,以及 HBM 与片上 VMEM 之间的数据搬运。高性能 Pallas 内核往往要手动安排异步拷贝与计算重叠。
- 编译器负责生成嵌套流水线:在简单 add 示例中,Helion 将 tile 循环映射为 host 侧 grid,并通过 Pallas 调用让数据加载和计算重叠;在 Flash Attention 中,外层处理 Q tile,内层围绕 K/V tile 生成或选择流水线策略。
- 自动调优是关键卖点:Helion 会在不同 block size、pipeline buffer size、内层循环生成方式之间搜索。对于某些形状,它可以选择 Pallas 的
emit_pipeline;在 VMEM 允许时,也可能预取更多数据并展开循环,以减少计算单元空泡。
意义与影响
这篇文章的真正信号,是 PyTorch 生态正在认真面对“GPU 之外的主流训练与推理硬件”。TPU 与 NVIDIA B200 等 GPU 在 BF16 算力和 HBM 带宽等指标上已具有可比性,但开发体验并不相同。过去,想写出高性能 TPU kernel 往往需要深入 Pallas 和 TPU 内存层级;Helion 试图把这部分复杂性收敛到编译器和自动调优器中。
对工程团队而言,这带来三类价值:第一,性能关键路径可以在不完全掌握 Pallas 的情况下开始优化;第二,同一套高层 kernel 有机会同时服务 TPU 与 GPU,降低多硬件维护成本;第三,面对 Flash Attention 这类形状敏感的算子,自动调优能减少人工试错。
当然,素材展示的是特定工作负载与形状下的结果,不能简单外推为所有 TPU kernel 都能自动获得同等效率。但方向已经很清晰:AI 框架的竞争不只在模型 API,也在谁能把复杂硬件的性能更稳定地交给普通开发者。
来源:PyTorch Blog
评论
正在确认登录状态……
正在加载评论……