Repository navigation
Conversation
…m_v2_transposed_rhs_dlhs`) On the Tokamax GMM v2 path the backward dlhs is computed as gmm_v2(dout, W.swapaxes(1, 2)) because gmm_v2 has no native transpose_rhs support. XLA materializes the swapaxes as a weight-sized transpose copy per grouped matmul inside the remat (on Gemma4-26B-A4B / v6e-128 ~87 ms/step of 507 MB (wi) / 253 MB (wo) copies). - kernels/megablox/gmm_v2_trhs.py: the GMM v2 Pallas kernel with a static `transpose_rhs` option. When set, rhs is given as [groups, n, k] (the untransposed forward weight), the rhs BlockSpec reads [tile_n, tile_k] tiles and the MXU contracts the lane dimension of both operands (NT matmul), so dlhs reads the gathered weight in place. Only the unquantized, no-bias, no-fused-activation path is supported with transpose_rhs=True. - kernels/megablox/ops.py: new `gmm(..., use_gmm_v2_transposed_rhs_dlhs=False)` non-diff argument; `_dlhs_run_tokamax_v2` dispatches to the transposed-RHS kernel when it is set, the forward was not `transpose_rhs` and the weight is not a QArray. Forward and drhs are unchanged. - layers/moe.py: passes `config.use_gmm_v2_transposed_rhs_dlhs` to the gmm call. - configs: new `use_gmm_v2_transposed_rhs_dlhs: false` (base.yml, types.py) with validation (requires `use_gmm_v2=True`, no quantization); docs entry; pyconfig test. - kernels/megablox/gmm_trhs_bench.py: single-chip correctness + speed gate comparing the NT kernel against gmm_v2 on the explicitly transposed weight, at kernel level and through the `ops.gmm` custom VJP. Measured on Gemma4-26B-A4B, v6e-128 (seq 16k, per-device batch 2, FSDP-128, full remat): 4.036 -> 3.957 s/step on top of the layer/model optimizations and fused combine / permute kernels.
There was a problem hiding this comment.
Code Review
This pull request introduces the use_gmm_v2_transposed_rhs_dlhs configuration option, which optimizes the backward dlhs path in Tokamax GMM v2 by using a transposed-RHS kernel (gmm_v2_trhs.py) to read forward weights in place, avoiding an expensive transpose copy. It also adds a benchmark script gmm_trhs_bench.py and updates configuration validation and unit tests. The review comments correctly identify a critical issue in the benchmark script where jax.block_until_ready is incorrectly called as a module function instead of a method on the JAX array.
| def _timeit(fn, args, iters): | ||
| out = fn(*args) | ||
| jax.block_until_ready(out) | ||
| t0 = time.perf_counter() | ||
| for _ in range(iters): | ||
| out = fn(*args) | ||
| jax.block_until_ready(out) | ||
| return out, (time.perf_counter() - t0) * 1e3 / iters |
There was a problem hiding this comment.
In JAX, block_until_ready() is a method on JAX arrays (e.g., out.block_until_ready()), not a function in the jax module. Calling jax.block_until_ready(out) will raise an AttributeError: module 'jax' has no attribute 'block_until_ready'. Please call the method directly on the array instead.
| def _timeit(fn, args, iters): | |
| out = fn(*args) | |
| jax.block_until_ready(out) | |
| t0 = time.perf_counter() | |
| for _ in range(iters): | |
| out = fn(*args) | |
| jax.block_until_ready(out) | |
| return out, (time.perf_counter() - t0) * 1e3 / iters | |
| def _timeit(fn, args, iters): | |
| out = fn(*args) | |
| out.block_until_ready() | |
| t0 = time.perf_counter() | |
| for _ in range(iters): | |
| out = fn(*args) | |
| out.block_until_ready() | |
| return out, (time.perf_counter() - t0) * 1e3 / iters |
| return ref_t(d, w.swapaxes(1, 2)) | ||
|
|
||
| o_nt, ms_nt = _timeit(jax.jit(nt), (dout, w), args.iters) | ||
| wt = jax.block_until_ready(jax.jit(lambda w: w.swapaxes(1, 2))(w)) |
There was a problem hiding this comment.
|
Note: This has already been implemented on |
I see, thanks!! |
|
Closing in favor of the existing |
Context
Workload: Gemma4-26B-A4B (128 experts, top-8), v6e-128, seq 16k, per-device batch 2, FSDP-128, full remat, bf16,
sparse_matmul+ Tokamax GMM v2. Step time went from 7.58 s to 3.96 s across the three PRs.On the Tokamax GMM v2 path the backward
dlhs = dout @ W^Tis computed asgmm_v2(dout, W.swapaxes(1, 2))because gmm_v2 has no nativetranspose_rhs. XLA materializes the swapaxes as a weight-sized transpose copy per grouped matmul inside the remat — on this workload ~87 ms/step of 507 MB (wi) / 253 MB (wo) copies.This PR adds
src/maxtext/kernels/megablox/gmm_v2_trhs.py: the GMM v2 Pallas kernel with a statictranspose_rhsoption. When set, the RHS is the untransposed forward weight[groups, n, k], the RHS BlockSpec reads[tile_n, tile_k]tiles and the MXU contracts the lane dimension of both operands (NT matmul), so the gathered weight is read in place.ops.gmmgets a non-diffuse_gmm_v2_transposed_rhs_dlhsargument and dispatches the dlhs to it from the plain path (forward and drhs unchanged; not used fortranspose_rhs=Trueforwards or quantized weights).Base:
4d291bf02(the commit the stack was developed and measured on). Independent of Parts 1 and 2 at the code level.Commits
use_gmm_v2_transposed_rhs_dlhs)ops.gmmdispatch + config keyHow to enable
Tests
tests/unit/pyconfig_test.py: validation of the key.src/maxtext/kernels/megablox/gmm_trhs_bench.py: single-chip correctness + speed gate comparing the NT kernel against gmm_v2 on the explicitly transposed weight (kernel level with/withoutgroup_offset, and through theops.gmmcustom VJP at thewi/woshapes). The kernel needs a TPU (it usespltpu.get_tpu_info()), so there is no CPU unit test; theops.gmmplumbing was checked on CPU through the Megablox interpret path.--pyink-indentation=2 --line-length=122) clean.Follow-up
With this PR alone the kernel is used for every
ops.gmmdlhs (i.e.wiandwoin the default fused layout). Part 1's split-expertwilayout (kernels/megablox/split_gmm.py) still computes its dlhs withgmm_v2on the transposed halves; once both PRs land, a small follow-up routes that path through the transposed-RHS kernel as well (the 4.036 → 3.957 s/step above was measured with the kernel applied to both the splitwihalves andwo).Related PRs