From 822f0d01e2303f4dc4dbff597f8bc1d0411dddfe Mon Sep 17 00:00:00 2001 From: Jintao Huang Date: Wed, 17 Jun 2026 16:19:28 +0800 Subject: [PATCH 1/5] update --- src/mcore_bridge/model/gpt_model.py | 39 ++++++++++++++++++++++++----- 1 file changed, 33 insertions(+), 6 deletions(-) diff --git a/src/mcore_bridge/model/gpt_model.py b/src/mcore_bridge/model/gpt_model.py index c7b7e0a..a1ea4c9 100644 --- a/src/mcore_bridge/model/gpt_model.py +++ b/src/mcore_bridge/model/gpt_model.py @@ -33,6 +33,7 @@ logger = get_logger() mcore_016 = version.parse(megatron.core.__version__) >= version.parse('0.16.0rc0') +mcore_019 = version.parse(megatron.core.__version__) >= version.parse('0.19.0rc0') class OutputLayerLinear(TELinear): @@ -507,12 +508,38 @@ def _postprocess(self, if self.training: mtp_loss_for_log = ( torch.sum(mtp_loss) / num_tokens if num_tokens > 0 else mtp_loss.new_tensor(0.0)) - MTPLossLoggingHelper.save_loss_to_tracker( - mtp_loss_for_log, - mtp_layer_number, - self.config.mtp_unroll_steps, - avg_group=parallel_state.get_data_parallel_group(with_context_parallel=True), - ) + if hasattr(MTPLossLoggingHelper, 'save_metrics_to_tracker'): + # mcore >= 0.19 main branch: save_metrics_to_tracker with correct/total + with torch.no_grad(): + preds = torch.argmax(mtp_logits, dim=-1) + valid = loss_mask_.bool() + correct = ((preds == mtp_labels) & valid).sum().float() + total = valid.sum().float() + MTPLossLoggingHelper.save_metrics_to_tracker( + mtp_loss_for_log, + correct, + total, + mtp_layer_number, + self.config.mtp_unroll_steps, + avg_group=parallel_state.get_data_parallel_group(with_context_parallel=True), + ) + elif mcore_019: + # mcore >= 0.19 dev branch: save_loss_to_tracker with num_tokens + MTPLossLoggingHelper.save_loss_to_tracker( + mtp_loss_for_log, + num_tokens, + mtp_layer_number, + self.config.mtp_unroll_steps, + avg_group=parallel_state.get_data_parallel_group(with_context_parallel=True), + ) + else: + # mcore < 0.19: original signature + MTPLossLoggingHelper.save_loss_to_tracker( + mtp_loss_for_log, + mtp_layer_number, + self.config.mtp_unroll_steps, + avg_group=parallel_state.get_data_parallel_group(with_context_parallel=True), + ) mtp_loss_scale = self.config.mtp_loss_scaling_factor / self.config.mtp_unroll_steps # Clamp to avoid 0/0=NaN when a CP rank's tokens are all rolled out of range. safe_num_tokens = num_tokens.clamp(min=1) From d11d3a5285496f1cb9c5cc7d3d437395a3ca197d Mon Sep 17 00:00:00 2001 From: Jintao Huang Date: Wed, 17 Jun 2026 16:29:42 +0800 Subject: [PATCH 2/5] fix --- src/mcore_bridge/model/gpt_model.py | 13 ++++++++----- 1 file changed, 8 insertions(+), 5 deletions(-) diff --git a/src/mcore_bridge/model/gpt_model.py b/src/mcore_bridge/model/gpt_model.py index a1ea4c9..87dd1cb 100644 --- a/src/mcore_bridge/model/gpt_model.py +++ b/src/mcore_bridge/model/gpt_model.py @@ -510,11 +510,14 @@ def _postprocess(self, torch.sum(mtp_loss) / num_tokens if num_tokens > 0 else mtp_loss.new_tensor(0.0)) if hasattr(MTPLossLoggingHelper, 'save_metrics_to_tracker'): # mcore >= 0.19 main branch: save_metrics_to_tracker with correct/total - with torch.no_grad(): - preds = torch.argmax(mtp_logits, dim=-1) - valid = loss_mask_.bool() - correct = ((preds == mtp_labels) & valid).sum().float() - total = valid.sum().float() + from megatron.core.transformer.multi_token_prediction import _compute_mtp_acceptance_counts + correct, total = _compute_mtp_acceptance_counts( + mtp_logits, + mtp_labels, + loss_mask_, + output_layer=None, + runtime_gather_output=True, + ) MTPLossLoggingHelper.save_metrics_to_tracker( mtp_loss_for_log, correct, From c672d7f089d21a27d095dc7f3d3e521f4bc9a275 Mon Sep 17 00:00:00 2001 From: Jintao Huang Date: Wed, 17 Jun 2026 16:38:54 +0800 Subject: [PATCH 3/5] fix --- src/mcore_bridge/model/gpt_model.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/src/mcore_bridge/model/gpt_model.py b/src/mcore_bridge/model/gpt_model.py index 87dd1cb..7780ae9 100644 --- a/src/mcore_bridge/model/gpt_model.py +++ b/src/mcore_bridge/model/gpt_model.py @@ -529,7 +529,7 @@ def _postprocess(self, elif mcore_019: # mcore >= 0.19 dev branch: save_loss_to_tracker with num_tokens MTPLossLoggingHelper.save_loss_to_tracker( - mtp_loss_for_log, + torch.sum(mtp_loss), num_tokens, mtp_layer_number, self.config.mtp_unroll_steps, From f61cab94992a854fc74492d8764001c7b9a3afd6 Mon Sep 17 00:00:00 2001 From: Jintao Huang Date: Wed, 17 Jun 2026 16:49:05 +0800 Subject: [PATCH 4/5] update --- src/mcore_bridge/model/gpt_model.py | 20 +++++--------------- 1 file changed, 5 insertions(+), 15 deletions(-) diff --git a/src/mcore_bridge/model/gpt_model.py b/src/mcore_bridge/model/gpt_model.py index 7780ae9..c33303e 100644 --- a/src/mcore_bridge/model/gpt_model.py +++ b/src/mcore_bridge/model/gpt_model.py @@ -508,24 +508,19 @@ def _postprocess(self, if self.training: mtp_loss_for_log = ( torch.sum(mtp_loss) / num_tokens if num_tokens > 0 else mtp_loss.new_tensor(0.0)) + avg_group = parallel_state.get_data_parallel_group(with_context_parallel=True) if hasattr(MTPLossLoggingHelper, 'save_metrics_to_tracker'): # mcore >= 0.19 main branch: save_metrics_to_tracker with correct/total from megatron.core.transformer.multi_token_prediction import _compute_mtp_acceptance_counts correct, total = _compute_mtp_acceptance_counts( - mtp_logits, - mtp_labels, - loss_mask_, - output_layer=None, - runtime_gather_output=True, - ) + mtp_logits, mtp_labels, loss_mask_, output_layer=None, runtime_gather_output=True) MTPLossLoggingHelper.save_metrics_to_tracker( mtp_loss_for_log, correct, total, mtp_layer_number, self.config.mtp_unroll_steps, - avg_group=parallel_state.get_data_parallel_group(with_context_parallel=True), - ) + avg_group=avg_group) elif mcore_019: # mcore >= 0.19 dev branch: save_loss_to_tracker with num_tokens MTPLossLoggingHelper.save_loss_to_tracker( @@ -533,16 +528,11 @@ def _postprocess(self, num_tokens, mtp_layer_number, self.config.mtp_unroll_steps, - avg_group=parallel_state.get_data_parallel_group(with_context_parallel=True), - ) + avg_group=avg_group) else: # mcore < 0.19: original signature MTPLossLoggingHelper.save_loss_to_tracker( - mtp_loss_for_log, - mtp_layer_number, - self.config.mtp_unroll_steps, - avg_group=parallel_state.get_data_parallel_group(with_context_parallel=True), - ) + mtp_loss_for_log, mtp_layer_number, self.config.mtp_unroll_steps, avg_group=avg_group) mtp_loss_scale = self.config.mtp_loss_scaling_factor / self.config.mtp_unroll_steps # Clamp to avoid 0/0=NaN when a CP rank's tokens are all rolled out of range. safe_num_tokens = num_tokens.clamp(min=1) From 1c75cc37bfe533267fd811125bb5bc71cd7053c7 Mon Sep 17 00:00:00 2001 From: Jintao Huang Date: Wed, 17 Jun 2026 17:33:50 +0800 Subject: [PATCH 5/5] fix --- src/mcore_bridge/model/gpt_model.py | 8 -------- 1 file changed, 8 deletions(-) diff --git a/src/mcore_bridge/model/gpt_model.py b/src/mcore_bridge/model/gpt_model.py index c33303e..2158b90 100644 --- a/src/mcore_bridge/model/gpt_model.py +++ b/src/mcore_bridge/model/gpt_model.py @@ -521,14 +521,6 @@ def _postprocess(self, mtp_layer_number, self.config.mtp_unroll_steps, avg_group=avg_group) - elif mcore_019: - # mcore >= 0.19 dev branch: save_loss_to_tracker with num_tokens - MTPLossLoggingHelper.save_loss_to_tracker( - torch.sum(mtp_loss), - num_tokens, - mtp_layer_number, - self.config.mtp_unroll_steps, - avg_group=avg_group) else: # mcore < 0.19: original signature MTPLossLoggingHelper.save_loss_to_tracker(