资讯动态

JAX Pallas TPU 内核编程指南:从硬件内存模型到流水线优化与 Megacore 并行

发布时间:2026/9/20 20:39:39 来源:尧图企业网站定制
机器学习深度学习【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址https://gitcode.com/gh_mirrors/jax/jax点击查看免费下载导读本文是 docs/pallas/tpu 系列文档的深度整理主题为如何使用 PallasJAX 的自定义内核语言在 Google TPU 上编写高性能内核。文章先剖析 TPU 与 GPU 截然不同的硬件特性HBM/VMEM/SMEM 内存层级、向量寄存器、顺序执行语义再逐步讲解BlockSpec/grid迭代规则、dimension_semantics多核并行、SMEM 标量预取以及用gridBlockSpec自动生成流水线来重叠内存拷贝与计算。读完你将掌握 Pallas TPU 内核的正确编写姿势、性能调优规则和常见陷阱并能直接复现文中的矩阵加法、归约求和等完整内核示例。Pallas TPU 后端实验状态与正确性承诺Pallas 是 JAX 的扩展用于为 GPU 与 TPU 编写自定义内核见 docs/pallas/index.rst 的定位说明。其中 TPU 后端由 Mosaic 编译管线支撑相关 API 集中在 jax/experimental/pallas/tpu.pyVMEM、SMEM、PrefetchScalarGridSpec、emit_pipeline等均由此导出。使用 TPU 后端前需要明确两点实验性该功能仍处于实验阶段目前只接受 JAX NumPy 的一个子集且官方仍在改进错误信息。因此编写内核时遇到 not implemented 错误并不罕见。正确性承诺虽然功能实验性但官方对正确性非常严肃——只要内核被编译器接受就必然返回预期结果。当发现意外输出时务必用interpretTrue传入pallas_call再跑一遍同一内核进行对照。从源码看interpret模式会把pallas_call实现为对grid的jax.jit扫描见 jax/_src/pallas/pallas_call.py 中 runs thepallas_callas ajax.jitof a scan over the grid 的说明它不需要 TPU 或 GPU是唯一能在 CPU 上运行 Pallas 内核的方式非常适合调试。若两种模式结果不一致则属于编译器 bug应提交 bug report。理解 TPU与 GPU 截然不同的硬件内存空间、寄存器与计算单元TPU 是 Google 开发的专用机器学习加速器。与 GPU 相比其核心差异在于TPU 是带有超宽向量寄存器的顺序执行机器类似 CPU同时允许软件把部分操作调度到后台异步执行。TPU 的 TensorCore 由三类构件组成示意图见 docs/pallas/tpu/pipelining.md内存空间HBM高带宽内存常被理解为设备内存数组的主要存放地VMEM向量内存用于存放向量/数组值的缓存SMEM标量内存用于存放标量值的缓存。寄存器向量寄存器VREGs存数组值标量寄存器SREGs存标量值数据从对应缓存加载进寄存器。计算单元标量单元scalar unit、向量单元VPU、矩阵单元MXU。计算单元只对寄存器中的值运算结果也写回寄存器。异步后台操作TPU 允许软件调度以下操作在后台执行使其与主指令流异步HBM 内存访问不能直接发起必须由 DMA 子单元预取到更低层级内存矩阵乘法由 MXU 单元支持矩阵转置与置换由 XLU 单元支持。一次向量化计算的完整链路对 HBM 中的x、y做向量加法硬件层面需要五步把x、y从 HBM 拷贝到 VMEM从 VMEM 加载到 VREGs用 VPU 或 MXU 执行计算结果写入 VREGs把输出 VREGs 存回 VMEM把 VMEM 中的输出拷回 HBM。BlockSpec 与 grid 迭代Pallas TPU 内核的核心约束BlockSpecblock_shapeindex_map定义于 jax/_src/pallas/core.py与grid的行为在 TPU 上大体符合 Pallas 的一般语义每次调用内核体都拿到输入的切片并负责初始化输出切片。但 TPU 后端有几条特殊规则窗口形状限制不是所有窗口形状都受支持。如果输入的最后两维分别大于 8 和 128那么这两个维度的窗口形状必须是 8 和 128 的倍数如果输入维度更小窗口应覆盖整个维度。内存空间HBM 与 VMEM/SMEM 的分工pallas_call的输入通常位于 HBM但传入内核体的引用Ref指向的是更低层内存VMEM 或 SMEM。这让内核体能以极高速读写它们而所有与 HBM 的通信延迟极高由编译器负责并与计算重叠。这正是 Pallas TPU 内核高性能的来源。顺序 grid 语义带来的三个推论与 GPU 不同TPU 高度顺序化grid 通常按字典序顺序执行而非并行Multicore 配置除外见后文。这带来三个重要能力HBM 传输复用当两个字典序相邻的 grid 索引使用同一输入切片时第二次迭代的 HBM 传输会被跳过——数据已就绪无竞争输出写入多次内核体调用可以写同一输出切片而不会产生竞争条件。但要求所有写同一切片的调用必须连续输出切片前缀-后缀结构输出的连续限制通常意味着grid 维度的某个前缀总在变化输出切片而输出窗口对剩余后缀保持恒定。以矩阵乘法内核为例一般用 3 维 grid——前两维分别对应左操作数第一轴、右操作数第二轴的切片第三个最后轴负责 tile reduction 维度。reduction 轴必须是最后一维因为输出窗口在该轴上不变化输出引用因此可以作为部分和累加器反复使用。VMEM 容量与窗口大小VMEM 对这么底层的存储层级来说相当大16MB所以可以用很大的窗口。经验上窗口越大硬件利用率通常越好。但如果窗口加上溢出的向量寄存器所需空间超过 VMEM 容量就会看到底层编译器报出的内存不足OOM错误。维度顺序有意义把最后两维做大在普通jax.jit程序中中间数组的维度顺序通常不影响性能——编译器可自由重排。但 Pallas 暴露的是底层能力维度顺序对生成代码质量影响巨大。原因在于TPU 的绝大部分计算在 2D 向量寄存器上进行而Pallas TPU 只会把中间数组的最后两维映射到向量寄存器维度分别对应 sublanes 和 lanes。形状为(n, 1, 1)的数组至少需要n个向量寄存器表示n过大时可能导致寄存器溢出和因内存占用过大而触发 VMEM OOM。虽然底层编译器很擅长重排指令以降低寄存器压力但稳妥的经验法则是保持最后两维尤其是最后一维较大前导维度较小。Multicore TPU 配置dimension_semantics 与 Megacore单芯片双核抽象在较新的 TPU 世代如 v4、v5p芯片上的两个 TensorCore 常被抽象为单一设备即Megacore模式。两个 TensorCore 各有独立的 VMEM、VREGs、SMEM、SREGs 和计算单元但共享 HBM。概念上Megacore 设备像一台只有两个线程的极简 GPU。要利用多核Pallas 必须打破顺序 grid 执行保证把某个 grid 轴并行化到各核上。这是**显式选择opt-in**的通过pallas_call的compiler_params传入dimension_semantics实现pallas_call( ..., compiler_paramsdict( mosaicdict( dimension_semantics[parallel, parallel, arbitrary] ) ), )dimension_semantics 的语义该参数是一个列表条目数与 grid 轴数相同。只有标记为parallel的维度可以在核间分区。经验法则输出窗口不变的维度才是 parallel 的否则无法无竞争地并行因此dimension_semantics永远是若干个parallel轴后跟若干个arbitrary轴。arbitrary表示该维度不可做任何假设因此不能并行化。从源码看该参数在 Mosaic 管线中经过严格校验未提供时默认全为arbitrary长度必须与grid一致见 jax/_src/pallas/mosaic/lowering.py 与 jax/_src/pallas/mosaic/pipeline.py最终以dimension_semantics属性附着到编译产物上。在pallas_call_registration.py中它从mosaic_params中提取并传入 lowering见 jax/_src/pallas/mosaic/pallas_call_registration.py。在流水线示例中启用 Megacore 只需一行注解def add_matrices_pipelined_megacore(x: jax.Array, y: jax.Array) - jax.Array: block_spec pl.BlockSpec((256, 512), lambda i: (i, 0)) return pl.pallas_call( add_matrices_kernel, out_shapex, in_specs[block_spec, block_spec], out_specsblock_spec, grid(2,), compiler_paramsdict(mosaicdict(dimension_semantics(parallel,))) )(x, y)指定dimension_semantics后Pallas 会自动把 grid 拆分并在两个 TensorCore 上同时执行。注意Megacore 目前仅对 TPU v4 和 v5p 生效。在其他平台上提供该注解是 no-op但不指定它会导致即使有多个核也只用一个 TensorCore。并行化的收益与风险在 2 核 TPU 上分区内核常带来约 2 倍加速但实际收益可能显著小于 2 倍——尤其是当不同内核体实例的计算成本差异很大时若所有昂贵步骤恰好被映射到一个核、廉价步骤全在另一个核第二个核会空转到第一个核完成。此外Pallas TPU 一般偏好分区大小为核数倍数的轴并优先分区前导 grid 轴。把操作数放进 SMEMPrefetchScalarGridSpecTPU 上大部分计算发生在向量单元但控制流等场景常需要标量运算。为此 TPU 配有独立的标量单元和标量内存SMEM。经验法则任何用于控制流决策的数据都应放在 SMEM。SMEM 是低延迟内存支持随机访问但单条指令只能读写 32 位值——相比 VMEM 事务 4KBi 的粒度小得多却因为没有对齐要求而灵活得多。当内核不以规则模式访问输入 tile 时例如块稀疏内核标量内存非常有用。在 Pallas 中把grid参数换成grid_spec为PrefetchScalarGridSpec、并设置非零的num_scalar_prefetch即可实现。规则如下若num_scalar_prefetch为n则pallas_call的前n个参数被放入 SMEM这些参数不应指定BlockSpec后续所有参数的BlockSpec的index_map会额外收到这些 SMEM 引用即 index_map 签名中多了前导标量参数。从源码看PrefetchScalarGridSpec定义于 jax/_src/pallas/mosaic/core.py它继承自GridSpec构造参数为num_scalar_prefetch、grid、in_specs、out_specs、scratch_shapes在get_grid_mapping中前num_scalar_prefetch个输入被切分出来并映射为TPUMemorySpace.SMEM的引用jax/_src/pallas/mosaic/core.py。配套的测试用例如 tests/pallas/tpu/pallas_call_test.py展示了标准用法标量索引s被预取进 SMEMx的index_map接收(i, s_ref)并在内核内用pl.load(s_ref, (i,))读取标量来决定切片位置。该类测试均可在interpretTrue下于 CPU 验证。支持的数据类型与计算放置目前 Pallas TPU仅支持以下数据类型类别支持情况jnp.float32支持jnp.bfloat16支持jnp.int*支持所有精度除jnp.int4jnp.uint*支持所有精度计算放置规则所有标量0D数组存放在标量寄存器中相关运算在标量核执行其余所有操作即使是单元素的 1D 数组都在向量核执行。支持的操作全景与性能特征矩阵乘法矩阵乘法结果始终是 float32。若输入不是 float32推荐用lax.dot并设preferred_element_typejnp.float32使用lax.dot_general时操作数最后两维的转置可以融合进乘法提升整体性能。精度控制Pallas TPU 的 lowering 感知jax.default_matmul_precision。追求最高性能和最低精度用bfloat16关心数值精度则设为float32。警告即使给矩阵乘法传入 32 位操作数除非显式请求float32精度它们仍会被舍入到bfloat16。转置值至少有 4 维时除最后两维外的任意轴转置是免费的否则只实现了最后两维的转置注意最后两维的某些转置可以融合进矩阵乘法。内存访问引用Ref的任意切片均可读可写但要受实现约束32 位宽的输入目前没有限制更窄的类型只支持部分切片模式最后两维上对齐到 8 和 128 的倍数、且长度为 8 和 128 的倍数的读写总是受支持。由于向量内存的读写通常发生在(8, 128)的 tile 上读写至少二维的引用时最佳性能条件是基址偏移可被 tiling 整除读取区域大小是 tile 大小的倍数。元素操作硬件一般只支持用32 位类型做逐元素计算。加载低精度操作数时通常应先 upcast 到 32 位类型再做元素操作。不同元素操作的成本差异显著官方将其分为三档操作成本jnp.add、 便宜jnp.sub、- 便宜jnp.mul、* 便宜/、//、% 中等jnp.max、jnp.min 便宜jnp.whereselect 便宜jnp.abs 便宜\|、^、、~ 便宜、 便宜比较等 便宜类型转换.astype 便宜jnp.exp 中等jnp.tanh 中等jnp.pow 中等jnp.sin 昂贵jnp.cos 昂贵许多 JAX 函数由其他原语组合实现故该表并非穷尽。例如jax.nn.relu由比较和jnp.where实现因此也能在 Pallas 内核中工作。数组构造器所有常量数组构造器均受支持jnp.ones、jnp.zeros、jnp.full。值得注意jax.random模块目前与 Pallas 不兼容。归约sum、max、min归约受支持但每次只能归约单个数组轴且性能差异明显最后一维归约一般最慢倒数第二维归约较快但仍慢于前导维。广播广播的性能特征与归约非常相似除最后两维外的广播总是受支持且免费沿倒数第二维广播较慢沿最后一维广播最慢。Reshape除最后两维外的 reshape 均受支持且免费能修改最后两维的 reshape 只有两种受支持情况(1)某些前导维被展平到倒数第二维上(2)添加一个刚被归约移除的维度。控制流TPU 后端目前对控制流支持有限支持cond、fori_loop和for_loop。但循环原语在编译期会被完全展开所以请把循环次数trip count控制在合理的小范围内。过度使用控制流会导致底层代码生成显著退化推荐尽可能把计算密集的操作挤进单个基本块。用流水线重叠内存 I/O 与计算VMEM/SMEM 的两大约束使用pallas_call直接搬运数组有两大约束容量VMEM 和 SMEM 很小v4 TPU 的 VMEM 只有 16MiBSMEM 只有几十到几百 KiB。作为参照一个f32[2048, 2048]数组恰好 16MiB——上面的朴素内核无法扩展到超出中等规模的数组带宽HBM↔VMEM 的拷贝比绝大多数计算指令慢得多。add_matrices大概率花在 HBM/VMEM 间拷贝上的时间多于加法本身。流水线的思想把拷贝藏进计算流水线的目标是在并行地做 HBM↔VMEM 拷贝的同时利用计算单元。朴素程序的问题在于先把所有x、y拷贝完才开始计算造成拷贝与计算的串行依赖。如果把计算切分成多个子计算例如把矩阵加法拆成若干块的加法就能用一块子计算的拷贝去重叠另一块子计算的计算。以把(512, 512)的x、y沿前导轴切成x1, x2、y1, y2各(256, 512)为例流水线执行序列为拷贝x1、y1进 VMEM开始异步拷贝x2、y2进 VMEM从 VMEM 加载x1, y1到 VREGs计算z1 x1 y1把z1存入 VMEM开始把z1从 VMEM 拷回 HBM等待x2, y2拷贝完成加载x2, y2到 VREGs计算z2 x2 y2存z2进 VMEM等待z1拷完开始拷z2回 HBM最后等待z2拷完。任何时刻只要在做计算就有异步拷贝在进行——拷贝的部分时间没有被浪费。判定流水线效率的两个关键数字需要执行的浮点运算量FLOPs和为此需要拷贝的字节数。二者之比FLOPs/内存字节称为算术强度arithmetic intensity它决定了流水线是计算受限还是内存受限。用 grid 与 BlockSpec 表达流水线Pallas 用grid和BlockSpec自动生成上述流水线无需手写异步序列。在流水线中BlockSpec.block_shape为(256, 512)第一次迭代取x1、第二次取x2def x_index_map(i): return (i, 0) block_spec pl.BlockSpec((256, 512), x_index_map)完整内核如下——pallas_call负责把x、y拷入 VMEM、分配输出 VMEM 缓冲区z_vmem_ref并在内核结束后把输出拷回 HBMdef add_matrices_kernel(x_vmem_ref, y_vmem_ref, z_vmem_ref): # 从 VMEM 加载到 VREGs x_vregs x_vmem_ref[:, :] y_vregs y_vmem_ref[:, :] # 执行向量加法 z_vregs x_vregs y_vregs # 把 VREGs 中的结果存回 VMEM z_vmem_ref[:, :] z_vregs def add_matrices_pipelined(x: jax.Array, y: jax.Array) - jax.Array: block_spec pl.BlockSpec((256, 512), lambda i: (i, 0)) return pl.pallas_call( add_matrices_kernel, out_shapex, in_specs[block_spec, block_spec], out_specsblock_spec, grid(2,) )(x, y)只加了很少的代码BlockSpec和grid就完成了大量工作BlockSpec提供了足够信息去预取输入块——例如迭代i时把i 1传入index_map得到下一迭代所需的块然后发起异步拷贝对输出则等待上一迭代输出拷贝完成再开始当前迭代输出的拷贝。pallas_call的完整签名含grid_spec、debug、interpret、compiler_params等参数见 jax/_src/pallas/pallas_call.py。参数化流水线块大小是最重要的调优旋钮块大小是优化 Pallas 内核性能时最重要的调优参数更小的块会给流水线循环增加更多迭代每次迭代工作量更少。还可以同时沿第二维切分输入输出def add_matrices_pipelined_2d( x: jax.Array, y: jax.Array, *, bm: int 256, bn: int 256 ) - jax.Array: m, n x.shape block_spec pl.BlockSpec((bm, bn), lambda i, j: (i, j)) return pl.pallas_call( add_matrices_kernel, out_shapex, in_specs[block_spec, block_spec], out_specsblock_spec, grid(m // bm, n // bn), )(x, y) np.testing.assert_array_equal( add_matrices_pipelined_2d(x, y, bm256, bn256), x y ) np.testing.assert_array_equal( add_matrices_pipelined_2d(x, y, bm128, bn128), x y ) np.testing.assert_array_equal( add_matrices_pipelined_2d(x, y, bm512, bn512), x y )二维grid在底层被降级为嵌套循环即多层流水线。处理归约初始化累加器是关键把(8, 512, 512)的数组沿轴 0 归约为(512, 512)可用大小为(8,)的grid每次迭代把x[i]累加到输出 VMEM 缓冲区。朴素实现是错误的# Warning: this implementation is incorrect! def naive_sum_kernel(x_ref, o_ref): o_ref[...] x_ref[...] def naive_sum(x: jax.Array) - jax.Array: grid, *out_shape x.shape return pl.pallas_call( naive_sum_kernel, gridgrid, # block_shape 中的 None 表示取大小为 1 并在内核中挤掉该维 in_specs[pl.BlockSpec((None, *out_shape), lambda i: (i, 0, 0))], out_specspl.BlockSpec(out_shape, lambda i: (0, 0)), out_shapejax.ShapeDtypeStruct(out_shape, x.dtype), )(x)这里in_specs把(512, 512)整块载入 VMEM该维度不流水线化index_map每次选x的第i维block_shape中的None表示选择一个单例维度并在内核里挤掉因此内核里的x_ref是(512, 512)。out_specs的index_map为lambda i: (0, 0)说明o_ref在流水线中保持不变每次迭代都可读写它。问题在于o_ref初始内容是垃圾累加会基于垃圾进行导致结果错误。因此在内核中做归约时必须初始化保存归约值的Ref。用pl.whenjax.lax.cond的便捷封装配合pl.program_id查询当前 grid 轴的迭代序号在迭代 0 初始化def sum_kernel(x_ref, o_ref): pl.when(pl.program_id(axis0) 0) def _(): o_ref[...] jnp.zeros_like(o_ref) o_ref[...] x_ref[...] def sum(x: jax.Array) - jax.Array: grid, *out_shape x.shape return pl.pallas_call( sum_kernel, gridgrid, in_specs[pl.BlockSpec((None, *out_shape), lambda i: (i, 0, 0))], out_specspl.BlockSpec(out_shape, lambda i: (0, 0)), out_shapejax.ShapeDtypeStruct(out_shape, x.dtype) )(x)归约必须放在 grid 的最右minormost维度Pallas 用BlockSpec、grid和内核函数生成的流水线不会从 HBM 读回输出——输出一旦写回 HBM 就无法再访问。因此不能跨越会被重新访问的 grid 维度做归约所有归约都必须发生在 grid 的最右维度。上例 grid 只有一维天然满足。排查与验证要点interpret 模式优先所有内核先用interpretTrue在 CPU 上验证逻辑正确性源码文档明确这是唯一的 CPU 运行方式再上真机。测试目录 tests/pallas/tpu/pallas_call_test.py 中的用例如标量预取、vmap 组合均以interpret参数双路径覆盖可作为编写参考对齐规则涉及向量内存访问时优先保证最后两维偏移可被(8, 128)tile 整除、长度是其倍数VMEM 预算窗口过大导致的 VMEM OOM 会以底层编译器错误出现注意评估窗口 寄存器溢出所需空间控制流预算循环会被完全展开保持 trip count 小把计算密集操作放进同一基本块。总结Pallas 让开发者在不必完全理解 TPU 硬件的情况下也能开始写内核但理解硬件显然更利于写出高性能内核。本文覆盖了 Pallas TPU 的核心知识体系硬件内存模型HBM/VMEM/SMEM与顺序执行语义、BlockSpec/grid的切片与连续输出规则、维度顺序对寄存器压力的影响、dimension_semantics多核Megacore并行、PrefetchScalarGridSpec的 SMEM 标量预取、支持的数据类型与各类操作的成本特征以及用gridBlockSpec自动生成流水线来重叠内存 I/O 与计算、并正确处理归约累加器初始化。文中内核均可在 docs/pallas/tpu/pipelining.md 中查看完整可运行版本。延伸阅读grid与BlockSpec的通用概念见 docs/pallas/grid_blockspec.md 与 docs/pallas/quickstart.mdPallas 的总体设计见 docs/pallas/design.md。动手练习建议实现一个把其他维度也流水线化的sum内核并为add和sum内核补充 Megacore 注解。赞分享机器学习深度学习【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址https://gitcode.com/gh_mirrors/jax/jax点击查看免费下载相关推荐JAX Pallas TPU 快速入门内存空间、Ref 内核与流水线并行emit_pipeline 实战JAX Pallas TPU 快速入门内存空间、Ref 内核与流水线并行emit_pipeline 实战 本篇指南以 JAX 开源仓库中的 docs/pa人工智能机器学习深度学习编译器高性能计算JAX Pallas TPU 流水线编程实战内存层级、多重缓冲、动态块形状与 Megacore 调度JAX Pallas TPU 流水线编程实战内存层级、多重缓冲、动态块形状与 Megacore 调度 本文围绕 JAX 仓库中 TPU 专属的 Pallas人工智能机器学习深度学习编译器高性能计算JAX Pallas SparseCore 内核编写指南在 TPU 稀疏核上实现 gather、scatter 与流水线内核JAX Pallas SparseCore 内核编写指南在 TPU 稀疏核上实现 gather、scatter 与流水线内核 SparseCore稀疏核是人工智能机器学习深度学习编译器高性能计算上一篇ScyllaDB 一致性级别Consistency Level实战演示基于 cqlsh 的 3 节点集群读写与故障演练下一篇auto-novel数据可视化用户行为与翻译统计图表设计创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价