Skip to content

Gemma4 26B MoE on TPU v6e: MaxText layer/model optimizations (one commit per change) - #5572

Draft
csgoogle wants to merge 7 commits into
AI-Hypercomputer:mainfrom
csgoogle:g4-v6e128-maxtext
Draft

csgoogle wants to merge 7 commits into
AI-Hypercomputer:mainfrom
csgoogle:g4-v6e128-maxtext

Conversation

@csgoogle

@csgoogle csgoogle commented Oct 6, 2026 •

Copy link
Copy Markdown
Contributor

Draft. Part 1 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).

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 (baseline) to 3.82 s across the three PRs; this PR holds the 7 pure MaxText layer / model changes — one commit per optimization, requiring no new config flags, no custom Pallas kernels, and no checkpoint layout changes.

Base: 4d291bf02 (the commit the stack was developed and measured on).

Commits

# Commit What / why s/step before → after
1 Gemma4: replicate pre_forward_scale_2 like the other norm scales Removes a per-layer all-gather of a norm scale that every other scale avoided 4.805 → 4.793
2 RMSNorm / Gemma4 router: multiply norm scales in fp32 Removes bf16 round trips (converts + copies) around the scale multiply → 4.773
3 Embed: shard the embedding table on the vocabulary dimension Avoids the per-step all-gather of the [vocab, emb] table under FSDP → 4.762
4 RotaryEmbedding: compute split-half RoPE directly One rotate + multiply instead of concat / slice copies → 4.716
5 Vocabulary tiling: tile along unsharded seq_len to avoid FSDP all-to-all The previous tiling axis forced an all-to-all per tile → 4.587 (after use_iota_embed=False, a config change measured at 4.716 → 4.613)
6 MoE: TensorCore top_k_weights, per-expert scale and group_size; single rsqrt in RMSNorm Keeps small routing ops off SparseCore, where they serialize behind the weight all-gathers → 4.527
7 MoE: gate/up split custom VJP Replaces two zero-pads + fp32 add in backward of the (gate, up) slice with a single concatenate 4.309 → 4.285

Commit messages carry the per-commit details and measurements. All 7 commits are active by default and need no configuration flags.

Tests

  • Existing unit tests for the touched layers; pyink 24.10.1 (--pyink-indentation=2 --line-length=122) clean.

Related PRs

pre_forward_scale_2 was the only RMSNorm scale sharded as ("embed",); all the others are
("norm",) (replicated). Its gradient reduction could not be batched with the others, so each
layer paid its own sharded collective plus relayout copies.

Change: models/gemma4.py Gemma4MoE pre_forward_scale_2 sharding -> ("norm",).

Result: 4.803 -> 4.780 s/step (-23 ms); all-reduce 117.7 -> 52.5 ms/step.

Measured on Gemma4-26B-A4B pre-training, TPU v6e-128, seq 16384, per_device_batch_size=2, FSDP=128, full remat, float32_weight_sum=true.
RMSNorm cast its fp32 scale to the compute dtype before the multiply. In backward, the bf16
scale-gradient all-reduce was followed by a convert into the fp32 scan carry, which blocked
XLA's while-loop all-reduce code motion, so 8 norm-gradient reductions per local layer stayed
inside the scanned loop (plus bf16<->fp32 relayout copies).

Change: layers/normalizations.py RMSNorm always multiplies the fp32 normalized activations by
the fp32 scale and rounds once; models/gemma4.py computes the router gate input in fp32.

Result: 4.780 -> 4.765 s/step (-15 ms); all-reduce 52.5 -> 20.3 ms/step, data
formatting -26.5 ms/step. Numerically more accurate (single rounding).

Measured on Gemma4-26B-A4B pre-training, TPU v6e-128, seq 16384, per_device_batch_size=2, FSDP=128, full remat, float32_weight_sum=true.
The table was declared ("vocab", "embed_vocab"), which under FSDP shards the 2816-wide
embedding dim (splits badly across 128 chips) and forced compaction reshapes each step.

Change: layers/embeddings.py Embed sharding -> ("embed_vocab", "vocab").

Result: 4.765 -> 4.758 s/step (-7 ms); data formatting -8.9 ms/step.

Measured on Gemma4-26B-A4B pre-training, TPU v6e-128, seq 16384, per_device_batch_size=2, FSDP=128, full remat, float32_weight_sum=true.
Building full-width sin/cos by concatenation and applying rotate_half made XLA split q/k
norm + RoPE into three kernels per tensor (and recompute the norm).

Change: layers/embeddings.py RotaryEmbedding computes out1 = x1*cos - x2*sin,
out2 = x2*cos + x1*sin on the half-width sin/cos and concatenates once (identical math).

Result: 4.758 -> 4.709 s/step (-49 ms); norm + RoPE + transpose fuse into one kernel.

Measured on Gemma4-26B-A4B pre-training, TPU v6e-128, seq 16384, per_device_batch_size=2, FSDP=128, full remat, float32_weight_sum=true.

@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 moe_split_expert_weight_layout configuration, which stores the prefused MoE wi weight in a split layout to eliminate concatenation and splitting overhead during forward and backward passes. It implements a new split_gmm kernel to chain grouped matrix multiplications and optimizes several MoE activation routing, sorting, and normalization paths to improve XLA fusion and performance. Feedback on the changes suggests handling None gradients defensively in the custom VJP _split_gate_up_bwd to prevent potential runtime errors during partial differentiation.

Comment thread src/maxtext/layers/moe.py
Comment on lines +91 to +93
def _split_gate_up_bwd(_residuals: None, grads: tuple[jax.Array, jax.Array]) -> tuple[jax.Array]:
grad_gate, grad_up = grads
return (jnp.concatenate([grad_gate, grad_up], axis=-1),)

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

medium

In JAX custom VJPs for multi-output functions, if one of the outputs is not differentiated (e.g., during testing, debugging, or partial differentiation), its corresponding cotangent in grads can be None. To prevent runtime errors when jnp.concatenate receives None, it is safer to handle None gradients defensively by replacing them with zeros of the same shape and dtype as the other gradient.

Suggested change
def _split_gate_up_bwd(_residuals: None, grads: tuple[jax.Array, jax.Array]) -> tuple[jax.Array]:
grad_gate, grad_up = grads
return (jnp.concatenate([grad_gate, grad_up], axis=-1),)
def _split_gate_up_bwd(_residuals: None, grads: tuple[jax.Array, jax.Array]) -> tuple[jax.Array]:
grad_gate, grad_up = grads
if grad_gate is None and grad_up is None:
return (None,)
if grad_gate is None:
grad_gate = jnp.zeros_like(grad_up)
elif grad_up is None:
grad_up = jnp.zeros_like(grad_gate)
return (jnp.concatenate([grad_gate, grad_up], axis=-1),)

vocab_tiling_nnx_loss flattened (batch, seq_len, emb) directly to
(num_tiles, batch * seq_len // num_tiles, emb), which split the FSDP-sharded batch
dimension across vocabulary tiles and forced cross-chip all-to-all + collective-permute
(214 MB each) in forward and backward.

- utils/vocabulary_tiling.py: when seq_len % num_tiles == 0, tile along the unsharded
  seq_len dimension (swapaxes of reshape(batch, num_tiles, seq_len // num_tiles, ...))
  in forward and backward.
- trainers/pre_train/train.py: compute total_weights before the model forward so its
  scalar all-reduce overlaps with layer 0 (always on; no env gate).

Result: 4.003 -> 3.966 s/step (-37 ms; 232.9 TFLOP/s/chip, 25.37% MFU).
Measured on Gemma4-26B-A4B pre-training, TPU v6e-128, seq 16384, per_device_batch_size=2, FSDP=128, full remat.
…e rsqrt in RMSNorm

Three tiny gathers in the Gemma4 MoE routing path were being offloaded by XLA to
SparseCore, where they serialize behind the weight all-gathers and token permutes:
- top_k_weights = take_along_axis(router_probs, top_k_indices) -> one-hot compare/select
  + reduce over num_experts on the TensorCore (exact: exactly one non-zero term).
- per_expert_scale lookup -> same one-hot select/reduce; its VJP is a reduction, not a scatter.
- group_size = bincount(selected_experts) -> masked sum over an iota of num_experts
  (num_experts <= 256), avoiding a SparseCore scatter-add.
RMSNorm: compute x * rsqrt(mean2 + eps) once in fp32 and reuse it for both the
with_scale and no-scale paths (previously the rsqrt was evaluated twice).
Always on; no env gate.
@csgoogle
csgoogle force-pushed the g4-v6e128-maxtext branch 2 times, most recently from 4414b23 to 1201eb5 Compare October 7, 2026 12:50
Differentiating the gate/up slice of the fused [rows, 2N] wi output emitted two
zero-pads + an fp32 add (~38 ms/step on Gemma4-26B-A4B, v6e-128). _split_gate_up
wraps the two half-slices in a custom VJP whose backward is a single concatenate
of the two half-gradients. Always on; no config option.

This branch has not been deployed

No deployments
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.

1 participant