Skip to content

Gemma4 26B MoE on TPU v6e: step-time optimizations (4.80 -> 4.00 s/step, one commit per optimization) - #5474

Closed
csgoogle wants to merge 11 commits into
AI-Hypercomputer:mainfrom
csgoogle:g4-v6e128-optimizations
Closed

csgoogle wants to merge 11 commits into
AI-Hypercomputer:mainfrom
csgoogle:g4-v6e128-optimizations

Conversation

@csgoogle

Copy link
Copy Markdown
Contributor

Summary

Performance optimizations for Gemma 4 26B-A4B (MoE) pre-training on TPU v6e-128, found through XProf/HLO headroom analysis. There is one commit per optimization, and each commit message gives the issue, the change and the measured result.

End to end: 4.803 → 4.003 s/step (−16.7%), 230.7 TFLOP/s/chip, 25.1% MFU, 8,186 tok/s/chip, with float32_weight_sum=true throughout (step-0 loss 13.004 unchanged). O15 (use_iota_embed=False) is config-only, so it has no commit here.

# Commit Gate Step time (s)
O10 Gemma4: replicate pre_forward_scale_2 like the other norm scales always on 4.803 → 4.780
O11 RMSNorm / Gemma4 router: multiply norm scales in fp32 always on → 4.765
O12 Embed: shard the embedding table on the vocab dim always on → 4.758
O13 RotaryEmbedding: split-half RoPE computed directly always on → 4.709
O14 MoE: reuse the routing permutation instead of re-argsorting always on → 4.707
(O15) config: use_iota_embed=False config → 4.611
O16a MoE: split-expert wi layout + all-gather scheduling fence MAXTEXT_G4_WLAYOUT=2, MAXTEXT_G4_AG_FENCE=all O16 total: → 4.206
O16b MoE: fused TensorCore Pallas combine kernel (+ unit test) MAXTEXT_G4_COMBINE=tc (_BLOCK=256) (in O16)
O17 MoE: bind the all-gather fence to the gmm weight views MAXTEXT_G4_FENCE_SRC=gmm → 4.155
O18 Megablox: transposed-RHS gmm_v2 for backward dlhs MAXTEXT_G4_GMM_TRHS=1 → 4.078
O19 MoE combine: column-oriented backward kernel v5 MAXTEXT_G4_COMBINE_BWD=v5 (default), _RQ=1024 → 4.044
O20 MoE: unpadded tgmm_v2 output, gate/up split VJP, 2-D row info MAXTEXT_G4_TGMM_NOSLICE=1, MAXTEXT_G4_GLU_CUSTOM_VJP=1, MAXTEXT_G4_ROWINFO_IOTA=1 → 4.003

Details and profiles: Google Doc "Gemma 4 26B-A4B Pre-training on TPU v6e: Optimizations & Headroom Analysis" (internal).

Behavior changes

  • O16–O20 are opt-in env gates, default off. With the env vars unset, the MoE path is unchanged.
  • O10–O14 are always on:
    • O10: Gemma4 only.
    • O11: changes RMSNorm numerics for all models. The scale multiply now happens in fp32 with a single rounding, which is more accurate. Previously this applied only when scale_offset != 0.
    • O12: changes the Embed table sharding from ("vocab","embed_vocab") to ("embed_vocab","vocab") for all models. This can affect checkpoint sharding and restore.
    • O13: same math as before, different fusion.
    • O14: same result, one fewer sort.
    • If you'd prefer gates for O11–O12, I can add them.
  • MAXTEXT_G4_WLAYOUT=2 stores wi as [2, E/2, emb, 2*mlp]. That is a different param layout, so checkpoints written with it are not interchangeable with the default layout.
  • The env gates are a stopgap to keep this PR low-risk. I'm happy to promote them to base.yml config flags.

Recipe used for the measurements

MAXTEXT_G4_COMBINE=tc MAXTEXT_G4_COMBINE_BLOCK=256 MAXTEXT_G4_WLAYOUT=2 MAXTEXT_G4_AG_FENCE=all \
MAXTEXT_G4_FENCE_SRC=gmm MAXTEXT_G4_GMM_TRHS=1 MAXTEXT_G4_COMBINE_BWD=v5 MAXTEXT_G4_COMBINE_BWD_RQ=1024 \
MAXTEXT_G4_TGMM_NOSLICE=1 MAXTEXT_G4_GLU_CUSTOM_VJP=1 MAXTEXT_G4_ROWINFO_IOTA=1 \
python3 -m maxtext.trainers.pre_train.train maxtext/configs/base.yml model_name=gemma4-26b \
  max_target_length=16384 per_device_batch_size=2 ici_fsdp_parallelism=128 shard_exp_on_fsdp=true \
  use_iota_embed=False num_vocab_tiling=4 vocab_tiling_ag_once=true moe_use_direct_token_gather=true \
  prefuse_moe_weights=true fused_mlp=true use_tokamax_gmm=true use_gmm_v2=true use_gmm_v2_heuristic_tiling=false \
  wi_tile_{fwd,dlhs,drhs}_{batch_seq,embed_dim,mlp_dim}=512,2816,1408 \
  wo_tile_{fwd,dlhs,drhs}_{batch_seq,embed_dim,mlp_dim}=1024,2816,768 \
  sa_block_*=1024  # full remat and float32_weight_sum=true come from the model/base defaults

XLA flags: --xla_tpu_scoped_vmem_limit_kib=81920, plus the standard v6e async-collective / SparseCore-offload flag set (--xla_tpu_enable_sparse_core_collective_offload_*, --xla_tpu_offload_gather_to_sparsecore=true, --xla_tpu_enable_sparse_core_offload_queuing_in_lhs=true, --xla_tpu_sparse_core_offload_queuing_overlap_limit=8, --xla_tpu_sparse_core_all_gather_latency_multiplier=3, and others).

Testing

  • Every commit byte-compiles.
  • tests/unit/moe_combine_tc_test.py: Pallas interpret-mode forward and grad checks against a jnp reference, run on CPU at O16b (v4 backward) and at O20 (v5 backward plus variants). 18/18 pass at O16b, 21/21 at O20.
  • End-to-end runs on v6e-128 (18 steps, synthetic data), with step times as in the table above. Step-0 loss matches the baseline.
  • The branch is based on 4d291bf, which is what I measured on. It rebases onto current main with no conflicts, but I haven't re-measured on main.

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 (O10): 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 (O11): 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 (O12): 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 (O13): 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.
The direct-gather routing path already has sorted_selected_experts; the unpermute and the
custom unsort backward ran jnp.argsort on it again (262,144-element sort per layer/pass).

Change: layers/moe.py adds _route_activations_precomputed and _unsort_activations, which
scatter by the existing permutation; unpermute and the local-permute path use them.

Result (O14): 4.709 -> 4.707 s/step; XProf sort 49.1 -> 41.7 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.
Slicing the prefused wi [E, emb, 2*mlp] into w0/w1 and re-concatenating it in gmm_up (plus
the backward split/concat) cost ~410 ms/step of copies. Passing the fused 1 GB wi directly
removes the copies but creates one 1 GB SparseCore all-gather that stalls token gathers.

Change (opt-in, default off):
- MAXTEXT_G4_WLAYOUT=2: store wi as [2, E/2, emb, 2*mlp] so the two expert halves are
  contiguous leading-dim slices (~507 MB each) - no copy to extract or all-gather.
- kernels/megablox/split_gmm.py: chains two gmm_v2 calls (group_offset 0 and E/2, the second
  accumulating into the first output) and restores the rows at the half boundary; backward
  uses two group-offset tgmm_v2 calls that write drhs_lo/drhs_hi directly.
- MAXTEXT_G4_AG_FENCE=wi|all: zero-valued data dependency from the gathered weights onto the
  router logits so the wi (and wo) all-gathers are issued before routing/token permutation.

Result: standalone 4.611 -> 4.545 s/step; combined with part 2 (next commit) O16 gives
4.611 -> 4.206 s/step (-405 ms).

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 SparseCore unpermute gather (bf16[262144, 2816]) + fp32 top-k weighted sum cost
~594 ms/step (379 ms exposed SC wait), and take_along_axis on per_expert_scale occupied both
SparseCores ~2.7 ms per call.

Change (opt-in, default off):
- kernels/moe_combine_tc.py: Pallas TC kernel fusing unpermute + weighted sum (fp32
  accumulation, honoring float32_weight_sum) with a custom VJP. Because argsort is stable, rows
  of each 256-token block routed to one expert are a contiguous range of the expert-sorted
  output; the kernel DMAs aligned row windows into VMEM and reduces on the MXU.
- layers/moe.py: MAXTEXT_G4_COMBINE=tc routes unpermute through it; per_expert_scale is
  gathered with an MXU one-hot dot instead of take_along_axis.
- tests/unit/moe_combine_tc_test.py: interpret-mode fwd/grad tests vs a jnp reference.

Result: standalone 4.611 -> 4.456 s/step; O16 (both parts) 4.611 -> 4.206 s/step (-405 ms,
219.5 TFLOP/s/chip, 23.9% MFU).

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 O16 fence read element [0,0,0] of the raw all-gather outputs before the reshape/cast
consumed by split_gmm, so XLA materialized an extra weight buffer across the fwd/remat
boundary (~50 ms/step copy).

Change (opt-in): MAXTEXT_G4_FENCE_SRC=gmm reads the fence scalar from the reshaped bf16 views.

Result (O17): 4.206 -> 4.155 s/step (-51 ms); data formatting 206 -> 155.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.
Backward dlhs (grad @ W^T) needed the RHS in [G, N, K] layout, so XLA emitted HBM transpose
copies of wi_lo/wi_hi/wo every layer (~77 ms/step).

Change (opt-in): MAXTEXT_G4_GMM_TRHS=1. kernels/megablox/gmm_v2_trhs.py adds a transpose_rhs
gmm_v2 that streams weight tiles in native [G, K, N] layout and contracts in VMEM; used by
split_gmm (wi dlhs) and the new split_gmm.gmm_single (wo dlhs). gmm_trhs_bench.py is a
microbenchmark.

Result (O18): 4.155 -> 4.078 s/step (-77 ms); data formatting 155.5 -> 62.7 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.
The v4 combine backward (80.1 ms/step) re-streamed dy for every 256-row chunk and transposed
[C, 256] fp32 selection matrices in VMEM each iteration.

Change: kernels/moe_combine_tc.py backward v5 (now the default when MAXTEXT_G4_COMBINE=tc;
MAXTEXT_G4_COMBINE_BWD=v4 selects the old one): RQ (default 1024) buffer rows per chunk, W^T
pre-transposed once per grid step, selection built directly in [RQ, C], dW written as
[T, 128]. Adds moe_combine_bench.py and test variants.

Result (O19): 4.078 -> 4.044 s/step (-34 ms); moe_combine_bwd 80.1 -> 47.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.
Three leftovers: tgmm_v2 allocated wo drhs as [128, 768, 2816] and sliced to 704 in HBM
(~22 ms/step + extra zero-init); _row_info's jnp.stack/pad created a relayout + pad
(~16 ms/step); differentiating the gate/up slice emitted pads + fp32 add (~38 ms/step).

Change (opt-in):
- MAXTEXT_G4_TGMM_NOSLICE=1: pallas_mosaic_tpu_v2_tgmm_kernel allocates the exact output when
  size_k is sublane-aligned and clamps the output block to min(tile_k, size_k).
- MAXTEXT_G4_ROWINFO_IOTA=1: moe_combine_tc builds row info directly in 2-D.
- MAXTEXT_G4_GLU_CUSTOM_VJP=1: custom VJP for the gate/up split (backward = concat).

Result (O20): 4.044 -> 4.003 s/step (-41 ms; 230.7 TFLOP/s/chip, 25.1% MFU).

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 several TPU-targeted optimizations and new Pallas kernels for Mixture of Experts (MoE) layers in MaxText, specifically for Google TPU v6e. Key additions include a transposed RHS grouped matmul (GMM) kernel to avoid HBM transpose copies, and a fused TensorCore MoE 'combine' kernel (unpermute + top-k weighted sum) with f32 accumulation. It also optimizes Rotary Position Embeddings (RoPE) by enabling split-half RoPE fusion, and adjusts RMSNorm scale casting to allow XLA to hoist scale-gradient all-reduces out of scanned loops. The review feedback identifies two critical issues: a mismatched tile size configuration in moe.py where wi_tile_fwd_batch_seq is incorrectly used for wo GMM padding, and an invalid JAX API call to jax.typeof(x) in moe_combine_tc.py which will cause a runtime AttributeError.

Comment thread src/maxtext/layers/moe.py
and self.mesh.devices.flat[0].platform == "tpu"
):
# MAXTEXT_G4_GMM_TRHS: wo gmm with NT-kernel dlhs (no wo^T copy).
_wo_in, _wo_pad = max_utils.maybe_pad(intermediate_layer, fwd_tile("wi_tile_fwd_batch_seq"))

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

For the wo GMM padding, wi_tile_fwd_batch_seq is incorrectly used instead of wo_tile_fwd_batch_seq. Since wi and wo can have different batch tile sizes (e.g., 512 vs 1024 in the recommended recipe), padding to wi_tile_fwd_batch_seq can result in an input size that is not a multiple of the wo GMM tile size, leading to runtime crashes or out-of-bounds errors. Please use wo_tile_fwd_batch_seq instead.

Suggested change
_wo_in, _wo_pad = max_utils.maybe_pad(intermediate_layer, fwd_tile("wi_tile_fwd_batch_seq"))
_wo_in, _wo_pad = max_utils.maybe_pad(intermediate_layer, fwd_tile("wo_tile_fwd_batch_seq"))



def _vma(x) -> frozenset:
mat = getattr(jax.typeof(x), "manual_axis_type", None)

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

high

The function jax.typeof(x) is not a valid JAX API and will raise an AttributeError at runtime. Please use the standard Python type(x) instead to retrieve the type of the array and check for the manual_axis_type attribute.

Suggested change
mat = getattr(jax.typeof(x), "manual_axis_type", None)
mat = getattr(type(x), "manual_axis_type", None)

@csgoogle

csgoogle commented Oct 6, 2026

Copy link
Copy Markdown
Contributor Author

Superseded by #5572, #5573, #5574 (split by area, rebuilt without the all-gather scheduling fence).

@csgoogle csgoogle closed this Oct 6, 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.

1 participant