Skip to content

Commit 54e0a9e

Browse files
authored
[bugfix] fix embedding sharded_state_dict (#119)
1 parent 4450669 commit 54e0a9e

1 file changed

Lines changed: 42 additions & 2 deletions

File tree

src/mcore_bridge/model/gpt_model.py

Lines changed: 42 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -12,13 +12,14 @@
1212
from megatron.core.extensions.transformer_engine import TELinear
1313
from megatron.core.inference.contexts import BaseInferenceContext
1414
from megatron.core.models.common.embeddings.rotary_pos_embedding import RotaryEmbedding
15+
from megatron.core.models.common.language_module.language_module import LanguageModule
1516
from megatron.core.models.gpt import GPTModel as McoreGPTModel
1617
from megatron.core.packed_seq_params import PackedSeqParams
1718
from megatron.core.tensor_parallel.mappings import (copy_to_tensor_model_parallel_region,
1819
gather_from_sequence_parallel_region,
19-
gather_from_tensor_model_parallel_region,
2020
reduce_from_tensor_model_parallel_region)
21-
from megatron.core.transformer.multi_token_prediction import MTPLossAutoScaler, MTPLossLoggingHelper
21+
from megatron.core.transformer.multi_token_prediction import (MTPLossAutoScaler, MTPLossLoggingHelper,
22+
tie_word_embeddings_state_dict)
2223
from megatron.core.transformer.spec_utils import ModuleSpec
2324
from megatron.core.utils import WrappedTensor, deprecate_inference_params
2425
from packaging import version
@@ -579,3 +580,42 @@ def _postprocess(self,
579580

580581
def get_input_tensor(self):
581582
return self.decoder.input_tensor
583+
584+
def sharded_state_dict(
585+
self,
586+
prefix: str = '',
587+
sharded_offsets: tuple = (),
588+
metadata: Optional[dict] = None,
589+
) -> ShardedStateDict:
590+
"""Override to support embedding-task models that have no output_layer.
591+
592+
``LanguageModule.sharded_state_dict`` and ``McoreGPTModel.sharded_state_dict``
593+
both assume an ``output_layer`` exists. For ``task_type == 'embedding'`` we set
594+
``self.output_layer = None``, so we bypass those output_layer-related branches
595+
by calling the grandparent (``MegatronModule``) directly, while preserving the
596+
MTP embedding tying behavior from ``LanguageModule``.
597+
"""
598+
if getattr(self, 'output_layer', None) is not None:
599+
return super().sharded_state_dict(prefix, sharded_offsets, metadata)
600+
601+
assert not sharded_offsets, 'Unexpected sharded offsets'
602+
if mcore_016:
603+
from megatron.core.transformer.utils import ensure_metadata_has_dp_cp_group
604+
metadata = ensure_metadata_has_dp_cp_group(metadata)
605+
kwargs = {'tp_group': self.tp_group, 'dp_cp_group': metadata['dp_cp_group']}
606+
else:
607+
kwargs = {}
608+
# Skip LanguageModule.sharded_state_dict to avoid KeyError on output_layer.weight.
609+
sharded_state_dict = super(LanguageModule, self).sharded_state_dict(prefix, sharded_offsets, metadata)
610+
611+
# Preserve MTP embedding tying behavior from LanguageModule.
612+
if getattr(self, 'mtp_process', False) and not self.pre_process:
613+
first_stage_word_emb_key = f'{prefix}embedding.word_embeddings.weight'
614+
emb_weight = self.embedding.word_embeddings.weight
615+
tie_word_embeddings_state_dict(
616+
sharded_state_dict,
617+
emb_weight,
618+
first_stage_word_emb_key,
619+
**kwargs,
620+
)
621+
return sharded_state_dict

0 commit comments

Comments
 (0)