Skip to content

Commit 626f895

Browse files
committed
Refactor state dict canonicalization
Do not modify `sharded_state_dict` method, instead use `generate_state_dict` to apply the transformation only when loading/saving.
1 parent d003ce5 commit 626f895

5 files changed

Lines changed: 238 additions & 121 deletions

File tree

‎megatron/core/models/hybrid/hybrid_block.py‎

Lines changed: 7 additions & 89 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,7 @@
1313
from torch import Tensor, nn
1414

1515
from 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
1717
from megatron.core.enums import Fp8Recipe
1818
from megatron.core.extensions.transformer_engine import TENorm
1919
from megatron.core.fp4_utils import get_fp4_context
@@ -32,25 +32,6 @@
3232
from megatron.core.transformer.utils import sharded_state_dict_default
3333
from 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
5637
class 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

‎megatron/core/models/hybrid/hybrid_layer_fusion.py‎

Lines changed: 128 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -16,6 +16,8 @@
1616

1717
from typing import TYPE_CHECKING
1818

19+
from megatron.core.dist_checkpointing.mapping import ShardedStateDict
20+
from megatron.core.dist_checkpointing.utils import apply_prefix_mapping
1921
from megatron.core.extensions.transformer_engine import TENorm
2022
from megatron.core.fusions.fused_bias_dropout import get_bias_dropout_add
2123
from megatron.core.models.hybrid.hybrid_layer_allocation import Symbols as LayerSymbols
@@ -222,3 +224,129 @@ def build_fused_layer(
222224
build_kwargs["pp_layer_offset"] = pp_layer_offset
223225

224226
return build_module(fused_spec, **build_kwargs)
227+
228+
229+
# Canonical slot name each layer symbol uses in its stand-alone
230+
# sharded-state-dict output – i.e., the attribute path under which the
231+
# primitive's weights live when the block contains just that symbol. For a
232+
# fused `[XY]` block, the same primitives sit under `self_attention` (for the
233+
# sequence mixer X) and `mlp` (for the channel mixer Y); canonicalization
234+
# uses this table to rewrite fused keys back into the stand-alone layout so
235+
# fused and unfused patterns produce the same checkpoint keys. Only Mamba
236+
# needs an intra-block rename (`self_attention.` -> `mixer.`); every other
237+
# sequence mixer is stand-alone-hosted in a TransformerLayer whose
238+
# `self_attention` slot already matches the fused layout.
239+
_CANONICAL_SLOT_FOR_SYMBOL: dict[str, str] = {
240+
LayerSymbols.MAMBA: "mixer",
241+
LayerSymbols.GDN: "self_attention",
242+
LayerSymbols.ATTENTION: "self_attention",
243+
LayerSymbols.DS_ATTENTION: "self_attention",
244+
LayerSymbols.MLP: "mlp",
245+
LayerSymbols.MOE: "mlp",
246+
}
247+
248+
249+
def canonicalize_hybrid_sharded_state_dict(
250+
sharded_state_dict: ShardedStateDict,
251+
layer_prefix: str,
252+
layer_type_list: list[str],
253+
physical_offset: int = 0,
254+
sub_layer_offset: int = 0,
255+
) -> None:
256+
"""Rewrite HybridStack layer keys into the canonical (unfused) layout, in place.
257+
258+
`HybridStack.sharded_state_dict` emits keys indexed by global physical
259+
block position within the model (a fused `[XY]` group still occupies a
260+
single physical block). Fused blocks are realized as `TransformerLayer`s
261+
whose `self_attention` slot holds the sequence mixer and `mlp` slot
262+
holds the channel mixer, so their keys do not match what a stand-alone
263+
`X` followed by stand-alone `Y` would produce. This function rewrites
264+
each fused block's keys into two sub-layer-indexed prefixes that
265+
do match: `layers.{sub_layer_offset + i}.mixer.*` for mamba sub-layers,
266+
`layers.{sub_layer_offset + i}.mlp.*` for MLP, etc. Stand-alone blocks
267+
are simply re-indexed from physical to sub-layer index.
268+
269+
The resulting keys are fusion-independent – a checkpoint written with
270+
`[*-]M` and one written with `*-M` end up with the same set of keys, so
271+
the dist_checkpointing layer can load either into either.
272+
273+
Args:
274+
sharded_state_dict: The sharded state dict to rewrite in place. Only
275+
entries whose keys start with `layer_prefix` are touched.
276+
layer_prefix: The full prefix up to and including `"layers."`, e.g.
277+
`"decoder.layers."`. Keys outside this prefix are left alone.
278+
layer_type_list: The per-physical-block layer-type symbols for this
279+
pipeline segment (a single `"M"`, `"*"`, etc. for a stand-alone
280+
block; a two-char string like `"*-"` for a fused group).
281+
physical_offset: The global physical-block index at which this
282+
pipeline segment starts (i.e., the value `HybridStack` uses to
283+
derive each layer's `layer_number`). Defaults to 0, which is
284+
correct for non-pipeline-parallel runs.
285+
sub_layer_offset: The global sub-layer index at which this pipeline
286+
segment starts. Accounts for sub-layers contributed by earlier
287+
pipeline segments so that fused groups in those segments are
288+
correctly counted. When the model is not pipeline-parallel (or
289+
no earlier segment contains a fusion group), this is equal to
290+
`physical_offset` and may be left at its default.
291+
292+
Notes:
293+
- Build up a combined prefix map across all layers and apply it in
294+
one `apply_prefix_mapping` pass. This keeps the rewrite narrow
295+
(sibling keys outside `layer_prefix` are untouched) and lets the
296+
function safely run on the full model state dict without tripping
297+
on embedding or output-layer entries.
298+
- Order matters inside the combined prefix map: `apply_prefix_mapping`
299+
picks the first matching prefix, so the specific sub-prefixes
300+
(e.g. `"input_layernorm."`) are inserted before the bare block
301+
prefix fallback.
302+
- Norms attached to the X sub-layer (e.g. `input_layernorm` for DSA)
303+
stay with X's sub-layer index; norms attached to Y (e.g.
304+
`pre_mlp_layernorm` for MoE) attach to Y's sub-layer index.
305+
"""
306+
prefix_map: dict[str, str] = {}
307+
sub_layer_cursor = sub_layer_offset
308+
309+
for local_layer_idx, layer_type in enumerate(layer_type_list):
310+
physical_prefix = f'{layer_prefix}{physical_offset + local_layer_idx}.'
311+
312+
if len(layer_type) == 1:
313+
# Stand-alone block: one physical block == one sub-layer, and the
314+
# block's attribute layout already matches the canonical
315+
# stand-alone layout. Only the outer block index needs to move
316+
# from the local module-list index to the global sub-layer index.
317+
canonical_prefix = f'{layer_prefix}{sub_layer_cursor}.'
318+
prefix_map[physical_prefix] = canonical_prefix
319+
sub_layer_cursor += 1
320+
else:
321+
# Fused block `[XY]`: split the single physical block's keys into
322+
# two sub-layer prefixes so the checkpoint looks exactly as it
323+
# would for stand-alone `X` followed by stand-alone `Y`.
324+
x_sym, y_sym = layer_type[0], layer_type[1]
325+
canonical_x_prefix = f'{layer_prefix}{sub_layer_cursor}.'
326+
canonical_y_prefix = f'{layer_prefix}{sub_layer_cursor + 1}.'
327+
slot_for_x = _CANONICAL_SLOT_FOR_SYMBOL[x_sym]
328+
slot_for_y = _CANONICAL_SLOT_FOR_SYMBOL[y_sym]
329+
330+
# Specific sub-prefixes before the bare block prefix fallback.
331+
prefix_map[f'{physical_prefix}input_layernorm.'] = (
332+
f'{canonical_x_prefix}input_layernorm.'
333+
)
334+
prefix_map[f'{physical_prefix}self_attention.'] = (
335+
f'{canonical_x_prefix}{slot_for_x}.'
336+
)
337+
prefix_map[f'{physical_prefix}self_attn_bda.'] = (
338+
f'{canonical_x_prefix}self_attn_bda.'
339+
)
340+
prefix_map[f'{physical_prefix}pre_mlp_layernorm.'] = (
341+
f'{canonical_y_prefix}pre_mlp_layernorm.'
342+
)
343+
prefix_map[f'{physical_prefix}mlp.'] = f'{canonical_y_prefix}{slot_for_y}.'
344+
prefix_map[f'{physical_prefix}mlp_bda.'] = f'{canonical_y_prefix}mlp_bda.'
345+
# Fallback for any stray top-level fused-block keys (e.g.,
346+
# `_extra_state` attached to the TransformerLayer itself);
347+
# attach them to X's sub-layer index by convention.
348+
prefix_map[physical_prefix] = canonical_x_prefix
349+
sub_layer_cursor += 2
350+
351+
if prefix_map:
352+
apply_prefix_mapping(sharded_state_dict, prefix_map)

‎megatron/core/models/hybrid/hybrid_model.py‎

Lines changed: 10 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -194,12 +194,16 @@ def __init__(
194194
first_stage_layers=self.config.num_layers_in_first_pipeline_stage,
195195
last_stage_layers=self.config.num_layers_in_last_pipeline_stage,
196196
)
197-
# Sub-layer offset mirrors `layer_offset` but counts each character in a
198-
# fused `[XY]` group separately. `HybridStack` uses it to emit
199-
# sharded-state-dict keys in the canonical (unfused) layout so
200-
# checkpoints saved under one fusion configuration load cleanly under
201-
# another.
202-
sub_layer_offset = get_sub_layer_offset(main_pattern, layer_offset)
197+
# Read at checkpoint save/load time by
198+
# `megatron.training.checkpointing._apply_hybrid_canonicalization_if_applicable`,
199+
# which rewrites the decoder's sharded keys into a fusion-independent
200+
# layout so checkpoints saved under one fusion configuration load
201+
# cleanly under another. `_decoder_physical_offset` mirrors `layer_offset`
202+
# (physical-block index where this pipeline segment starts);
203+
# `_decoder_sub_layer_offset` is its sub-layer counterpart, counting
204+
# each character of a fused `[XY]` group separately.
205+
self._decoder_physical_offset = layer_offset
206+
self._decoder_sub_layer_offset = get_sub_layer_offset(main_pattern, layer_offset)
203207

204208
# Determine if MTP is needed (based on pattern parsing)
205209
self.mtp_process = (
@@ -263,7 +267,6 @@ def __init__(
263267
pre_process=self.pre_process,
264268
layer_type_list=layer_type_list,
265269
pp_layer_offset=layer_offset,
266-
sub_layer_offset=sub_layer_offset,
267270
post_process=self.post_process,
268271
dtype=config.params_dtype,
269272
pg_collection=self.pg_collection,

0 commit comments

Comments
 (0)