Skip to content

MoE: fused TensorCore Pallas permute/combine kernels and split-expert wi layout for TPU v6e - #5573

Draft
csgoogle wants to merge 3 commits into
AI-Hypercomputer:mainfrom
csgoogle:g4-v6e128-tc-gather-kernels
Draft

csgoogle wants to merge 3 commits into
AI-Hypercomputer:mainfrom
csgoogle:g4-v6e128-tc-gather-kernels

Conversation

@csgoogle

@csgoogle csgoogle commented Oct 6, 2026 •

Copy link
Copy Markdown
Contributor

Draft. Part 2 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, moe_use_direct_token_gather=true. Step time went from 7.58 s to 3.82 s across the three PRs.

On v6e:

  1. The MoE token gathers (the permute into expert-sorted order, the unpermute + Top-K weighted sum "combine", and their backward scatters) are offloaded by XLA to SparseCore, where they serialize behind the expert-weight all-gathers and show up as exposed SparseCore / collective wait around the grouped matmuls. Commits 1 and 2 add two fused TensorCore Pallas kernels (src/maxtext/kernels/moe_combine_tc.py) that keep both on the TensorCore and never materialize the [tokens * top_k, emb] unsorted copy.
  2. With prefuse_moe_weights=true, wi is stored as [E, emb, 2*mlp] (~1 GB across FSDP), which either incurs ~400 ms/step of concat/split copies around gmm_v2 or stalls SparseCore behind a single ~1 GB all-gather. Commit 3 adds moe_split_expert_weight_layout=true, storing wi as [2, E, emb/2, 2*mlp] so its two leading slices (W[:E/2] and W[E/2:], ~507 MB each) are all-gathered separately and consumed by two chained gmm_v2 calls (src/maxtext/kernels/megablox/split_gmm.py) into one shared output buffer.

Base: 4d291bf02 (the commit the stack was developed and measured on); independent of Part 1 except for the measured numbers.

Commits

# Commit What / why s/step before → after
1 MoE: fused TensorCore Pallas combine kernel (moe_combine_kernel="tc") Fused unpermute + Top-K weighted sum; backward is a window-DMA permutation for dx and a masked lane reduction for dw 4.527 → 4.233 (−294 ms standalone, works with or without #3)
2 MoE: fused TensorCore Pallas permute kernel (moe_permute_kernel="tc") Token gather into expert-sorted order with unit weights; backward = combine forward with unit weights. Moves the group_size histogram above the gather (no behaviour change) 4.233 → 4.309 standalone, but −492 ms of exposed collective wait once #3 is on
3 MoE: split-expert wi layout with chained gmm_v2 (moe_split_expert_weight_layout) Stores the prefused wi as two expert halves ([2, E, emb/2, 2*mlp]) that are all-gathered separately and consumed by two chained GMM v2 calls (kernels/megablox/split_gmm.py), so the ~1 GB fused weight is never concatenated or split 4.285 → 4.036 (−249 ms)

Only gated code is added to layers/moe.py; with the defaults (moe_combine_kernel="xla", moe_permute_kernel="xla", moe_split_expert_weight_layout=false) the MoE path is unchanged.

How to enable

# Recipe 1 (Standard checkpoint layout, works at any FSDP size including FSDP=256; saves ~294 ms/step):
moe_combine_kernel: "tc"               # requires sparse_matmul=true, use_ring_of_experts=false, use_ragged_sort=false
moe_permute_kernel: "xla"
moe_split_expert_weight_layout: false

# Recipe 2 (Maximum throughput with split-expert layout, saves an extra ~197 ms/step; requires fsdp * cp <= num_experts):
moe_combine_kernel: "tc"
moe_permute_kernel: "tc"               # requires moe_use_direct_token_gather=true, use_ring_of_experts=false
moe_split_expert_weight_layout: true   # requires prefuse_moe_weights=true, shard_exp_on_fsdp=true, use_gmm_v2=true

When using Recipe 2, combine moe_permute_kernel="tc" + moe_split_expert_weight_layout=true with the XLA flags
--xla_tpu_offload_gather_to_sparsecore=false --xla_tpu_offload_all_supported_gathers_to_sparsecore=false
so that no MoE gather is left serialized on SparseCore behind the expert-weight all-gathers.

Note

With moe_split_expert_weight_layout=true (Recipe 2), the checkpoint shape of wi changes to [2, *scan_dims, num_experts, emb/2, 2*mlp] (a reshape of the fused weight). Use split_gmm.convert_checkpoint_tree(state, to_split=True/False) (included in src/maxtext/kernels/megablox/split_gmm.py) to convert checkpoints in either direction. With moe_split_expert_weight_layout=false (Recipe 1), checkpoint shapes are unchanged.

Tests

  • tests/unit/moe_combine_tc_test.py (new, Pallas interpret mode on CPU with race detection, compiled on TPU): combine forward (≤ 1 bf16 ulp vs reference) and backward (dx exact, dw to 1e-2) over block sizes / top_k / expert counts / skewed routing / shared boundary granules / align=16 / f32 weights; jax.checkpoint + lax.scan gradient check; permute forward (exact) and backward; fallbacks. Also checks the per-block scatter staging table (which boundary granules are read from HBM vs carried from the previous block's output buffer, never twice). 18 tests, ~55 s on CPU.
  • tests/unit/split_gmm_test.py (new): forward + gradients of the chained split GMM against a dense reference, using a pure-JAX model of gmm_v2 that reproduces the boundary-tile masking (CPU).
  • tests/unit/pyconfig_test.py: validation of moe_combine_kernel, moe_permute_kernel, and moe_split_expert_weight_layout.
  • src/maxtext/kernels/moe_combine_bench.py: single-chip correctness + benchmark script.
  • pyink 24.10.1 (--pyink-indentation=2 --line-length=122) clean.

Related PRs

@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 fused TensorCore Pallas kernels for MoE "combine" and "permute" operations to optimize TPU performance by avoiding large unsorted memory copies. It integrates these kernels into the MoE layer, exposes them via new configuration options, and adds corresponding benchmarks, unit tests, and validation checks. The feedback suggests adding a safety check in kernel_supported to ensure that the number of experts per token (k) does not exceed _LANES (128), preventing potential layout failures in the Pallas kernel.

Comment on lines +190 to +203
def kernel_supported(x_shape, x_dtype, num_tokens, k, block_tokens, align=8) -> bool:
"""True iff the Pallas kernels can handle this problem (see the module docstring)."""
if len(x_shape) != 2:
return False
r, e = x_shape
return (
jnp.dtype(x_dtype) == jnp.bfloat16
and e % _LANES == 0
and r == num_tokens * k
and r % align == 0
and block_tokens % _LANES == 0
and num_tokens % block_tokens == 0
and num_tokens <= _MAX_TOKENS
)

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

The Pallas kernel implementation relies on _LANES (128) for the layout of top-k weights and row info in VMEM. If k (the number of experts per token) exceeds _LANES, operations like _padded_weights (which uses _LANES - w.shape[1]) will fail or produce incorrect results. Adding a check for k <= _LANES in kernel_supported ensures that the kernel safely falls back to the XLA path for unsupported shapes.

Suggested change
def kernel_supported(x_shape, x_dtype, num_tokens, k, block_tokens, align=8) -> bool:
"""True iff the Pallas kernels can handle this problem (see the module docstring)."""
if len(x_shape) != 2:
return False
r, e = x_shape
return (
jnp.dtype(x_dtype) == jnp.bfloat16
and e % _LANES == 0
and r == num_tokens * k
and r % align == 0
and block_tokens % _LANES == 0
and num_tokens % block_tokens == 0
and num_tokens <= _MAX_TOKENS
)
def kernel_supported(x_shape, x_dtype, num_tokens, k, block_tokens, align=8) -> bool:
"""True iff the Pallas kernels can handle this problem (see the module docstring)."""
if len(x_shape) != 2:
return False
r, e = x_shape
return (
jnp.dtype(x_dtype) == jnp.bfloat16
and e % _LANES == 0
and r == num_tokens * k
and r % align == 0
and block_tokens % _LANES == 0
and num_tokens % block_tokens == 0
and num_tokens <= _MAX_TOKENS
and k <= _LANES
)

@csgoogle
csgoogle force-pushed the g4-v6e128-tc-gather-kernels branch from f6b2dcb to 80be58d Compare October 7, 2026 14:35
@csgoogle csgoogle changed the title MoE: fused TensorCore Pallas permute / combine (gather + gather-sum) kernels for TPU v6e MoE: fused TensorCore Pallas permute/combine kernels and split-expert wi layout for TPU v6e Oct 7, 2026
@csgoogle
csgoogle force-pushed the g4-v6e128-tc-gather-kernels branch from 80be58d to dce4481 Compare October 7, 2026 14:57
csgoogle and others added 3 commits October 7, 2026 18:38
The MoE "combine" (unpermute of the expert outputs back into token order followed by the
Top-K weighted sum) is a [tokens * top_k, emb] gather plus an einsum. On v6e the gather is
offloaded to SparseCore, where it queues behind the expert-weight all-gathers, and the
unsorted copy is materialized in HBM in both forward and backward.

- kernels/moe_combine_tc.py: fused TensorCore kernel with a custom VJP. Tokens are processed
  in blocks of 256; because the expert sort is stable, the rows of one token block that went
  to one expert form a single contiguous row range, so each grid step DMAs at most one row
  window per expert into VMEM and does the reorder + weighting as an MXU matmul against a
  0/1-structured selection matrix built from a per-row (token, k) info array. Forward:
  y = S @ buf (f32 accumulation, the same arithmetic as `float32_weight_sum=True`); backward:
  dx = S^T @ dy (a permutation, written back with window DMAs and read-modify-write of the
  shared boundary granules) and dw via a masked lane reduction. Non-bf16 routing weights use
  an exact bf16 hi + lo split. Falls back to a pure-jnp reference when the preconditions do
  not hold (see `kernel_supported`).
- layers/moe.py: `moe_combine_kernel="tc"` routes the local (no expert parallelism, no
  ring-of-experts, no ragged sort) unpermute through the kernel; `"xla"` (default) keeps the
  gather + einsum.
- configs: new `moe_combine_kernel: "xla" | "tc"` (base.yml, types.py) with validation;
  docs/reference/core_concepts/moe_configuration.md entry.
- tests/unit/moe_combine_tc_test.py: interpret-mode forward/backward checks against the
  reference (exact dx, 1-ulp y), f32 weights, shared granules, align=16, scan + remat.
- kernels/moe_combine_bench.py: single-chip correctness check + benchmark (argparse flags).

Measured on Gemma4-26B-A4B, v6e-128 (seq 16k, per-device batch 2, FSDP-128, full remat):
4.527 -> 4.233 s/step.
With `moe_use_direct_token_gather=True` the token permute (gather of the [tokens, emb]
activations into expert-sorted [tokens * top_k, emb] order) and its backward scatter-add are
offloaded by XLA to SparseCore on v6e, where they queue behind the expert-weight all-gathers
and show up as exposed SparseCore wait in front of the first grouped matmul.

- kernels/moe_combine_tc.py: `permute(x, sort_idx, group_sizes, num_experts_per_tok)` reuses
  the combine kernel's block / row-window DMA scheme with unit weights: the forward is a
  one-hot MXU selection per token block, the backward (dx[t] = sum of the top_k rows of
  token t) is the combine forward with unit weights. Custom VJP; returns None when
  `kernel_supported` does not hold so the caller can fall back.
- layers/moe.py: `moe_permute_kernel="tc"` routes the direct-token-gather path through the
  kernel when there is no expert parallelism / ring-of-experts and the activations are not
  quantized; `"xla"` (default) keeps the XLA gather. The expert histogram (`group_size`) is
  now computed before the gather so it can be passed to the kernel (no behaviour change).
- configs: new `moe_permute_kernel: "xla" | "tc"` (base.yml, types.py) with validation
  (requires `moe_use_direct_token_gather=True`, `use_ring_of_experts=False`); docs entry.
- tests/unit/moe_combine_tc_test.py: interpret-mode permute forward/backward (uniform and
  skewed routing) and the unsupported-shape fallback.

Intended configuration: `moe_permute_kernel="tc"` together with the XLA flags
`--xla_tpu_offload_gather_to_sparsecore=false --xla_tpu_offload_all_supported_gathers_to_sparsecore=false`,
so that no MoE gather is serialized on SparseCore behind the expert-weight all-gathers.

Measured on Gemma4-26B-A4B, v6e-128 (seq 16k, per-device batch 2, FSDP-128, full remat):
4.233 -> 4.309 s/step standalone (the kernel is slightly slower than the SparseCore gather in
isolation), but it removes ~492 ms of exposed collective wait once the split expert-weight
layout and the transposed-RHS backward kernel are enabled, for a net 4.036 -> 3.957 s/step
in the full stack.
…ght_layout)

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 them but creates a single 1 GB all-gather.

New config option `moe_split_expert_weight_layout` (default false = unchanged layout):
- wi is stored as [2, E, emb/2, 2*mlp] (a row-major reshape of the fused weight, sharded
  on dim 1) so the two expert halves W[:E/2] and W[E/2:] are contiguous leading-dim slices
  (~507 MB each) that are all-gathered separately and consumed without any concat.
- 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
  that the second kernel's sublane masking zeroes; backward chains dlhs the same way and
  runs two group-offset tgmm_v2 calls that write drhs_lo / drhs_hi directly.
- configs/types.py validates the preconditions (prefuse_moe_weights, shard_exp_on_fsdp,
  sparse_matmul, use_gmm_v2, no quantization / emb chunks / EP, even E and emb).
- tests/unit/split_gmm_test.py checks fwd + grads against a dense reference using a
  pure-JAX model of gmm_v2 that reproduces the boundary-tile masking.
Note: the checkpoint shape of `wi` changes with the option (a reshape of the fused layout).

Measured on Gemma4-26B-A4B, v6e-128: 4.611 -> 4.545 s/step standalone (with the fused
combine kernel 4.611 -> 4.206 s/step).

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