Skip to content

Megablox: transposed-RHS gmm_v2 kernel for the backward dlhs - #5574

Closed
csgoogle wants to merge 1 commit into
AI-Hypercomputer:mainfrom
csgoogle:g4-v6e128-gmm-trhs
Closed

csgoogle wants to merge 1 commit into
AI-Hypercomputer:mainfrom
csgoogle:g4-v6e128-gmm-trhs

Conversation

@csgoogle

@csgoogle csgoogle commented Oct 6, 2026 •

Copy link
Copy Markdown
Contributor

Draft. Part 3 of 3 of the Gemma4-26B-A4B / TPU v6e optimization stack. Supersedes draft PR #5474 (split by area and rebuilt without the all-gather scheduling fence, with every environment-variable gate replaced by a config key).

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^T is computed as gmm_v2(dout, W.swapaxes(1, 2)) because gmm_v2 has no native transpose_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 static transpose_rhs option. 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.gmm gets a non-diff use_gmm_v2_transposed_rhs_dlhs argument and dispatches the dlhs to it from the plain path (forward and drhs unchanged; not used for transpose_rhs=True forwards 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

# Commit What / why s/step before → after
1 Megablox: transposed-RHS gmm_v2 kernel for the backward dlhs (use_gmm_v2_transposed_rhs_dlhs) NT gmm_v2 kernel + ops.gmm dispatch + config key 4.036 → 3.957 (on top of Parts 1 and 2)

How to enable

use_gmm_v2_transposed_rhs_dlhs: true   # requires use_gmm_v2=true, no quantization

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/without group_offset, and through the ops.gmm custom VJP at the wi / wo shapes). The kernel needs a TPU (it uses pltpu.get_tpu_info()), so there is no CPU unit test; the ops.gmm plumbing was checked on CPU through the Megablox interpret path.
  • pyink 24.10.1 (--pyink-indentation=2 --line-length=122) clean.

Follow-up

With this PR alone the kernel is used for every ops.gmm dlhs (i.e. wi and wo in the default fused layout). Part 1's split-expert wi layout (kernels/megablox/split_gmm.py) still computes its dlhs with gmm_v2 on 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 split wi halves and wo).

Related PRs

…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.

@gemini-code-assist gemini-code-assist Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Comment on lines +63 to +70
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

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

critical

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.

Suggested change
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))

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

critical

Similarly, jax.block_until_ready does not exist. Please call .block_until_ready() directly on the JAX array returned by the JITted function.

Suggested change
wt = jax.block_until_ready(jax.jit(lambda w: w.swapaxes(1, 2))(w))
wt = jax.jit(lambda w: w.swapaxes(1, 2))(w).block_until_ready()

@shuningjin

Copy link
Copy Markdown
Collaborator

Note: This has already been implemented on main in #4741 (pallas_mosaic_tpu_v2_gmm_kernel.gmm_v2(..., transpose_rhs=True), gated by moe_gmm_v2_dlhs_transpose_rhs: true).

@csgoogle

csgoogle commented Oct 8, 2026

Copy link
Copy Markdown
Contributor Author

Note: This has already been implemented on main in #4741 (pallas_mosaic_tpu_v2_gmm_kernel.gmm_v2(..., transpose_rhs=True), gated by moe_gmm_v2_dlhs_transpose_rhs: true).

I see, thanks!!

@csgoogle

Copy link
Copy Markdown
Contributor Author

Closing in favor of the existing moe_gmm_v2_dlhs_transpose_rhs: true config from #4741 — thank you @shuningjin!

@csgoogle csgoogle closed this Oct 10, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants