JAX Pallas Mosaic GPU 快速上手:从 GPU 内存空间到 Tensor Core 流水线矩阵乘法
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
本篇快速上手指南围绕 JAX 实验性 Pallas 扩展中的Mosaic GPU后端展开,教你如何在 NVIDIA GPU(示例以 Hopper/H100 为目标,但内存空间、网格(grid)与流水线(pipelining)等核心概念适用于所有受支持的 GPU 代际)上直接编写内核(kernel)。读完本文,你将掌握 Pallas 在 GPU 上的编程模型(warpgroup 抽象)、GMEM/SMEM/TMEM 三种内存空间的正确用法,并能从零写出两个可运行的示例——一个填充常量的极简内核和一个完整覆盖 Tensor Core 矩阵乘法的流水线内核。
Mosaic GPU 的编程模型:以 warpgroup 为单位
在进入代码之前,先理解 Mosaic GPU 最核心的抽象:Pallas 的每个 "thread" 对应一个 warpgroup(即 4 个 warp,共 128 条 CUDA 线程)。你编写的是一段直接操作数组的直线式(straight-line)代码,整个 warpgroup 以锁步(lockstep)方式同步执行,因此不需要像裸 CUDA 那样管理单独的线程——没有 threadIdx 的显式分支,也没有手工的线程协作。
这一点与 Triton 的编程模型有显著差异:在 Triton 中流水线(pipelining)是编译器自动完成的优化;而在 Pallas/Mosaic GPU 中,流水线是显式编程的(详见 Mosaic GPU Pipelining)。这意味着你能精确控制数据搬运与计算的重叠方式,代价是必须自己理解内存空间和异步指令。
本文所有示例都需要如下导入:
import jax import jax.numpy as jnp from jax.experimental import pallas as pl from jax.experimental.pallas import mosaic_gpu as plgpuplgpu是 Mosaic GPU 后端的 Python 入口,从源码看,它从 jax/_src/pallas/mosaic_gpu/core.py 与 jax/_src/pallas/mosaic_gpu/pipeline.py 等模块统一导出了kernel、emit_pipeline、BlockSpec、SwizzleTransform、TilingTransform、wgmma、commit_smem等全套 API(见 jax/experimental/pallas/mosaic_gpu.py)。该模块在文件头也明确标注:这些 API 高度不稳定,可能每周变动,使用时需自担风险。
GPU 内存空间:GMEM、SMEM 与 TMEM
Pallas 内核通过Ref(JAX 的可变数组引用)访问内存。在 GPU 上,每个 Ref 都隶属于一个特定的内存空间(memory space):
| 内存空间 | 全称 | 容量/速度 | 用途 |
|---|---|---|---|
| GMEM | Global Memory / HBM | 大、慢 | 内核输入与输出 |
| SMEM | Shared Memory | 小、快、每个 SM 独占 | 块内线程共享,用于为 Tensor Core 运算暂存数据 |
| TMEM | Tensor Memory | 快、每个 SM 独占 | Tensor Core 运算专用(仅 Blackwell 及之后代际可用) |
在 jax/_src/pallas/mosaic_gpu/core.py 中,MemorySpace枚举还额外定义了第 4 个成员REGS(寄存器),注释明确指出TMEM 是 Blackwell 新增的成员,Hopper 上不可用。内核中所有标量/数组值(即 JAX 数组)默认都位于寄存器中;若编译器寄存器耗尽,会插入 spill(反复存取),导致性能下降。
典型的数据流决定了内核的写法:
- Tensor Core 工作负载(Hopper):
GMEM → SMEM → Tensor Cores → registers → SMEM → GMEM;在 Blackwell 上,TMEM 会取代 SMEM 承担 Tensor Core 输入/输出的暂存角色。 - 纯 ALU 工作负载(如逐元素运算):完全绕过 SMEM 与 Tensor Core,走
GMEM → registers → GMEM即可。只有块内线程需要交换或复用数据时,才需要显式使用共享内存。
关于显式分配:SMEM 与 TMEM 可以通过plgpu.kernel的scratch_shapes参数分配,也可以用pl.run_scoped在作用域内分配;直接调用内存空间对象即可,例如plgpu.SMEM((128, 128), jnp.float16)会在共享内存中分配一个 128×128 的 float16 数组。如果要让某个BlockSpec显式访问 GMEM,可以设置BlockSpec(memory_space=plgpu.GPUMemorySpace.GMEM)(详见 Mosaic GPU Reference)。
你的第一个内核:填充常量
最简单的内核是向输出数组填充一个常量:
@plgpu.kernel(out_type=jax.ShapeDtypeStruct((128,), jnp.float32)) def fill_42(o_ref): o_ref[...] = jnp.full_like(o_ref, 42.0) result = fill_42() # [42.0, 42.0, ...]plgpu.kernel会替你完成三件事:分配输出缓冲区 → 在设备上运行计算 → 返回一个 JAX 数组。装饰器中的out_type用jax.ShapeDtypeStruct声明输出的形状与 dtype;内核函数接收的参数o_ref就是输出 Ref,对其整体赋值即完成写入。
注意这里没有声明任何 grid——单个 warpgroup 顺序执行完整个 128 元素数组即可。这种写法适合小规模、单块的运算。
用 grid 并行处理大数组
要处理更大的数组,需要引入grid(网格)。每个 grid 点会作为一个独立的CUDA block并行运行在不同的 SM(流式多处理器)上:
@plgpu.kernel( out_type=jax.ShapeDtypeStruct((1024,), jnp.float32), grid=(8,), grid_names=('i',), ) def iota(o_ref): i = jax.lax.axis_index('i') o_ref[pl.ds(i * 128, 128)] = jnp.arange(128, dtype=jnp.float32) + i * 128 result = iota() # [0.0, 1.0, ..., 1023.0]这里有两个关键 API:
jax.lax.axis_index(name):返回当前 grid 块在该轴上的编号,用于确定本块负责输出数组的哪一段。grid=(8,)声明了 8 个并行块,grid_names=('i',)给这个轴起名i,与axis_index('i')对应。pl.ds(start, size):构造一个大小为size的动态切片(dynamic slice),等价于start:start+size的索引写法。本例中第i块负责写入i*128 : i*128+128区间,8 个块合起来恰好覆盖 1024 个元素。
需要强调的是,grid上的并行块之间没有顺序保证,因此每个块必须通过axis_index自行定位自己负责的输出区间,这是所有基于 grid 的 Mosaic GPU 内核的基本模式。
为什么需要流水线:Tensor Core 的饥饿问题
上面两个例子都不涉及流水线。但任何命中 Tensor Core 的运算——矩阵乘法、注意力等——都必须把 GMEM↔SMEM 的数据搬运与计算重叠起来。原因很直接:如果不重叠,Tensor Core 会在等待数据到达期间完全闲置。
plgpu.emit_pipeline正是为此设计的,它接收三部分:
- sequential grid:要执行的流水线步数(通常沿收缩维,即 K 维);
BlockSpecs:描述每一步如何从输入中切片出所需的数据块;- body 函数:每一步要执行的计算。
整体分工是:外层的plgpu.kernelgrid 负责并行(把输出切块、每个 CUDA block 算一块),emit_pipeline负责块内的顺序归约(沿 K 维迭代累加)。
与 Triton 中编译器自动插入多级缓冲不同,emit_pipeline的所有参数都是显式的。从 jax/_src/pallas/mosaic_gpu/pipeline.py 的源码可以看到其校验逻辑:grid的所有维度必须严格为正;max_concurrent_steps必须大于所有BlockSpec的delay_release值,否则直接抛出ValueError。源码还会在max_concurrent_steps大于总步数时自动将其收缩到总步数,以避免过度分配 SMEM 缓冲。
Hopper 上的矩阵乘法内核
下面是针对 Hopper GPU 的完整 matmul 内核。它使用wgmma(warpgroup matrix multiply accumulate)指令,该指令由单个 Mosaic GPU thread 发出、在 Tensor Core 上异步执行:
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)注意:
wgmma是 Hopper 专用指令。Blackwell 用户应改用tcgen05指令,参见 Blackwell Matrix Multiplication。
逐段拆解这个内核:
并行网格(parallel grid)。plgpu.kernel(..., grid=(m // tile_m, n // tile_n))把输出[M, N]切成tile_m × tile_n的块,每个输出块对应一个 CUDA block,并行地在不同 SM 上执行。pid_m/pid_n通过jax.lax.axis_index取得当前块的行列编号。
顺序网格(sequential grid)。emit_pipeline(..., grid=(k // tile_k,))是沿 K 维的流水线循环。每一步从两个输入中各加载一个tile_k宽的切片(BlockSpec的(block_shape, index_map)分别声明块形状与切片位置),执行一次wgmma累加。
scratch_types。它声明每个并行 grid 点所需的临时内存分配,字典中的每个 key 会作为关键字参数传入内核函数。本例分配了两块:
o_smem:位于 SMEM 的输出暂存缓冲区;acc:plgpu.ACC,即Tensor Core 累加器。wgmma异步地向其中累加结果;它通常驻留在寄存器中,是 Mosaic GPU 特有的 Ref 类型(源码中即WGMMAAccumulatorRef,见 jax/experimental/pallas/mosaic_gpu.py 中ACC的别名导出)。注意累加器使用float32,即使输入输出都是 float16——这是 matmul 精度稳定性的关键。
delay_release=1。告诉流水线额外多保留一个缓冲的生命周期。如果不设置,流水线会在某次迭代的输入块使用完毕后立刻释放缓冲,下一次迭代就可能覆盖这块数据——而此时wgmma可能还在异步读取它,从而引发静默数据竞争。结合plgpu.wgmma_wait(1)(等待在途的 wgmma 数量不超过 1,即当前迭代发出的 wgmma 将在下一轮被等待),可以始终保留一个 wgmma 在飞行中,保持 Tensor Core 满载。正如 Mosaic GPU Pipelining 中强调的:省略delay_release会产生静默数据竞争,务必小心使用。
收尾(drain)阶段。流水线结束后,累加器里是完整的输出块:
o_smem[...] = acc[...].astype(out_dtype):把累加结果从寄存器写入 SMEM(同时从 float32 转回 float16);plgpu.commit_smem():让 SMEM 写入对 TMA(Tensor Memory Accelerator)引擎可见;plgpu.copy_smem_to_gmem(...):用pl.ds构造的目标切片把 SMEM 数据异步拷回 GMEM 中当前块负责的区域;plgpu.wait_smem_to_gmem(0):等待所有拷贝完成后再退出内核。
流水线示意如上图:外层 grid 将输出划分为多个块并行计算,每个块内部沿 K 维顺序迭代,每一步的 TMA 数据搬运(GMEM→SMEM)与上一步的wgmma计算相重叠。这种"搬运与计算重叠"正是emit_pipeline的价值所在——异步的 GMEM/SMEM 拷贝延迟很长,而 Tensor Core 计算必须基于寄存器或 SMEM 中的 Ref,两者不同步重叠就会互相等待(详见 Mosaic GPU Pipelining)。
深入emit_pipeline的关键参数
从 jax/_src/pallas/mosaic_gpu/pipeline.py 的emit_pipeline签名可以看到,它支持的参数包括body、grid、in_specs、out_specs、max_concurrent_steps(默认 1)与init_carry。结合 Mosaic GPU Pipelining 与 CompilerParams 的说明,两个最值得调优的参数是:
max_concurrent_steps:控制并发内存传输的最大数量。
- 增大该值会占用更多 SMEM 存放临时缓冲,但能提高内存子系统的利用率;
- 官方建议对该参数做自动调优(autotune);
- 较小值(如 2)由于 SMEM 占用低,可能获得更高 occupancy,对 ALU 密集内核的吞吐有利,但因硬件调度会产生更多噪声;
- 较大值(4~6)最适合无法从额外 occupancy 中获益的内核。
delay_release:延迟缓冲复用。
- 以
max_concurrent_steps=2、delay_release=1为例:第 0 次迭代拷入 SMEM 的缓冲要到第 3 次迭代才会被复用,而标准的双缓冲策略在第 2 次迭代就会复用; - 当你在 body 中不等待
plgpu.wgmma(即不做wgmma_wait)时,delay_release=1是必需的——否则流水线会在 WGMMA 仍在读取时就开始覆盖缓冲; - 该技巧常用于让多个异步 matmul 同时在飞行中以喂满 Tensor Core 流水线,但代价是重叠的传输变少,
emit_pipeline的效率会下降。
兼容 API:通过pl.pallas_call使用流水线
为了与 Pallas TPU 保持兼容,Mosaic GPU 也实现了既有的pl.pallas_callAPI。默认情况下,Mosaic GPU 上的pl.pallas_call会把内核沿 CUDA grid 并行切分;要开启流水线,需要传入一个plgpu.CompilerParams对象作为compiler_params参数,其中与流水线相关的选项是:
dimension_semantics:一个由'parallel'/'sequential'组成的元组,声明每个 grid 维度的迭代语义。'parallel'维被切分到 CUDA grid 上,'sequential'维被顺序流水线化。注意:如果没有维度被标记为'sequential',就不会发生任何流水线化!max_concurrent_steps:与emit_pipeline中的同名参数一致。delay_release:与emit_pipeline中的同名参数一致。
从 jax/_src/pallas/mosaic_gpu/core.py 的CompilerParams定义看,它还包含approx_math(允许使用近似数学实现,默认 False)、unsafe_no_auto_barriers(关闭自动插入的 barrier,需满足严格条件才安全)、reduction_scratch_bytes(跨 warp 归约预留的 SMEM 字节数,H100/B200 上2*128*6*4=6144字节通常是较好的取值)、skip_device_barrier(跳过内核启动前的跨设备 barrier,误用会导致竞争)等参数,且校验规则要求profile_space与profile_dir必须同时设置或同时不设置。
不过官方文档明确建议:优先使用plgpu.kernel而非pl.pallas_call,因为plgpu.kernel支持更多特性——例如指定 warpgroup 数量(num_threads)与 warp 特化(详见 Mosaic GPU Pipelining 中的emit_pipeline_warp_specialized与compute_context用法)。两种方式下,pallas_call/emit_pipeline都支持使用plgpu.BlockSpec代替pl.BlockSpec,从而指定 GPU 特有的内存变换(如TilingTransform与SwizzleTransform,用于把 SMEM 数据排布成wgmma要求的寄存器片段与共享内存矩阵布局)。
与仓库实现的对应关系
本文涉及的核心 API 均可在仓库源码中找到实现:
plgpu.kernel、plgpu.ACC、BlockSpec、CompilerParams、MemorySpace(含 GMEM/SMEM/TMEM/REGS 四成员)定义于 jax/_src/pallas/mosaic_gpu/core.py;emit_pipeline(含max_concurrent_steps与delay_release的约束校验、SMEM 缓冲自动收缩逻辑)实现于 jax/_src/pallas/mosaic_gpu/pipeline.py;wgmma、wgmma_wait、commit_smem、copy_smem_to_gmem、wait_smem_to_gmem等 GPU 原语导出自 jax/_src/pallas/mosaic_gpu/primitives.py;- 全部 GPU 专用 API 的公开入口在 jax/experimental/pallas/mosaic_gpu.py,其中
GMEM、SMEM、TMEM、REGS是MemorySpace成员的便捷别名。
仓库中的 GPU 测试(如 tests/pallas/mosaic_gpu_test.py、tests/pallas/mgpu_examples_test.py)大量使用了emit_pipeline、plgpu.kernel与wgmma,是查看真实用法与边界行为的参考;pipelining.md中还提供了带np.testing.assert_allclose(result, a @ b)校验的完整可运行示例,可直接对照验证。
下一步
- Mosaic GPU Pipelining——流水线深入讲解,包括 warp 特化(
emit_pipeline_warp_specialized、num_compute_wgs、memory_registers、wg_axis等参数); - Mosaic GPU Reference——完整 API 参考(内存空间、布局、Tensor Core 运算、内存引用变换);
- Blackwell Matrix Multiplication——使用
tcgen05指令的 Blackwell 矩阵乘法; - Collective Matrix Multiplication——GPU 集合通信矩阵乘法。
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考