ML Kernel 的性能工作,常常被两件事绑住:一是要理解底层硬件的执行模型,二是同一算法换到另一种加速器后,往往又要重写一遍实现。Helion 是面向 PyTorch 的高层 DSL,目标是让开发者以更高层的方式编写性能可移植的 Kernel。现在,Helion 与 Google 合作构建了 TPU 后端,可将 Helion Kernel 编译到 Pallas,为 PyTorch 工作流带来一条面向 TPU 的路径。
这件事的重点不只是“支持 TPU”。更有价值的是:Kernel 作者可以把算法表达、分块策略和硬件后端的适配尽量分层处理,减少框架边界和硬件边界同时侵入业务模型代码的情况。
从 PyTorch Kernel 描述到 Pallas
根据公布的信息,这条编译链路可以概括为:
PyTorch 调用方
-> Helion 高层 Kernel DSL
-> Helion TPU 后端
-> Pallas
-> TPU 执行
Helion 所在的位置很关键。它面向的是希望继续保持 PyTorch 开发体验、同时需要编写定制高性能算子的工程师;Pallas 则成为 TPU 方向的目标后端。这样,调用方不必把 TPU 专用实现直接散落在模型层或训练循环中。
这里的“可移植”不应理解为完全不需要调优。同一份 Kernel 的正确性表达可以更容易复用,但 tile 大小、内存访问方式、并行映射和数值行为仍可能随硬件而变化。可移植 DSL 降低的是实现和维护成本,不会自动消除性能工程本身。
异构 Kernel 开发的边界更清楚了
在 GPU 与 TPU 并存的环境中,常见的困难不是单独写出某个算子,而是长期维护多个版本:
- 模型代码要在不同设备上选择不同实现。
- 算法修复后,多个后端都要同步修改和验证。
- 基准测试容易只覆盖一个设备,另一个后端在规模扩大后才暴露问题。
- 专用实现进入训练代码后,部署环境的变化会变得昂贵。
Helion 的 TPU 后端提供了一种更清晰的职责划分:Kernel 作者在 Helion 层描述算子,后端负责将其导向相应目标。在架构上,可以把设备分发收敛在一个小入口中,而不是让训练代码到处出现设备判断。
下面的 Python 脚本可以作为项目接入前的环境检查。它不假定特定版本的 Helion API,因此可直接运行;把它保存为 check_kernel_env.py 后执行即可。
import importlib.util
import platform
packages = ("torch", "helion", "jax", "jaxlib")
print(f"Python platform: {platform.platform()}")
for package in packages:
available = importlib.util.find_spec(package) is not None
print(f"{package:8} {'available' if available else 'missing'}")
if importlib.util.find_spec("torch") is not None:
import torch
print(f"PyTorch version: {torch.__version__}")
print(f"CUDA available: {torch.cuda.is_available()}")
if importlib.util.find_spec("jax") is not None:
import jax
print("JAX devices:")
for device in jax.devices():
print(f" - {device}")
python check_kernel_env.py
这个检查的目的不是要求项目同时依赖 PyTorch 与 JAX。Pallas 是 TPU 后端的重要上下文,而实际依赖、安装方式和支持的版本应以 Helion TPU 后端的发布说明为准。尤其在托管 TPU 环境中,应先确认运行时暴露的设备与编译栈,而不是只根据本地机器的包安装结果判断可用性。
设计 Kernel 时,把“语义”与“落地方式”分开
如果准备把现有算子迁到这种多后端模型,可以先检查代码是否混合了三类逻辑:算子数学语义、分块与并行策略、特定设备的调用路径。前两者通常适合放在 Kernel 实现和其参数化配置附近,第三类则应留在很薄的后端选择层。
下面是一个可改造的项目组织示例。kernel_registry.py 是可运行的纯 Python 代码,用来说明分发边界;实际项目中,把占位函数替换为 Helion 编译后的 Kernel 调用即可。
# kernel_registry.py
from collections.abc import Callable
Kernel = Callable[[object], object]
def select_kernel(target: str, gpu_kernel: Kernel, tpu_kernel: Kernel) -> Kernel:
targets = {
"gpu": gpu_kernel,
"tpu": tpu_kernel,
}
try:
return targets[target]
except KeyError as exc:
supported = ", ".join(sorted(targets))
raise ValueError(f"unsupported target {target!r}; use: {supported}") from exc
def reference_scale(x):
return [value * 0.5 for value in x]
if __name__ == "__main__":
kernel = select_kernel("tpu", reference_scale, reference_scale)
print(kernel([2.0, 4.0, 8.0]))
python kernel_registry.py
# [1.0, 2.0, 4.0]
真实接入时,gpu_kernel 与 tpu_kernel 不一定需要是两份不同的源码。理想状态是它们来自同一个 Helion Kernel 定义,只是由不同后端编译或加载。上例刻意没有伪造 Helion 的具体 DSL 语法;具体的 Kernel 定义、编译入口和 Pallas 相关配置,应以当前 TPU 后端 API 为准。
验证不能只看“能跑”
把 Kernel 编译到新后端后,至少要把验证分为三层:
- 数值正确性:使用参考 PyTorch 实现,对随机输入、边界形状、非连续布局和不同数据类型做比对。
- 设备语义:确认 dtype、精度策略、归约顺序以及异常输入的行为符合训练或推理需求。
- 端到端性能:分别测量编译成本、单 Kernel 延迟、吞吐量和真实模型中的占比,避免只在微基准中得出结论。
对于归约、softmax、归一化这类数值敏感算子,还应预先定义容差,而不是用逐元素完全相等作为跨硬件验收标准。对于动态形状场景,则需要额外观察编译缓存与 shape 变化带来的开销。
采用建议
Helion 的 TPU 后端适合那些已经在 PyTorch 中维护定制算子、同时需要覆盖 TPU 的团队。第一步不应是把所有 Kernel 一次性迁移,而应选择一个热点算子做试点:建立参考实现,确认 TPU 编译路径,测量正确性与端到端收益,再决定是否扩大覆盖范围。
需要持续保留的工程纪律是:让模型代码依赖稳定的 Kernel 接口,让硬件特定编译细节停留在后端层,并把每个后端的性能数据纳入持续基准。这样,硬件异构才不会演变成多套难以同步的 Kernel 代码库。