JAX Pallas 在 TPU 上编写高效矩阵乘法内核:从分块算法、性能建模到算子融合
JAX Pallas 在 TPU 上编写高效矩阵乘法内核从分块算法、性能建模到算子融合【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax本文基于 docs/pallas/tpu/matmul.md 编写。该文档演示了如何用 JAX 的 Pallas 低级内核编写语言在 TPUTensorCore/MXU上编写高效的矩阵乘法内核涵盖块矩阵乘法与流水线pipelining原理、FLOPs 与内存带宽的性能建模方法、bfloat16数据类型的优化以及如何把内核模板化以融合 RHS 转置和激活函数。读完本文你将掌握用pl.pallas_callBlockSpecPrefetchScalarGridSpec写出可运行的 TPU matmul 内核并能用算术强度arithmetic intensity判断自己的 matmul 是受计算还是受内存搬运限制。背景为什么在 TPU 上写 matmul 要先理解分块与流水线矩阵乘法是现代深度学习与大语言模型的核心线性代数原语我们希望用 TPU、GPU 这类拥有专用矩阵乘法单元TPU 上称为 MXU即 Matrix Multiply Unit的加速器把它跑得尽可能快。要高效利用 TPU 的 MXU需要先掌握三个概念块矩阵乘法block matrix multiplication、分块tiling与流水线pipelining。块矩阵乘法假设我们要实现matmul(x, y)它把(m, k)与(k, n)两个数组相乘但我们只能使用一个只支持小矩阵比如三个维度均 ≤ 256的原语matmul_small。矩阵乘法有一个很好的性质输出的每个块都可以表示为输入的行块与列块的若干次小矩阵乘法之和。形式化地说设输入 $x \in \mathbb{R}^{m \times k}$、$y \in \mathbb{R}^{k \times n}$输出 $z \in \mathbb{R}^{m \times n}$我们沿各个维度把它们按块大小 $b_m, b_k, b_n$ 分解。例如 $x$ 可分解为$$ \begin{bmatrix} x_{0, 0} \cdots x_{0, i_k} \ x_{1, 0} \cdots x_{1, i_k} \ \vdots \ddots \vdots \ x_{i_m, 0} \cdots x_{i_m, i_k} \ \end{bmatrix} $$其中每个块 $x_{ik} \in \mathbb{R}^{b_m \times b_k}$$y$、$z$ 同理分解。对于某个输出块 $z_{ij}$有$$ z_{ij} \sum_k x_{ik} y_{kj} $$即每个输出块是若干个小块矩阵乘法 $x_{ik} y_{kj}$ 的累加。用 NumPy 实现如下def matmul_small(x: np.ndarray, y: np.ndarray) - np.ndarray: m, k, n x.shape[0], x.shape[1], y.shape[0] assert m 256 assert k 256 assert n 256 return np.matmul(x, y) def block_matmul( x: np.ndarray, y: np.ndarray, *, bm: int 256, bk: int 256, bn: int 256, ) - np.ndarray: m, k x.shape _, n y.shape z np.zeros((m, n), dtypex.dtype) for m_i in range(m // bm): for n_i in range(n // bn): for k_i in range(k // bk): m_slice slice(m_i * bm, (m_i 1) * bm) k_slice slice(k_i * bk, (k_i 1) * bk) n_slice slice(n_i * bn, (n_i 1) * bn) x_block x[m_slice, k_slice] y_block y[k_slice, n_slice] z[m_slice, n_slice] matmul_small(x_block, y_block) return z上面的实现假设输入维度能被bm/bk/bn整除。block_matmul把一个大矩阵乘法拆解成许多小矩阵乘法每个(bm, bn)大小的输出块由若干个(bm, bk) × (bk, bn)的小矩阵乘法累加得到。可以用一个 4096³ 的随机输入验证它与x y一致容差atol1e-6, rtol1e-6。TPU 和 GPU 正是这样工作的硬件原生支持类似matmul_small的小矩阵乘法因此做大规模 matmul 时就是把上述block_matmul分解应用到硬件上。Pallas 会自动把较大的块再细分到 MXU 上执行——文档中明确指出while MXUs are only capable of multiplying small blocks, Pallas will automatically take bigger blocks and automatically tile them over the MXUs.Tiling 与 Pipelining在 docs/pallas/tpu/pipelining.md 中已经覆盖了 Pallas 中如何对计算进行分块与流水线化。核心动机是TPU 的计算单元MXU只读 VMEM靠近计算单元的片上存储数据需要先从 HBM高带宽内存拷贝到 VMEM。如果计算单元总是等内存搬运完成才开始干活效率会很低。流水线化的关键是用BlockSpec与grid描述任务让第 $i1$ 次迭代的内存搬运与第 $i$ 次迭代的计算重叠保证计算单元始终忙碌。注意我们在block_matmul里已经有一个三重嵌套循环它天然对应 Pallas 的grid循环里的切片逻辑则对应BlockSpec。你的第一个矩阵乘法内核把上面的思想落地就得到第一个 Pallas matmul 内核创建一个三维grid对应 NumPy 代码中的三重循环grid的最后一维对应矩阵乘法的收缩contraction维它是一个reduction 维度因此必须初始化累加器z_ref。def matmul_kernel(x_ref, y_ref, z_ref): pl.when(pl.program_id(2) 0) def _(): z_ref[...] jnp.zeros_like(z_ref) z_ref[...] x_ref[...] y_ref[...] def matmul( x: jax.Array, y: jax.Array, *, bm: int 128, bk: int 128, bn: int 128, ): m, k x.shape _, n y.shape return pl.pallas_call( matmul_kernel, out_shapejax.ShapeDtypeStruct((m, n), x.dtype), in_specs[pl.BlockSpec((bm, bk), lambda i, j, k: (i, k)), pl.BlockSpec((bk, bn), lambda i, j, k: (k, j))], out_specspl.BlockSpec((bm, bn), lambda i, j, k: (i, j)), grid(m // bm, n // bn, k // bk), compiler_paramspltpu.CompilerParams( dimension_semantics(parallel, parallel, arbitrary)), )(x, y)关键点逐一说明grid(m // bm, n // bn, k // bk)三个维度分别对应输出行块索引i、输出列块索引j、收缩块索引k。pl.program_id(2)即第三个维度收缩维的程序 ID。pl.when(pl.program_id(2) 0)只有收缩维第一步时把累加器清零随后每次迭代用累积。BlockSpec的 index 函数in_specs中第一个BlockSpec((bm, bk), lambda i, j, k: (i, k))说明每个(bm, bk)的输入块取自x[i, k]第二个取自y[k, j]out_specs的(bm, bn)块写入z[i, j]。dimension_semantics(parallel, parallel, arbitrary)前两个维度标注为parallel可任意顺序并行执行收缩维标注为arbitrary必须顺序执行因为存在累加依赖。关于CompilerParams其完整定义位于 jax/_src/pallas/mosaic/core.py。从源码可以看到dimension_semantics的可选值包括parallel、core_parallel、subcore_parallel、arbitrary等core.py 第 70 行parallel表示可任意顺序执行、arbitrary表示必须顺序执行。此外它还支持vmem_limit_bytes覆盖内核默认 VMEM 上限需配合--xla_tpu_scoped_vmem_limit_kibN标志、collective_id、has_side_effects防止内核被 XLA 做 CSE 消除等参数读者可按需查阅。验证方式np.testing.assert_array_equal(x y, matmul(x, y))m, k, n 4096, 4096, 4096 k1, k2 random.split(random.key(0), 2) x random.normal(k1, (m, k), dtypejnp.float32) y random.normal(k2, (k, n), dtypejnp.float32) np.testing.assert_array_equal(x y, matmul(x, y))仓库中已经有一个与本文档配套的官方示例实现jax/experimental/pallas/ops/tpu/matmul.py。它的matmul_kernel使用acc_ref作为显式 f32 累加器、用preferred_element_typeacc_ref.dtype调用jnp.dotmatmul函数则通过PrefetchScalarGridSpec分配scratch_shapes[pltpu.VMEM((l, r), acc_dtype)]并针对int8/int4/uint8/uint4输入把累加器 dtype 自动切换为int32matmul.py 第 63-65 行。这可以作为你在实践中参照/扩展的起点。矩阵乘法性能分析FLOPs 与内存带宽分析 matmul 性能通常关注两件事浮点运算总量FLOPs与内存带宽用量。注意区分FLOPs指浮点运算的次数数量FLOP/s指每秒执行的浮点运算次数速率。一次(m, k) × (k, n)矩阵乘法的 FLOPs 约为2 * m * k * n严格说是n * m * (2k - 1)但k足够大时近似成立。最小内存带宽用量假设 float32输入从 HBM 拷贝进 VMEM 的总大小加上输出写回 HBM 的大小即(m * k k * n m * n) * 4 bytes/float32。如果同一份输入被多次重读实际用量会更大。一个关键的观察matmul 的 FLOPs 随输入规模三次方增长而最小带宽用量只随输入规模二次方增长。这意味着 matmul 越大计算相对拷贝的比例越高。def matmul_flops(m: int, k: int, n: int): return 2 * m * k * n def matmul_membw(m: int, k: int, n: int, dtype: jnp.dtype): return (m * k k * n m * n) * np.dtype(dtype).itemsize print(matmul_flops(1024, 1024, 1024)) # 2147483648 print(matmul_membw(1024, 1024, 1024, jnp.float32)) # 12582912算术强度与 compute bound / memory bound接下来把理论数字和真实芯片对照。原始文档的 notebook 运行在TPU v5e上如果你自己跑数字可能不同。TPU v5e 拥有约197 TFLOP/s 的 bf16/f32 计算能力与819 GB/s 的内存带宽。两者的比值称为算术强度arithmetic intensity它给出一个临界点当 FLOPs / 内存带宽用量 低于该比值时芯片就会变成memory bound内存受限——计算单元在等待数据搬运时空转高于该比值则compute bound计算受限。TPU v5e 上这个临界点约为 240 FLOPs/byte。v5e_flops 197e12 v5e_membw 819e9 v5e_op_intensity v5e_flops / v5e_membw # ~240.5 def matmul_flops_intensity(m: int, k: int, n: int, dtype: jnp.dtype): flops matmul_flops(m, k, n) membw matmul_membw(m, k, n, dtype) return flops / membw粗略地说一次 matmul 的计算耗时约2 * m * k * n / (197 TFLOP/s)秒VMEM 搬运耗时约(m*k k*n m*n) * 4 bytes / 819GB/s秒。例如(1024, 1024) × (1024, 1024)的 float32 matmulprint(f{matmul_flops_intensity(1024, 1024, 1024, jnp.float32)} flops/byte) # 约 170.7 flops/byte低于 v5e 的能力 → memory bound这个强度低于芯片能力所以它是 memory bound 的。当矩阵变大后会从 memory bound 跨越到 compute bound。对于m k n的方阵在 TPU v5e 上跨过临界点2m**3 / 12m**2 240即m k n 1440之后就是 compute bound 了。bfloat16 矩阵乘法更小 dtype 更易 compute bound另一个让 matmul 更容易 compute bound 的办法是使用更小的 dtype。前面的例子用 float32 输入输出但 TPU v5e 也原生支持bfloat16bf16矩阵乘法FLOP/s 不变但内存带宽用量减半因此小矩阵也更容易变成 compute bound。(1024, 1024, 1024)的 bf16 matmul 强度约为 341 flops/byte此时已经 compute bound。MXU 原生的 bf16 matmul 例程接受两个 bf16 输入矩阵、并以f32 累积。在 Pallas 中触发它的方式向jnp.matmul或jnp.dot传入preferred_element_typejnp.float32累加器Ref使用 f32 dtype写回 HBM 前把结果下转换downcast回 bf16。这样既不损失精度、又不多做类型转换还保住了 bf16 的内存带宽收益。完整内核如下注意目前分配 scratch 空间的唯一方式是通过pltpu.PrefetchScalarGridSpec它允许你在 VMEM 中分配 scratch 空间其签名可在 jax/_src/pallas/mosaic/core.py 中查看参数包括num_scalar_prefetch、grid、in_specs、out_specs、scratch_shapesdef matmul_kernel(x_ref, y_ref, z_ref, acc_ref, *, nsteps): pl.when(pl.program_id(2) 0) def _(): acc_ref[...] jnp.zeros_like(acc_ref) acc_ref[...] jnp.dot( x_ref[...], y_ref[...], preferred_element_typejnp.float32 ) pl.when(pl.program_id(2) nsteps - 1) def _(): z_ref[...] acc_ref[...].astype(z_ref.dtype) jax.jit(static_argnames[bm, bk, bn]) def matmul( x: jax.Array, y: jax.Array, *, bm: int 128, bk: int 128, bn: int 128, ): m, k x.shape _, n y.shape return pl.pallas_call( functools.partial(matmul_kernel, nstepsk // bk), grid_specpltpu.PrefetchScalarGridSpec( num_scalar_prefetch0, in_specs[ pl.BlockSpec((bm, bk), lambda i, j, k: (i, k)), pl.BlockSpec((bk, bn), lambda i, j, k: (k, j)), ], out_specspl.BlockSpec((bm, bn), lambda i, j, k: (i, j)), scratch_shapes[pltpu.VMEM((bm, bn), jnp.float32)], grid(m // bm, n // bn, k // bk), ), out_shapejax.ShapeDtypeStruct((m, n), x.dtype), compiler_paramspltpu.CompilerParams( dimension_semantics(parallel, parallel, arbitrary)), )(x, y)与第一个版本相比这里发生了三个变化累加器移入显式 scratch 空间acc_ref由scratch_shapes[pltpu.VMEM((bm, bn), jnp.float32)]分配在 VMEM 中而不是直接累加进输出z_ref因此z_ref可以保持 bf16 dtype。nsteps参数通过functools.partial(matmul_kernel, nstepsk // bk)把收缩维的步数传给内核用于判断最后一次迭代pl.program_id(2) nsteps - 1才做输出转换。jax.jit(static_argnames[bm, bk, bn])把块大小作为静态参数避免因 Python 常量变化导致重复编译。验证np.testing.assert_array_equal(x y, matmul(x, y))m, k, n 4096, 4096, 4096 k1, k2 random.split(random.key(0), 2) x random.normal(k1, (m, k), dtypejnp.bfloat16) y random.normal(k2, (k, n), dtypejnp.bfloat16) np.testing.assert_array_equal(x y, matmul(x, y))流水线内核的性能为什么块大小如此重要前面关于 FLOPs 与内存用量的分析是在整个矩阵乘法的粗粒度上进行的。但在实践中我们流水线化执行的其实是分块后的矩阵乘法——内核里有一个用小块做 matmul 的循环。因此真正重要的是内核每次实例化的 FLOPs 与内存带宽用量之比而不是全局比例。并且分块后同一个值可能被从内存中多次读取。具体来说第一个操作数的带宽为(bm * bk)乘以 grid 维度后为(bm * bk) * m // bm * n // bn * k // bk m * k * n // bn第二个操作数同理。总带宽用量为(m * k * n // bn k * n * m // bm m * n) * element_size因此bm、bk、bn对性能极其关键。即使矩阵是全世界最大的只要块尺寸选得很小每次调用内核时 FLOPs 太少、不足以掩盖后台的内存搬运就会 memory bound。直觉结论要 compute bound就把块做得尽可能大。但有两个主要约束VMEM 用量块越大VMEM 占用越多块足够大时就会耗尽 VMEM。流水线气泡pipeline bubbles块相对矩阵越大流水线循环迭代次数越少流水线首尾气泡相对总体的占比越大这部分开销不可忽视。在 Pallas 中把 matmul 性能调好本质上就是选择合适的块大小来平衡这个优化问题。实践中常见做法是对一大批候选块大小做扫描sweep、逐一 profile再挑选最优者。简单的计时实验下面用timeit做简单的计时实验并计算 FLOP/s 与相对芯片峰值的利用率百分比。注意timeit测到的还包括 Python 派发等开销因此是内核实际运行时间的上界import timeit def benchmark(f, ntrials: int 100): def run(*args, **kwargs): # Compile function first jax.block_until_ready(f(*args, **kwargs)) # Time function result timeit.timeit(lambda: jax.block_until_ready(f(*args, **kwargs)), numberntrials) time result / ntrials # print(fTime: {time}) return time return run def analyze_matmul(m: int, k: int, n: int, dtype: np.dtype, mm_func): x jnp.ones((m, k), dtypedtype) y jnp.ones((k, n), dtypedtype) time benchmark(mm_func)(x, y) print(f----- {m} x {k} x {n} -----) print(Matmul time: , time) mm_flops matmul_flops(m, k, n) / time print(Matmul FLOP/s: , mm_flops) print(fFLOP/s utilization: {mm_flops / v5e_flops * 100:.4f}%) print() print(bm128, bk128, bn128) mm functools.partial(matmul, bm128, bk128, bn128) analyze_matmul(1024, 1024, 1024, jnp.bfloat16, mm) analyze_matmul(4096, 4096, 4096, jnp.bfloat16, mm) analyze_matmul(8192, 8192, 8192, jnp.bfloat16, mm) print(bm512, bk1024, bn1024) mm functools.partial(matmul, bm512, bk1024, bn1024) analyze_matmul(1024, 1024, 1024, jnp.bfloat16, mm) analyze_matmul(4096, 4096, 4096, jnp.bfloat16, mm) analyze_matmul(8192, 8192, 8192, jnp.bfloat16, mm)原始文档中的实验结果指出更大的块尺寸帮助很大大矩阵能拿到 80–90% 的利用率但最小的那个 matmul1024³很难获得好的性能。与 XLA 生成的 matmul 对比print( XLA matmul ) mm jnp.matmul analyze_matmul(1024, 1024, 1024, jnp.bfloat16, mm) analyze_matmul(4096, 4096, 4096, jnp.bfloat16, mm) analyze_matmul(8192, 8192, 8192, jnp.bfloat16, mm)XLA 在生成 matmul 方面非常擅长文档原话XLA isverygood at generating matmuls不必期望 Pallas 超过它但 Pallas 经过非常基础的块大小调优后已经能逼近 XLA 的性能进一步的块尺寸扫描有望完全追平。这里的数值依赖具体硬件与运行环境请以你自己机器上的测量为准。模板化矩阵乘法把算子融合进内核现在我们已经有了一个基本的 matmul 内核可以尝试把其他算子**融合fuse**进去。融合的意义在于一个高效的 compute bound matmul 内核之后如果跟一个 memory bound 的独立算子如 transpose、activation会拖累整体性能把算子融进内核则省去额外的内存往返。融合 RHS 转置x y.T假设我们要计算x y.T而不是x y。朴素做法是先算y.T再喂给高效的 matmul 内核——但y.T本身不是免费的它要拷贝 $O(n^2)$ 的数据。理想情况是在一个内核里边做矩阵乘法边完成转置。加速器通常原生支持融合 RHS 转置的矩阵乘法例程TPU v5e 的 MXU 支持小块的x y.T可通过jax.lax.dot_general触发这比先转置再 matmul更高效。def matmul_kernel(x_ref, y_ref, z_ref, acc_ref, *, nsteps, transpose_rhs): pl.when(pl.program_id(2) 0) def _(): acc_ref[...] jnp.zeros_like(acc_ref) # dot_general expects a data structure (contraction_dims, batch_dims), # where contraction_dims are the set of dimensions for LHS and RHS that will # be contracted (reduced) in the matmul; batch_dims, on the other hand, are # looped over. The remaining dimensions will be the input and output dimension # of the matmul. if transpose_rhs: dims ((1,), (1,)), ((), ()) else: dims ((1,), (0,)), ((), ()) acc_ref[...] jax.lax.dot_general( x_ref[...], y_ref[...], dims, preferred_element_typejnp.float32, ) pl.when(pl.program_id(2) nsteps - 1) def _(): z_ref[...] acc_ref[...].astype(z_ref.dtype) jax.jit(static_argnames[bm, bk, bn, transpose_rhs]) def matmul( x: jax.Array, y: jax.Array, *, bm: int 128, bk: int 128, bn: int 128, transpose_rhs: bool False, ): if transpose_rhs: y y.swapaxes(0, 1) y_block_spec pl.BlockSpec((bn, bk), lambda i, j, k: (j, k)) else: y_block_spec pl.BlockSpec((bk, bn), lambda i, j, k: (k, j)) m, k x.shape _, n y.shape return pl.pallas_call( functools.partial(matmul_kernel, nstepsk // bk, transpose_rhstranspose_rhs), grid_specpltpu.PrefetchScalarGridSpec( num_scalar_prefetch0, in_specs[ pl.BlockSpec((bm, bk), lambda i, j, k: (i, k)), y_block_spec, ], out_specspl.BlockSpec((bm, bn), lambda i, j, k: (i, j)), scratch_shapes[pltpu.VMEM((bm, bn), jnp.float32)], grid(m // bm, n // bn, k // bk), ), out_shapejax.ShapeDtypeStruct((m, n), x.dtype), compiler_paramspltpu.CompilerParams( dimension_semantics(parallel, parallel, arbitrary)), )(x, y)几个值得深入理解的细节dot_general的dims约定dims ((1,), (1,)), ((), ())表示 LHS 的第 1 维与 RHS 的第 1 维收缩对应x y.Tdims ((1,), (0,)), ((), ())表示 LHS 第 1 维与 RHS 第 0 维收缩对应普通x y。剩余维度成为 matmul 的输入/输出维。逻辑转置 vs 物理转置matmul函数内部执行y y.swapaxes(0, 1)。因为在 JIT 后的 JAX 计算里维度顺序是逻辑的而非物理的重排维度不意味着物理布局变化但把数组传入pallas_call时会强制 major-to-minor 的维度顺序约束。通过在matmul内部转置y我们请求y以转置布局(n, k)进入内核而调用方仍然传入逻辑(k, n)形状的数组。基准测试的注意事项为了公平地 benchmark 转置我们希望y传入内核时已经是物理转置布局从而不把 relayout 时间算进去。因此 wrapper 里先逻辑上把它转回(k, n)再传给matmul因为matmul期望逻辑(k, n)顺序def analyze_matmul(m: int, k: int, n: int, dtype: np.dtype, mm_func, transpose_rhs: bool False): x jnp.ones((m, k), dtypedtype) if transpose_rhs: y jnp.ones((n, k), dtypedtype) jax.jit def _wrapper(x, y): y y.swapaxes(0, 1) return mm_func(x, y, transpose_rhsTrue) else: y jnp.ones((k, n), dtypedtype) _wrapper mm_func time benchmark(_wrapper)(x, y) print(f----- {m} x {k} x {n} -----) print(Matmul time: , time) mm_flops matmul_flops(m, k, n) / time print(Matmul FLOP/s: , mm_flops) print(fFLOP/s utilization: {mm_flops / v5e_flops * 100:.4f}%) print() print(bm128, bk128, bn128) mm functools.partial(matmul, bm128, bk128, bn128) analyze_matmul(1024, 1024, 1024, jnp.bfloat16, mm, transpose_rhsTrue) analyze_matmul(4096, 4096, 4096, jnp.bfloat16, mm, transpose_rhsTrue) analyze_matmul(8192, 8192, 8192, jnp.bfloat16, mm, transpose_rhsTrue) print(bm512, bk1024, bn1024) mm functools.partial(matmul, bm512, bk1024, bn1024) analyze_matmul(1024, 1024, 1024, jnp.bfloat16, mm, transpose_rhsTrue) analyze_matmul(4096, 4096, 4096, jnp.bfloat16, mm, transpose_rhsTrue) analyze_matmul(8192, 8192, 8192, jnp.bfloat16, mm, transpose_rhsTrue)原始文档的实验结论是多做了这个转置利用率却基本不变we get the same utilization despite the extra transpose这正是融合的价值。顺带一提编译器层面也有类似的融合能力CompilerParams中的fuse_transposed_lhs_in_matmul字段jax/_src/pallas/mosaic/core.py 第 114-122 行是给编译器的提示用于在 matmul 中融合转置后的 LHS例如jnp.einsum(km,kn-mn, lhs, rhs)但它只是尽力而为best-effort是否融合由编译器判断且并非总是有收益。融合激活函数融合激活函数同样常见其动机是不要用一个高效的 compute bound matmul 内核紧跟一个缓慢的 memory bound 激活内核。激活只在最后一次迭代pl.program_id(2) nsteps - 1时施加在 f32 累加结果上然后一次性 downcast 写回def matmul_kernel( x_ref, y_ref, z_ref, acc_ref, *, nsteps, transpose_rhs, activation ): pl.when(pl.program_id(2) 0) def _(): acc_ref[...] jnp.zeros_like(acc_ref) if transpose_rhs: dims ((1,), (1,)), ((), ()) else: dims ((1,), (0,)), ((), ()) acc_ref[...] jax.lax.dot_general( x_ref[...], y_ref[...], dims, preferred_element_typejnp.float32, ) pl.when(pl.program_id(2) nsteps - 1) def _(): z_ref[...] activation(acc_ref[...]).astype(z_ref.dtype) jax.jit(static_argnames[bm, bk, bn, activation]) def matmul( x: jax.Array, y: jax.Array, *, bm: int 128, bk: int 128, bn: int 128, transpose_rhs: bool False, activation: Callable[[jax.Array], jax.Array] lambda x: x, ): if transpose_rhs: y y.swapaxes(0, 1) y_block_spec pl.BlockSpec((bn, bk), lambda i, j, k: (j, k)) else: y_block_spec pl.BlockSpec((bk, bn), lambda i, j, k: (k, j)) m, k x.shape _, n y.shape return pl.pallas_call( functools.partial( matmul_kernel, nstepsk // bk, transpose_rhstranspose_rhs, activationactivation, ), grid_specpltpu.PrefetchScalarGridSpec( num_scalar_prefetch0, in_specs[ pl.BlockSpec((bm, bk), lambda i, j, k: (i, k)), y_block_spec, ], out_specspl.BlockSpec((bm, bn), lambda i, j, k: (i, j)), scratch_shapes[pltpu.VMEM((bm, bn), jnp.float32)], grid(m // bm, n // bn, k // bk), ), out_shapejax.ShapeDtypeStruct((m, n), x.dtype), compiler_paramspltpu.CompilerParams( dimension_semantics(parallel, parallel, arbitrary)), )(x, y)配套的基准测试函数支持同时开启转置与激活def analyze_matmul(m: int, k: int, n: int, dtype: np.dtype, mm_func, transpose_rhs: bool False, activation lambda x: x): x jnp.ones((m, k), dtypedtype) if transpose_rhs: y jnp.ones((n, k), dtypedtype) jax.jit def _wrapper(x, y): y y.swapaxes(0, 1) return mm_func(x, y, transpose_rhsTrue, activationactivation) else: y jnp.ones((k, n), dtypedtype) _wrapper functools.partial(mm_func, activationactivation) time benchmark(_wrapper)(x, y) print(f----- {m} x {k} x {n} -----) print(Matmul time: , time) mm_flops matmul_flops(m, k, n) / time print(Matmul FLOP/s: , mm_flops) print(fFLOP/s utilization: {mm_flops / v5e_flops * 100:.4f}%) print() activation jax.nn.relu print(bm128, bk128, bn128) mm functools.partial(matmul, bm128, bk128, bn128) analyze_matmul(1024, 1024, 1024, jnp.bfloat16, mm, activationactivation) analyze_matmul(4096, 4096, 4096, jnp.bfloat16, mm, activationactivation) analyze_matmul(8192, 8192, 8192, jnp.bfloat16, mm, activationactivation) print(bm512, bk1024, bn1024) mm functools.partial(matmul, bm512, bk1024, bn1024) analyze_matmul(1024, 1024, 1024, jnp.bfloat16, mm, activationactivation) analyze_matmul(4096, 4096, 4096, jnp.bfloat16, mm, activationactivation) analyze_matmul(8192, 8192, 8192, jnp.bfloat16, mm, activationactivation)原始文档的实验结论是融合激活函数几乎不影响利用率The additional fused activation barely affects our utilization at all。值得注意的是这里的activation是通过jax.jit(static_argnames[..., activation])作为静态参数传入的Callable类型来自typing模块的导入见文档开头的from typing import Callable。小结与实践路线本文覆盖了在 TPU 上用 Pallas 编写高效矩阵乘法的完整路径理解块矩阵乘法把大 matmul 分解为小块 matmulmatmul_small的累加TPU/GPU 硬件正是如此工作。tiling 与 pipelining用grid表达嵌套循环、用BlockSpec表达切片让内存搬运与计算重叠。第一个内核pl.pallas_callpl.when(pl.program_id(2) 0)初始化累加器 z_ref[...] x_ref[...] y_ref[...]。性能建模FLOPs2*m*k*n与最小带宽(m*k k*n m*n) * itemsize之比即算术强度与芯片的 FLOP/s ÷ 带宽v5e 约 240 FLOPs/byte比较判断 memory bound 还是 compute bound用 bf16 输入 f32 累加preferred_element_typejnp.float32显著降低带宽。块大小调优流水线化后实际看的是每次内核实例的 FLOPs/带宽bm/bk/bn越大越易 compute bound但受 VMEM 容量与流水线气泡约束实践上做块尺寸扫描 profile。算子融合模板通过jax.lax.dot_general的(contraction_dims, batch_dims)融合 RHS 转置x y.T在最后一次迭代施加activation融合激活函数二者几乎不损失利用率。原始文档还给读者留下了三个进阶练习输入融合有时需要把某个算子融合到 matmul 的输入上尝试进一步模板化 matmul 内核以支持输入融合CompilerParams.allow_input_fusion字段也可作参考见 jax/_src/pallas/mosaic/core.py 第 92-93 行。int8矩阵乘法TPU v5 原生支持int8matmul其 FLOPs 是 bf16 的两倍尝试加入支持并测量能达到的利用率仓库示例 jax/experimental/pallas/ops/tpu/matmul.py 已展示了 int8/uint8 输入时累加器切换为 int32 的做法。反向传播支持用jax.custom_vjp为matmul函数添加反向传播。继续阅读本文的配套 Jupyter Notebookdocs/pallas/tpu/matmul.ipynb前序指南《TPU Pipelining》讲解 TPU 内存层级与流水线 APIdocs/pallas/tpu/pipelining.md仓库内可直接运行的官方示例 matmul 内核jax/experimental/pallas/ops/tpu/matmul.pypltpu.CompilerParams与pltpu.PrefetchScalarGridSpec的源码定义jax/_src/pallas/mosaic/core.pyPallas TPU 相关文档索引docs/pallas/tpu/index.rst【免费下载链接】jaxComposable transformations of PythonNumPy programs: differentiate, vectorize, JIT to GPU/TPU, and more项目地址: https://gitcode.com/GitHub_Trending/ja/jax创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考