Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 29 additions & 1 deletion docs/user-guide/hybrid-model-migration.md
Original file line number Diff line number Diff line change
Expand Up @@ -46,7 +46,8 @@ the first `-` or `E`. Source GPT layer 1 maps to the next pair, and so on.
The pattern can also describe execution layout. A `|` marks a pipeline segment
boundary, and `/` introduces a repeated Multi-Token Prediction (MTP) pattern.
For example, `*-*-|*-*-` places four GPT-equivalent blocks across two pipeline
segments. Separators do not count as layers.
segments. Separators do not count as layers. Square brackets group layers into
one logical layer, as described in [Bracketed layer groups](#bracketed-layer-groups).

HybridModel provides the following benefits:

Expand All @@ -68,6 +69,33 @@ An architecture-preserving `*-` or `*E` migration should be validated for
numerical equivalence, and a pattern that adds another layer family should be
treated as a new architecture and benchmarked independently.

### Bracketed layer groups

Square brackets group consecutive layers into one logical layer, for example
`[*E][*E]` or `M[M*E]`. HybridModel builds each group as one nested
`HybridStack`, and the group occupies a single layer index in checkpoint keys:
`[*-]` stores its attention and MLP under one `layers.<i>.` prefix, as a GPT
transformer block does. Adding or removing brackets therefore changes the
checkpoint keys of the affected layers. Each symbol inside a group still counts
as one layer for layer numbering, so layer-indexed settings such as FP8 layer
ranges and hash-routed MoE thresholds are unchanged.

Use groups with MoE expert-parallel communication overlap
(`--overlap-moe-expert-parallel-comm`). The overlap schedule interleaves one
microbatch's MoE all-to-all with another microbatch's compute one logical layer
at a time, so group each MoE layer with the layers that run before it, as in
`[M*E]`.

Groups have the following constraints:

- A group cannot be empty or nested, cannot span a `|` pipeline boundary, and
cannot appear in an MTP pattern.
- An MoE layer must be the last layer of its group.
- The layers of a group must use different checkpoint namespaces: at most one
Mamba layer, one attention or GDN layer, and one MLP or MoE layer.
- Groups cannot be combined with mHC connections, wide residual streams, or MoE
shortcut connections.

## 2. How to Convert a Checkpoint

There are two ways to bring `GPTModel` weights into a `HybridModel` run. Both
Expand Down
18 changes: 16 additions & 2 deletions megatron/core/models/common/fine_grained_callables.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,11 @@

from megatron.core import tensor_parallel
from megatron.core.models.gpt.fine_grained_callables import build_transformer_layer_callables
from megatron.core.models.hybrid.fine_grained_callables import (
build_hybrid_stack_callables,
get_hybrid_stack_moe_metadata,
)
from megatron.core.models.hybrid.hybrid_block import HybridStack
from megatron.core.transformer.moe.moe_layer import MoELayer
from megatron.core.transformer.multi_token_prediction import (
MultiTokenPredictionLayer,
Expand Down Expand Up @@ -67,8 +72,12 @@ def submodule_mtp_pre_dispatch_forward(node, hidden_states):
offset = get_mtp_layer_offset(
layer.config, model.vp_stage, pp_rank=model.pg_collection.pp.rank()
)
node.chunk_state.mtp_hidden_states = list(torch.chunk(hidden_states, 1 + offset, dim=0))
hidden_states = node.chunk_state.mtp_hidden_states[offset]
chunks = list(torch.chunk(hidden_states, 1 + offset, dim=0))
# These chunks cross from pre-dispatch to MTP post-process. Detach
# their stored views so the slots do not traverse final_norm's graph
# twice; node.backward_impl merges their gradients into this slot.
node.chunk_state.mtp_hidden_states = [node.detach(chunk) for chunk in chunks]
hidden_states = chunks[offset]

input_ids, position_ids, padding_mask, mtp_input_mask, decoder_input, hidden_states = (
layer._get_embeddings(
Expand Down Expand Up @@ -154,6 +163,8 @@ def get_layer_moe_metadata(layer):

if isinstance(layer, MultiTokenPredictionLayer):
return get_layer_moe_metadata(layer.mtp_model_layer)
if isinstance(layer, HybridStack):
return get_hybrid_stack_moe_metadata(layer)
if isinstance(layer, TransformerLayer):
is_moe = isinstance(layer.mlp, MoELayer)
num_local_experts = layer.mlp.num_local_experts if is_moe else None
Expand All @@ -170,6 +181,9 @@ def build_layer_callables(layer):

if isinstance(layer, MultiTokenPredictionLayer):
return build_mtp_layer_callables(layer)
if isinstance(layer, HybridStack):
forward_funcs, backward_dw, _, _ = build_hybrid_stack_callables(layer)
return forward_funcs, backward_dw
if isinstance(layer, TransformerLayer):
return build_transformer_layer_callables(layer)

Expand Down
Loading
Loading