JAX Pallas Mosaic GPU 软件流水线实战:emit_pipeline 与 Warp Specialization 构建高吞吐 GPU 内核

发布时间:2026/9/10 6:22:17
JAX Pallas Mosaic GPU 软件流水线实战:emit_pipeline 与 Warp Specialization 构建高吞吐 GPU 内核 JAX Pallas Mosaic GPU 软件流水线实战emit_pipeline 与 Warp Specialization 构建高吞吐 GPU 内核【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax本文围绕 JAX 仓库中的 Pallas Mosaic GPU 流水线指南展开讲解在 Hopper 及更新的 GPU 上如何用plgpu.emit_pipeline显式编排 GMEM/SMEM 数据搬运与 TensorCore 计算的重叠如何通过pl.pallas_call兼容 API 启用流水线以及如何使用emit_pipeline_warp_specialized做 warp 特化warp specialization的矩阵乘法内核。读完本文你可以掌握max_concurrent_steps、delay_release等关键调优参数的语义、wgmma 指令与流水线缓冲区的正确配合方式并能将仓库中的两个完整 matmul 示例作为模板迁移到自己的 GPU 内核开发中。1. 显式编程的流水线模型Pallas 对软件流水线software pipelining采取显式编程的路线用户自己在内核里写出流水线循环明确控制每一步的拷贝与计算。官方建议先阅读面向 TPU 的通用流水线指南标签pallas_software_pipelining见docs/pallas/design/目录下的设计文档建立整体认知然后再学习 Mosaic GPU 后端的差异。与 Triton 的编程模型相比这是一个显著区别在 Triton 中流水线是编译器自动完成的优化而在 Pallas 中流水线由程序员通过plgpu.emit_pipeline对顺序循环做流水线化配合plgpu.kernel把问题并行切分到 CUDA grid 上显式构建。显式编程换来的是对多级缓冲、释放时机、warpgroup 分工的完全控制权这也是后文一系列调优参数存在的意义。2. GPU 内存空间SMEM 与 GMEM流水线的前提是理解 Pallas 在 GPU 上的两层内存空间。Ref 主要存在于两种内存空间之一通过BlockSpec的memory_space参数显式指定例如BlockSpec(memory_spaceplgpu.GPUMemorySpace.GMEM)plgpu.GPUMemorySpace.SMEM共享内存SMEM Ref 可以通过数组索引语法解引用到寄存器参与计算即x y_ref[...]。使用emit_pipeline时Ref 就落在这个内存空间。plgpu.GPUMemorySpace.GMEM全局内存/HBM分配在 GMEM 的 Ref不参与流水线且不能直接用数组索引访问。GMEM 必须经由 SMEM 中转读取用plgpu.copy_gmem_to_smem写回用plgpu.copy_smem_to_gmem或者直接用plgpu.emit_pipeline流水线进 SMEM。emit_pipeline的核心价值在于用异步的 GMEM↔SMEM 数据搬运去掩盖 TensorCore 计算或反过来用计算掩盖搬运。异步拷贝延迟很长而所有 TensorCore 计算必须发生在寄存器上矩阵乘法的输入来自 SMEM Ref。把两者重叠起来是 GPU 内核吞吐量的关键。3.plgpu.emit_pipeline接口与参数详解推荐做法是用plgpu.emit_pipeline对顺序循环做流水线并用plgpu.kernel在 CUDA grid 上并行切分问题。emit_pipeline的 API 与pl.pallas_call相似但额外暴露了几个 GPU 专属选项参数语义body流水线单步函数接收循环索引与各输入的 SMEM Ref若启用 carry 则还有上一轮携带值。gridbody的调用次数流水线步数。与 CUDA grid 不同流水线 grid 保证顺序执行。in_specs/out_specs与pl.pallas_call类似但接受plgpu.BlockSpec可指定 GPU 专属转换如 swizzling详见参考文档中的 memory reference transforms。max_concurrent_steps最大并发内存搬运步数。值越大占用越多 SMEM 存放临时缓冲但能提升内存子系统利用率。官方建议对该参数做自动调优小值如 2在 ALU 密集型内核中可能因 SMEM 占用低而获得更高 occupancy但硬件调度的噪声更大4–6 之间的大值更适合无法利用额外 occupancy 的内核。delay_release缓冲被流水线复用前额外等待的迭代数。例如第 0 迭代拷入 SMEM 的缓冲在delay_release1、max_concurrent_steps2时不会在第 2 迭代被复用而是等到第 3 迭代标准双缓冲策略是第 2 迭代。如果没有对流水线操作数 awaitplgpu.wgmma必须设置delay_release1否则流水线会在 WGMMA 还在读取缓冲时就开始覆写造成静默数据竞争silent data races。该参数对允许多个异步 matmul 同时在飞、保持 TensorCore 流水线满载这类优化很关键但会略微降低emit_pipeline自身重叠内存搬运的效率。从源码结构看emit_pipeline的实现位于 jax/_src/pallas/mosaic_gpu/pipeline.py可以印证几点文档中没有明说、但对实操很重要的约束grid 各维度必须严格为正否则直接抛ValueError当静态 grid 的总步数小于max_concurrent_steps时实现会自动收缩到总步数以减少 SMEM 中 Ref 的分配大小pipeline.py#L326-L329max_concurrent_steps必须大于所有delay_release值实现会做显式校验pipeline.py#L312-L319在 Hopper 之前的 GPU 上走的是cp.async路径仅支持输入型流水线不允许out_specs且要求输入访问可证明在界内或显式设置OOBFillMode.PROMISE_IN_BOUNDSpipeline.py#L331-L351。因此在 Ampere 及更早架构上使用 TMA 风格的流水线示例需要额外注意需要注意 API 的演进文档示例将delay_release1作为emit_pipeline的关键字参数传入而当前源码中emit_pipeline的签名不再直接接收该参数而是从每个输入BlockSpec的delay_release属性读取见 pipeline.py#L312-L314 的getattr(s, delay_release, 0)。如果你使用的 JAX 版本不再接受该关键字参数应在对应的plgpu.BlockSpec上设置delay_release语义不变。兼容 APIpl.pallas_callCompilerParams作为emit_pipeline的替代为了与 Pallas TPU 保持兼容Mosaic GPU 也实现了pl.pallas_call。默认情况下Mosaic GPU 上的pl.pallas_call只在 CUDA grid 上并行切分内核要启用流水线需要传入plgpu.CompilerParams作为compiler_params其中与流水线相关的选项有dimension_semanticsLiteral[parallel, sequential]的元组逐个指定 grid 维度的迭代语义。parallel维度在 CUDA grid 上并行切分sequential维度做顺序流水线。注意如果没有任何维度被标记为sequential就不会发生流水线max_concurrent_steps与plgpu.emit_pipeline中的同名参数语义一致。delay_release同上。流水线让你在顺序迭代间复用 scratch 缓冲例如实现归约。此外使用 Mosaic GPU 后端时pallas_call也支持用plgpu.BlockSpec替代pl.BlockSpec以指定 GPU 专属内存转换。源码佐证pallas_call.py#L45-L48 中当dimension_semantics为None时默认取(parallel,) * len(grid)——这正对应无sequential维度则无流水线的行为CompilerParams字段定义在 core.py#L135-L149。官方明确建议能用plgpu.kernel就用plgpu.kernel因为它支持更多特性如指定 warpgroup 数量、warp 特化。4. 示例Hopper GPU 上的 Matmul 内核下面这个矩阵乘法内核面向 Hopper GPU使用 Hopper 专属的wgmmawarpgroup matrix multiply accumulate指令。wgmma由单个 Mosaic GPU 线程发出在 TensorCore 上异步执行。内核实现[M, K] [K, N] [M, N]的分块矩阵乘法每个输出块在 CUDA grid 上并行计算。该 grid 由外层plgpu.kernel的grid参数指定在矩阵乘法非收缩维度 M、N上并行。在单个 program instance 内部用plgpu.emit_pipeline沿收缩维度 K运行顺序流水线。每次迭代从两个输入矩阵各加载一个 tile、相乘并把结果累加进累加器 Refplgpu.ACC。plgpu.ACC是一种特殊的、驻留在寄存器上的 Ref保存 WGMMA 的中间结果。累加完整个收缩维度后把结果写出到输出 Ref。执行真正的矩阵乘法时调用plgpu.wgmma(accumulator, LHS, RHS)把操作数推进 TensorCore 流水线。所有 WGMMA 操作按序执行可以看作向一个队列里压操作由于wgmma是异步指令用plgpu.wgmma_wait(N)等待直到队列中至多剩 N 个未完成操作。本实现等待 1 个在飞 WGMMA即当前迭代压入的 WGMMA 会在下一个迭代才被等待——这正是与delay_release1配合、始终让一个 WGMMA 留在飞、避免每次迭代冲刷 TensorCore 流水线的关键。其他要点wgmma要求操作数满足特定布局见 CUDA 文档中的 register fragments / shared-memory matrix layouts。本例用输入 BlockSpec 上的TilingTransform与SwizzleTransform实现当前需要手工指定文档注明未来 Mosaic GPU 会推断 transforms届时可省去手工指定delay_release参数与plgpu.wgmma_wait(1)联用始终保持一个 WGMMA 在飞以维持 TensorCore 利用率。完整代码继承自原文档def matmul(a, b, tile_m128, tile_n128, swizzle128): dtype jnp.float16 swizzle_elems swizzle // jnp.dtype(dtype).itemsize tile_k swizzle_elems grid_m m // tile_m grid_k k // tile_k grid_n n // tile_n assert tile_m % swizzle_elems 0 # Note: Transforms will be inferred automatically # by Mosaic GPU in the future. transforms ( plgpu.TilingTransform((8, swizzle_elems)), plgpu.SwizzleTransform(swizzle), ) def kernel(a_gmem, b_gmem, o_gmem, o_smem, acc): def pipeline_step(_, a_smem, b_smem): plgpu.wgmma(acc, a_smem, b_smem) plgpu.wgmma_wait(1) # pl.program_id obtains the index into the grid. pid_m pl.program_id(0) pid_n pl.program_id(1) pipeline plgpu.emit_pipeline( pipeline_step, in_specs[ plgpu.BlockSpec( (tile_m, tile_k), lambda k: (pid_m, k), transformstransforms ), plgpu.BlockSpec( (tile_k, tile_n), lambda k: (k, pid_n), transformstransforms ), ], grid(grid_k,), max_concurrent_steps2, delay_release1, ) pipeline(a_gmem, b_gmem) # Store WGMMA accumulator to SMEM and then to GMEM. o_smem[...] acc[...].astype(dtype) plgpu.commit_smem() m_slice pl.ds(pid_m * tile_m, tile_m) n_slice pl.ds(pid_n * tile_n, tile_n) plgpu.copy_smem_to_gmem(o_smem, o_gmem.at[m_slice, n_slice]) plgpu.wait_smem_to_gmem(0) return plgpu.kernel( kernel, out_shapejax.ShapeDtypeStruct((m, n), jnp.float16), scratch_shapesdict( o_smemplgpu.SMEM((tile_m, tile_n), jnp.float16), accplgpu.ACC((tile_m, tile_n), jnp.float32) ), # grid specifies the CUDA grid. # Instances of kernel will be executed in parallel over this grid. grid(grid_m, grid_n), grid_names(m, n), )(a, b) m 132 * 128 n 4 * 128 k 10 * 64 key1, key2 jax.random.split(jax.random.key(42), 2) a jax.random.uniform(key1, shape(m, k), dtypejnp.float16) b jax.random.uniform(key2, shape(k, n), dtypejnp.float16) result matmul(a, b) np.testing.assert_allclose(result, a b)对应前面的 import 约定import jax from jax import lax from jax import numpy as jnp from jax.experimental.pallas import mosaic_gpu as plgpu from jax.experimental import pallas as pl import numpy as np几个值得注意的细节scratch_shapes声明了o_smemSMEM 输出暂存与accplgpu.ACCfloat32 累加器它们不是 GMEM 输入输出而是内核的局部缓冲输出路径是寄存器累加器 → SMEM → GMEMo_smem[...] acc[...].astype(dtype)先落 SMEMplgpu.commit_smem()提交再用plgpu.copy_smem_to_gmem异步写回最后plgpu.wait_smem_to_gmem(0)等待完成pl.program_id(0/1)获取的是外层 CUDA grid 索引M/N 分块号而emit_pipeline的grid(grid_k,)是内层顺序流水线的 K 维步数两层 grid 语义完全不同——这是读懂 Mosaic GPU 内核的关键。5. Warp Specializationwarp 特化warp 特化是指让每个 warp/warpgroup 只负责单一任务把调度灵活性交给 GPU 硬件。回顾一下硬件背景每个流式多处理器SM内都有可以互相切换 warp 的 warp scheduler例如某个 warp 停顿时可以立刻调度另一个 warp 执行。实践中这种多指令流方式往往比单指令流 编译器静态调度并尽力重叠更快。在 Hopper 上特别有价值的是warpgroup 特化用独立的 warpgroup 专门发射 TMAGMEM/SMEM 拷贝其余 warpgroup 专门做算术。索引计算和发射 TMA 本身也要花时间若与算术挤在一条指令流里可能让 TensorCore 空转特化后通信与算术分离两个 warpgroup 之间用一个consumed barrier同步——向内存 warpgroup 发出可以开始下一个 TMA 了的信号。plgpu.emit_pipeline_warp_specialized接口Pallas 通过plgpu.emit_pipeline_warp_specialized辅助函数启用 warp 特化。该 pipeline emitter 处理了内存线程里的全部逻辑用户只需描述计算线程的工作。它与普通emit_pipeline共享相似的 API支持的参数为plgpu.emit_pipeline_warp_specialized( body: Callable, * grid: tuple[int, ...], in_specs: Sequence[pallas_core.BlockSpec] (), out_specs: Sequence[pallas_core.BlockSpec] (), max_concurrent_steps: int, compute_context: Callable num_compute_wgs: int, memory_registers: int wg_axis: str, memory_thread_idx: int | None None, )GPU 专属参数说明num_compute_wgs计算线程/warpgroup 数量。pipeline emitter 固定使用单个内存线程所以在plgpu.kernel中应设num_threadsnum_compute_wgs1memory_registers分配给内存线程的寄存器数其余寄存器在计算线程间均分。默认值 40需要根据是否出现寄存器溢出register spill上下调整。源码 docstring 中明确写道在 H100 上 40 是个合理的起点pipeline.py#L697-L698wg_axis线程/warpgroup 轴的轴名即plgpu.kernel的thread_name参数所指定的名字memory_thread_idx指定哪个 Pallas 线程担任内存线程默认取最后一个线程compute_context允许指定只在计算线程执行的流水线 prologue/epilogue并支持定义贯穿流水线的 loop carry 的初始化与消费。所有计算线程专属的数组都应在这里实例化避免内存线程在寄存器中物化它们否则会因寄存器溢出而变慢。流水线 body 会并行地在所有计算线程上运行且由于这些计算线程调度在同一个 CUDA block 内它们共享 SMEM。内核内可以用lax.axis_index获取 Pallas 线程索引从而在计算线程之间分工。源码层面还有两处值得了解的补充见 emit_pipeline_warp_specialized 实现当前实现额外提供manual_produced_barriers/manual_consumed_barriers开关允许把 barrier 直接传入 body 由用户手动 arrive以及pipeline_state参数——当多个参数几乎相同的流水线要顺序执行时用START → STEADY* → STOP状态序列衔接可以避免流水线之间的 bubble且文档提示需配合get_allocations手动分配模式使用。6. 示例带 warp 特化的矩阵乘法下例把前面的 matmul 扩展为 warp 特化版本使用 2 个计算线程分别处理 RHS 矩阵的不同列但共享同一 LHS因此每次流水线调用计算输出矩阵中的 2 个相邻块。实现要点用compute_context模式下例的compute_thread函数初始化 WGMMA 累加器并把最终累加器从寄存器拷入 SMEM。累加器必须创建在compute_thread内部否则会被分配到内存线程上、白白浪费寄存器执行 WGMMA 时用pl.run_state包住wgmma从而创建一个以 carry 值初始化的 accumulator refpl.run_state(do_wgmma)(plgpu.ACC.init(acc))不使用pl.pallas_call而是用 GPU 专属的plgpu.kernel入口通过num_threads指定每个 CUDA block 的线程数通过thread_name指定内核内可查询 Pallas 线程索引的轴名本例为wg。完整代码继承自原文档def matmul_warp_specialized(a, b, tile_m128, tile_n128, swizzle128, compute_wgs2): dtype jnp.float16 elems_128b swizzle // jnp.dtype(dtype).itemsize tile_k elems_128b grid_m m // tile_m grid_k k // tile_k grid_n n // tile_n assert tile_m % elems_128b 0 transforms ( plgpu.TilingTransform((8, elems_128b)), plgpu.SwizzleTransform(128), ) def kernel(a_gmem, b_gmem, o_gmem, o_smem): wg_idx lax.axis_index(wg) wg_slice pl.ds(wg_idx * tile_n, tile_n) # pl.program_id obtains the index into the pallas_call grid. pid_m pl.program_id(0) pid_n pl.program_id(1) def compute_thread(pipeline): acc plgpu.layout_cast( jnp.full((tile_m, tile_n), 0, dtypejnp.float32), plgpu.Layout.WGMMA, ) # yield marks the place where the pipelined loop will be inserted. # Its argument are the initial carry values, and its result is the carry # value after the loop completes. final_acc pipeline(acc) o_smem[:, wg_slice] final_acc[...].astype(dtype) def kernel_body(_, a_smem, b_smem, carry): acc carry b_smem_wg b_smem.at[:, wg_slice] def do_wgmma(acc_ref): plgpu.wgmma(acc_ref, a_smem, b_smem_wg) acc pl.run_state(do_wgmma)( plgpu.ACC.init(acc)) return acc pipeline plgpu.emit_pipeline_warp_specialized( kernel_body, in_specs[ plgpu.BlockSpec( (tile_m, tile_k), lambda k: (pid_m, k), transformstransforms ), plgpu.BlockSpec( (tile_k, tile_n * 2), lambda k: (k, pid_n),transformstransforms ), ], grid(grid_k,), compute_contextcompute_thread, max_concurrent_steps2, num_compute_wgscompute_wgs, memory_registers40, memory_thread_idx2, wg_axiswg, ) # Call the pipeline pipeline(a_gmem, b_gmem) # Copy the output from SMEM to GMEM. plgpu.commit_smem() m_slice pl.ds(pid_m * tile_m, tile_m) n_slice pl.ds(pid_n * tile_n * 2, tile_n * 2) plgpu.copy_smem_to_gmem(o_smem, o_gmem.at[m_slice, n_slice]) plgpu.wait_smem_to_gmem(0) return plgpu.kernel( kernel, out_shapejax.ShapeDtypeStruct((m, n), jnp.float16), scratch_shapesdict( o_smemplgpu.SMEM((tile_m, tile_n * 2), jnp.float16) ), grid(grid_m, grid_n // 2), grid_names(m, n), num_threads3, # 2 compute, 1 memory. thread_namewg )(a, b) m 132 * 128 n 4 * 128 k 10 * 64 key1, key2 jax.random.split(jax.random.key(42), 2) a jax.random.uniform(key1, shape(m, k), dtypejnp.float16) b jax.random.uniform(key2, shape(k, n), dtypejnp.float16) result matmul_warp_specialized(a, b) np.testing.assert_allclose(result, a b)与非特化版本对比可以读出特化带来的结构变化grid(grid_m, grid_n // 2)N 维 grid 减半因为每个 program instance 一次算 2 个相邻输出块对应的输出写回切片也变成tile_n * 2宽n_slice pl.ds(pid_n * tile_n * 2, tile_n * 2)RHS 的in_specs一次性加载(tile_k, tile_n * 2)整块 B 矩阵 tile由内存线程统一搬进 SMEM计算线程再用b_smem.at[:, wg_slice]按wg_idx各取自己那一半列——共享 LHS 正是特化版省掉一半 A 矩阵搬运的原因num_threads3对应 2 compute 1 memory 的布局memory_thread_idx2把最后一个线程指定为内存线程与默认最后一个是内存线程一致累加器acc在compute_thread中用plgpu.layout_cast(..., plgpu.Layout.WGMMA)构造pipeline(acc)调用即把流水线循环插入yield标记处并拿到循环结束后的最终 carry随后写入o_smem的对应列区。7. 实践要点小结显式流水线意味着显式责任Pallas 不会像 Triton 那样替你隐藏流水线逻辑max_concurrent_steps、delay_release、warpgroup 分工都暴露给用户也意味着像忘记 await wgmma 又忘设 delay_release这类错误会以静默数据竞争的形式出现调试时应格外警惕调优路径先跑通非特化emit_pipeline版本示例中max_concurrent_steps2起步再按内核属性ALU 密集 vs 访存密集自动调优max_concurrent_stepsSMEM 用量随并发步数增长会直接影响 occupancyTensorCore 利用率wgmma_wait(1)delay_release的组合目标是让队列中始终有一个 WGMMA 在飞避免每个 K 迭代都冲刷 TensorCore 流水线进阶到 warp 特化当 TMA 发射与索引计算开始挤占计算线程时间时切换到emit_pipeline_warp_specialized按num_threadsnum_compute_wgs1配置线程并根据寄存器溢出情况微调memory_registersH100 上 40 为起点架构适用性示例基于 Hopper 的 wgmma/TMA 路径从源码看Hopper 之前的 GPU 走cp.async路径且有仅输入型流水线、访问须可证明在界内的限制移植到 Ampere 时需要相应调整。相关延伸阅读与源码入口Mosaic GPU 参考文档transforms、wgmma 等、emit_pipeline 实现、emit_pipeline_warp_specialized 实现、pallas_call 兼容路径的参数解析、CompilerParams 定义以及同目录下的 Mosaic GPU 快速上手 与 collective matmul 指南。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考