JAX Pallas TPU 流水线编程实战:内存层级、多重缓冲、动态块形状与 Megacore 调度

发布时间:2026/9/10 16:05:01
JAX Pallas TPU 流水线编程实战:内存层级、多重缓冲、动态块形状与 Megacore 调度 JAX Pallas TPU 流水线编程实战内存层级、多重缓冲、动态块形状与 Megacore 调度【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax本文围绕 JAX 仓库中 TPU 专属的 Pallas 流水线文档展开系统讲解 TPU 的内存层级HBM/VMEM/SMEM与计算单元结构并深入介绍 Pallas TPU 流水线的全部平台特性手动内存空间分配、多重缓冲Multiple Buffering、pltpu.emit_pipeline内核内流水线、Lookahead 预取、动态块形状以及 Megacore 双 TensorCore 并行调度。读完本文你将掌握在 TPU 上用pl.pallas_call编写 HBM-VMEM 流水线内核的完整参数配置方法并能通过dimension_semantics与core_axis让内核同时跑满两颗 TensorCore。以下内容以 TPU Pipelining 文档 为主体骨架并结合同仓库源码实现进行纵深补充。# 通用导入 import jax from jax.experimental import pallas as pl from jax.experimental.pallas import tpu as pltpu import jax.numpy as jnp import numpy as np关于流水线Software Pipelining的通用概念——例如为什么需要流水线、HBM 带宽与 TensorCore 计算吞吐之间的 gap 如何靠 DMA 重叠来弥补——可参考仓库中的通用文档 Pallas Pipelining 总览。本文聚焦 TPU 平台特有的部分。TPU 及其内存空间HBM、VMEM、SMEM 与寄存器在讨论 TPU 流水线之前必须先理解 TPU 的硬件布局一块 TPU 及其 TensorCore 由三部分构成——内存空间数组可以驻留的地方、寄存器临时存放标量与数组值和计算单元对寄存器中的值执行运算。文档中以 HBM 中两个数组x、y为起点描述了一张 TPU 内存层级示意原文档中以内嵌图片形式给出这里用文字描述其结构内存空间Memory spacesHBMHigh-Bandwidth Memory即我们通常所说的设备显存所有 JAX Array 默认驻留于此VMEMVector Memory用于存放向量和数组值的 SRAM 缓存SMEMScalar Memory专为标量值设计的 SRAM 缓存。寄存器RegistersTensorCore 有两类寄存器——VREG向量寄存器存放数组值SREG标量寄存器存放标量值。值可以从各自的缓存加载到寄存器VREG 从 VMEM 加载SREG 从 SMEM 加载。计算单元Compute unitsTensorCore 包含标量单元、向量单元VPU和矩阵单元MXU三者都能执行数值计算。它们可以异步工作但这一异步性由 TPU 编译器管理——因此从程序员视角看TPU 程序是单线程的。计算单元操作 SREG 与 VREG 中的值输出也写回这些寄存器。理解这张图的关键结论是TPU 上的流水线本质上就是让 HBM→VMEM 的数据搬运DMA与 TensorCore 上的计算VPU/MXU在时间上重叠。TPU 专属内存空间枚举Pallas 将 TPU 内存层级的所有级别都暴露给用户。仓库源码 jax/experimental/pallas/tpu.py 将这些枚举转发自 Mosaic 核心实现其中第 87~93 行直接导出了ACC、CMEM、SMEM、VMEM、VMEM_SHARED、HBM、SEMAPHORE等内存空间常量。文档给出的 Pallas 内存空间与标准内存类型DRAM/SRAM的对应关系如下Pallas 枚举TPU 内存空间类型DRAM/SRAMpl.ANY通常是 HBM也可能是 VMEMDRAMpltpu.VMEMVMEMSRAMpltpu.SMEMSMEMSRAMpltpu.SEMAPHORE信号量SRAM各空间的语义要点pltpu.VMEM向量 SRAM。若不指定VMEM 就是默认的内存空间。pltpu.SMEM标量 SRAM。只有标量类型的 load/store 可以访问 SMEM。pl.ANY向编译器提示该内存空间不受约束。多数情况下 XLA 会把它放到 HBM。分配给ANY空间的 buffer 无法用普通数组索引语法如x[...]直接解引用必须先通过pltpu.sync_copy或pltpu.async_copy把值复制到 VMEM 或 SMEM 的 buffer 中才能使用。pltpu.SEMAPHORE用于分配信号量可构建 barrier 或跟踪异步操作。还可以把信号量从内核返回以构建异步内核——这是实验特性参见 异步流水线设计文档。默认行为TPU 上的流水线通常是 HBMDRAM↔ VMEM向量 SRAM之间进行。pallas_call在 TPU 上的默认约定是pallas_call的实参默认位于 HBM而传入用户内核体的输入 buffer 存放在 VMEM。也就是说你不需要为常规的内核显式配置任何内存空间Pallas 会自动完成 HBM↔VMEM 的搬运与流水线编排。手动指定内存空间memory_space与scratch_shapes虽然不专属流水线但你可以通过BlockSpec上的memory_space参数手动控制输入/输出 buffer 的内存空间。文档特别强调两条约束除非memory_space标记为VMEM否则不允许对该 buffer 进行流水线化内存空间还可以用于通过pallas_call的scratch_shapes参数为内核指定 scratch工作区参数。scratch buffer 会在内核迭代之间持久存在适合存放部分累加、归约结果等中间量。scratch buffer 必须位于VMEM、SMEM或SEMAPHORE之一。文档给出的完整示例把一个 HBM bufferx_hbm_ref的一片切到 scratch VMEM buffer 中做算术运算再把结果写回输出 VMEM buffer——def hbm_vmem_kernel(x_hbm_ref, out_vmem_ref, scratch_vmem_ref): pltpu.sync_copy(x_hbm_ref.at[0:1], scratch_vmem_ref) out_vmem_ref[...] scratch_vmem_ref[...] 1 x jax.random.uniform(jax.random.key(0), (8, 128), jnp.float32) out pl.pallas_call(hbm_vmem_kernel, in_specs[pl.BlockSpec(memory_spacepl.ANY)], # 输入 x 留在 HBM out_shapejax.ShapeDtypeStruct((1, 128), jnp.float32), scratch_shapes(pltpu.VMEM(shape(1, 128), dtypejnp.float32),) )(x) np.testing.assert_allclose(out, x[0:1] 1)注意三个要点输入用pl.BlockSpec(memory_spacepl.ANY)声明留在 HBM输出走默认的 VMEM由out_shape推断scratch 通过scratch_shapes传入一个pltpu.VMEM(shape..., dtype...)元组。pltpu.sync_copy的实现位于 jax/_src/pallas/mosaic/helpers.py第 25 行def sync_copy而pltpu模块只是把jax/_src/pallas/mosaic/下的实现按需转发出来见 jax/experimental/pallas/tpu.py 第 32、44 行的sync_copy/async_copy导入。多重缓冲Multiple Buffering流水线的深度由 buffer 数量决定。pl.BlockSpec的pipeline_mode选项支持按实参粒度指定多重缓冲传入一个pl.Buffered对象即可为某个特定实参分配指定数量的 bufferpl.BlockSpec( pipeline_modepl.Buffered(buffer_countbuffer_count) )所有输入和输出的默认 buffer 数量是 2。从源码看pl.Buffered是一个 frozen dataclass定义在 jax/_src/pallas/core.py第 211~236 行除buffer_count外还带有use_lookahead是否启用 lookahead、revisit输出块重复访问策略RevisitMode.IMMEDIATE为输出默认值与prefetched_count等字段——其中use_lookahead的注释明确写道启用 lookahead 后流水线可以在任何槽位一空出来就提前开始获取下一个变化的块这正是下一节 TPU 特性中 Lookahead 预取的底层开关。pltpu.emit_pipeline在内核内部构建流水线pl.pl.pallas_call的流水线只发生在内核入口处而pltpu.emit_pipeline允许你在内核体内部构建流水线。文档列出的典型使用场景嵌套流水线例如外层流水线负责芯片间通信内层流水线负责 HBM↔VMEM 搬运使用emit_pipeline专属特性如Lookahead 预取与动态块形状见下文两节。它的签名与pl.pallas_call类似需要指定内核体kernel、grid 以及输入输出的 block specsdef emit_pipeline( kernel: Callable, grid: tuple[int], in_specs: PyTree[BlockSpec] None, out_specs: PyTree[BlockSpec] None, dimension_semantics: tuple[GridDimensionSemantics] None, core_axis: int | None None, ) - Callable: ... # 给定内层内核与 BlockSpecs返回一个自定义流水线dimension_semantics和core_axis两个参数用于把内核 grid 划分到 Megacore 的两颗 TensorCore 上见后文 Megacore 一节。从源码看真实实现位于 jax/_src/pallas/mosaic/pipeline.py第 2000 行起def emit_pipeline当前版本签名额外提供了tiling、core_axis_name、trace_scopes、no_pipelining等参数并且当dimension_semantics未指定时会默认填充(ARBITRARY,) * len(grid)第 2025~2026 行——即不声明就不可并行这与文档中不指定dimension_semantics只会使用单颗 TensorCore的说明相互印证。Lookahead 预取提前拉取数据以覆盖变长计算Lookahead prefetch 是流水线的一种调度策略一旦某个缓冲槽位空出来流水线就立即尝试拉取下一个输入块而不是等到它即将被使用的前一个迭代。文档给出的例子非常直观假设内核 grid 为(8,)各迭代需要获取的块索引序列是0, 0, 0, 0, 1, 1, 1, 1——标准调度第 0 次迭代取块0要等到第 3 次迭代才开始取块1Lookahead 调度第 0 次迭代就同时开始取块0和块1。文档解释了为什么值得引入这点控制流开销当每个块的计算量可变时例如某些块包含跳过的或减少的工作量块真正被使用的前一个迭代里可能没有足够的计算来完全掩盖内存搬运时间。此时希望更早地把块拉进缓冲区。Lookahead 开销较小但默认关闭。启用方式与多重缓冲一致在pipeline_mode中把pl.Buffered的use_lookahead置为Truepl.BlockSpec( pipeline_modepl.Buffered(buffer_countbuffer_count, use_lookaheadTrue) )在源码中可以看到 lookahead 的落地路径jax/_src/pallas/mosaic/pipeline.py 中BufferedRef.use_lookahead属性第 539 行附近以及第 1297~1301 行针对use_lookahead的迭代调度分支而 Mosaic GPU 一侧的 lowering第 1270 行也消费同一标志——可见这是流水线基础设施层的统一特性。动态块形状pl.BoundedSlice与pl.dspltpu.emit_pipeline支持对动态但有界bounded形状的块做流水线。用法约定BlockSpec 的动态维度要用pl.BoundedSlice(max_size)标注而不是静态整数max_size是该块的最大尺寸index_map返回的对应索引必须是元素element索引而非块索引构造的动态切片pl.ds(start, size)且start、size都允许是动态值。一个最小声明示例pl.BlockSpec( block_shape(pl.BoundedSlice(32), 256), index_maplambda *grid_idxs: (pl.ds(start, end), 0), )文档给出的端到端例子内核按slices中描述的动态尺寸块把x逐块拷贝到输出——# 内核按 slices 中给出的动态尺寸块将 x 拷贝到输出 def dynamic_block_example_kernel(x_hbm, slices_hbm, o_hbm, slices_smem): pltpu.sync_copy(slices_hbm, slices_smem) # 把 slices 拷入 SMEM def pipeline_body(x_vmem, o_vmem): o_vmem[...] x_vmem[...] def index_map(i): start slices_smem[i, 0] size slices_smem[i, 1] - slices_smem[i, 0] return (pl.ds(start, size), 0) block_spec pl.BlockSpec(block_shape(pl.BoundedSlice(8), 128), index_mapindex_map) pltpu.emit_pipeline( pipeline_body, grid(slices.shape[0],), in_specs[block_spec], out_specsblock_spec )(x_hbm, o_hbm) x jax.random.uniform(jax.random.key(0), (8, 128), jnp.float32) slices jnp.array([[0, 2], [2, 3], [3, 5], [5, 8]], dtypejnp.int32) hbm_block_spec pl.BlockSpec(memory_spacepl.ANY) out pl.pallas_call(dynamic_block_example_kernel, in_specs[hbm_block_spec, hbm_block_spec], out_specshbm_block_spec, out_shapejax.ShapeDtypeStruct((8, 128), jnp.float32), scratch_shapes(pltpu.SMEM(slices.shape, jnp.int32),) )(x, slices) np.testing.assert_allclose(x, out)这个例子里有几处值得注意的 TPU 细节动态边界slices本身放在 HBM通过pltpu.sync_copy拷入 SMEM 供index_map逐迭代读取SMEM 只能存标量/标量级小数组slices通过scratch_shapes(pltpu.SMEM(slices.shape, jnp.int32),)声明为内核 scratch。源码侧pl.BoundedSlice定义在 jax/_src/pallas/core.py第 415 行起而 mosaic/pipeline.py 第 162~165 行有对应约束BoundedSlice的index_map必须返回ds动态切片否则报 Must return a ds from the index_map for a BoundedSlice——与文档中必须用pl.ds的要求完全一致。Megacore 配置下的 TPUdimension_semantics与core_axis部分 TPU 芯片带有两颗 TensorCore 但对 JAX 用户呈现为单个设备这就是所谓Megacore。两颗 TensorCore 各自拥有独立的 VMEM、VREG、SMEM、SREG 与计算单元但共享 HBM。文档中的示意描述为Megacore 的 TPU 在概念上像一个非常简单的 GPU——只有两条线程。要让内核同时利用两颗 TensorCore基本思路是如果计算中存在**可尴尬并行embarrassingly parallel**的维度就把它拆分到两颗 TensorCore 上。通过给pallas_call提供dimension_semantics注解来指示哪些维度可以并行。文档给出的矩阵加法示例def add_matrices_kernel(x_vmem_ref, y_vmem_ref, z_vmem_ref): # 从 VMEM 加载 x、y 到 VREG x_vregs x_vmem_ref[:, :] y_vregs y_vmem_ref[:, :] # 执行向量化加法 z_vregs x_vregs y_vregs # 把 VREG 中的结果写回 VMEM z_vmem_ref[:, :] z_vregs 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_shapejax.ShapeDtypeStruct.like(x), in_specs[block_spec, block_spec], out_specsblock_spec, grid(2,), compiler_paramspltpu.CompilerParams( dimension_semantics(parallel,)) )(x, y) x, y jnp.ones((512, 512)), jnp.ones((512, 512)) add_matrices_pipelined_megacore(x, y)dimension_semantics语义规定如下它必须是与grid等长的元组每一项取值parallel或arbitraryparallel告诉 Pallas 该维度的 for 循环各次迭代可以独立执行而不影响正确性arbitrary该 grid 维度不能做任何并行假设因此不可被并行化。指定dimension_semantics后内核会在每颗 TensorCore 上同时执行Pallas 会自动拆分 grid。文档特别提醒平台限制Megacore 目前仅在 TPU v4 和 TPU v5p 上可用在其他平台上提供dimension_semantics注解是 no-op但不提供则意味着即便存在多颗 TensorCore 也只会用到其中一颗。仓库中这一机制已有大量生产级使用例如 TPU MatMul 内核在 jax/experimental/pallas/ops/tpu/matmul.py第 82 行使用dimension_semantics(parallel, parallel, arbitrary)FlashAttention 在 jax/experimental/pallas/ops/tpu/flash_attention.py 多处标注各 grid 维度的并行性SplashAttention、PagedAttention 等稀疏注意力内核也都遵循同一模式可视为本文示例的直接印证。在emit_pipeline中划分 corecore_axis当使用pltpu.emit_pipeline构建内核内流水线时应把core_axis传入emit_pipeline。core_axis是用于划分 grid 的某个并行 grid 轴的下标。文档给出的模板把内核沿一个前置的并行 grid 维度划分——def kernel_body(...): def inner_pipeline_body(...): ... pltpu.emit_pipeline(inner_pipeline_body, grid(4, 4), core_axis0, dimension_semantics(parallel, sequential)) pl.pallas_call( kernel_body, grid(num_cores,), compiler_paramspltpu.CompilerParams( dimension_semantics(parallel,)) )模式总结为外层pallas_call的 grid 对应 core 数并标注parallel内层emit_pipeline通过core_axis指定其 grid 的哪个维度按 TensorCore 切分。小结TPU 流水线参数速查能力关键 API / 参数默认行为HBM↔VMEM 自动流水线pl.pallas_call常规用法实参默认在 HBM内核输入 buffer 在 VMEM手动内存空间pl.BlockSpec(memory_space...)、pallas_call(scratch_shapes...)scratch 必须位于 VMEM/SMEM/SEMAPHORE非 VMEM 不可流水线多重缓冲pipeline_modepl.Buffered(buffer_countn)默认 buffer 数为 2内核内流水线 / 嵌套流水线pltpu.emit_pipeline(kernel, grid, in_specs, out_specs, ...)未声明dimension_semantics时视为全arbitrary源码默认值Lookahead 预取pl.Buffered(buffer_countn, use_lookaheadTrue)默认关闭动态块形状pl.BoundedSlice(max_size)pl.ds(start, size)元素索引仅限emit_pipelineMegacore 双核并行pltpu.CompilerParams(dimension_semantics(parallel, ...))emit_pipeline(core_axis...)未声明时仅用单颗 TensorCoreMegacore 目前限 TPU v4 / v5p掌握以上配置组合即可覆盖文档中描述的 TPU 流水线全部场景从最省心的自动 HBM-VMEM 流水到 lookahead 应对变长块负载、动态形状应对不规则分块再到 Megacore 双 TensorCore 的自动 grid 划分。更多 TPU 侧实践可继续参阅同目录的 硬件介绍、MatMul 教程与 快速上手。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考