Add SparseCore collective offloading for MoE with custom VJP for FSDP All-Gather to Reduce-Scatter. - #5125
Open
copybara-service[bot] wants to merge 1 commit into
Open
Add SparseCore collective offloading for MoE with custom VJP for FSDP All-Gather to Reduce-Scatter.#5125copybara-service[bot] wants to merge 1 commit into
copybara-service[bot] wants to merge 1 commit into
Conversation
copybara-service
Bot
requested review from
A9isha,
NuojCheng,
RissyRan,
SurbhiJainUSC,
abhinavclemson,
aireenmei,
bvandermoon,
darisoy,
dipannita08,
gagika,
gobbleturk,
hengtaoguo,
huytransformer,
igorts-git,
jiangjy1982,
khatwanimohit,
michelle-yooh,
richjames0,
shralex,
shuningjin,
vipannalla,
xibinliu and
zxhe-sean
as code owners
September 3, 2026 05:21
… All-Gather to Reduce-Scatter. - Split MoE SparseCore collective offloading controls into `moe_pin_sparse_core_fsdp_all_gather` and `moe_pin_sparse_core_ep_all_gather`, while retaining `moe_pin_sparse_core_all_gathers` as an umbrella flag for backward compatibility. - Implement custom VJP `_fsdp_all_gather_with_rs` on MoE weights to ensure the forward pass executes All-Gather on SparseCore and the backward pass executes Reduce-Scatter on SparseCore, eliminating the backward scheduling bottleneck. - Automatically disable SparseCore collective offloading if the backend or target compile topology is TPU but not Gen7. - Add unit tests in `tests/unit/moe_test.py` covering flag propagation, separate flag control, non-Gen7 TPU disabling, Gen7 TPU retention, and custom VJP autodiff gradient correctness. PiperOrigin-RevId: 975451266
copybara-service
Bot
force-pushed
the
test_975451266
branch
from
September 4, 2026 01:30
b29dd21 to
fff09d4
Compare
Codecov Report❌ Patch coverage is
📢 Thoughts on this report? Let us know! |
4 tasks
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Add SparseCore collective offloading for MoE with custom VJP for FSDP All-Gather to Reduce-Scatter.
moe_pin_sparse_core_fsdp_all_gatherandmoe_pin_sparse_core_ep_all_gather, while retainingmoe_pin_sparse_core_all_gathersas an umbrella flag for backward compatibility._fsdp_all_gather_with_rson MoE weights to ensure the forward pass executes All-Gather on SparseCore and the backward pass executes Reduce-Scatter on SparseCore, eliminating the backward scheduling bottleneck.tests/unit/moe_test.pycovering flag propagation, separate flag control, non-Gen7 TPU disabling, Gen7 TPU retention, and custom VJP autodiff gradient correctness.