|
12 | 12 | from megatron.core.extensions.transformer_engine import TELinear |
13 | 13 | from megatron.core.inference.contexts import BaseInferenceContext |
14 | 14 | from megatron.core.models.common.embeddings.rotary_pos_embedding import RotaryEmbedding |
| 15 | +from megatron.core.models.common.language_module.language_module import LanguageModule |
15 | 16 | from megatron.core.models.gpt import GPTModel as McoreGPTModel |
16 | 17 | from megatron.core.packed_seq_params import PackedSeqParams |
17 | 18 | from megatron.core.tensor_parallel.mappings import (copy_to_tensor_model_parallel_region, |
18 | 19 | gather_from_sequence_parallel_region, |
19 | | - gather_from_tensor_model_parallel_region, |
20 | 20 | 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) |
22 | 23 | from megatron.core.transformer.spec_utils import ModuleSpec |
23 | 24 | from megatron.core.utils import WrappedTensor, deprecate_inference_params |
24 | 25 | from packaging import version |
@@ -579,3 +580,42 @@ def _postprocess(self, |
579 | 580 |
|
580 | 581 | def get_input_tensor(self): |
581 | 582 | 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