Skip to content

[feat] FSDP support for HybridStack EP-overlap (3/4 of #4798) - #4943

Open
Connor-XY wants to merge 8 commits into
NVIDIA:pull-request/4942from
Connor-XY:pr4798-3-fsdp-hybrid
Open

Connor-XY wants to merge 8 commits into
NVIDIA:pull-request/4942from
Connor-XY:pr4798-3-fsdp-hybrid

Conversation

@Connor-XY

@Connor-XY Connor-XY commented May 22, 2026 •

Copy link
Copy Markdown
Contributor

Review this follow-up's changes against #4942: scoped diff. The GitHub base is pull-request/4942, verified at the current parent head; this scoped diff excludes the parent changes.

Before merging #4942, retarget this PR to main so GitHub preserves the dependent review discussion.

What does this PR do?

Add Megatron FSDP v1 support for bracketed HybridStack groups in the 1F1B EP-overlap schedule. Each bracket group becomes one FSDP unit; the outer decoder stack is excluded.

Apply the same instance filter to parameter buckets and lifecycle hooks, so separate groups never share the outer stack's bucket identity. Preserve the existing constructor's positional arguments. Regression coverage includes grouped, mixed, and ungrouped unit selection, plus three-step loss/parameter parity against non-overlapped FSDP for grouped attention and Mamba patterns.

Dependencies and review scope

Validation

Refreshed on current main through #4942. Black, isort, repository-pinned pylint, Ruff, and diff checks pass for this slice. A limited CPU harness using PyTorch 2.14.0 passes all six unit-selection cases and fails the filtered case when the bucketing fix is removed. It executes the source grouping function with ordinary CPU parameters; this does not validate CUDA/FSDP collectives.

An internal GB200 run tested the integrated stack at #6960 c331df5ef3f2d6e057ad11cee28cf6341d9cdd8b plus #7192 d2385590f5ed4e0189e19e434d262ed5c7ac02e7 on four GB200 GPUs (PyTorch 2.12.0a0+0291f960b6.nv26.04.48445190, CUDA 13.2, Transformer Engine 2.14.0+f031cf87). All 18 FSDP tests passed on every rank: 12 three-step loss/final-parameter parity cases across attention/Mamba/MLA, alltoall/flex dispatchers, shared experts off/on, plus six unit/bucket tests. This validates the integrated source combination; standalone GitHub CI has been triggered at this exact head.

@copy-pr-bot

copy-pr-bot Bot commented May 22, 2026

Copy link
Copy Markdown

This pull request requires additional validation before any workflows can run on NVIDIA's runners.

Pull request vetters can view their responsibilities here.

Contributors can view more details about this message here.

@Connor-XY

Copy link
Copy Markdown
Contributor Author

Rebased onto the updated #4942 branch to pick up the two review fixes there (ffcae6d, 08e141e): the output_processor branch in HybridModel._postprocess, and gating the final_norm → final_layernorm sharded-key rename on bracketed-group patterns so non-grouped hybrid checkpoints keep loading. No changes to this PR's own commits.

@Connor-XY
Connor-XY force-pushed the pr4798-3-fsdp-hybrid branch from 0b44daf to bc5a0f7 Compare July 28, 2026 00:22
@Connor-XY

Copy link
Copy Markdown
Contributor Author

Rebased onto current main via the updated #4942 branch. No conflicts in this PR's own commits — the conflicts were all in #4942's files against the MLA-in-HybridModel port (#4452); details there.

@Connor-XY
Connor-XY marked this pull request as ready for review September 29, 2026 22:32
@Connor-XY
Connor-XY requested review from a team as code owners September 29, 2026 22:32
@Connor-XY

Copy link
Copy Markdown
Contributor Author

/ok to test 166d10e

@Connor-XY
Connor-XY changed the base branch from main to pull-request/4942 September 30, 2026 04:19
ddp_config, ["optim_grads_params"]
):
supported_fsdp_unit_modules = [TransformerLayer, MoETransformerLayer, MambaLayer]
from megatron.core.models.hybrid.hybrid_block import HybridStack

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

top-level

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Done in 1715f7e: HybridStack is now imported at the top of mcore_fsdp_adapter.py. It doesn't create an import cycle. I checked fresh-process imports of megatron.core, the adapter, hybrid_block/hybrid_model and megatron.training, plus a full unit-test collection, on our container.

Comment on lines +293 to +294
fsdp_unit_modules=self.fsdp_unit_modules,
fsdp_unit_filter=_fsdp_unit_filter,

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

These two now overlap quite a bit. Keep only fsdp_unit_filter because it's more general?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Agreed. In 1715f7e fsdp_unit_filter is now the only thing MegatronFSDP and BucketingPolicy use to pick FSDP units. It's a standalone predicate now, not a post-filter on the class match. I kept fsdp_unit_modules only as shorthand, lowered to an equivalent FSDPUnitTypeFilter, because it's the public megatron-fsdp / fully_shard argument (it also takes class-path strings) and many callers use it. Passing both now raises. The MCore adapter passes just fsdp_unit_filter: its unit types minus the outer HybridStack of a grouped pattern.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

many callers use it

How bad is it really? We could provide a converter from fsdp_module_types to a filter to ease the migration.

Connor-XY and others added 8 commits September 30, 2026 15:16
Rewrite comments that only made sense relative to the diff or to a past
incident: drop a line-number reference in favor of the function name,
state the recompute-hook invariant instead of the symptom it caused, and
describe the checkpoint keys each pattern publishes instead of calling
them historical or pre-existing.

Signed-off-by: Yan Xu <yxu1@nvidia.com>
The schedule plan checked static config on every build_schedule_plan()
call, i.e. every microbatch, walked all modules each time, and mixed
assert with ValueError. Move the checks to HybridModel.__init__ and raise
ValueError:

- CUDA graphs (the message no longer claims the check is specific to
  grouped patterns; it applies to any HybridModel with EP overlap),
- mHC connections and MoE shortcut connections,
- mixed context-parallel layouts, checked on every HybridStack including
  nested bracketed groups.

Drop the hash-routed MoE and wide-residual checks: TransformerConfig
already rejects both with overlap_moe_expert_parallel_comm.

Signed-off-by: Yan Xu <yxu1@nvidia.com>
Read known attributes directly and pass arguments unconditionally:

- HybridStack sets final_norm to None when it has no final norm (like
  TransformerBlock.final_layernorm), so the final_layernorm alias and the
  overlap callables read it directly. The flextron hook checks for a None
  final norm instead of a missing attribute.
- The schedule plan dispatches on type: HybridStack for the layer-type
  symbol, and TransformerLayer / MultiTokenPredictionLayer for the
  layer-level quantization context. Its imports move to module level.
- The MoE preprocess slot always sets shared_expert_output; the combine
  slot reads it and mlp_norm_manager directly. The Mamba pre-layer call
  no longer probes the chunk state for an inference context the training
  schedule never sets.
- HybridStack.forward drops rotary_pos_cos / rotary_pos_sin /
  rotary_pos_cos_sin, which nothing passes in, and HybridModel._postprocess
  drops is_spec_decode, which no caller passes.
- A bracketed group stack never applies full recomputation itself; gate
  on the static is_layer_group_stack instead of threading the private
  _checkpointed_forward_in_parent flag through forward. checkpointed_forward
  identifies group layers with isinstance(layer, HybridStack) and passes
  input_ids, packed_seq_params_by_layout and cp_layout_plan unconditionally.
- MambaLayer.backward_dw and MambaMixer.backward_dw call their delayed
  wgrad unconditionally. Building the overlap callables with
  delay_wgrad_compute now raises for a Mamba layer whose mixer cannot
  defer its wgrad (e.g. GatedDeltaProductMixer) instead of silently
  skipping those weight gradients.

Signed-off-by: Yan Xu <yxu1@nvidia.com>
Adjust the mcore-FSDP adapter and the megatron-FSDP core so HybridStack
(including nested grouped HybridStack instances) is a valid FSDP unit and
participates in the EP-overlap schedule plan. Add the
``test_fsdp_hybrid_overlap`` integration test exercising the FSDP +
grouped HybridModel forward/backward path.

Part 3/4 of splitting NVIDIA#4798 (original changes by @Wohox). Depends on the
HybridStack changes in part 2/4 (#TBD).

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
Signed-off-by: Yan Xu <yxu1@nvidia.com>
Signed-off-by: Yan Xu <yxu1@nvidia.com>
Address review feedback on the FSDP HybridStack support:

- Import HybridStack at module level in the MCore FSDP adapter instead of
  inside __init__.
- Make fsdp_unit_filter the one mechanism that selects FSDP units.
  MegatronFSDP and BucketingPolicy now consult only the filter; the public
  fsdp_unit_modules argument is kept as shorthand for a class-based filter
  (FSDPUnitTypeFilter), and passing both raises a ValueError.
- The MCore adapter passes only fsdp_unit_filter, built from its unit
  types, and leaves out the outer HybridStack of a grouped hybrid pattern.

Signed-off-by: Yan Xu <yxu1@nvidia.com>
@Connor-XY

Copy link
Copy Markdown
Contributor Author

/ok to test 1715f7e

This branch was successfully deployed

2 active deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

3 participants