Pallas
也称为 Pallas、Pallas kernel、JAX Pallas、TPU 自定义 kernel
先用大白话
Pallas 是工程师在编译器调度不够好时,为 TPU 手写调优 kernel 的方式。
技术定义
Pallas 是 JAX 的扩展,用于编写 TPU 和 GPU 的自定义 kernel,可以显式控制切分、HBM 与 VMEM 之间的数据搬运,以及 MXU 和向量单元的工作调度。
工程细节
XLA 处理模型的大部分,但 ragged paged attention、分组 MoE 矩阵乘法和 Gated DeltaNet 循环等性能关键操作需要显式控制块大小、双缓冲和 lane 布局,Pallas 在 Python 中提供这种控制。Official Preview 中的 TPU 推理优化大多是 Pallas kernel:GDN v3 把 Conv1D 和 GDN 融合为一个 kernel,grouped matmul 对专家权重做三缓冲,attention kernel 采用 sequence-on-lane 布局。Helion 的 TPU 后端也生成 Pallas。TorchTPU 可以从原生 PyTorch 调用 Pallas kernel,因此为 TorchAX 写的 kernel 只需调整 wrapper 和布局即可迁移。
为什么重要
Pallas 在 TPU 上的角色相当于 NVIDIA 上的 CUDA C++、CUTLASS 和 Triton。原生 PyTorch 支持并不消除对它的需求,张量形状和布局仍然要为 MXU 调优。生态编写新模型 Pallas kernel 的速度,决定了 TPU 外部化的节奏。
如何在 InferenceX 中解读
Ironwood 上报告的 Pallas kernel 收益包括:GDN v3 在 kernel 层面 decode 加速 1.41 倍、prefill 1.60 倍、混合 batch 2.14 倍;异步状态传输在并发 512 下吞吐量提高 11.3%;拆分 attention 预取块与计算块使 decode 吞吐量提高 49%。