From 80e5a20c1c499ae10011e26d2dec278bbb03cc10 Mon Sep 17 00:00:00 2001 From: Teodor-Dumitru Ene Date: Fri, 6 Feb 2026 13:17:42 -0600 Subject: [PATCH 1/2] Enable inference for MTP models --- megatron/core/models/gpt/gpt_model.py | 1 + megatron/core/models/mamba/mamba_model.py | 1 + megatron/core/transformer/multi_token_prediction.py | 7 ++++++- 3 files changed, 8 insertions(+), 1 deletion(-) diff --git a/megatron/core/models/gpt/gpt_model.py b/megatron/core/models/gpt/gpt_model.py index b0d3f085240..4170ec27e17 100644 --- a/megatron/core/models/gpt/gpt_model.py +++ b/megatron/core/models/gpt/gpt_model.py @@ -621,6 +621,7 @@ def _postprocess( config=self.config, cp_group=self.pg_collection.cp, packed_seq_params=packed_seq_params, + in_inference_mode=in_inference_mode, ) sequence_parallel_override = False diff --git a/megatron/core/models/mamba/mamba_model.py b/megatron/core/models/mamba/mamba_model.py index 8dd614fdaaa..764a3e36664 100644 --- a/megatron/core/models/mamba/mamba_model.py +++ b/megatron/core/models/mamba/mamba_model.py @@ -335,6 +335,7 @@ def forward( config=self.config, cp_group=self.pg_collection.cp, packed_seq_params=packed_seq_params, + in_inference_mode=in_inference_mode, ) sequence_parallel_override = False diff --git a/megatron/core/transformer/multi_token_prediction.py b/megatron/core/transformer/multi_token_prediction.py index 1c431491ca2..3ea31eafb8e 100755 --- a/megatron/core/transformer/multi_token_prediction.py +++ b/megatron/core/transformer/multi_token_prediction.py @@ -617,6 +617,7 @@ def process_mtp_loss( config: TransformerConfig, cp_group: Optional[torch.distributed.ProcessGroup] = None, packed_seq_params: Optional[PackedSeqParams] = None, + in_inference_mode: bool = False, ) -> Tensor: """Process Multi-Token Prediction (MTP) loss computation. @@ -635,14 +636,18 @@ def process_mtp_loss( config (TransformerConfig): Model configuration containing mtp_num_layers etc. cp_group (Optional[ProcessGroup]): Context parallelism process group. packed_seq_params (Optional[PackedSeqParams]): Packed sequence parameters. + in_inference_mode (bool): Whether the model is in inference mode. Affects loss computation. Returns: Tensor: Updated hidden states after MTP loss processing (first chunk only). """ - mtp_labels = labels.clone() hidden_states_list = torch.chunk(hidden_states, 1 + config.mtp_num_layers, dim=0) hidden_states = hidden_states_list[0] + if in_inference_mode: + return hidden_states + + mtp_labels = labels.clone() if loss_mask is None: loss_mask = torch.ones_like(mtp_labels) From 44c116511878ef8ad50cb6cf0c9b56323a9cb17a Mon Sep 17 00:00:00 2001 From: Teodor-Dumitru Ene Date: Mon, 9 Feb 2026 16:32:26 -0600 Subject: [PATCH 2/2] Fix outlier cases --- megatron/core/models/gpt/gpt_model.py | 1 - megatron/core/models/mamba/mamba_model.py | 1 - megatron/core/transformer/multi_token_prediction.py | 4 +--- 3 files changed, 1 insertion(+), 5 deletions(-) diff --git a/megatron/core/models/gpt/gpt_model.py b/megatron/core/models/gpt/gpt_model.py index 4170ec27e17..b0d3f085240 100644 --- a/megatron/core/models/gpt/gpt_model.py +++ b/megatron/core/models/gpt/gpt_model.py @@ -621,7 +621,6 @@ def _postprocess( config=self.config, cp_group=self.pg_collection.cp, packed_seq_params=packed_seq_params, - in_inference_mode=in_inference_mode, ) sequence_parallel_override = False diff --git a/megatron/core/models/mamba/mamba_model.py b/megatron/core/models/mamba/mamba_model.py index 764a3e36664..8dd614fdaaa 100644 --- a/megatron/core/models/mamba/mamba_model.py +++ b/megatron/core/models/mamba/mamba_model.py @@ -335,7 +335,6 @@ def forward( config=self.config, cp_group=self.pg_collection.cp, packed_seq_params=packed_seq_params, - in_inference_mode=in_inference_mode, ) sequence_parallel_override = False diff --git a/megatron/core/transformer/multi_token_prediction.py b/megatron/core/transformer/multi_token_prediction.py index 98aa8f41801..393accec5df 100755 --- a/megatron/core/transformer/multi_token_prediction.py +++ b/megatron/core/transformer/multi_token_prediction.py @@ -620,7 +620,6 @@ def process_mtp_loss( config: TransformerConfig, cp_group: Optional[torch.distributed.ProcessGroup] = None, packed_seq_params: Optional[PackedSeqParams] = None, - in_inference_mode: bool = False, ) -> Tensor: """Process Multi-Token Prediction (MTP) loss computation. @@ -639,7 +638,6 @@ def process_mtp_loss( config (TransformerConfig): Model configuration containing mtp_num_layers etc. cp_group (Optional[ProcessGroup]): Context parallelism process group. packed_seq_params (Optional[PackedSeqParams]): Packed sequence parameters. - in_inference_mode (bool): Whether the model is in inference mode. Affects loss computation. Returns: Tensor: Updated hidden states after MTP loss processing (first chunk only). @@ -647,7 +645,7 @@ def process_mtp_loss( hidden_states_list = torch.chunk(hidden_states, 1 + config.mtp_num_layers, dim=0) hidden_states = hidden_states_list[0] - if in_inference_mode: + if labels is None: return hidden_states mtp_labels = labels.clone()