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: 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.
There was a problem hiding this comment.
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.
| 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),) |
There was a problem hiding this comment.
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.
| 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.
4414b23 to
1201eb5
Compare
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.
273c96f to
102a322
Compare
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
pre_forward_scale_2like the other norm scales[vocab, emb]table under FSDPseq_lento avoid FSDP all-to-alluse_iota_embed=False, a config change measured at 4.716 → 4.613)top_k_weights, per-expert scale andgroup_size; single rsqrt in RMSNorm(gate, up)slice with a single concatenateCommit messages carry the per-commit details and measurements. All 7 commits are active by default and need no configuration flags.
Tests
--pyink-indentation=2 --line-length=122) clean.Related PRs
wilayout — MoE: fused TensorCore Pallas permute/combine kernels and split-expert wi layout for TPU v6e #5573