Skip to content

[TC sort 4/4] All-gather / combine MoE tokens in the TC kernels' 3D layout - #5507

Open
gobbleturk wants to merge 2 commits into
mattdavidow/tc-sort-3-gmm3dfrom
mattdavidow/tc-sort-4-3d-dispatch
Open

gobbleturk wants to merge 2 commits into
mattdavidow/tc-sort-3-gmm3dfrom
mattdavidow/tc-sort-4-3d-dispatch

Conversation

@gobbleturk

Copy link
Copy Markdown
Collaborator

Adds moe_tc_ragged_3d_dispatch (default False; requires moe_tc_ragged_sort and moe_tc_ragged_3d_gmm under ring of experts, no TP / mlp_bias).

The local tokens are reshaped to (batch, seq, emb // 128, 128) before the expert-parallel all-gather (on the fp8 qvalue when moe_quantize_token_all_gather), so the TC ragged sort consumes the gathered tokens without a relayout. The TC unsort keeps its output in the same layout, the combine reduce-scatter runs on it, and only the local result is reshaped back to (batch, seq, emb). The custom VJPs keep the cotangents in the 3D layout, so the backward all-gather / reduce- scatter are 3D as well. This removes the full-size (N, D/128, 128) <-> (N, D) relayout kernels around the gathered tokens and the combine output in both passes; only reshapes of the 16x smaller local tensors remain. Calls that do not use the TC sort (e.g. dropless fallback) reshape the gathered tokens back to 2D.

512 chips (v7x 8x8x8, FSDP 64, EP 16, real checkpoint + routing), on top of the TC sort + weights-on-activation + 3D gmm stack: 3.471 -> 3.341 s/step (-3.7%), loss matches (lm_loss @39: 4.032 vs 3.997). Flatten/unflatten kernels (~106 ms/step) are gone from the profile.

Description

Start with a short description of what the PR does and how this is a change from
the past.

The rest of the description includes relevant details and context, examples:

  • why is this change being made,
  • the problem being solved and any relevant context,
  • why this is a good solution,
  • some information about the specific implementation,
  • shortcomings of the solution and possible future improvements.

If the change fixes a bug or a Github issue, please include a link, e.g.,:
FIXES: b/123456
FIXES: #123456

You can also provide a comma-separated list. If you don't want to close a bug but
simply to reference it, use BUGS, e.g.:
BUGS: b/123456

Notice 1: Once all tests pass, the "pull ready" label will automatically be assigned.
This label is used for administrative purposes. Please do not add it manually.

Notice 2: For external contributions, our settings currently require an approval from a MaxText maintainer to trigger CI tests.

Notice 3: Tests marked scheduled_only are skipped on PRs and only run in the periodic scheduled CI.
If you need them (plus the TPU7X suite) on your PR, add the scheduled-only label; the full pipeline
re-runs as a scheduled run. Remove and re-add the label to trigger it again.

Tests

Please describe how you tested this change, and include any instructions and/or
commands to reproduce.

Checklist

Before submitting this PR, please make sure (put X in square brackets):

  • I have performed a self-review of my code. For an optional AI review, add the gemini-review label.
  • I have necessary comments in my code, particularly in hard-to-understand areas.
  • I have run end-to-end tests tests and provided workload links above if applicable.
  • I have made or will make corresponding changes to the doc if needed, including adding new documentation pages to the relevant Table of Contents (toctree directive) as explained in our documentation.

@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 the moe_tc_ragged_3d_dispatch configuration and associated logic to optimize Mixture of Experts (MoE) routing by keeping token-side tensors in a 3D layout (batch, seq, emb // 128, 128) during expert-parallel all-gather, sorting, and combine operations, thereby avoiding expensive relayouts. Feedback on the changes highlights a critical issue where x_is_3d evaluates to False because the gathered tensor x has 4 dimensions, which silently disables the 3D layout optimization during the combine and reduce-scatter passes. It is recommended to update x_is_3d to support 4D layouts and directly set tc_out_3d=use_tc_3d_dispatch() to ensure the optimization is correctly applied.

Comment thread src/maxtext/layers/moe.py
topk_argsort_indices=routing.topk_argsort_indices,
tc_routing=routing.tc_routing,
tc_prescaled=tc_prescaled,
tc_out_3d=x_is_3d and use_tc_3d_dispatch(),

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

When use_tc_3d_dispatch() is active, the expert-parallel all-gather keeps the tokens in the 4D layout (batch, seq, emb // 128, 128). Consequently, the gathered tensor x has ndim == 4.

However, x_is_3d is defined on line 3308 as:

x_is_3d = self.config.num_moe_emb_chunks <= 0 and (x.qvalue if isinstance(x, qpl.QArray) else x).ndim == 3

Because x.ndim is 4, x_is_3d evaluates to False. This causes tc_out_3d to be False as well, which silently disables the 3D layout optimization for the combine and reduce-scatter passes, falling back to 2D and re-introducing the relayout overhead.

To resolve this, x_is_3d should be updated to allow ndim == 4 when use_tc_3d_dispatch() is active:

x_is_3d = self.config.num_moe_emb_chunks <= 0 and (
    (x.qvalue if isinstance(x, qpl.QArray) else x).ndim in (3, 4)
)

Since line 3308 is outside the current diff hunk, we can temporarily bypass it here by setting tc_out_3d=use_tc_3d_dispatch(), but the definition of x_is_3d on line 3308 must also be updated to ensure the wo GMM (gmm_fn) correctly outputs in 3D layout without relayout.

Suggested change
tc_out_3d=x_is_3d and use_tc_3d_dispatch(),
tc_out_3d=use_tc_3d_dispatch(),

@gobbleturk gobbleturk changed the title [DRAFT] All-gather / combine MoE tokens in the TC kernels' 3D layout [TC sort 4/4] All-gather / combine MoE tokens in the TC kernels' 3D layout Oct 2, 2026
@RissyRan

RissyRan commented Oct 2, 2026

Copy link
Copy Markdown
Collaborator

Let's have tests guard this feature?

Shall we add some assertions?

The flag can silently do nothing. use_tc_3d_dispatch() checks nine conditions, including TC sort, 3D GMM, ring of experts, ragged sort, no emb chunks, tp == 1, no mlp_bias and emb % 128 == 0, and falls back to the normal layout without a message.

@RissyRan

RissyRan commented Oct 2, 2026

Copy link
Copy Markdown
Collaborator

For one comment:

Rowwise combine-gradient quantization behaves differently. With moe_quantize_combine_bwd_method="rowwise", the combine psum_scatter now runs on 4D tensors. _all_gather_quantized_payload uses channelwise_axes=tuple(range(x.ndim - 1)), so scales come out per (b, s, emb//128) slice, one per 128 values, instead of one per token. That's numerically finer, so not wrong, but:
   - it gathers about emb/128 times more scale data;
   - it no longer matches the documented "per-token, matching megablox _bwd_quantize_gradient"

Suggested with no slowdown:

Confirmed on CPU. With channelwise_axes=(0, 1) on the 4D gradient, the scale shape is (b, s, 1, 1), one per token, and the dequantized result is bit-identical to today's 3D version. The PR's current code gives (b, s, 8, 1), eight scales per token at emb=1024.

The fix is to make the number of token axes explicit instead of assuming "everything but the last axis". It goes in moe.py, in the PR's branch:

def _all_gather_quantized_payload(
    x: jax.Array,
    axis_name: str | tuple[str, ...],
    *,
    axis: int,
    tiled: bool,
    method: str,
    qtype: jnp.dtype = jnp.float8_e5m2,
    num_feature_axes: int = 1,
) -> jax.Array:
  """...
    num_feature_axes: Number of trailing axes that form one token's features. 'rowwise' computes one scale
      per token over these axes, e.g. 2 for the (batch, seq, emb // 128, 128) layout of moe_tc_ragged_3d_dispatch.
  """
  static = method.startswith("fixed")
  token_axes = tuple(range(x.ndim - num_feature_axes))
  if not static and method != "rowwise":
    raise ValueError(...)
  if not static and axis % x.ndim not in token_axes:
    raise ValueError(f"'rowwise' requires gathering along a token axis, got axis={axis}.")
  ...
  else:  # dynamic rowwise with per-token scale
    x_q = qpl.quantize(x_f32, qtype=qtype, channelwise_axes=token_axes, calibration_method="absmax")

Then have the combine backward pass the layout. The combine tensor is always (batch, seq, *features), so the feature axis count is grads.ndim - 2:

  gathered_grads = _all_gather_quantized_payload(
      grads, axis_name, axis=scatter_dimension, tiled=tiled, method=bwd_method, num_feature_axes=grads.ndim - 2
  )

Why this approach:
- No relayout. The gradient stays 4D, so the PR's speedup is kept. Reshaping to 3D around the psum_scatter would put back the full-size relayout the PR removes.
- Same numbers as the 2D path. The scales are one per token again, and the result matches the 3D path exactly, as shown above. Scale traffic goes back to one float32 per token instead of emb/128 per token.
- No change for existing callers. num_feature_axes defaults to 1, so other callers behave as before. For today's 3D input, grads.ndim - 2 is also 1.
- fixed,<b> is unaffected, because it uses a per-tensor scale.

@gobbleturk
gobbleturk force-pushed the mattdavidow/tc-sort-3-gmm3d branch from fec6edb to c410112 Compare October 3, 2026 01:13
@gobbleturk
gobbleturk force-pushed the mattdavidow/tc-sort-4-3d-dispatch branch from f88aada to 1adbfff Compare October 3, 2026 01:13
@gobbleturk
gobbleturk force-pushed the mattdavidow/tc-sort-3-gmm3d branch from c410112 to d32c858 Compare October 3, 2026 01:34
@gobbleturk
gobbleturk force-pushed the mattdavidow/tc-sort-4-3d-dispatch branch 2 times, most recently from 7908c8b to 1ba5424 Compare October 3, 2026 02:38
@gobbleturk
gobbleturk force-pushed the mattdavidow/tc-sort-3-gmm3d branch from d32c858 to d67b699 Compare October 3, 2026 02:38
Adds moe_tc_ragged_3d_dispatch (default False; requires moe_tc_ragged_sort and
moe_tc_ragged_3d_gmm under ring of experts, no TP / mlp_bias).

The local tokens are reshaped to (batch, seq, emb // 128, 128) before the
expert-parallel all-gather (on the fp8 qvalue when moe_quantize_token_all_gather),
so the TC ragged sort consumes the gathered tokens without a relayout. The TC
unsort keeps its output in the same layout, the combine reduce-scatter runs on
it, and only the local result is reshaped back to (batch, seq, emb). The custom
VJPs keep the cotangents in the 3D layout, so the backward all-gather / reduce-
scatter are 3D as well. This removes the full-size (N, D/128, 128) <-> (N, D)
relayout kernels around the gathered tokens and the combine output in both
passes; only reshapes of the 16x smaller local tensors remain. Calls that do not
use the TC sort (e.g. dropless fallback) reshape the gathered tokens back to 2D.

512 chips (v7x 8x8x8, FSDP 64, EP 16, real checkpoint + routing), on top of the
TC sort + weights-on-activation + 3D gmm stack: 3.471 -> 3.341 s/step (-3.7%),
loss matches (lm_loss @39: 4.032 vs 3.997). Flatten/unflatten kernels
(~106 ms/step) are gone from the profile.
- ragged_sort_tc_test: 3D-layout token input and combine output (out_3d) vs the
  reference, with 3D and 2D sorted buffers, prescaled weights, truncated buffer,
  flatten and DeepSeek-V3 hidden size.
- moe_test: moe_tc_ragged_3d_dispatch loss/grad parity with the 2D dispatch.
@gobbleturk
gobbleturk force-pushed the mattdavidow/tc-sort-3-gmm3d branch from d67b699 to 84a83ea Compare October 3, 2026 02:50
@gobbleturk
gobbleturk force-pushed the mattdavidow/tc-sort-4-3d-dispatch branch from 1ba5424 to 4fde613 Compare October 3, 2026 02:50

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.

3 participants