Skip to content

[feat] HybridStack grouped syntax and EP-overlap (2/4 of #4798) - #4942

Open
Connor-XY wants to merge 8 commits into
NVIDIA:mainfrom
Connor-XY:pr4798-2-hybrid-stack-grouped
Open

Connor-XY wants to merge 8 commits into
NVIDIA:mainfrom
Connor-XY:pr4798-2-hybrid-stack-grouped

Conversation

@Connor-XY

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

Copy link
Copy Markdown
Contributor

What does this PR do?

Adds bracketed HybridStack groups and EP communication overlap to HybridModel. Patterns such as [*-] and [*E] make several physical layers one logical layer for scheduling and checkpoint storage. Groups reject nesting, nonterminal MoE layers, and colliding checkpoint namespaces. This is part 2 of the original #4798 split by @Wohox, @Connor-XY, and @guihong-nv; the common scheduler prerequisite #4941 has merged.

Grouping is what makes EP overlap possible for hybrid models. The overlap scheduler overlaps one microbatch's MoE all-to-all with another microbatch's compute at logical-layer granularity, so an attention (or Mamba/GDN) layer and the MoE layer after it must form one logical layer, e.g. [*E].

The overlap callables support attention, MLA, Mamba, GDN, and MoE layers, with per-microbatch norm-offload state and correct delayed-weight-gradient hook ownership. Hybrid MTP is not supported with EP overlap: as on main, HybridModel raises for an MTP pattern with overlap_moe_expert_parallel_comm. Existing ungrouped Hybrid checkpoint names and direct-stack offsets remain compatible. HybridModel.sharded_state_dict omits the output layer's (empty) TE extra state, as GPTModel does. GPT checkpoint import into grouped models is isolated in #7192.

Change to the common MTP callable (affects GPTModel). build_mtp_layer_callables in models/common/fine_grained_callables.py now detaches the per-depth hidden-state chunks it stores for the MTP post-process slot (node.detach), and the schedule node merges their gradients. Before, the post-process slot backpropagated through the pre-dispatch slot's graph, so the final norm's backward graph was traversed twice whenever the pre-dispatch slot applies the final norm. The code path runs for GPTModel with mtp_num_layers >= 1 and overlap_moe_expert_parallel_comm=True. The existing GPT chunk-level 1F1B tests with MTP (tests/unit_tests/a2a_overlap/test_schedule_chunk_1f1b.py) pass unchanged.

Current-main integration and boundaries

Rebased on b4d72b79c6b7e30010894ea0827a06c3477387fd. The refresh preserves eager hash routing, sequence-parallel token alignment, precomputed MTP embeddings, CP layout inputs, inference caches, and tensor observation hooks. Group parsing accounts for the current CSA/HCA/window attention symbols. MTP metrics retain main's explicit depth mapping. Ungrouped shortcut checkpoints keep physical storage indices, and static inference accepts caller-provided embeddings without token IDs.

With EP overlap, HybridModel rejects CUDA graphs, mixed CP-layout conversion, mHC, and MoE shortcuts at construction; TransformerConfig already rejects hash-routed MoE and wide residuals with overlap. These features require dedicated schedule integration before they can be enabled together.

Review and dependent PRs

The dependent PRs target main and contain this PR's commits until it merges, so this PR merges first.

Validation

  • Required formatter, isort, pylint, Ruff, Python parsing, and whitespace checks pass. Mypy remains advisory in the repository script and reports missing optional dependencies and typing findings in the local environment.
  • CPU source-isolation checks pass for the MTP LayerNorm gradient regression, grouped hash routing and inference handling, postprocess hooks, and existing MTP embedding/CP contracts. Removing the MTP fix makes its gradient regression fail. These checks use real PyTorch with CUDA/import scaffolding replaced and do not validate distributed GPU execution.
  • The kernel determinism coverage gate passes without exemptions. The registered Mamba replay test now exercises deferred projection gradients, including packed layouts.
  • GPU validation of the full stack ran on the top of the stack (Test MLA grouped HybridStack FSDP EP-overlap #6960 c331df5ef plus Support GPT checkpoint loading into grouped HybridModel #7192 d2385590f). Those heads contain this PR's commits unchanged, and the dependent PRs don't modify any file this PR touches. Each rank passed 18 FSDP, 19 regression, 31 checkpoint, and two un-packed Mamba replay cases (70 total). The four packed replay cases failed in that container before compute because its selective_scan_cuda build didn't match the platform. A rerun of all six ordinary/deferred Mamba replay cases on unchanged source in a matching container passed on every rank, including both packed layouts. No test expectations or dependency code were changed.
  • Static-inference and shortcut-checkpoint regression tests pass against the fixes and fail against the prior implementations.
  • After the pre-merge cleanup commits (14a16f198, 305c0f9a7, ad603d21c), the unit tests for the files this PR changes and tests/unit_tests/a2a_overlap were rerun at ad603d21c on an internal GB200 node. Results match the previous head's run, with no new failures. CI at the previous head 3b24fcadc passed apart from the code-owner approval gate.

@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
Connor-XY force-pushed the pr4798-2-hybrid-stack-grouped branch 7 times, most recently from 5e5dee3 to 91e86bb Compare June 3, 2026 19:21
@Victarry Victarry mentioned this pull request Jun 10, 2026
44 of 72 tasks
@Connor-XY
Connor-XY force-pushed the pr4798-2-hybrid-stack-grouped branch from 91e86bb to 7a0106e Compare June 29, 2026 16:29
@Connor-XY
Connor-XY force-pushed the pr4798-2-hybrid-stack-grouped branch 2 times, most recently from 360afee to b3467e6 Compare July 16, 2026 00:36
@Connor-XY

Copy link
Copy Markdown
Contributor Author

/claude review

Comment thread megatron/core/models/hybrid/hybrid_model.py Outdated
@Connor-XY
Connor-XY force-pushed the pr4798-2-hybrid-stack-grouped branch from eb80207 to 3242e57 Compare July 20, 2026 16:33
@Connor-XY

Copy link
Copy Markdown
Contributor Author

/claude review

@Connor-XY
Connor-XY force-pushed the pr4798-2-hybrid-stack-grouped branch from 3242e57 to b80f2c9 Compare July 20, 2026 16:38

@claude claude Bot left a comment

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.

Light review complete. The hybrid grouping / bracketed-pattern feature is well-covered by new unit tests (layer allocation, sharded state dict, forward equivalence, EP-overlap callables, postprocess). One potential bug flagged inline in router.py: the aux_loss MTP-slot indexing fix (line 595) was not applied to the identical z_loss path (line 698), which can still index past the metrics tracker when MTP wraps a HybridStack and z_loss is enabled.

@Connor-XY

Copy link
Copy Markdown
Contributor Author

/claude review

@Connor-XY

Copy link
Copy Markdown
Contributor Author

/ok to test 4aa1be7

@claude claude Bot left a comment

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.

LGTM

@Connor-XY

Copy link
Copy Markdown
Contributor Author

/ok to test 12407a4

Refresh PR NVIDIA#4942 against main while preserving hash routing, MTP, CP layouts, and ungrouped checkpoint compatibility. Integrate the cross-slot MTP gradient fix and reject unsupported direct-overlap execution modes.

Original grouped HybridStack work by @Wohox, @Connor-XY, and @guihong-nv in NVIDIA#4798/NVIDIA#4942.

Signed-off-by: Yan Xu <yxu1@nvidia.com>
@Connor-XY
Connor-XY force-pushed the pr4798-2-hybrid-stack-grouped branch from 12407a4 to 9f08c2f Compare September 29, 2026 22:19
@svcnvidia-nemo-ci svcnvidia-nemo-ci added the Final Review PR is in the "final review" stage label Sep 29, 2026
Connor-XY added a commit to Connor-XY/Megatron-LM that referenced this pull request Sep 29, 2026
Map logical grouped checkpoint keys and nested FSDP paths to physical GPT
attention and MLP layers. Detect homogeneous versus indexed source keys
and retarget model and optimizer state consistently.

Carry the grouped checkpoint regression tests with this implementation,
including parallel-layout roundtrips and model, optimizer, and FSDP loads.
This is the checkpoint slice of e8db9dd, extracted from PR NVIDIA#4942.

Signed-off-by: Yan Xu <yxu1@nvidia.com>
Connor-XY added a commit to Connor-XY/Megatron-LM that referenced this pull request Sep 29, 2026
Map logical grouped checkpoint keys and nested FSDP paths to physical GPT
attention and MLP layers. Detect homogeneous versus indexed source keys
and retarget model and optimizer state consistently.

Carry the grouped checkpoint regression tests with this implementation,
including parallel-layout roundtrips and model, optimizer, and FSDP loads.
This is the checkpoint slice of e8db9dd, extracted from PR NVIDIA#4942.

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

Copy link
Copy Markdown
Contributor Author

/ok to test 3b24fca

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>
@Connor-XY

Copy link
Copy Markdown
Contributor Author

/ok to test ad603d2

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

Labels

complexity: high Final Review PR is in the "final review" stage

Projects

None yet

Development

Successfully merging this pull request may close these issues.

6 participants