用 Pallas 编写 TPU 内核:JAX 的 TPU 后端细节、限制与最佳实践
发布时间:2026/9/11 10:34:55来源:尧图网络
用 Pallas 编写 TPU 内核JAX 的 TPU 后端细节、限制与最佳实践【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jaxPallas 是 JAX 内置的低层内核编写框架允许你用 Python NumPy 语义直接编写在 TPU/GPU 上运行的 kernel。本文基于仓库中的 Pallas TPU 内核编写指南系统讲解在 Google TPU 上运行 Pallas 内核时需要注意的关键细节网格grid与BlockSpec的约束、内存层级HBM/VMEM/SMEM与数组布局、多核并行配置、SMEM 标量预取以及逐元素运算、矩阵乘法、归约、控制流等受支持操作的能力边界与成本模型。读完本文你将掌握如何写出既被 TPU 编译器接受、又能贴近硬件特性的高性能 Pallas TPU 内核。实验性状态与正确性保证在深入细节之前必须明确一点Pallas 的 TPU 后端仍处于实验阶段。当前仓库中的文档details.rst明确指出TPU 后端只接受 JAX NumPy 的一个子集且错误信息仍在持续改进中。不过实验性并不等于不可靠。项目对正确性保持严肃态度因此写 TPU 内核时遇到 not implemented 错误并不罕见但反过来只要内核被编译器接受它就必定返回预期结果。如果你观察到意外输出最有效的调试手段是在pallas_call中传入interpretTrue用解释模式interpreter运行同一个内核将解释模式的结果与真实硬件编译运行的结果进行对比若两者出现分歧则说明编译器存在 bug应向 JAX 项目提交 bug report仓库中的相关源码位于 jax/experimental/pallas/。interpretTrue是排查 Pallas TPU 内核正确性问题的第一道防线它绕过了 Mosaic 编译器直接在解释器中模拟内核语义。什么是 TPU与 GPU 的本质差异TPU 是 Google 开发、专为机器学习负载优化的硬件加速器。你可以把它想象成专精于 ML 的 GPU但两者的架构差异非常显著。Pallas 的价值在于即使你不完全理解底层硬件也能开始编写 TPU 内核而深入理解硬件则更容易写出高性能内核。从软件视角看TPU 与 GPU 最核心的区别可以概括为一句话TPU 是带有极宽向量寄存器的顺序执行机器类似 CPU同时允许软件将特定操作调度到后台异步执行。这些可异步执行的操作包括HBM 内存访问无法直接发起必须由 DMA 子单元预取到更低层级的内存层次中矩阵乘法由 MXUMatrix Multiply Unit单元支持矩阵转置与置换由 XLU 单元支持。这种主指令流 后台加速单元的架构决定了 Pallas TPU 内核的很多设计约束例如后文会讲到的顺序网格执行、内存层级预取与重叠、以及输出窗口必须连续的限制。值得注意的属性和限制BlockSpec 与网格迭代BlockSpec在 Pallas 中的语义与预期一致内核体的每次调用都会拿到输入的一个切片并负责初始化输出的一个切片。但 TPU 后端对 block 的形状有额外约束block 形状约束TPU 上只支持 rank 至少为 1 的 block。此外block 形状的最后两个维度必须分别能被 8 和 128 整除或者等于整个数组对应维度的尺寸。这一约束直接来源于 TPU 的向量寄存器组织方式8 条 sublanes × 128 lanes详见下文数组布局一节。Pallas TPU 内核还有一个独特的内存空间处理方式pallas_call的输入通常驻留在 HBMTPU 主内存但传入内核体的引用references指向更低层级内存中的缓冲区VMEM 或 SMEM。这使得内核体可以高速读写这些缓冲区而与 HBM 之间高延迟的通信完全由编译器处理并与计算重叠。这也是 Pallas TPU 内核能够隐藏 HBM 延迟的关键机制。与 GPU 相比TPU 是高度顺序化的机器网格通常不是并行处理的而是按字典序lexicographic order顺序执行多核配置除外见下文。这解锁了两个有趣的能力当两个字典序上相邻的网格索引使用输入的同一切片时第二次迭代的 HBM 传输会被跳过——数据已经可用了内核体的多次调用可以写入输出的同一切片而不会有竞态风险但要求写入同一切片的所有调用在网格上是连续的。连续限制意味着通常网格维度的某个前缀会负责变化输出窗口而剩余的后缀维度保持输出窗口不变。最典型的例子就是矩阵乘法内核通常使用 3 维网格前两维分别对应左操作数第一轴和右操作数第二轴的切片第三维最后一维负责平铺归约维度。归约维必须放在最后一维因为输出窗口沿该轴不变化输出引用可以充当部分和的累加器accumulator。仓库中的官方矩阵乘法示例 jax/experimental/pallas/ops/tpu/matmul.py 完美展示了这一模式def matmul_kernel(x_tile_ref, y_tile_ref, o_tile_ref, acc_ref): pl.when(pl.program_id(2) 0) def init(): acc_ref[...] jnp.zeros_like(acc_ref) acc_ref[...] acc_ref[...] jnp.dot( x_tile_ref[...], y_tile_ref[...], preferred_element_typeacc_ref.dtype, ) o_tile_ref[...] acc_ref[...].astype(o_tile_ref.dtype)其中网格为(x.shape[0] // l, y.shape[1] // r, x.shape[1] // block_k)归约维度block_k对应的网格轴是最后一个轴acc_refVMEM 中的 float32 累加器在所有归约迭代中复用从而避免了重复初始化与部分和合并的开销。该内核还使用scratch_shapes[pltpu.VMEM((l, r), acc_dtype)]在 VMEM 中显式分配累加器缓冲区。VMEM 容量提示VMEM 对于如此低层的内存层次来说相当大16MB因此可以使用较大的窗口尺寸window size。通常窗口越大硬件利用率越高。但如果窗口尺寸加上寄存器溢出所需的临时空间超过了 VMEM 容量你会看到低层编译器报出的 out-of-memory 错误——这也提示我们不要无上限地放大窗口。数组布局维度顺序是有意义的在 Pallas 中数组的维度顺序layout是有意义的。在普通的jax.jit程序中中间数组的排列顺序通常不影响性能因为编译器可以自由重排但 Pallas 暴露的是低层能力维度顺序会显著影响生成代码的质量。TPU 的大部分计算发生在 2D 向量寄存器上典型的寄存器尺寸是8×128以 32 位值为例对应 TPU 的 sublanes 和 lanes。当一个向量值从 VMEM 加载到寄存器例如x x_ref[...]时数组的最后两个维度会被平铺到寄存器中。Pallas 只会考虑将中间数组的最后两维映射到 8×128 的向量寄存器维度上。下图展示了 12×320 的数组如何用 6 个 8×128 的 tile 进行平铺这种 tile 化布局对内核编写者有几个重要推论数组的最后两轴与其他轴待遇不同。涉及最后两轴的归约、reshape、转置通常更昂贵某些涉及最后两维的 reshape 不受支持会直接导致编译器报错而其他维度的 reshape 是免费的在编译期完成。最后两轴上的单例singleton维度是浪费的。它们会占据整个 tile 维度中的 1 个元素位置而其余位置被白白填充。过多的寄存器消耗还可能导致寄存器溢出spill到 VMEM进而降低内核性能。所有向量计算都会补齐到 tile 尺寸。两个 1×1 数组相加的代价与两个 8×128 数组相加相同两个 8×128×1×1 数组相加的代价是 8×128 数组相加的 1024 倍因为前者会被补齐到 8×128×8×128。换句话说写 Pallas TPU 内核时应避免在最后两维出现单例维度并尽量让归约/reshape/转置发生在除最后两维之外的轴上。多核 TPU 配置dimension_semantics在新一代 TPU 中一个芯片chip上的两个核心cores通常被抽象为单个设备。要利用多核能力Pallas 必须打破顺序网格执行的保证将某个网格轴并行分配到多个核心上。这是一个**主动选择opt-in**的过程需要给pallas_call传入额外的compiler_params参数指定dimension_semanticspallas_call( ..., compiler_paramspltpu.CompilerParams( dimension_semantics[parallel, parallel, arbitrary] ), )dimension_semantics是一个列表长度与网格轴数一致每个元素取值为parallel该维度可以按任意顺序执行可以被划分到多个核心上arbitrary该维度必须顺序执行。经验法则输出窗口不随之变化的维度是 parallel 的输出窗口随之变化的维度是 arbitrary 的。因此dimension_semantics通常呈现为若干个parallel轴 若干个arbitrary轴的形式。在源码层面这一参数在 jax/_src/pallas/mosaic/core.py 的CompilerParams数据类中定义其 docstring 明确说明dimension_semantics是内核每个网格维度的语义列表parallel表示可以任意顺序执行arbitrary表示必须顺序执行。同文件还列出了其他常用编译器参数例如vmem_limit_bytes覆盖 VMEM 上限需配合--xla_tpu_scoped_vmem_limit_kibN使用、collective_id选择 barrier 信号量、disable_bounds_checks等。上述矩阵乘法内核的完整调用正是这样配置的见 jax/experimental/pallas/ops/tpu/matmul.pyreturn pl.pallas_call( matmul_kernel, out_shapejax.ShapeDtypeStruct((x.shape[0], y.shape[1]), out_dtype), grid_specpltpu.PrefetchScalarGridSpec( num_scalar_prefetch0, in_specs[ pl.BlockSpec((l, block_k), lambda i, _, k: (i, k)), pl.BlockSpec((block_k, r), lambda _, j, k: (k, j)), ], out_specspl.BlockSpec((l, r), lambda i, j, k: (i, j)), grid(x.shape[0] // l, y.shape[1] // r, x.shape[1] // block_k), scratch_shapes[pltpu.VMEM((l, r), acc_dtype)], ), compiler_paramspltpu.CompilerParams( dimension_semantics(parallel, parallel, arbitrary)), debugdebug, )(x, y)注意这里的dimension_semantics(parallel, parallel, arbitrary)前两维对应输出窗口的两个方向可以并行划分到核心最后一维归约维输出窗口不变保持顺序执行。这正印证了文档中的规则。关于并行收益的提醒将内核划分到 2 核 TPU 设备上通常能带来 2 倍加速但实际收益可能显著小于此——特别是当不同内核体实例的成本差异很大时。如果所有昂贵的步骤都被映射到一个核心而廉价步骤都在另一个核心上那么后者会空转到前者完成。Pallas TPU 一般倾向于划分尺寸为核心数整数倍的网格轴并且优先划分靠前的网格轴。将操作数放入 SMEMPrefetchScalarGridSpecTPU 上大部分计算发生在向量单元vector unit但很多场景需要执行标量操作例如控制流判断。为此 TPU 配备了独立的标量单元scalar unit和附属的标量内存 SMEM。经验法则任何用于控制流决策的数据都应放在 SMEM 中。SMEM 是低延迟、支持随机访问的内存但单条指令只能读写 32 位值相比 VMEM 事务 4KiB 的粒度小得多但因无对齐要求而更灵活。标量内存在实现对输入 tile 访问模式不规则的内核例如块稀疏内核时非常有用。在 Pallas 中做法是把pallas_call的grid参数换成PrefetchScalarGridSpec并设置非零的num_scalar_prefetch若num_scalar_prefetch n则pallas_call的前 n 个参数会被放入 SMEM这些参数不应指定BlockSpec其余参数照常指定BlockSpec但它们的 BlockSpec 回调除了网格索引外还会收到指向前导操作数的 SMEM 引用。PrefetchScalarGridSpec与num_scalar_prefetch字段定义在 jax/_src/pallas/mosaic/core.py并经由 jax/experimental/pallas/tpu.py 导出为pltpu.PrefetchScalarGridSpecSMEM、VMEM、HBM等内存空间常量也在此导出jax/experimental/pallas/tpu.py中的SMEM MemorySpace.SMEM等。支持的数据类型目前 Pallas TPU 支持以下数据类型jnp.float32jnp.bfloat16jnp.int*所有精度除jnp.int4jnp.uint*所有精度jnp.bool_计算放置标量与向量所有标量即 0D数组都会存放在标量寄存器中相关操作在标量核心上执行其他所有操作即使是单元素的 1D 数组都在向量核心上执行。这一规则会影响你的代码组织需要 0D 标量语义的中间结果应保持 0D避免意外变成 1D 而被提升到向量单元。支持的操作详解矩阵乘法矩阵乘法总是产生 float32 格式的结果。如果输入不是 float32推荐使用lax.dot并设置preferred_element_typejnp.float32正如上述matmul_kernel示例中preferred_element_typeacc_ref.dtype的用法。使用lax.dot_general时可以将矩阵乘法操作数最后两维的转置融合进操作本身从而提升内核整体性能。精度控制Pallas TPU 的 lowering 会感知jax.default_matmul_precision追求最佳性能同时接受最低精度时使用bfloat16在意数值精度时可将精度设为float32。重要警告即使你向矩阵乘法传入 32 位操作数除非显式请求float32精度否则它们会被舍入到bfloat16。如果你的内核依赖 32 位精度务必显式设置精度否则可能静默引入精度损失。转置若数组至少 4 维除最后两轴之外的任意转置都是免费的否则只实现了最后两轴的转置注意最后两维的某些转置可以融合进矩阵乘法见上文。访问内存引用reference的任意切片都可以读或更新但受实现约束限制对 32 位宽的输入目前没有限制对更窄的类型只支持部分切片模式最后两维中偏移量对齐到 8 的倍数、长度是 128 的倍数分别对应两个维度的读写总是受支持的。由于对向量内存的读写通常以(8, 128)的 tile 为单位进行对至少 2 维引用的读写最佳性能条件是访问的基准偏移可被 tiling 整除且读区尺寸是 tile 尺寸的倍数。逐元素操作大量逐元素操作受支持但需注意硬件通常只支持使用 32 位类型进行逐元素计算。加载低精度操作数后通常应先将其提升upcast到 32 位类型再做逐元素运算。不同逐元素操作的成本差异巨大文档将其分为三类便宜、中等和昂贵。操作成本jnp.add、jnp.sub、-jnp.mul、*/、//、%jnp.max、jnp.minjnp.whereselectjnp.abs|、^、、~、比较运算等类型转换.astypejnp.expjnp.tanhjnp.powjnp.sinjnp.cos注意很多 JAX 函数是由其他 primitive 组合实现的因此该列表未必全面。例如jax.nn.relu由比较运算和jnp.where实现因此同样可以在 Pallas 内核中工作。以这个成本表为参照可以指导内核设计中的廉价原语组合策略。数组构造所有常量数组构造器均受支持jnp.ones、jnp.zeros、jnp.full。归约浮点值支持sum、max、min归约布尔值支持any和all整数归约不受支持。性能方面沿最后一维的归约通常最慢沿倒数第二维的归约更快但仍慢于沿前导维度的归约。因此归约轴的选择对内核性能有直接影响。广播广播的性能特征与归约非常相似沿除最后两维之外的所有维度的广播总是受支持且免费的沿倒数第二维的广播较慢沿最后一维的广播最慢。Reshape与布局一节呼应除最后两维之外的所有维度的 reshape 都受支持且免费。只有两种情况下 reshape 可以修改最后两维某些前导维度被扁平化flatten到倒数第二维上它添加了一个刚被归约移除的维度。随机数生成Pallas 支持jax.random模块中最常用的函数如uniform、normal、bernoulli。key 必须是threefry2x32类型的 key这也是 JAX 的默认设置。key 既可以直接传入内核也可以在内核内部生成。更多细节可参考仓库中的 Pallas TPU 随机数指南。控制流TPU 后端目前对控制流的支持有限支持以下函数condfori_loopfor_loop但循环原语在编译期会被完全展开fully unrolled所以循环次数trip count应保持在合理的小范围内。过度使用控制流会导致低层代码生成质量显著退化——推荐的做法是尽可能把更多计算密集型操作塞进单个基本块basic block中。综合建议如何写出高性能 Pallas TPU 内核综合文档与仓库源码details.rst、matmul.py、core.py可以提炼出以下可落地的内核编写清单形状约束先行block 形状的 rank ≥ 1最后两维分别满足 8 与 128 的整除性或等于数组对应维度。网格设计遵循输出窗口前缀原则让网格前缀维驱动输出窗口变化后缀维保持窗口不变如归约维放最后并利用窗口复用跳过重复的 HBM 传输。布局优先最后两维避免单例维度归约、reshape、转置尽量放到前导维度理解所有向量计算都会 padding 到 8×128 tile。多核显式启用只有dimension_semantics标为parallel的轴才可能被划分到多个核心注意负载不均衡风险优先划分核心数整数倍的网格轴。控制流数据放 SMEM用PrefetchScalarGridSpec(num_scalar_prefetchn)将前 n 个参数预取到标量内存这些参数不指定BlockSpec。精度显式声明矩阵乘法默认结果是 float32 但对输入有 bfloat16 舍入风险需要精度时显式设置preferred_element_type或jax.default_matmul_precision低精度操作数做逐元素运算前先 upcast 到 32 位。控制流克制cond/fori_loop/for_loop会全展开循环次数保持小规模尽量把热路径收敛到单个基本块。正确性兜底遇到意外结果用interpretTrue对比解释模式输出再决定是否向 JAX 项目提交 bug report。Pallas TPU 后端虽然实验性较强但其接受即正确的契约、对内存层级与硬件单元的显式暴露使得它成为在 JAX 生态内编写贴近硬件、可预测性能的 TPU 内核的有力工具。理解本文所讲的硬件属性与操作成本是在这一后端上写出高性能内核的前提。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
网站建设高端定制企业官网