news 2026/9/12 8:46:35

Mosaic GPU 如何用 emit_pipeline 给 Pallas kernel 写软件流水线,重叠计算与访存

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Mosaic GPU 如何用 emit_pipeline 给 Pallas kernel 写软件流水线,重叠计算与访存

Mosaic GPU 如何用 emit_pipeline 给 Pallas kernel 写软件流水线,重叠计算与访存

【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax

用 JAX Pallas 为 NVIDIA GPU 写 TensorCore kernel(矩阵乘法、attention 等)时,有一个绕不开的性能问题:GMEM(HBM)与 SMEM(shared 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(下文缩写为plgpu),kernel 通过plgpu.kernel启动,每个 Pallas thread 对应一个 warpgroup(4 个 warp = 128 个 CUDA thread),代码按 lockstep 执行,不需要管理单个 CUDA thread。

kernel 里的数据通过 Ref(JAX 的可变数组引用)访问,每个 Ref 位于特定内存空间:

  • GMEMplgpu.GMEM):全局内存/HBM,容量大、延迟高,kernel 的输入输出都在这里。GMEM Ref 不能直接用下标索引访问,只能通过plgpu.copy_gmem_to_smem/plgpu.copy_smem_to_gmememit_pipeline经 SMEM 中转。
  • SMEMplgpu.SMEM):SM 内的 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 np

emit_pipeline 的参数含义

在 Mosaic GPU Pipelining 指南中,推荐用plgpu.emit_pipeline顺序循环做流水线,同时用plgpu.kernel把问题在 CUDA grid 上并行划分。emit_pipeline的 API 与pl.pallas_call类似,但额外暴露了几个 GPU 专用选项:

  • bodygrid:语义同pl.pallas_callgrid表示body会运行多少次,与 CUDA grid 不同,pipeline grid 保证顺序执行。
  • in_specs/out_specs:同pl.pallas_call,但额外接受plgpu.BlockSpec实例,可以指定 GPU 专用的内存 reference transform(如 swizzling),transform 的完整说明见 Mosaic GPU Reference。
  • max_concurrent_steps:控制最大并发内存传输数。更大的值会消耗更多 SMEM 存放临时缓冲区,但可以提高内存子系统利用率。文档建议对该参数做 autotune。
  • delay_release:指定缓冲区被流水线复用前要额外等待的迭代数。例如delay_release=1max_concurrent_steps=2时,第 0 次迭代拷入 SMEM 的缓冲区要到第 3 次迭代才被复用(标准双缓冲是第 2 次)。如果你的 pipeline 操作数上还挂着未 await 的plgpu.wgmma,就必须设置delay_release=1,否则流水线会在 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.ACC,K 维全部累加完后把结果写回输出。

def matmul(a, b, tile_m=128, tile_n=128, tile_k=64, out_dtype=jnp.float16): m, k = a.shape _, n = b.shape @plgpu.kernel( out_type=jax.ShapeDtypeStruct((m, n), out_dtype), scratch_types=dict( o_smem=plgpu.SMEM((tile_m, tile_n), out_dtype), acc=plgpu.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_release=1 ), plgpu.BlockSpec( (tile_k, tile_n), lambda ki: (ki, pid_n), delay_release=1 ), ], max_concurrent_steps=2, )(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)

按执行顺序看这段代码里每个关键点的职责:

  1. 两级 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 内部的顺序归约,不产生额外并行度。
  2. scratch_types。为每个并行 grid 点声明临时内存:plgpu.SMEM((tile_m, tile_n), out_dtype)是结果暂存用的 shared memory,plgpu.ACC((tile_m, tile_n), jnp.float32)是 TensorCore 累加器,wgmma会异步累加到它上面。dict 里的每个 key 会作为同名关键字参数传给 kernel 函数。
  3. wgmma_wait(1)wgmma是异步指令,所有 WGMMA 操作按序执行,可以理解为往队列里压操作;plgpu.wgmma_wait(N)等待到 in-flight 的 WGMMA 不超过 N 个。这里 wait for 1,意味着当前迭代发出的 WGMMA 会在下一次迭代才被等待,保证 TensorCore pipeline 里始终有活干,否则每次迭代都会 flush TensorCore pipeline。
  4. delay_release=1。写在两个输入plgpu.BlockSpec上(对应上面「若操作数上有未 await 的 WGMMA 就必须设置」的要求)。没有它,流水线会立即释放 SMEM 缓冲区,下一次迭代覆盖数据时wgmma可能还在读,产生静默数据竞争。
  5. Drain 阶段。K 维累加完成后,把寄存器里的累加器转存到 SMEM(o_smem[...] = acc[...].astype(out_dtype)),plgpu.commit_smem()让 SMEM 写入对 TMA 引擎可见,再用plgpu.copy_smem_to_gmem异步写回 GMEM,最后plgpu.wait_smem_to_gmem(0)等待全部拷贝完成。

关于 transform:wgmma要求操作数满足 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、Kswizzle // 元素宽度的倍数;目前支持jnp.float32jnp.bfloat16jnp.float16和 FP8 类型,累加器一般是jnp.float32

验证 kernel 输出正确

用文档示例的输入规模(m = 132 * 128n = 4 * 128k = 10 * 64,float16)生成随机数据,跑完 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), dtype=jnp.float16) b = jax.random.uniform(key2, shape=(k, n), dtype=jnp.float16) result = matmul(a, b) np.testing.assert_allclose(result, a @ b)

assert_allclose通过即为该示例的正确性判据;注意这里输入维度都与默认 tile 尺寸整除(m是 128 的倍数、ktile_k=64的倍数、n是 128 的倍数),改成其他尺寸前先确认整除,否则 grid 计算会静默丢块。

调参:max_concurrent_steps 与 delay_release

max_concurrent_steps是流水线里最值得调的旋钮。Pipelining 指南给出的调参依据是:

  • 值越大,并发传输越多,但每个额外并发 step 都要占用 SMEM 存放临时缓冲区;
  • 较小的值(例如 2)有时能获得更高 occupancy(SMEM 占用低),对 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_threads=num_compute_wgs+1
  • memory_registers:分给内存线程的寄存器数,其余寄存器在计算线程间均分。默认 40,出现 register spill 时向上或向下调整;
  • wg_axis:线程/warpgroup 轴的名字,即plgpu.kernelthread_name参数;
  • memory_thread_idx:指定哪个 Pallas thread 作为内存线程,默认最后一个;
  • compute_context:定义只在计算线程里执行的 pipeline 前/后置逻辑,并定义 loop carry 的初始化与消费。所有计算线程专属的数组都应在这里实例化,避免内存线程在寄存器里物化它们(否则会因 register spill 变慢)。

文档中的 warp-specialized matmul 示例用 2 个计算线程分别处理 RHS 的不同列、共享同一个 LHS,每次 pipeline 调用计算输出矩阵的 2 个相邻 block。启动 kernel 的关键部分如下(其中mngrid_mgrid_ntile_mtile_n是原示例顶部定义的矩阵尺寸与 grid 值,kernel为示例中定义的 kernel 函数):

return plgpu.kernel( kernel, out_shape=jax.ShapeDtypeStruct((m, n), jnp.float16), scratch_shapes=dict( o_smem=plgpu.SMEM((tile_m, tile_n * 2), jnp.float16) ), grid=(grid_m, grid_n // 2), grid_names=("m", "n"), num_threads=3, # 2 compute, 1 memory. thread_name="wg" )(a, b)

num_threads=3对应num_compute_wgs=2加 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_stepsdelay_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 spill:spill 会带来显著性能退化。编译期ptxas的消息中能看到 spill 警告,设置环境变量MOSAIC_GPU_DUMP_PTXAS=1可把这些日志打到标准输出。使用 warp specialization 时,spill 是判断memory_registers该调大还是调小的依据。
  • 静默数据竞争delay_release缺失导致的竞争不会报错,只能靠np.testing.assert_allclose(result, a @ b)这类数值对照暴露;改动 pipeline 操作数的等待策略后,先跑一遍数值验证再谈性能。

延伸阅读

  • 通用(平台无关)的流水线概念推导、双缓冲展开过程:Software Pipelining 教程
  • emit_pipeline_warp_specializedplgpu.kernel的完整示例与参数:Mosaic GPU Pipelining
  • 内存空间、transform、Barrier、commit_smem的语义细节:Mosaic GPU Reference

【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/12 8:45:49

MIMO雷达DOA估计:波形正交性与虚拟阵列构建实战

简介:本资源是一套面向雷达信号处理初学者与进阶开发者的MIMO雷达波形设计与DOA估计MATLAB实现方案,聚焦于多输入多输出雷达系统中波形合成、频谱共享及到达角估计等核心问题,适用于通信与雷达交叉领域学习、课程设计或科研原型验证。压缩包仅…

作者头像 李华
网站建设 2026/9/12 8:43:09

CNN-LSTM-Attention实现Matlab时间序列预测与负荷回归

简介:一份基于卷积神经网络-长短期记忆网络结合注意力机制的多变量时间序列预测Matlab实现,涵盖CNN-LSTM-Attention、CNN-GRU-Attention、CNN-BILSTM-Attention三套可运行方案。资源面向需要完成课程设计、毕业设计或快速入门时序预测的在校学生与科研人…

作者头像 李华
网站建设 2026/9/12 8:42:59

几分钟免费把网页打包成应用:PakePlus桌面应用打包实战

几分钟免费把网页打包成应用:PakePlus桌面应用打包实战 【免费下载链接】PakePlus Turn any webpage/HTML/Vue/React and so on into desktop and mobile app under 5M with easy in few minutes. 轻松将任意网站/HTML/Vue/React等项目构建为轻量级(小于5M)多端桌面…

作者头像 李华
网站建设 2026/9/12 8:39:40

Mojo 内核自动调优结果分析:kprofile 与 tuning_codegen 实战指南

Mojo 内核自动调优结果分析:kprofile 与 tuning_codegen 实战指南 【免费下载链接】mojo The Modular Platform (includes MAX & Mojo) 项目地址: https://gitcode.com/GitHub_Trending/mo/mojo kprofile 与 tuning_codegen 是 Mojo/MAX 仓库中 max/kern…

作者头像 李华
网站建设 2026/9/12 8:38:35

153 本免费极客时间电子书:Python 核心教材直接拿走

153 本免费极客时间电子书:Python 核心教材直接拿走 【免费下载链接】geektime-books :books: 极客时间电子书 项目地址: https://gitcode.com/GitHub_Trending/ge/geektime-books 找 Python 免费电子书,网盘链接死一片、广告夹一堆,翻…

作者头像 李华
网站建设 2026/9/12 8:36:27

低功耗开发实战:从寄存器配置到系统级功耗治理

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华