JAX Pallas 软件流水线(Software Pipelining)完全指南:通信-计算重叠的原理、API 与实战陷阱
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
本文以 docs/pallas/pipelining.md 为核心骨架展开,系统讲解 JAX Pallas 中软件流水线的概念基础、双缓冲手工推导、
pl.pallas_call流水线 API、分块内核写法,以及缓冲重用与归约累加两大易错点,并结合仓库源码(jax/_src/pallas/core.py、jax/_src/pallas/mosaic/pipeline.py、jax/_src/pallas/mosaic_gpu/pipeline.py)与平台专属文档(docs/pallas/tpu/pipelining.md、docs/pallas/gpu/pipelining.md)做纵深补充。读完本文,你将掌握:内存层级与带宽瓶颈的直觉模型、如何手工推导双缓冲流水线、如何使用 Pallas 的grid/BlockSpec/kernel三要素编写可重叠通信与计算的流水线内核,以及如何避开缓冲重访与归约初始化这两类"看起来对、结果错"的经典陷阱。
1. 为什么需要软件流水线
软件流水线(Software Pipelining)是性能优化中的一项重要技术:即使操作之间存在数据依赖,也可以通过重叠多个异步操作来隐藏延迟。在编写内核(kernel)的语境下,最常见的形式是让通信与内存搬移和计算相互重叠,从而让硬件加速器在等待数据到达时不再空转。本教程聚焦于"通信-计算流水线"(communication-compute pipelining)这一类问题:先建立概念模型,再介绍 Pallas 的流水线 API,最后给出若干真实可运行的示例。
本文只覆盖流水线的概念基础;平台专属的细节可参考 TPU 流水线参考 与 Mosaic GPU 流水线参考。
2. 先理解内存层级(Memory Hierarchies)
理解流水线的前提,是先弄清楚加速器上不同类型的存储空间及其容量、延迟/带宽之间的权衡。大多数硬件架构(包括 CPU、GPU 和 TPU)都提供了多种存储空间,在容量与延迟/带宽之间做出取舍。对 Pallas 而言,我们通常关心四类存储:
- 寄存器(Registers):物理上离处理器最近的内存。任何计算执行之前,值通常必须先加载到寄存器中。
- SRAM(在 GPU 上称为共享内存/Shared Memory、L1/L2 缓存,在 TPU 上称为 VMEM):同样离处理器较近,但容量比寄存器大。现代 ML 加速器的 SRAM 通常在 10–100 MB 量级(例如 TPU v5p 拥有 96 MB VMEM,H100 GPU 拥有约 30 MB L1 缓存与 50 MB L2 缓存)。访问 SRAM 的延迟大约是访问寄存器延迟的 10 倍量级。
- DRAM(又称 HBM):容量远大于 SRAM,现代 ML 加速器通常有 10–100 GB 量级。访问延迟大约比 SRAM 再高 10 倍量级。
- 网络(Network)通信:当单个设备的 DRAM 容量不足,或需要利用并行计算时,网络通信变得关键。本文不涉及分布式流水线,跨设备写流水线可参考多设备分布式 TPU 内核指南。
2.1 一次完整的 HBM → 计算 → HBM 数据流
要对存放在 HBM 中的值X、Y执行计算,硬件需要依次完成:
- 把
x和y从 HBM 拷贝到 SRAM; - 把值从 SRAM 加载到寄存器;
- 执行计算并把结果存入寄存器;
- 把输出寄存器中的值写回 SRAM;
- 把 SRAM 中的输出值拷回 HBM。
下面就是一个忠实地完成上述流程的 Pallas 函数(注意:这是 TPU 示例):
def add_matrices_kernel(x_sram_ref, y_sram_ref, z_sram_ref): # Load x and y from SRAM into registers x_regs = x_sram_ref[:, :] y_regs = y_sram_ref[:, :] # Execute a vectorized add z_regs = x_regs + y_regs # Store the output values in registers back into SRAM z_sram_ref[:, :] = z_regs def add_matrices(x: jax.Array, y: jax.Array) -> jax.Array: # pallas_call will first allocate scratch buffers for `x` and `y` in SRAM. # It will then copy `x` and `y` from HBM into SRAM. z = pl.pallas_call( add_matrices_kernel, out_shape=jax.ShapeDtypeStruct.like(x) )(x, y) # pallas_call will also copy the output from SRAM back into HBM. return z x, y = jnp.ones((512, 512)), jnp.ones((512, 512)) add_matrices(x, y)这里定义了两个函数:
add_matrices_kernel操作的是位于 SRAM 中的Ref。从 SRAMRef加载产生的是寄存器中的值;寄存器中的值行为类似jax.Array,可以对其使用jnp和jax.lax运算产生新的寄存器值;当需要返回结果时,将值存入输出的 SRAMRef。add_matrices操作的是jax.Array。它把x、y传入pallas_call:pallas_call负责把x、y拷贝进 SRAM,并分配内核操作所需的 SRAM 缓冲(包括输出缓冲);内核执行完毕后,pallas_call再把输出缓冲中的值拷回 HBM,得到输出jax.Array。
2.2 两个必须正视的约束:容量与带宽
Pallas 暴露了 SRAM 等底层存储空间,但要写出高性能内核,必须更精细地利用各类存储,尤其要同时考虑:
- 内存容量(Memory capacity):SRAM 很小!如果数组太大,上面的内核根本无法运行,因为输入放不进 SRAM。作为参考,一个
f32[2048, 2048]的数组就有 16 MiB,因此上述朴素内核只能处理中等规模以下的数组。 - 内存带宽(Memory bandwidth):在 HBM 与 SRAM 之间拷贝很耗时,至少比绝大多数计算指令慢得多。上面的
add_matrices很可能把大部分时间花在 HBM↔SRAM 的拷贝上,而不是加法本身。
带着这两个约束,我们需要重新思考如何榨取加速器的性能——这正是流水线的用武之地。
3. 流水线基础:把大问题切小并重叠
如何既利用内存层级中各类存储的优势,又能操作存放在 HBM 中的大数组、同时用快速的 SRAM 做计算?流水线是一种非常通用的编程模式,它要求把问题拆成可以并行重叠的更小子问题。
流水线的第一步,是把问题划分成能放进 SRAM 的小子问题。以逐元素(elementwise)运算为例,可以简单地对源数组每次处理一个切片,得到如下 3 个步骤(又称 3 个阶段 / stages):
- copy_in:把切片
A[i]从 HBM 拷入 SRAMX; - compute:把
X加载到寄存器,计算结果并存回 SRAMY; - copy_out:把结果
Y拷回 HBM 的A[i]。
注意步骤 1–3 之间存在数据依赖,必须先完成步骤 1 才能开始步骤 2,因此不能简单重叠。然而不同子问题实例之间没有数据依赖——也就是说,可以在执行块A[i+1]的步骤 1 的同时,执行块A[i]的步骤 2 和块A[i-1]的步骤 3。
上图描绘了一个理想化的流水线程序如何随时间调度。关键洞察是:在内核运行的大部分时间里,拷贝操作与计算操作并行执行,从而可以用计算"隐藏" HBM/SRAM 之间的搬移开销,让处理器保持尽可能高的占用率。
调度图两端各有一段启动(startup)与收尾(teardown)时间,称为"气泡"(bubbles)——此时流水线正在"填充"或"排空",只有部分阶段在执行。绝大部分时间花在流水线的稳态阶段(steady-state),此时每个流水线阶段都在不同子问题迭代上并行执行。更通用的流水线目标是在 N 个阶段上实现 N 路并行;但对内核流水线而言,瓶颈通常是内存带宽或处理速度,因此目标往往是实现处理器 FLOP/s 的完全利用——即任意时刻总有一个compute块在执行。上图中 compute 块在 8 个时隙中活跃了 6 个,假设每个计算时隙处理器都被完全利用,则实现了 75% 的处理器利用率。
4. 手工推导一个双缓冲(Double-Buffered)流水线
先看一段伪代码形式的逐元素程序:从 HBM 加载A[i](copy_in),加 1 后把结果写回 HBM(copy_out):
for i in range(N): copy_in(A[i], X) Y = X + 1 copy_out(Y, A[i])问题在于copy_in和copy_out通常是阻塞操作:GPU/TPU 在等待拷贝完成时空闲,然后内存又空闲着等计算。我们希望"预取"(pre-fetch)下一次循环迭代所需的输入,在当前迭代执行计算的同时异步发起拷贝,让计算与内存通信同时发生。
为了推演这个代码变换,先把循环按 N=4 展开,并把拷贝指令拆成copy_start(发起异步拷贝)与copy_wait(等待拷贝完成)两部分来表达异步性:
# Itr 1 copy_in_start(A[0], X) copy_in_wait(X) Y = X + 1 copy_out_start(Y, A[0]) copy_out_wait(Y) # Itr 2 copy_in_start(A[1], X) copy_in_wait(X) Y = X + 1 copy_out_start(Y, A[1]) copy_out_wait(Y) # Itr 3 copy_in_start(A[2], X) copy_in_wait(X) Y = X + 1 copy_out_start(Y, A[2]) copy_out_wait(Y) # Itr 4 copy_in_start(A[3], X) copy_in_wait(X) Y = X + 1 copy_out_start(Y, A[3]) copy_out_wait(Y)展开之后,流水线变换的本质就清晰了:尽可能早地发出copy_start,尽可能晚地执行copy_wait(恰好在使用该值之前)。但当前循环状态对X存在一个"假数据依赖"——不能在异步拷贝数据进X的同时又用X做计算,否则可能产生竞态(race condition)。因此引入**多缓冲(multiple-buffering)**技术:为每个输入X和每个输出Y各保留 2 个缓冲。有了 2 个缓冲,可以把copy_in_start提前一个迭代(3 个缓冲则可以提前 2 个迭代,依此类推),循环被改写为:
# Prologue copy_in_start(A[0], X[0]) # Itr 1 copy_in_start(A[1], X[1]) copy_in_wait(X[0]) Y[0] = X[0] + 1 copy_out_start(Y[0], A[0]) copy_out_wait(Y[0]) # Itr 2 - Steady state copy_in_start(A[2], X[0]) copy_in_wait(X[1]) Y[1] = X[1] + 1 copy_out_start(Y[1], A[1]) copy_out_wait(Y[1]) # Itr 3 - Steady state copy_in_start(A[3], X[1]) copy_in_wait(X[0]) Y[0] = X[0] + 1 copy_out_start(Y[0], A[2]) copy_out_wait(Y[0]) # Itr 4 - No copy-in copy_in_wait(X[1]) Y[1] = X[1] + 1 copy_out_start(Y[1], A[3]) copy_out_wait(Y[1])接下来,把copy_out_wait尽量推迟——推迟到下一次循环迭代写Y之前:
# Prologue copy_in_start(A[0], X[0]) # Itr 1 copy_in_start(A[1], X[1]) copy_in_wait(X[0]) Y[0] = X[0] + 1 copy_out_start(Y[0], A[0]) # Itr 2 - Steady state copy_in_start(A[2], X[0]) copy_in_wait(X[1]) Y[1] = X[1] + 1 copy_out_start(Y[1], A[1]) copy_out_wait(Y[0]) # 推迟到此 # Itr 3 - Steady state copy_in_start(A[3], X[1]) copy_in_wait(X[0]) Y[0] = X[0] + 1 copy_out_start(Y[0], A[2]) copy_out_wait(Y[1]) # 推迟到此 # Itr 4 - No copy-in copy_in_wait(X[1]) Y[1] = X[1] + 1 copy_out_start(Y[1], A[3]) copy_out_wait(Y[0]) # 推迟到此 # Epilogue copy_out_wait(Y[1]) # 排空最后把循环重新卷回for循环,就得到下面的流水线化循环:
# Prologue copy_in_start(A[0], X[0]) # Main loop for i in range(N): cur_slot = i % 2 next_slot = (i + 1) % 2 if i+1 < N: copy_in_start(A[i+1], X[next_slot]) copy_in_wait(X[cur_slot]) Y[cur_slot] = X[cur_slot] + 1 copy_out_start(Y[cur_slot], A[i]) if i > 0: copy_out_wait(Y[next_slot]) # Epilogue copy_out_wait(Y[1])4.1 泛化:流水线的三要素
若要把上述循环推广到更广泛的计算,本质上需要向流水线指定 3 条信息:
- grid(网格):
for循环的边界,指明子问题的个数。本例中是大小为(N,)的一维网格。 - kernel(内核):输入加载到 SRAM 后真正执行的计算。本例中是逐元素加法
Y = X + 1。 - data_slices(数据切片):把子问题映射到 HBM 缓冲中相应切片的规则。本例中数据切片是恒等函数
lambda i: i。
只要用户能指定这三者,就可以按照该模式写出各种各样的程序:
def double_buffered_pipeline( grid: tuple[int, ...], kernel: Callable, in_slices: Callable, out_slices: Callable): # Prologue copy_in_start(in_hbm[in_slices(0)], in_sram[0]) # Main loop grid_size = prod(grid) for i in range(grid_size): cur_slot = i % 2 next_slot = (i + 1) % 2 if (i + 1) < grid_size: copy_in_start(in_hbm[in_slices(i+1)], in_sram[next_slot]) copy_in_wait(in_sram[cur_slot]) kernel(in_sram[cur_slot], out_sram[cur_slot]) copy_out_start(out_sram[cur_slot], out_hbm[out_slices(i)]) if i > 0: copy_out_wait(out_sram[next_slot]) # Epilogue last_slot = (grid_size - 1) % 2 copy_out_wait(out_sram[last_slot])至此我们看到了如何手工实现一个流水线循环。接下来看看如何使用 Pallas 现成的 API——它把"维护多个缓冲、重叠异步通信与计算"的样板代码都抽象掉了。
5. Pallas 流水线 API
Pallas 提供了一套流水线 API,把维护多缓冲、重叠异步通信与计算的样板代码抽象出来。API 的基础知识在 Pallas 快速入门 中已有覆盖,这里简要回顾以保持完整性,并重点讨论流水线带来的几个"锋利的边角"(sharp edges)。
5.1 Grid(网格)
程序grid是一个整数元组,按数组的方式指明子问题的个数。流水线的结构可以理解为一个嵌套for循环,循环边界即 grid 的每个分量:
# For grid (N, M, K) for n in range (N): for m in range(M): for k in range(K): kernel()内核总共会被调用prod(grid)次。更详细的说明见 grid 与 blockspec 文档。
5.2 BlockSpecs(块规格)
BlockSpec指明每次子问题迭代要拷贝的数据块的大小与切片。pl.BlockSpec的基本构造参数是:
block_shape:一个数据切片的大小;index_map:接收当前子问题的 program id,输出源缓冲的分块索引(blocked indices)。分块索引指明每次迭代拷贝哪个块——假设源缓冲已被按block_shape切分成若干块;memory_space:指定输入被拷贝到哪种存储空间,默认是 SRAM。
pl.BlockSpec( block_shape: tuple[int, ...], index_map: Callable, memory_space: pl.MemorySpace )内核的每个输入和每个输出都各需要一个BlockSpec。
从源码看,BlockSpec定义在 jax/_src/pallas/core.py#L548,字段为block_shape、index_map、memory_space与pipeline_mode。其中block_shape除了int | None,还支持更精细的BlockDim类型(如pl.Element、pl.Squeezed、pl.Blocked、pl.BoundedSlice、pl.Indirect):None表示该维度被 squeeze 掉、不出现在内核里;pl.BoundedSlice(定义于 jax/_src/pallas/core.py#L415)则允许对某维度指定"有界但动态"的切片大小(详见第 9 节)。memory_space使用pl.MemorySpace枚举(jax/_src/pallas/core.py#L271),包含ANY(不限定,通常落到 HBM)、DEFAULT(后端决定)、ERROR(checkify 错误空间)、INDEX(标量预取参数)、KEY(PRNG key)等。
5.3 Kernel(内核)
内核函数指明每个子问题要做的计算。内核不应返回任何输出,所有输出都应写入传入内核的输出缓冲。默认情况下所有输入、输出缓冲都是 SRAM 缓冲(除非用户在对应BlockSpec上通过memory_space覆盖了行为)。
def kernel(*input_buffers, *output_buffers): # ... perform compute # ... store result into output buffers当前子问题的索引可以在内核内部通过pl.program_id(grid_axis: int)查询(对应实现见 jax/_src/pallas/primitives.py#L61)。
5.4 Pallas Call(主入口)
pl.pallas_call是 Pallas 的主入口,当提供grid与BlockSpec时执行流水线调度。其签名如下:
def pallas_call( kernel, grid: tuple[int, ...], in_specs: Sequence[PyTree[BlockSpec]], out_specs: PyTree[BlockSpec], out_shape: PyTree[jax.ShapeDtypeStruct], ) -> Callable:pallas_call返回一个可调用对象:用输入值调用它,会返回与out_shape形状一致的输出。in_specs、out_specs、out_shape都是各自元素类型的 PyTree:in_specs与传给内核的输入缓冲的 PyTree 结构要匹配,out_specs与out_shape的 PyTree 结构也要匹配。后端实现位于 jax/_src/pallas/pallas_call.py。
6. 实战示例:分块逐元素内核
回到教程开头那个朴素add_matrices_kernel,这次改用流水线。我们将两个存放在 HBM 中、形状为f32[4096, 4096]的输入数组,按block_shape=(512, 512)切成子问题,在内核中每次只把两个块相加。由于加法是逐元素的,每个index_map都是相同的:在第i, j次迭代选中第i, j个块。
# Note: This is a TPU example. total_shape = (4096, 4096) block_shape = (512, 512) def add_matrices_pipelined_kernel(x_ref, y_ref, o_ref): o_ref[...] = x_ref[...] + y_ref[...] def add_matrices_pipelined(x: jax.Array, y: jax.Array): return pl.pallas_call( add_matrices_pipelined_kernel, grid=tuple(total // block for (total, block) in zip(total_shape, block_shape)), in_specs=[ pl.BlockSpec(block_shape, index_map=lambda i, j: (i, j)), pl.BlockSpec(block_shape, index_map=lambda i, j: (i, j)) ], out_specs=pl.BlockSpec(block_shape, index_map=lambda i, j: (i, j)), out_shape=jax.ShapeDtypeStruct(total_shape, dtype=jnp.float32), )(x, y) x = jax.random.uniform(jax.random.key(0), total_shape, dtype=jnp.float32) y = jax.random.uniform(jax.random.key(1), total_shape, dtype=jnp.float32) result = add_matrices_pipelined(x, y) np.testing.assert_array_equal( result, x + y )可以看到,用这套 API 写一个流水线内核,代码量并不比最初的朴素加法内核多多少!
6.1 参数化块大小
把块形状参数化是常见需求。块大小可能是调优 Pallas 内核性能时最重要的参数:它让我们控制流水线的形态——例如选更小的块,会给流水线循环增加更多迭代,而每次迭代做的工作更少。下面是参数化版本:
def add_matrices_pipelined_param( x: jax.Array, y: jax.Array, *, bm: int = 256, bn: int = 256 ) -> jax.Array: m, n = x.shape block_spec = pl.BlockSpec((bm, bn), lambda i, j: (i, j)) return pl.pallas_call( add_matrices_kernel, out_shape=x, in_specs=[block_spec, block_spec], out_specs=block_spec, grid=(m // bm, n // bn), )(x, y) np.testing.assert_array_equal( add_matrices_pipelined_param(x, y, bm=256, bn=256), x + y ) np.testing.assert_array_equal( add_matrices_pipelined_param(x, y, bm=128, bn=128), x + y ) np.testing.assert_array_equal( add_matrices_pipelined_param(x, y, bm=512, bn=512), x + y )7. 锋利的边角(Sharp Edges)
虽然流水线在心理模型上非常接近"在一个循环里反复调用内核函数",但中间缓冲并没有被完全隐藏,会带来几个微妙的 bug 来源。
7.1 缓冲重访(Buffer Revisiting)
一个通用的经验法则是:传入内核的输入缓冲应视为只读,输出缓冲应视为只写。
绝大多数情况下,向输入写、从输出读都会导致错误结果。原因是传入内核的 SRAM 缓冲只是底层 HBM 缓冲中数据的副本:如果更新了输入 SRAM 缓冲,更新结果永远不会被写回 HBM;如果读输出缓冲,读到的也永远不会是 SRAM 里最新写入的值。这与使用通用缓存时的"陈旧数据"(staleness)问题类似。
缓冲支持同时读写的情况只有两种:一是归约累加(见下文),二是通过给pallas_call传input_output_aliases参数,把一对输入/输出缓冲标记为输入-输出别名(aliased)。
7.2 归约与累加(Reductions and accumulation)
归约/累加只能沿 grid 的最后一维(最内层维度)进行,并且缓冲必须首先手动初始化。
归约是流水线少数支持对输出缓冲"边读边写"的场景之一,但它能工作的原因很微妙:Pallas 的流水线发射器(pipeline emitter)做了一项优化——如果连续两次迭代的数据切片相同,流水线就不会对该缓冲发起copy_in/copy_out。这意味着上一次迭代用过的 SRAM 缓冲会原样传给下一次迭代的内核,因此对输出缓冲的写入会在下一次迭代可见;一旦数据切片发生变化,最终累加好的 SRAM 缓冲才会被写回 HBM。这也是归约必须沿 grid 最后一维进行的原因——我们希望在最内层循环中、输出缓冲还在 SRAM 时完成全部累加,然后一次性写回 HBM,之后再也不碰那个输出块。
作为具体例子,考虑把(8, 1024, 1024)的数组沿第一个轴归约成(1024, 1024):
x = jnp.ones((8, 1024, 1024)) jnp.sum(x, axis=0)用pallas_call实现时,可以用大小为(8,)的 grid,每次迭代把x[i]加载进 SRAM,然后把它累加进输出 SRAM 缓冲。先看一个错误的朴素实现:
# Note: This is a TPU example. # Warning: this implementation is incorrect! def incorrect_sum_kernel(x_ref, o_ref): o_ref[...] += x_ref[...] def incorrect_sum(x: jax.Array, block_size: tuple[int, ...] = (256, 256)) -> jax.Array: reduction_size, *out_shape = x.shape grid = (reduction_size, *(out // blk for out, blk in zip(out_shape, block_size))) return pl.pallas_call( incorrect_sum_kernel, grid=grid, # None in `block_shape` means we pick a size of 1 and squeeze it away in_specs=[pl.BlockSpec((None, *block_size), lambda i, j, k: (i, j, k))], out_specs=pl.BlockSpec(block_size, lambda i, j, k: (j, k)), out_shape=jax.ShapeDtypeStruct(out_shape, x.dtype), )(x) result = incorrect_sum(x) print(result)结果是完全错误的!这个内核里有两处错误:
- 我们是沿第一个grid 维度累加,而不是沿最后一个grid 维度;
o_ref初始包含垃圾值,因此在开始累加前必须把它初始化为零。
修复这两点后,得到修正版内核。新内核用@pl.when创建一个条件:当沿归约轴的 program id 为0时,说明开始累加一个新的输出块,先将其清零;同时把归约维度移到了grid的最后一维:
# Note: This is a TPU example. def correct_sum_kernel(x_ref, o_ref): @pl.when(pl.program_id(2) == 0) def _(): o_ref[...] = jnp.zeros_like(o_ref) o_ref[...] += x_ref[...] def correct_sum(x: jax.Array, block_size: tuple[int, ...] = (256, 256)) -> jax.Array: reduction_size, *out_shape = x.shape # We moved the reduction to the last axis of the grid. grid = (*(out // blk for out, blk in zip(out_shape, block_size)), reduction_size) return pl.pallas_call( correct_sum_kernel, grid=grid, # None in `block_shape` means we pick a size of 1 and squeeze it away in_specs=[pl.BlockSpec((None, *block_size), lambda i, j, k: (k, i, j))], out_specs=pl.BlockSpec(block_size, lambda i, j, k: (i, j)), out_shape=jax.ShapeDtypeStruct(out_shape, x.dtype), )(x) result = correct_sum(x) print(result)这里有两个值得记住的细节:
block_shape中的None表示取大小为 1 并在传给内核时 squeeze 掉,因此输入块规格(None, *block_size)实际把归约维当成单元素维度处理;- 归约维必须位于 grid 的最后一维,并且用
pl.program_id(2) == 0(本例中归约轴是第 2 个网格轴)判断"是否为该输出块的第一次累加",从而手动把缓冲初始化为零。
8. 分析流水线的性能
流水线内核的性能如何?答案取决于硬件的瓶颈在哪里。通常关心 3 个量:
- 内存延迟(Memory latency)$\alpha$:一次内存传输的最小延迟。
- 内存带宽(Memory bandwidth)$\beta$:从 HBM 到 SRAM 的传输速率(字节/秒)。
- FLOP/s$F$:处理器每秒可执行的浮点运算次数。
如果处理速度 FLOP/s 是瓶颈,称程序为计算受限(compute-bound);如果带宽或延迟是瓶颈,则称为内存受限(memory-bound)。一般而言,我们的优化目标就是让内核成为计算受限,即充分利用硬件全部的处理能力。
假设程序每次内核迭代需要传输 $X$ 字节、执行 $Y$ 次浮点运算。$X$ 与 $Y$ 的比值取决于计算类型:逐元素运算(如加法、乘法)中两者同比例增长;而矩阵乘法中,计算量随问题规模立方增长,内存量只随规模平方增长。
在计算受限场景下,运行 $N$ 次迭代的流水线大约耗时 $(\alpha + X/\beta) + N (Y/F)$ 秒:第一项是初始气泡的成本(若末尾也有气泡则乘以 2),第二项是流水线稳态阶段的总时间。当 $N$ 足够大、流水线足够长时,运行时间的主导项是 $F$——加速器的处理速度。
在内存受限场景下,还需要进一步区分瓶颈是延迟还是带宽:
- 如果瓶颈是带宽,总运行时间约为 $\alpha + N(X / \beta)$ 秒。与延迟受限场景相反,由于带宽已饱和,内存拷贝是串行进行的。内存受限通常不理想:处理器会有空闲间隙,而且在大多数硬件配置中,内存带宽 $\beta$ 比处理速度 $F$ 慢几个数量级。
- 如果瓶颈特指延迟而非带宽,可以通过插入更多流水线阶段来修复,代价是需要更多 SRAM 存放额外的缓冲。阶段足够多之后,问题会重新变为计算受限或带宽受限——取决于稳态阶段先撞上哪个瓶颈。多级流水线的缺点是:气泡的大小与阶段数成正比,因此务必保证流水线足够长,让气泡不占据总运行时间的可观比例。
平台支持方面:Pallas 在TPU 上只支持双缓冲——因为 TPU 程序可以使用较大的块尺寸,双缓冲通常已足以覆盖延迟;在GPU上,流水线阶段数既可以在 Triton 后端(通过CompilerParams)指定,也可以在 Mosaic GPU 后端(通过流水线发射器的参数)指定。
9. 平台特化:TPU 与 Mosaic GPU 的流水线进阶
主文档在第 5 节末把平台细节指向了对应文档,这里基于仓库内两份平台文档做纵深补充,帮助你按平台选择正确的入口。
9.1 TPU 流水线(docs/pallas/tpu/pipelining.md)
TPU 内存空间。Pallas 暴露了 TPU 完整的存储层级,下表把 Pallas 的 TPU 内存空间映射到标准内存类型(DRAM/SRAM):
| Pallas 枚举 | TPU 存储空间 | 类型(DRAM/SRAM) |
|---|---|---|
pl.ANY | HBM(通常)或 VMEM | DRAM |
pltpu.VMEM | VMEM | SRAM |
pltpu.SMEM | SMEM | SRAM |
pltpu.SEMAPHORE | 信号量 | SRAM |
要点:
MemorySpace.VMEM表示向量 SRAM,是未指定时的默认内存空间;MemorySpace.SMEM表示标量 SRAM,只能对 SMEM 做标量读写;MemorySpace.ANY是给编译器的"内存空间不受限"提示,多数情况下 XLA 会把它放到 HBM;ANY缓冲不能用数组索引语法(如x[...])直接解引用,必须先通过pltpu.sync_copy或pltpu.async_copy把值拷入 VMEM/SMEM 缓冲;MemorySpace.SEMAPHORE用于分配信号量,构造屏障或跟踪异步操作。
TPU 上的流水线通常发生在 HBM(DRAM)↔ VMEM(向量 SRAM)之间:pallas_call在 TPU 上的默认行为是参数假定存放在 HBM,内核体输入存放在 VMEM。注意:只有memory_space标记为VMEM时流水线才被允许。memory_space也可通过pallas_call的scratch_shapes参数给内核指定持久化的 scratch 缓冲(必须位于VMEM/SMEM/SEMAPHORE),用于存放部分累加、归约等中间结果。
多缓冲(Multiple Buffering)。TPU 上可以按参数粒度指定缓冲份数,通过pl.BlockSpec的pipeline_mode传入pl.Buffered对象:
pl.BlockSpec( pipeline_mode=pl.Buffered(buffer_count=buffer_count) )所有输入输出的默认缓冲份数为 2。源码中Buffered定义于 jax/_src/pallas/core.py#L212,除buffer_count外还支持use_lookahead(前瞻预取)、revisit(RevisitMode.IMMEDIATE/ANY,控制输出块在非连续迭代被重访时的处理方式)、prefetched_count(进入流水线前已预填充的窗口槽数)。
pltpu.emit_pipeline。这是 Pallas 内实现的流水线 API,允许在内核内部构造流水线而不只是在内核入口。典型用途:构造嵌套流水线(外层芯片间通信流水线 + 内层 HBM-VMEM 流水线)、使用 lookahead 预取与动态块形状等特性。签名与pl.pallas_call类似:
def emit_pipeline( kernel: Callable, grid: tuple[int], in_specs: PyTree[BlockSpec] = None, out_specs: PyTree[BlockSpec] = None, dimension_semantics: tuple[GridDimensionSemantics] = None, core_axis: int | None = None, ) -> Callable:前瞻预取(Lookahead Prefetch)。开启后,流水线会在缓冲槽一有空闲就立刻预取下下一个输入块,而不是等到该块被使用的前一个迭代。例如 grid 为(8,)、每迭代取块索引为0,0,0,0,1,1,1,1时,lookahead 会在第 0 次迭代就同时开始取块0和1,而标准调度要到第 3 次迭代才开始取块1。它主要适用于各块计算量不均衡(有些块被跳过或工作量较少)的场景,此时前一个迭代可能没有足够的计算量来完全重叠内存传输。lookahead 有一点控制流开销,默认关闭,可通过pl.Buffered(buffer_count=..., use_lookahead=True)开启。
动态块形状(Dynamic Block Shapes)。pltpu.emit_pipeline支持对有界动态形状的块做流水线:动态维在block_shape中标记为pl.BoundedSlice(max_size),index_map返回的对应索引应是pl.ds(start, size)构造的动态切片(start与size都是元素索引,且可以动态):
pl.BlockSpec( block_shape=(pl.BoundedSlice(32), 256), index_map=lambda *grid_idxs: (pl.ds(start, end), 0), )Megacore 配置。部分 TPU 芯片拥有两个 TensorCore,但对 JAX 用户表现为一个设备,即 megacore:两个 TensorCore 各自拥有独立的 VMEM/VREG/SMEM/SREG 与计算单元,但共享 HBM。通过给pallas_call传compiler_params=pltpu.CompilerParams(dimension_semantics=("parallel", ...))可以把 embarrassingly-parallel 的维度切分到两个 TensorCore 上并行执行;使用pltpu.emit_pipeline时则把core_axis(一个并行 grid 轴的索引)传入emit_pipeline。dimension_semantics每个元素取"parallel"或"arbitrary":"parallel"表示该维迭代可独立执行、互不影响正确性;"arbitrary"表示不可并行化。注意megacore 目前仅 TPU v4 与 v5p 可用:在其他平台上传dimension_semantics是 no-op,但不传它只会用到一个 TensorCore(即使有多个可用)。
9.2 Mosaic GPU 流水线(docs/pallas/gpu/pipelining.md)
Mosaic GPU 后端显式编程流水线,这与 Triton 的编程模型有显著差异——Triton 中流水线是编译器自动做的优化。推荐入口是plgpu.emit_pipeline(对顺序循环做流水线),配合plgpu.kernel(按 CUDA grid 并行切分问题)。emit_pipeline与pl.pallas_call的 API 类似,但有几个 GPU 特有选项(源码见 jax/_src/pallas/mosaic_gpu/pipeline.py#L261):
max_concurrent_steps:控制最大并发内存传输数。更多的并发步数会占用更多 SMEM 存放临时缓冲,但能提升内存子系统利用率,建议做 autotune。较低的值(如 2)由于 SMEM 占用更少,有时能获得更高占用率(occupancy),对 ALU 密集型内核有利,但硬件调度会引入更多噪声;较大的值(4–6)最适合无法从额外占用率获益的内核。delay_release:指定缓冲在被流水线重新使用前额外等待的迭代数。例如迭代 0 拷入 SMEM 的缓冲,在delay_release=1、max_concurrent_steps=2时直到迭代 3 才被复用(标准双缓冲策略是迭代 2)。如果不对流水线操作数await一次plgpu.wgmma,就必须设delay_release=1,否则流水线会在 WGMMA 还在读缓冲时就开始覆写它——省略该参数会造成静默数据竞争。
兼容 API:pl.pallas_call+CompilerParams。为保持与 Pallas TPU 兼容,Mosaic GPU 也实现了pl.pallas_call。默认它在 CUDA grid 上并行切分内核;通过compiler_params=plgpu.CompilerParams(...)传入与流水线相关的选项:
dimension_semantics:每个 grid 维取Literal['parallel', 'sequential']。parallel把对应维切分到 CUDA grid,sequential维被顺序流水线化。注意:如果没有任何维被标记为sequential,就不会发生任何流水线化!max_concurrent_steps、delay_release:与plgpu.emit_pipeline中的同名参数一致。
GPU 内存空间。BlockSpec(memory_space=...)可以指定 Ref 所在空间:plgpu.GPUMemorySpace.SMEM分配在共享内存(SMEM),SMEM Ref 可用数组索引语法解引用,emit_pipeline使用的正是这一空间;plgpu.GPUMemorySpace.GMEM分配在全局内存(GMEM/HBM),GMEM 中的 Ref 不做流水线处理,也不能直接用数组索引访问,必须通过plgpu.copy_gmem_to_smem/plgpu.copy_smem_to_gmem或plgpu.emit_pipeline流水线化到 SMEM。emit_pipeline的核心价值就是把 TensorCore 计算与 GMEM↔SMEM 数据传输重叠——异步拷贝延迟长,而 TensorCore 计算必须操作寄存器(矩阵乘则是 SMEM Ref)。
Hopper matmul 示例。GPU 文档中的示例内核用 Hopper 特有的wgmma(warpgroup matrix multiply accumulate)指令:wgmma由单个 Mosaic GPU 线程发出,在 TensorCore 上异步执行。外层plgpu.kernel的 grid 并行化矩阵乘的非收缩维 M、N,每个程序实例内部用plgpu.emit_pipeline对收缩维 K 做顺序流水线;每次流水线迭代加载两个输入 tile、调用plgpu.wgmma累加到plgpu.ACC(一种存放在寄存器、保存 WGMMA 中间结果的特殊 Ref),K 维累加完毕后写回输出。plgpu.wgmma_wait(N)等待在途 WGMMA 数量不超过 N;示例中delay_release=1与wgmma_wait(1)配合,始终让一个 WGMMA 在途,以保持 TensorCore 利用率和避免每迭代冲刷流水线。
Warp Specialization(Warp 特化)。Hopper+ 上可以把 TMA(GMEM/SMEM 拷贝)的发出工作交给独立的 memory warpgroup,与做算术的 compute warpgroup 分离,避免索引计算与 TMA 发出占据大量时间导致 TensorCore 空闲。Pallas 通过plgpu.emit_pipeline_warp_specialized(jax/_src/pallas/mosaic_gpu/pipeline.py#L660)支持:该辅助函数接管 memory thread 的全部逻辑,用户只需指定 compute thread 的工作。关键参数包括num_compute_wgs(compute warpgroup 数,总线程数需设为num_compute_wgs+1)、memory_registers(分配给 memory thread 的寄存器数,默认 40,出现寄存器溢出时上下调整)、wg_axis(线程轴名)、memory_thread_idx(指定哪个 Pallas 线程充当 memory thread,默认最后一个线程)以及compute_context(只在 compute thread 运行的前言/尾声,用于定义流水线 carry 的初始化与消费——所有 compute thread 专属数组都应在这里实例化,否则会在 memory thread 中物化、浪费寄存器并可能因寄存器溢出而变慢)。lax.axis_index可在内核中取得 Pallas 线程索引,用于在 compute threads 间划分工作。
10. 小结
软件流水线是内核优化的核心手段,其本质是把"切块子问题"与"异步通信"结合,用计算隐藏 HBM↔SRAM 的搬移延迟。从本文可以提炼出四条可复用的实践准则:
- 先想清楚瓶颈:用 $\alpha$(延迟)、$\beta$(带宽)、$F$(FLOP/s)三个量给内核定位,目标是把内核推向计算受限(compute-bound)区;
- 用三要素描述流水线:
grid(子问题个数)、kernel(SRAM 上的计算)、data_slices(BlockSpec的index_map),Pallas 会负责多缓冲与异步重叠的样板逻辑; - 块大小是头号调优旋钮:小块的流水线迭代更多、单次工作更少,应根据平台(TPU 双缓冲、GPU 多级缓冲)与硬件容量权衡;
- 警惕两个经典陷阱:输入缓冲只读、输出缓冲只写(除非显式
input_output_aliases);归约只能沿 grid 最后一维并先手动初始化缓冲。
继续深入可阅读 grid 与 blockspec、Pallas 快速入门,以及平台专属的 TPU 流水线 与 Mosaic GPU 流水线 文档。
【免费下载链接】jaxComposable transformations of Python+NumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考