使用 Pallas 编写 TPU 内核:从硬件架构到实战约束的完整指南
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/gh_mirrors/jax/jax
Pallas 是 JAX 提供的一套可编程内核编写接口,允许用户绕过jax.jit的自动抽象,直接以类似 NumPy 的语法描述在加速器上执行的内核逻辑。本文以当前仓库中 docs/pallas/tpu/details.rst 为核心,系统讲解在 Google TPU 上运行 Pallas 内核所必须掌握的硬件背景、内存模型、网格语义、多核配置以及当前后端所支持的操作集与限制。读完本文,你将理解 TPU 与 GPU 的本质差异,学会用pallas_call写出能在 TPU 上正确、高效运行的内核,并掌握调试与性能调优的关键要点。
Pallas TPU 后端现状:实验性但不牺牲正确性
Pallas 的 TPU 后端仍处于实验阶段,目前只接受 JAX NumPy 的一个子集。官方文档在 docs/pallas/tpu/details.rst 中明确给出了两条重要承诺:
- 正确性优先:虽然该功能整体实验性(尤其是错误提示信息仍在完善中),但 JAX 团队对正确性非常认真。写 TPU 内核时看到 "not implemented" 错误并不罕见,但只要一个内核被编译器接受,它必须返回预期结果。
- 解释模式兜底:如果看到意外输出,请将结果与在
pallas_call中传入interpret=True运行的内核进行对比。若两者结果不一致,说明存在 bug,应当提交 bug report。interpret参数在 jax/_src/pallas/pallas_call.py 中作为pallas_call的命名参数暴露,它让内核在解释器模式下运行,绕开真实硬件编译路径,是隔离编译器问题与用户逻辑问题的第一手段。
从 jax/_src/pallas/mosaic/pallas_call_registration.py 可以看到,TPU 后端的调用注册函数接收interpret: bool与compiler_params: dict[str, Any]两个关键参数,后者正是多核配置等高级参数的入口。
理解 TPU:与 GPU 截然不同的计算模型
顺序机器 + 超宽向量寄存器
TPU 是 Google 开发的机器学习专用硬件加速器。与 GPU 的大规模线程并行模型不同,TPU 本质上是顺序执行的机器,配备极宽的向量寄存器(类似 CPU)。这意味着在 Pallas 中写的网格(grid)通常不会被并行处理,而是按字典序(lexicographic order)顺序执行——这一特性会直接影响内核设计(详见下文"网格迭代"一节)。
后台异步执行的硬件单元
TPU 允许软件将特定操作调度到后台,使其与主指令流异步执行。涉及的主要单元包括:
- DMA 子单元:负责 HBM(主存)访问。HBM 无法被直接读写,必须先由 DMA 预取到更低层的内存层级;
- MXU 单元:执行矩阵乘法;
- XLU 单元:执行矩阵转置与置换(permute)。
这些异步能力意味着:只要合理安排,HBM 的高延迟通信可以被编译器调度到与计算重叠,这正是 Pallas TPU 内核性能的关键来源之一。
进一步阅读
文档推荐了多篇研究论文,涵盖从 TPUv1 到 TPUv4 的架构演进(详见 docs/pallas/tpu/details.rst)。虽然每篇论文针对特定代际,但其中大部分思想可迁移到后续代际。
值得注意的属性与限制
BlockSpec 与网格迭代:顺序执行带来的红利与约束
BlockSpec(参见 docs/pallas/grid_blockspec.md)在 Pallas 中的行为符合预期:内核体的每次调用都会获得输入的一个切片,并负责初始化输出的一个切片。
窗口形状限制(重要):并非所有窗口形状都被支持。如果输入的最后两个维度分别大于 8 和 128,则这两个维度上的窗口形状必须是相应因子的倍数;如果输入维度更小,则窗口应覆盖整个维度。这一限制直接来自 TPU 向量寄存器与内存传输以(8, 128)为基本 tile 的硬件事实(见下文"访问内存"一节)。
内存空间映射:Pallas TPU 内核最有趣的一点是内存空间的处理方式——pallas_call的输入通常驻留在 HBM(TPU 主存),但传给内核体的引用(references)指向的是更低层内存层级(VMEM 或 SMEM)中的缓冲区。这让内核体可以以极高速度读写这些引用,而与 HBM 的所有通信由编译器负责,并与计算重叠。
顺序网格带来的两个能力:
- HBM 传输去重:当两个字典序相邻的网格索引使用同一个输入切片时,第二次迭代的 HBM 传输会被跳过,因为数据已经可用;
- 安全的多写输出:内核体的多次调用可以写同一个输出切片而无需担心竞态条件。但要求写同一切片的所有调用是连续的。
"连续"限制的实践含义:通常这意味着网格维度的某个前缀总是变化着某次调用需要访问的输出切片,而输出窗口在剩余的后缀维度上保持不变。
矩阵乘法的典型网格设计:实现 Pallas TPU 矩阵乘法内核时,一般使用三维网格:前两维分别对应左操作数第一轴、右操作数第二轴的切片;第三维(最后一个网格轴)对归约维进行分块。归约维对应的网格轴必须是最后一个,因为输出窗口沿该轴不变化,输出引用可以充当部分和的累加器。
VMEM 容量提示:VMEM 对于如此低层级的内存来说相当大(16MB+),因此可以使用大窗口。通常窗口越大,硬件利用率越高。但窗口大小(加上溢出向量寄存器所需空间)若超过 VMEM 容量,会看到低层编译器报出的内存不足(OOM)错误。
维度顺序是有意义的
在普通jax.jit程序中,中间数组的维度顺序通常不影响性能,因为编译器可以自由重排。但 Pallas 旨在暴露底层能力,因此维度顺序会极大影响生成代码的质量。
原因在于:TPU 的大部分计算发生在二维向量寄存器上,而 Pallas TPU 只会将中间数组的最后两个维度映射到向量寄存器的维度(分别是 sublanes 和 lanes)。一个形状为(n, 1, 1)的数组至少需要n个向量寄存器来表示;当n过大时可能导致寄存器溢出,进而因内存占用过大触发 VMEM OOM。当然,低层编译器很擅长重排指令以降低寄存器压力,所以也可能不会溢出。但经验法则是:让最后两个维度尽量大(尤其是最后一维),让前导维度尽量小。
多核 TPU 配置:dimension_semantics 参数
在较新的 TPU 代际中,芯片上的两个核常常被抽象为单个设备。为了利用多核,Pallas 必须打破顺序网格执行保证,将一个网格轴并行化到多个核上。这是显式 opt-in的过程:pallas_call需要额外的compiler_params参数:
pallas_call( ..., compiler_params=dict( mosaic=dict( dimension_semantics=["parallel", "parallel", "arbitrary"] ) ), )该参数是一个列表,条目数与网格轴数相同。只有parallel维度可以被划分到多个核上。经验法则是:除输出窗口不变化的维度外,其余维度都是 parallel 的。因此dimension_semantics总是由一串parallel轴后跟一串arbitrary轴组成。
从源码看,jax/_src/pallas/mosaic/lowering.py 中MosaicGridMapping会对该参数做校验:未提供时默认全部为("arbitrary",);若列表长度与用户网格轴数不符会直接抛出ValueError。另外,来自vmap的维度会被自动映射为parallel(见lowering.py第 277-283 行的注释与实现)。随后get_dimension_semantics()(jax/_src/pallas/mosaic/lowering.py)会将其序列化为#tpu.dimension_semantics<...>属性写入生成的 IR。在 tests/pallas/tpu/pallas_pipeline_test.py 中可以看到多种(PARALLEL, ARBITRARY)组合的真实用法(常量pltpu.PARALLEL与pltpu.ARBITRARY由 jax/experimental/pallas/tpu.py 导出)。
性能预期:将内核划分到 2 核 TPU 设备上通常能带来约 2 倍加速,但也可能显著小于 2 倍——尤其当不同内核体实例的计算代价差异很大时:若所有昂贵步骤都被映射到一个核,而廉价步骤全在另一个核,后者就会空转等待。Pallas TPU 一般倾向于划分大小为核数整数倍的轴,并且优先划分前导网格轴。
将操作数放入 SMEM:PrefetchScalarGridSpec
TPU 上的大部分计算发生在向量单元,但许多场景(如控制流)需要执行标量运算。为此 TPU 配有独立的标量单元和独立标量内存SMEM。经验法则:任何用于控制流决策的数据都应放入 SMEM。
SMEM 是低延迟内存,支持随机访问,但单条指令只能读写 32 位值(相比 VMEM 事务 4KBi 的粒度小得多,但因无对齐要求而更加灵活)。
标量内存在实现不规则访问模式的内核(如 block-sparse 内核)时尤其有用。Pallas 中的做法是:将pallas_call的grid参数替换为PrefetchScalarGridSpec,并设置非零的num_scalar_prefetch参数:
- 若
num_scalar_prefetch为n,则pallas_call的前 n 个参数会被放入 SMEM,且不应为这些参数指定BlockSpec; - 其余所有参数的
BlockSpec不仅会接收到网格索引,还会接收到前导操作数的 SMEM 引用。
PrefetchScalarGridSpec的实现位于 jax/_src/pallas/mosaic/core.py:它继承自pallas_core.GridSpec,构造函数签名为__init__(num_scalar_prefetch, grid=None, in_specs=..., out_specs=..., scratch_shapes=())。在get_grid_mapping中,前num_scalar_prefetch个参数被切分出来,其引用 aval 被标记为TPUMemorySpace.SMEM(见第 183-197 行)。内存空间常量SMEM、VMEM、CMEM、ANY以及PrefetchScalarGridSpec均从 jax/experimental/pallas/tpu.py 公开导出。
支持的数据类型
目前 Pallas TPU 仅支持以下数据类型:
jnp.float32jnp.bfloat16jnp.int*(所有精度,除jnp.int4外)jnp.uint*(所有精度)
计算放置
所有标量(0D)数组存储在标量寄存器中,相关运算在标量核上执行;所有其他运算(即使是单元素但 1D+ 的数组)都在向量核上执行。
支持的操作详解
矩阵乘法
- 矩阵乘法总是以 float32 格式产生结果。若输入不是 float32,建议使用
lax.dot并设置preferred_element_type=jnp.float32; - 使用
lax.dot_general时,可以将矩阵乘法操作数最后两个维度的转置融合进操作,从而提升整体内核性能。
精度控制
Pallas TPU 的 lowering 会感知jax.default_matmul_precision:
- 追求最佳性能(和最低精度)时,使用
bfloat16; - 关心数值精度时,将精度设为
float32。
警告:即使向矩阵乘法传入 32 位操作数,除非显式请求float32精度,它们也会被舍入为bfloat16。
转置
- 若数组至少有 4 个维度,则除最后两个轴外任意轴的转置都是免费的;
- 否则只实现了最后两个轴的转置;
- 注意,最后两个维度的某些转置可以融合进矩阵乘法。
访问内存
引用(references)的任意切片都可以读取或更新,但要受实现约束:
- 目前对 32 位宽的输入没有任何限制;
- 更窄类型只支持部分切片模式;
- 最后两个维度上对齐到 8 和 128 的倍数、长度也是 8 和 128 倍数的读写总是被支持。
由于向量内存的读写通常以(8, 128)的 tile 为单位发生,当读写至少两维的引用时,最佳性能条件是:内存访问的基础偏移量可被 tile 整除,且读取区域的大小是 tile 大小的倍数。这与前述 BlockSpec 窗口形状的 8/128 限制同源。
逐元素运算
硬件一般只支持使用 32 位类型做逐元素计算。加载低精度操作数时,应先将它们升位(upcast)到 32 位类型再做逐元素运算。
不同逐元素运算的成本差异非常显著,文档将其分为三档:便宜(🟢)、中等(🌕)、昂贵(🔴):
| 操作 | 成本 |
|---|---|
jnp.add、+ | 🟢 |
jnp.sub、- | 🟢 |
jnp.mul、* | 🟢 |
/、//、% | 🌕 |
jnp.max、jnp.min | 🟢 |
jnp.where(select) | 🟢 |
jnp.abs | 🟢 |
\|、^、&、~ | 🟢 |
<<、>> | 🟢 |
比较运算(==等) | 🟢 |
类型转换(.astype) | 🟢 |
jnp.exp | 🌕 |
jnp.tanh | 🌕 |
jnp.pow | 🌕 |
jnp.sin | 🔴 |
jnp.cos | 🔴 |
许多 JAX 函数由其他 JAX 原语组合实现,因此该表并非穷尽。例如jax.nn.relu基于比较和jnp.where实现,因此在 Pallas 内核中也能工作。
数组构造函数
所有常量数组构造函数都受支持(jnp.ones、jnp.zeros、jnp.full)。值得注意的是,jax.random模块目前与 Pallas 不兼容。
归约
- 支持 sum、maximum、minimum 归约,但一次只能对一个数组轴进行归约;
- 沿最后一个维度的归约通常最慢;沿倒数第二维的归约较快,但仍慢于沿前导维度的归约。
广播
广播的性能特征与归约非常相似:
- 沿除最后两个维度外的广播始终受支持且免费;
- 沿倒数第二维的广播较慢;
- 沿最后一维的广播最慢。
重塑
- 除最后两个维度外的重塑受支持且免费;
- 重塑可以修改最后两个维度的仅有的两种受支持情况:
- 某些前导维度被展平到倒数第二维上;
- 增加一个刚被归约移除的维度。
控制流
TPU 后端目前对控制流的支持有限,当前支持cond、fori_loop和for_loop。但循环原语目前在编译时会被完全展开(unroll),因此应尽量控制循环趟数在合理范围内。
过度使用控制流会导致低层代码生成的显著退化,推荐尽可能将更多计算密集型操作塞进单个基本块中。
实战要点总结
综合文档与仓库源码,编写高性能 Pallas TPU 内核时值得记住的实践原则:
- 先确认正确性,再优化性能:任何内核先用
interpret=True验证语义正确,再关闭解释模式跑真实硬件; - 顺应硬件的 tile 结构:窗口、切片、读写保持
(8, 128)的倍数对齐,这是向量寄存器与内存传输的基本单元; - 网格设计以输出窗口不变轴为最后轴:归约维放最后,让输出引用充当累加器,并享受相邻网格索引的内存重用;
- 维度排序遵循"后两维大、前导维小":避免
(n, 1, 1)这类形状造成寄存器溢出; - 多核是显式选择:通过
compiler_params=dict(mosaic=dict(dimension_semantics=[...]))声明parallel轴,并优先划分前导网格轴; - 标量数据进 SMEM:控制流相关数据用
PrefetchScalarGridSpec+num_scalar_prefetch放入 SMEM; - 注意数据类型与运算成本的暗坑:矩阵乘法默认向 bfloat16 舍入、
jax.random不可用、归约/广播沿最后一维最慢、循环会被完全展开。
当前仓库中 jax/experimental/pallas/ops/tpu 目录(含 flash attention、paged attention、splash attention、megablox 等)提供了大量真实 TPU 内核实现,是学习这些约束在实战中如何落地的绝佳参考;相关测试可参见 tests/pallas/tpu 目录。
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/gh_mirrors/jax/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考