
JAX Pallas Mosaic GPU 快速上手从 GPU 内存空间到 Tensor Core 流水线矩阵乘法【免费下载链接】jaxComposable transformations of PythonNumPy 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 与 TMEMPallas 内核通过RefJAX 的可变数组引用访问内存。在 GPU 上每个 Ref 都隶属于一个特定的内存空间memory space内存空间全称容量/速度用途GMEMGlobal Memory / HBM大、慢内核输入与输出SMEMShared Memory小、快、每个 SM 独占块内线程共享用于为 Tensor Core 运算暂存数据TMEMTensor Memory快、每个 SM 独占Tensor Core 运算专用仅 Blackwell 及之后代际可用在 jax/_src/pallas/mosaic_gpu/core.py 中MemorySpace枚举还额外定义了第 4 个成员REGS寄存器注释明确指出TMEM 是 Blackwell 新增的成员Hopper 上不可用。内核中所有标量/数组值即 JAX 数组默认都位于寄存器中若编译器寄存器耗尽会插入 spill反复存取导致性能下降。典型的数据流决定了内核的写法Tensor Core 工作负载HopperGMEM → 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_spaceplgpu.GPUMemorySpace.GMEM)详见 Mosaic GPU Reference。你的第一个内核填充常量最简单的内核是向输出数组填充一个常量plgpu.kernel(out_typejax.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_typejax.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, dtypejnp.float32) i * 128 result iota() # [0.0, 1.0, ..., 1023.0]这里有两个关键 APIjax.lax.axis_index(name)返回当前 grid 块在该轴上的编号用于确定本块负责输出数组的哪一段。grid(8,)声明了 8 个并行块grid_names(i,)给这个轴起名i与axis_index(i)对应。pl.ds(start, size)构造一个大小为size的动态切片dynamic slice等价于start:startsize的索引写法。本例中第i块负责写入i*128 : i*128128区间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 内核。它使用wgmmawarpgroup matrix multiply accumulate指令该指令由单个 Mosaic GPU thread 发出、在 Tensor Core 上异步执行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)注意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 的输出暂存缓冲区accplgpu.ACC即Tensor Core 累加器。wgmma异步地向其中累加结果它通常驻留在寄存器中是 Mosaic GPU 特有的 Ref 类型源码中即WGMMAAccumulatorRef见 jax/experimental/pallas/mosaic_gpu.py 中ACC的别名导出。注意累加器使用float32即使输入输出都是 float16——这是 matmul 精度稳定性的关键。delay_release1。告诉流水线额外多保留一个缓冲的生命周期。如果不设置流水线会在某次迭代的输入块使用完毕后立刻释放缓冲下一次迭代就可能覆盖这块数据——而此时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 转回 float16plgpu.commit_smem()让 SMEM 写入对 TMATensor 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 密集内核的吞吐有利但因硬件调度会产生更多噪声较大值46最适合无法从额外 occupancy 中获益的内核。delay_release延迟缓冲复用。以max_concurrent_steps2、delay_release1为例第 0 次迭代拷入 SMEM 的缓冲要到第 3 次迭代才会被复用而标准的双缓冲策略在第 2 次迭代就会复用当你在 body 中不等待plgpu.wgmma即不做wgmma_wait时delay_release1是必需的——否则流水线会在 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*46144字节通常是较好的取值、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.pyemit_pipeline含max_concurrent_steps与delay_release的约束校验、SMEM 缓冲自动收缩逻辑实现于 jax/_src/pallas/mosaic_gpu/pipeline.pywgmma、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 PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考