[TC sort 4/4] All-gather / combine MoE tokens in the TC kernels' 3D layout - #5507
gobbleturk wants to merge 2 commits into
Conversation
There was a problem hiding this comment.
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.
| 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(), |
There was a problem hiding this comment.
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 == 3Because 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.
| tc_out_3d=x_is_3d and use_tc_3d_dispatch(), | |
| tc_out_3d=use_tc_3d_dispatch(), |
|
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. |
|
For one comment: Suggested with no slowdown: |
fec6edb to
c410112
Compare
f88aada to
1adbfff
Compare
c410112 to
d32c858
Compare
7908c8b to
1ba5424
Compare
d32c858 to
d67b699
Compare
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.
d67b699 to
84a83ea
Compare
1ba5424 to
4fde613
Compare
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:
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_onlyare 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-onlylabel; the full pipelinere-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):
gemini-reviewlabel.