十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

JAX Pallas 软件流水线(Software Pipelining)完全指南:通信-计算重叠的原理、API 与实战陷阱

JAX Pallas 软件流水线(Software Pipelining)完全指南:通信-计算重叠的原理、API 与实战陷阱 JAX Pallas 软件流水线Software Pipelining完全指南通信-计算重叠的原理、API 与实战陷阱【免费下载链接】jaxComposable transformations of PythonNumPy 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 VMEMH100 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_shapejax.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_callpallas_call负责把x、y拷贝进 SRAM并分配内核操作所需的 SRAM 缓冲包括输出缓冲内核执行完毕后pallas_call再把输出缓冲中的值拷回 HBM得到输出jax.Array。2.2 两个必须正视的约束容量与带宽Pallas 暴露了 SRAM 等底层存储空间但要写出高性能内核必须更精细地利用各类存储尤其要同时考虑内存容量Memory capacitySRAM 很小如果数组太大上面的内核根本无法运行因为输入放不进 SRAM。作为参考一个f32[2048, 2048]的数组就有 16 MiB因此上述朴素内核只能处理中等规模以下的数组。内存带宽Memory bandwidth在 HBM 与 SRAM 之间拷贝很耗时至少比绝大多数计算指令慢得多。上面的add_matrices很可能把大部分时间花在 HBM↔SRAM 的拷贝上而不是加法本身。带着这两个约束我们需要重新思考如何榨取加速器的性能——这正是流水线的用武之地。3. 流水线基础把大问题切小并重叠如何既利用内存层级中各类存储的优势又能操作存放在 HBM 中的大数组、同时用快速的 SRAM 做计算流水线是一种非常通用的编程模式它要求把问题拆成可以并行重叠的更小子问题。流水线的第一步是把问题划分成能放进 SRAM 的小子问题。以逐元素elementwise运算为例可以简单地对源数组每次处理一个切片得到如下 3 个步骤又称 3 个阶段 / stagescopy_in把切片A[i]从 HBM 拷入 SRAMXcompute把X加载到寄存器计算结果并存回 SRAMYcopy_out把结果Y拷回 HBM 的A[i]。注意步骤 1–3 之间存在数据依赖必须先完成步骤 1 才能开始步骤 2因此不能简单重叠。然而不同子问题实例之间没有数据依赖——也就是说可以在执行块A[i1]的步骤 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 后把结果写回 HBMcopy_outfor 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下一次循环迭代所需的输入在当前迭代执行计算的同时异步发起拷贝让计算与内存通信同时发生。为了推演这个代码变换先把循环按 N4 展开并把拷贝指令拆成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 i1 N: copy_in_start(A[i1], 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(i1)], 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 流水线 APIPallas 提供了一套流水线 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.IndirectNone表示该维度被 squeeze 掉、不出现在内核里pl.BoundedSlice定义于 jax/_src/pallas/core.py#L415则允许对某维度指定有界但动态的切片大小详见第 9 节。memory_space使用pl.MemorySpace枚举jax/_src/pallas/core.py#L271包含ANY不限定通常落到 HBM、DEFAULT后端决定、ERRORcheckify 错误空间、INDEX标量预取参数、KEYPRNG 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都是各自元素类型的 PyTreein_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, gridtuple(total // block for (total, block) in zip(total_shape, block_shape)), in_specs[ pl.BlockSpec(block_shape, index_maplambda i, j: (i, j)), pl.BlockSpec(block_shape, index_maplambda i, j: (i, j)) ], out_specspl.BlockSpec(block_shape, index_maplambda i, j: (i, j)), out_shapejax.ShapeDtypeStruct(total_shape, dtypejnp.float32), )(x, y) x jax.random.uniform(jax.random.key(0), total_shape, dtypejnp.float32) y jax.random.uniform(jax.random.key(1), total_shape, dtypejnp.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_shapex, in_specs[block_spec, block_spec], out_specsblock_spec, grid(m // bm, n // bn), )(x, y) np.testing.assert_array_equal( add_matrices_pipelined_param(x, y, bm256, bn256), x y ) np.testing.assert_array_equal( add_matrices_pipelined_param(x, y, bm128, bn128), x y ) np.testing.assert_array_equal( add_matrices_pipelined_param(x, y, bm512, bn512), 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, axis0)用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, gridgrid, # 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_specspl.BlockSpec(block_size, lambda i, j, k: (j, k)), out_shapejax.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, gridgrid, # 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_specspl.BlockSpec(block_size, lambda i, j, k: (i, j)), out_shapejax.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.mdTPU 内存空间。Pallas 暴露了 TPU 完整的存储层级下表把 Pallas 的 TPU 内存空间映射到标准内存类型DRAM/SRAMPallas 枚举TPU 存储空间类型DRAM/SRAMpl.ANYHBM通常或 VMEMDRAMpltpu.VMEMVMEMSRAMpltpu.SMEMSMEMSRAMpltpu.SEMAPHORE信号量SRAM要点MemorySpace.VMEM表示向量 SRAM是未指定时的默认内存空间MemorySpace.SMEM表示标量 SRAM只能对 SMEM 做标量读写MemorySpace.ANY是给编译器的内存空间不受限提示多数情况下 XLA 会把它放到 HBMANY缓冲不能用数组索引语法如x[...]直接解引用必须先通过pltpu.sync_copy或pltpu.async_copy把值拷入 VMEM/SMEM 缓冲MemorySpace.SEMAPHORE用于分配信号量构造屏障或跟踪异步操作。TPU 上的流水线通常发生在 HBMDRAM↔ 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_modepl.Buffered(buffer_countbuffer_count) )所有输入输出的默认缓冲份数为 2。源码中Buffered定义于 jax/_src/pallas/core.py#L212除buffer_count外还支持use_lookahead前瞻预取、revisitRevisitMode.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_lookaheadTrue)开启。动态块形状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_maplambda *grid_idxs: (pl.ds(start, end), 0), )Megacore 配置。部分 TPU 芯片拥有两个 TensorCore但对 JAX 用户表现为一个设备即 megacore两个 TensorCore 各自拥有独立的 VMEM/VREG/SMEM/SREG 与计算单元但共享 HBM。通过给pallas_call传compiler_paramspltpu.CompilerParams(dimension_semantics(parallel, ...))可以把 embarrassingly-parallel 的维度切分到两个 TensorCore 上并行执行使用pltpu.emit_pipeline时则把core_axis一个并行 grid 轴的索引传入emit_pipeline。dimension_semantics每个元素取parallel或arbitraryparallel表示该维迭代可独立执行、互不影响正确性arbitrary表示不可并行化。注意megacore 目前仅 TPU v4 与 v5p 可用在其他平台上传dimension_semantics是 no-op但不传它只会用到一个 TensorCore即使有多个可用。9.2 Mosaic GPU 流水线docs/pallas/gpu/pipelining.mdMosaic GPU 后端显式编程流水线这与 Triton 的编程模型有显著差异——Triton 中流水线是编译器自动做的优化。推荐入口是plgpu.emit_pipeline对顺序循环做流水线配合plgpu.kernel按 CUDA grid 并行切分问题。emit_pipeline与pl.pallas_call的 API 类似但有几个 GPU 特有选项源码见 jax/_src/pallas/mosaic_gpu/pipeline.py#L261max_concurrent_steps控制最大并发内存传输数。更多的并发步数会占用更多 SMEM 存放临时缓冲但能提升内存子系统利用率建议做 autotune。较低的值如 2由于 SMEM 占用更少有时能获得更高占用率occupancy对 ALU 密集型内核有利但硬件调度会引入更多噪声较大的值4–6最适合无法从额外占用率获益的内核。delay_release指定缓冲在被流水线重新使用前额外等待的迭代数。例如迭代 0 拷入 SMEM 的缓冲在delay_release1、max_concurrent_steps2时直到迭代 3 才被复用标准双缓冲策略是迭代 2。如果不对流水线操作数await一次plgpu.wgmma就必须设delay_release1否则流水线会在 WGMMA 还在读缓冲时就开始覆写它——省略该参数会造成静默数据竞争。兼容 APIpl.pallas_callCompilerParams。为保持与 Pallas TPU 兼容Mosaic GPU 也实现了pl.pallas_call。默认它在 CUDA grid 上并行切分内核通过compiler_paramsplgpu.CompilerParams(...)传入与流水线相关的选项dimension_semantics每个 grid 维取Literal[parallel, sequential]。parallel把对应维切分到 CUDA gridsequential维被顺序流水线化。注意如果没有任何维被标记为sequential就不会发生任何流水线化max_concurrent_steps、delay_release与plgpu.emit_pipeline中的同名参数一致。GPU 内存空间。BlockSpec(memory_space...)可以指定 Ref 所在空间plgpu.GPUMemorySpace.SMEM分配在共享内存SMEMSMEM Ref 可用数组索引语法解引用emit_pipeline使用的正是这一空间plgpu.GPUMemorySpace.GMEM分配在全局内存GMEM/HBMGMEM 中的 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 特有的wgmmawarpgroup matrix multiply accumulate指令wgmma由单个 Mosaic GPU 线程发出在 TensorCore 上异步执行。外层plgpu.kernel的 grid 并行化矩阵乘的非收缩维 M、N每个程序实例内部用plgpu.emit_pipeline对收缩维 K 做顺序流水线每次流水线迭代加载两个输入 tile、调用plgpu.wgmma累加到plgpu.ACC一种存放在寄存器、保存 WGMMA 中间结果的特殊 RefK 维累加完毕后写回输出。plgpu.wgmma_wait(N)等待在途 WGMMA 数量不超过 N示例中delay_release1与wgmma_wait(1)配合始终让一个 WGMMA 在途以保持 TensorCore 利用率和避免每迭代冲刷流水线。Warp SpecializationWarp 特化。Hopper 上可以把 TMAGMEM/SMEM 拷贝的发出工作交给独立的 memory warpgroup与做算术的 compute warpgroup 分离避免索引计算与 TMA 发出占据大量时间导致 TensorCore 空闲。Pallas 通过plgpu.emit_pipeline_warp_specializedjax/_src/pallas/mosaic_gpu/pipeline.py#L660支持该辅助函数接管 memory thread 的全部逻辑用户只需指定 compute thread 的工作。关键参数包括num_compute_wgscompute warpgroup 数总线程数需设为num_compute_wgs1、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子问题个数、kernelSRAM 上的计算、data_slicesBlockSpec的index_mapPallas 会负责多缓冲与异步重叠的样板逻辑块大小是头号调优旋钮小块的流水线迭代更多、单次工作更少应根据平台TPU 双缓冲、GPU 多级缓冲与硬件容量权衡警惕两个经典陷阱输入缓冲只读、输出缓冲只写除非显式input_output_aliases归约只能沿 grid 最后一维并先手动初始化缓冲。继续深入可阅读 grid 与 blockspec、Pallas 快速入门以及平台专属的 TPU 流水线 与 Mosaic GPU 流水线 文档。【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表