
Mosaic GPU 如何用 emit_pipeline 给 Pallas kernel 写软件流水线重叠计算与访存【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax用 JAX Pallas 为 NVIDIA GPU 写 TensorCore kernel矩阵乘法、attention 等时有一个绕不开的性能问题GMEMHBM与 SMEMshared memory之间的异步拷贝延迟很长而 Hopper 的wgmma指令又只能在 SMEM/寄存器上的数据上执行。不做处理的话TensorCore 会在数据到达前空转。Mosaic GPU 后端提供的plgpu.emit_pipeline就是解决这个问题的显式软件流水线 API它把「按顺序搬运输入 tile」和「每步执行一次计算」重叠起来让硬件加速器在等数据的时候不停下来。本文以 Hopper GPU 上的分块矩阵乘法为例走一遍从写 kernel、到验证输出、到调参的完整过程。与 Triton 的关键区别要先说明Pallas 中的流水线是显式编程的而 Triton 的流水线是编译器自动完成的优化。这意味着缓冲区管理、并发拷贝数量、延迟释放这些决策都由你通过emit_pipeline的参数指定。准备Mosaic GPU 与内存空间Mosaic GPU 的入口模块是jax.experimental.pallas.mosaic_gpu下文缩写为plgpukernel 通过plgpu.kernel启动每个 Pallas thread 对应一个 warpgroup4 个 warp 128 个 CUDA thread代码按 lockstep 执行不需要管理单个 CUDA thread。kernel 里的数据通过 RefJAX 的可变数组引用访问每个 Ref 位于特定内存空间GMEMplgpu.GMEM全局内存/HBM容量大、延迟高kernel 的输入输出都在这里。GMEM Ref 不能直接用下标索引访问只能通过plgpu.copy_gmem_to_smem/plgpu.copy_smem_to_gmem或emit_pipeline经 SMEM 中转。SMEMplgpu.SMEMSM 内的 shared memory可用x y_ref[...]解引用到寄存器参与计算。emit_pipeline会把 pipeline 输入 pipelined 进 SMEM。ACCplgpu.ACC驻留在寄存器中的 TensorCore 累加器 Ref持有wgmma的中间结果。Hopper 上 TensorCore 工作的典型数据流是GMEM → SMEM → Tensor Cores → 寄存器 → SMEM → GMEM。emit_pipeline的主要用途就是让 TensorCore 计算与 GMEM/SMEM 之间的数据搬运重叠起来因为异步拷贝延迟很长而所有 TensorCore 计算必须在寄存器或矩阵乘法的 SMEM Ref上进行。import jax from jax import numpy as jnp from jax.experimental import pallas as pl from jax.experimental.pallas import mosaic_gpu as plgpu import numpy as npemit_pipeline 的参数含义在 Mosaic GPU Pipelining 指南中推荐用plgpu.emit_pipeline对顺序循环做流水线同时用plgpu.kernel把问题在 CUDA grid 上并行划分。emit_pipeline的 API 与pl.pallas_call类似但额外暴露了几个 GPU 专用选项body、grid语义同pl.pallas_call。grid表示body会运行多少次与 CUDA grid 不同pipeline grid 保证顺序执行。in_specs/out_specs同pl.pallas_call但额外接受plgpu.BlockSpec实例可以指定 GPU 专用的内存 reference transform如 swizzlingtransform 的完整说明见 Mosaic GPU Reference。max_concurrent_steps控制最大并发内存传输数。更大的值会消耗更多 SMEM 存放临时缓冲区但可以提高内存子系统利用率。文档建议对该参数做 autotune。delay_release指定缓冲区被流水线复用前要额外等待的迭代数。例如delay_release1、max_concurrent_steps2时第 0 次迭代拷入 SMEM 的缓冲区要到第 3 次迭代才被复用标准双缓冲是第 2 次。如果你的 pipeline 操作数上还挂着未 await 的plgpu.wgmma就必须设置delay_release1否则流水线会在 WGMMA 还在读缓冲区时就开始覆盖——文档原话省略这个参数会产生 silent data races静默数据竞争。主路径Hopper matmul kernel 的完整写法下面是一个针对 Hopper GPU 的分块矩阵乘法[M, K] [K, N] [M, N]来自 GPU Quickstart。外层plgpu.kernel的 grid 在 M、N非收缩维上并行每个输出 block 由一个 CUDA block 计算块内用plgpu.emit_pipeline在收缩维 K 上做顺序流水线每次迭代加载两个输入 tile、执行一次wgmma、把结果累加进plgpu.ACCK 维全部累加完后把结果写回输出。def matmul(a, b, tile_m128, tile_n128, tile_k64, out_dtypejnp.float16): m, k a.shape _, n b.shape plgpu.kernel( out_typejax.ShapeDtypeStruct((m, n), out_dtype), scratch_typesdict( o_smemplgpu.SMEM((tile_m, tile_n), out_dtype), accplgpu.ACC((tile_m, tile_n), jnp.float32), ), grid(m // tile_m, n // tile_n), grid_names(m, n), ) def kernel(a_gmem, b_gmem, o_gmem, o_smem, acc): pid_m jax.lax.axis_index(m) pid_n jax.lax.axis_index(n) def body(_, a_smem, b_smem): plgpu.wgmma(acc, a_smem, b_smem) plgpu.wgmma_wait(1) # Keep one wgmma in flight. plgpu.emit_pipeline( body, grid(k // tile_k,), in_specs[ plgpu.BlockSpec( (tile_m, tile_k), lambda ki: (pid_m, ki), delay_release1 ), plgpu.BlockSpec( (tile_k, tile_n), lambda ki: (ki, pid_n), delay_release1 ), ], max_concurrent_steps2, )(a_gmem, b_gmem) # Drain: move the accumulated result to GMEM via SMEM. o_smem[...] acc[...].astype(out_dtype) plgpu.commit_smem() # Make the SMEM write visible to the TMA engine. plgpu.copy_smem_to_gmem( o_smem, o_gmem.at[pl.ds(pid_m * tile_m, tile_m), pl.ds(pid_n * tile_n, tile_n)], ) plgpu.wait_smem_to_gmem(0) # Wait for all copies to finish. return kernel(a, b)按执行顺序看这段代码里每个关键点的职责两级 grid 的分工。plgpu.kernel(..., grid(m // tile_m, n // tile_n), grid_names(m, n))是并行 grid每个 grid 点是一个独立 CUDA block用jax.lax.axis_index查询自己在 grid 中的坐标emit_pipeline(..., grid(k // tile_k,))是顺序 grid是 K 维上的流水线循环。emit_pipeline只负责每个 block 内部的顺序归约不产生额外并行度。scratch_types。为每个并行 grid 点声明临时内存plgpu.SMEM((tile_m, tile_n), out_dtype)是结果暂存用的 shared memoryplgpu.ACC((tile_m, tile_n), jnp.float32)是 TensorCore 累加器wgmma会异步累加到它上面。dict 里的每个 key 会作为同名关键字参数传给 kernel 函数。wgmma_wait(1)。wgmma是异步指令所有 WGMMA 操作按序执行可以理解为往队列里压操作plgpu.wgmma_wait(N)等待到 in-flight 的 WGMMA 不超过 N 个。这里 wait for 1意味着当前迭代发出的 WGMMA 会在下一次迭代才被等待保证 TensorCore pipeline 里始终有活干否则每次迭代都会 flush TensorCore pipeline。delay_release1。写在两个输入plgpu.BlockSpec上对应上面「若操作数上有未 await 的 WGMMA 就必须设置」的要求。没有它流水线会立即释放 SMEM 缓冲区下一次迭代覆盖数据时wgmma可能还在读产生静默数据竞争。Drain 阶段。K 维累加完成后把寄存器里的累加器转存到 SMEMo_smem[...] acc[...].astype(out_dtype)plgpu.commit_smem()让 SMEM 写入对 TMA 引擎可见再用plgpu.copy_smem_to_gmem异步写回 GMEM最后plgpu.wait_smem_to_gmem(0)等待全部拷贝完成。关于 transformwgmma要求操作数满足 CUDA 文档 中定义的特定 SMEM 布局通常由plgpu.TilingTransform((8, swizzle_elems))plgpu.SwizzleTransform(swizzle_bytes)组合实现swizzle_elems swizzle 字节数除以元素宽度。上例的 quickstart 版本没有显式传transforms而 Pipelining 指南的完整示例在in_specs上手动指定了它们并备注「未来 Mosaic GPU 会自动推断 transform届时无需手动指定」。如果你扩展这个 kernel 时遇到 wgmma 参数校验报错先对照 Mosaic GPU Reference 的 Hopperwgmma一节支持的形状要求M可被 64 整除、N可被 8 整除且不超过 256、K是swizzle // 元素宽度的倍数目前支持jnp.float32、jnp.bfloat16、jnp.float16和 FP8 类型累加器一般是jnp.float32。验证 kernel 输出正确用文档示例的输入规模m 132 * 128、n 4 * 128、k 10 * 64float16生成随机数据跑完 kernel 后与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)assert_allclose通过即为该示例的正确性判据注意这里输入维度都与默认 tile 尺寸整除m是 128 的倍数、k是tile_k64的倍数、n是 128 的倍数改成其他尺寸前先确认整除否则 grid 计算会静默丢块。调参max_concurrent_steps 与 delay_releasemax_concurrent_steps是流水线里最值得调的旋钮。Pipelining 指南给出的调参依据是值越大并发传输越多但每个额外并发 step 都要占用 SMEM 存放临时缓冲区较小的值例如 2有时能获得更高 occupancySMEM 占用低对 ALU 占比高的 kernel 可能反而提升吞吐代价是硬件调度带来更多噪声较大的值4 到 6最适合无法从额外 occupancy 中获益的 kernel典型如受 TensorCore 吞吐限制的 matmul。文档的结论是「We recommend autotuning this parameter」即围绕 2/4/6 实测选择而不是套用固定值。delay_release则不是性能旋钮而是正确性开关只要 pipeline 操作数上存在未 await 的wgmma本文示例就是这种模式就必须设为 1它同时会让流水线重叠的内存传输变少所以只在确有多个异步 matmul 需要同时在飞时才值得用。可选进阶warp specialization上面的 kernel 中TMA 拷贝GMEM/SMEM 搬运和矩阵乘法由同一条指令流发出。而索引计算和 TMA 发射本身很耗时可能让 TensorCore 空等。Hopper GPU 上可以把一部分 warpgroup 专职发 TMA、其余 warpgroup 专职计算用consumed barrier在两组 warpgroup 之间同步通知内存组何时可以发下一批 TMA。Pallas 中用plgpu.emit_pipeline_warp_specialized实现它处理全部内存线程逻辑用户只需写计算线程的工作API 与emit_pipeline类似特有参数引自 Pipelining 指南num_compute_wgs计算线程/warpgroup 数量。流水线发射器始终使用单个内存线程所以在plgpu.kernel里应设置num_threadsnum_compute_wgs1memory_registers分给内存线程的寄存器数其余寄存器在计算线程间均分。默认 40出现 register spill 时向上或向下调整wg_axis线程/warpgroup 轴的名字即plgpu.kernel的thread_name参数memory_thread_idx指定哪个 Pallas thread 作为内存线程默认最后一个compute_context定义只在计算线程里执行的 pipeline 前/后置逻辑并定义 loop carry 的初始化与消费。所有计算线程专属的数组都应在这里实例化避免内存线程在寄存器里物化它们否则会因 register spill 变慢。文档中的 warp-specialized matmul 示例用 2 个计算线程分别处理 RHS 的不同列、共享同一个 LHS每次 pipeline 调用计算输出矩阵的 2 个相邻 block。启动 kernel 的关键部分如下其中m、n、grid_m、grid_n、tile_m、tile_n是原示例顶部定义的矩阵尺寸与 grid 值kernel为示例中定义的 kernel 函数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)num_threads3对应num_compute_wgs2加 1 个内存线程。文档特别强调WGMMA 累加器必须在compute_thread函数内创建用compute_context模式如果在内存线程里分配会白白浪费寄存器每步的wgmma则包在pl.run_state里把 carry 值初始化为 accumulator ref。替代路径用 pl.pallas_call CompilerParams如果代码要同时兼容 Pallas TPU 后端可以用pl.pallas_call而不是emit_pipeline。Mosaic GPU 也实现了该 API默认情况下它只在 CUDA grid 上并行划分 kernel要开启流水线需传入plgpu.CompilerParams作为compiler_params参数其中dimension_semantics每个 grid 维是parallel划分到 CUDA grid还是sequential顺序流水的 tuple。注意如果没有任何维度标记为sequential就不会发生任何流水线max_concurrent_steps、delay_release与plgpu.emit_pipeline同名选项含义相同。流水线的另一个收益是允许在顺序迭代之间复用 scratch 缓冲区例如实现 reduction。pallas_call在 Mosaic GPU 后端下也接受plgpu.BlockSpec替代pl.BlockSpec从而可以指定 GPU 专用 transform。不过文档的推荐是优先使用plgpu.kernel因为它支持更多特性如指定 warpgroup 数量、warp specialization。限制与排查硬件边界本文主路径用的wgmma是 Hopper 特有的指令wgmma_wait配套Blackwell 改用tcgen05指令与 TMEM 内存空间写法不同参考 Blackwell Matrix Multiplication。Quickstart说明核心概念内存空间、grid、pipelining适用于所有受支持的 GPU 代际但具体 TensorCore 指令会换。GMEM 偏移对齐当 SMEM reference 上应用了plgpu.TilingTransform时GMEM↔SMEM 拷贝中 GMEM 侧的偏移必须与 tile 尺寸对齐否则传输可能产生错误结果见 Mosaic GPU Reference 的 note。register spillspill 会带来显著性能退化。编译期ptxas的消息中能看到 spill 警告设置环境变量MOSAIC_GPU_DUMP_PTXAS1可把这些日志打到标准输出。使用 warp specialization 时spill 是判断memory_registers该调大还是调小的依据。静默数据竞争delay_release缺失导致的竞争不会报错只能靠np.testing.assert_allclose(result, a b)这类数值对照暴露改动 pipeline 操作数的等待策略后先跑一遍数值验证再谈性能。延伸阅读通用平台无关的流水线概念推导、双缓冲展开过程Software Pipelining 教程emit_pipeline_warp_specialized、plgpu.kernel的完整示例与参数Mosaic GPU Pipelining内存空间、transform、Barrier、commit_smem的语义细节Mosaic GPU Reference【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考