Repository navigation
Conversation
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.
There was a problem hiding this comment.
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.
| 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")) |
There was a problem hiding this comment.
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.
| _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) |
There was a problem hiding this comment.
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.
| mat = getattr(jax.typeof(x), "manual_axis_type", None) | |
| mat = getattr(type(x), "manual_axis_type", None) |
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=truethroughout (step-0 loss 13.004 unchanged). O15 (use_iota_embed=False) is config-only, so it has no commit here.pre_forward_scale_2like the other norm scalesuse_iota_embed=Falsewilayout + all-gather scheduling fenceMAXTEXT_G4_WLAYOUT=2,MAXTEXT_G4_AG_FENCE=allMAXTEXT_G4_COMBINE=tc(_BLOCK=256)MAXTEXT_G4_FENCE_SRC=gmmgmm_v2for backward dlhsMAXTEXT_G4_GMM_TRHS=1MAXTEXT_G4_COMBINE_BWD=v5(default),_RQ=1024tgmm_v2output, gate/up split VJP, 2-D row infoMAXTEXT_G4_TGMM_NOSLICE=1,MAXTEXT_G4_GLU_CUSTOM_VJP=1,MAXTEXT_G4_ROWINFO_IOTA=1Details and profiles: Google Doc "Gemma 4 26B-A4B Pre-training on TPU v6e: Optimizations & Headroom Analysis" (internal).
Behavior changes
scale_offset != 0.Embedtable sharding from("vocab","embed_vocab")to("embed_vocab","vocab")for all models. This can affect checkpoint sharding and restore.MAXTEXT_G4_WLAYOUT=2storeswias[2, E/2, emb, 2*mlp]. That is a different param layout, so checkpoints written with it are not interchangeable with the default layout.base.ymlconfig flags.Recipe used for the measurements
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
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.4d291bf, which is what I measured on. It rebases onto currentmainwith no conflicts, but I haven't re-measured onmain.