1313from torch import Tensor , nn
1414
1515from megatron .core .dist_checkpointing .mapping import ShardedStateDict
16- from megatron .core .dist_checkpointing .utils import apply_prefix_mapping , replace_prefix_for_sharding
16+ from megatron .core .dist_checkpointing .utils import replace_prefix_for_sharding
1717from megatron .core .enums import Fp8Recipe
1818from megatron .core .extensions .transformer_engine import TENorm
1919from megatron .core .fp4_utils import get_fp4_context
3232from megatron .core .transformer .utils import sharded_state_dict_default
3333from megatron .core .utils import WrappedTensor , deprecate_inference_params , make_viewless_tensor
3434
35- # Canonical slot name each layer symbol uses in its stand-alone
36- # sharded-state-dict output – i.e., the attribute path under which the
37- # primitive's weights live when the block contains just that symbol. For a
38- # fused `[XY]` block, the same primitives sit under `self_attention` (for the
39- # sequence mixer X) and `mlp` (for the channel mixer Y); `sharded_state_dict`
40- # uses this table to rewrite fused keys back into the stand-alone layout so
41- # fused and unfused patterns produce the same checkpoint keys. Only Mamba
42- # needs an intra-block rename (`self_attention.` -> `mixer.`); every other
43- # sequence mixer is stand-alone-hosted in a TransformerLayer whose
44- # `self_attention` slot already matches the fused layout.
45- _CANONICAL_SLOT_FOR_SYMBOL : dict [str , str ] = {
46- LayerSymbols .MAMBA : "mixer" ,
47- LayerSymbols .GDN : "self_attention" ,
48- LayerSymbols .ATTENTION : "self_attention" ,
49- LayerSymbols .DS_ATTENTION : "self_attention" ,
50- LayerSymbols .MLP : "mlp" ,
51- LayerSymbols .MOE : "mlp" ,
52- }
53-
5435
5536@dataclass
5637class HybridStackSubmodules :
@@ -103,13 +84,6 @@ class HybridStack(GraphableMegatronModule, MegatronModule):
10384 pp_layer_offset (int, optional): the global layer offset for this pipeline
10485 segment, measured in physical blocks (fused groups count as one).
10586 Defaults to 0.
106- sub_layer_offset (int, optional): the global sub-layer offset for this
107- pipeline segment, measured in unfused sub-layers (a fused `[XY]`
108- contributes 2). Used only to emit canonical-form sharded-state-dict
109- keys so fused and unfused patterns produce the same checkpoint
110- layout. When None (default), falls back to `pp_layer_offset`, which
111- is correct whenever no earlier pipeline segment contained a fusion
112- group.
11387 post_layer_norm (bool, optional): whether to include a final layer norm.
11488 Defaults to True.
11589 post_process (bool, optional): whether to include an output layer.
@@ -128,7 +102,6 @@ def __init__(
128102 pre_process : bool = True ,
129103 layer_type_list : Optional [list [str ]] = None ,
130104 pp_layer_offset : int = 0 ,
131- sub_layer_offset : Optional [int ] = None ,
132105 post_layer_norm : bool = True ,
133106 post_process : bool = True ,
134107 device = None ,
@@ -156,12 +129,6 @@ def __init__(
156129 "--hybrid-layer-pattern by HybridModel."
157130 )
158131 self .layer_type_list = layer_type_list
159- # Sub-layer offset defaults to pp_layer_offset, which is correct for
160- # any pattern whose earlier pipeline segments contain no fusion groups
161- # (i.e., physical-block count == sub-layer count up to that point).
162- self .sub_layer_offset = (
163- sub_layer_offset if sub_layer_offset is not None else pp_layer_offset
164- )
165132
166133 # Build layers from the pre-selected segment
167134 self .layers = nn .ModuleList ()
@@ -479,70 +446,21 @@ def sharded_state_dict(
479446 sharded_state_dict = {}
480447 layer_prefix = f'{ prefix } layers.'
481448
482- # Sub-layer cursor tracks the running sub-layer index across this
483- # segment; combined with `sub_layer_offset` it yields the same global
484- # sub-layer indices whether the current pattern contains fused groups
485- # or not. That is what gives checkpoints a fusion-independent layout –
486- # a pattern saved unfused can be loaded into a fused model (and vice
487- # versa) without any external translation step.
488- sub_layer_cursor = self .sub_layer_offset
449+ for local_layer_idx , layer in enumerate (self .layers ):
489450
490- for local_layer_idx , (layer_type , layer ) in enumerate (
491- zip (self .layer_type_list , self .layers )
492- ):
451+ global_layer_offset = layer .layer_number - 1 # self.layer_number starts at 1
493452 state_dict_prefix = (
494453 f'{ layer_prefix } { local_layer_idx } .' # module list index in HybridStack
495454 )
455+
456+ sharded_prefix = f'{ layer_prefix } { global_layer_offset } .'
496457 sharded_pp_offset = []
458+
497459 layer_sharded_state_dict = layer .sharded_state_dict (
498460 state_dict_prefix , sharded_pp_offset , metadata
499461 )
500462
501- if len (layer_type ) == 1 :
502- # Stand-alone block: one physical block == one sub-layer, and
503- # the block's attribute layout already matches the canonical
504- # stand-alone layout. Only the outer block index needs to move
505- # from the local module-list index to the global sub-layer
506- # index.
507- canonical_prefix = f'{ layer_prefix } { sub_layer_cursor } .'
508- replace_prefix_for_sharding (
509- layer_sharded_state_dict , state_dict_prefix , canonical_prefix
510- )
511- sub_layer_cursor += 1
512- else :
513- # Fused block `[XY]`: split the single physical block's keys
514- # into two sub-layer prefixes so the checkpoint looks exactly
515- # as it would for stand-alone `X` followed by stand-alone `Y`.
516- # Norms attached to X (e.g. `input_layernorm` for DSA) attach
517- # to X's sub-layer index; norms attached to Y (e.g.
518- # `pre_mlp_layernorm` for MoE) attach to Y's sub-layer index.
519- x_sym , y_sym = layer_type [0 ], layer_type [1 ]
520- canonical_x_prefix = f'{ layer_prefix } { sub_layer_cursor } .'
521- canonical_y_prefix = f'{ layer_prefix } { sub_layer_cursor + 1 } .'
522- slot_for_x = _CANONICAL_SLOT_FOR_SYMBOL [x_sym ]
523- slot_for_y = _CANONICAL_SLOT_FOR_SYMBOL [y_sym ]
524-
525- # Order matters: `apply_prefix_mapping` picks the first
526- # matching prefix, so list each specific sub-prefix before
527- # the bare block prefix fallback.
528- prefix_map = {
529- f'{ state_dict_prefix } input_layernorm.' : (
530- f'{ canonical_x_prefix } input_layernorm.'
531- ),
532- f'{ state_dict_prefix } self_attention.' : (f'{ canonical_x_prefix } { slot_for_x } .' ),
533- f'{ state_dict_prefix } self_attn_bda.' : (f'{ canonical_x_prefix } self_attn_bda.' ),
534- f'{ state_dict_prefix } pre_mlp_layernorm.' : (
535- f'{ canonical_y_prefix } pre_mlp_layernorm.'
536- ),
537- f'{ state_dict_prefix } mlp.' : f'{ canonical_y_prefix } { slot_for_y } .' ,
538- f'{ state_dict_prefix } mlp_bda.' : f'{ canonical_y_prefix } mlp_bda.' ,
539- # Fallback for any stray top-level fused-block keys (e.g.,
540- # `_extra_state` attached to the TransformerLayer itself);
541- # attach them to X's sub-layer index by convention.
542- state_dict_prefix : canonical_x_prefix ,
543- }
544- apply_prefix_mapping (layer_sharded_state_dict , prefix_map )
545- sub_layer_cursor += 2
463+ replace_prefix_for_sharding (layer_sharded_state_dict , state_dict_prefix , sharded_prefix )
546464
547465 sharded_state_dict .update (layer_sharded_state_dict )
548466
0 commit comments