Repository navigation
Conversation
There was a problem hiding this comment.
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.
| 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 | ||
| ) |
There was a problem hiding this comment.
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.
| 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 | |
| ) |
f6b2dcb to
80be58d
Compare
80be58d to
dce4481
Compare
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).
b505267 to
5d05c55
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,moe_use_direct_token_gather=true. Step time went from 7.58 s to 3.82 s across the three PRs.On v6e:
src/maxtext/kernels/moe_combine_tc.py) that keep both on the TensorCore and never materialize the[tokens * top_k, emb]unsorted copy.prefuse_moe_weights=true,wiis stored as[E, emb, 2*mlp](~1 GBacross FSDP), which either incurs ~400 ms/step of concat/split copies aroundgmm_v2or stalls SparseCore behind a single ~1 GB all-gather. Commit 3 addsmoe_split_expert_weight_layout=true, storingwias[2, E, emb/2, 2*mlp]so its two leading slices (W[:E/2]andW[E/2:], ~507 MB each) are all-gathered separately and consumed by two chainedgmm_v2calls (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
moe_combine_kernel="tc")dxand a masked lane reduction fordwmoe_permute_kernel="tc")group_sizehistogram above the gather (no behaviour change)wilayout with chained gmm_v2 (moe_split_expert_weight_layout)wias 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 splitOnly 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
When using Recipe 2, combine
moe_permute_kernel="tc"+moe_split_expert_weight_layout=truewith the XLA flags--xla_tpu_offload_gather_to_sparsecore=false --xla_tpu_offload_all_supported_gathers_to_sparsecore=falseso 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 ofwichanges to[2, *scan_dims, num_experts, emb/2, 2*mlp](a reshape of the fused weight). Usesplit_gmm.convert_checkpoint_tree(state, to_split=True/False)(included insrc/maxtext/kernels/megablox/split_gmm.py) to convert checkpoints in either direction. Withmoe_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 (dxexact,dwto 1e-2) over block sizes /top_k/ expert counts / skewed routing / shared boundary granules /align=16/ f32 weights;jax.checkpoint+lax.scangradient 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 ofmoe_combine_kernel,moe_permute_kernel, andmoe_split_expert_weight_layout.src/maxtext/kernels/moe_combine_bench.py: single-chip correctness + benchmark script.--pyink-indentation=2 --line-length=122) clean.Related PRs