From 0d3b93b1a2cbd8e9a22b56f0077f98d577b6ed8a Mon Sep 17 00:00:00 2001 From: xiaoyao0115 <1804647152@qq.com> Date: Mon, 29 Jun 2026 05:26:50 -0700 Subject: [PATCH 1/2] Restore per-microbatch MTP loss logging Signed-off-by: xiaoyao0115 <1804647152@qq.com> --- .../transformer/multi_token_prediction.py | 43 +++++++------ megatron/training/training.py | 8 +-- .../test_multi_token_prediction.py | 62 ++++++++++++------- 3 files changed, 65 insertions(+), 48 deletions(-) diff --git a/megatron/core/transformer/multi_token_prediction.py b/megatron/core/transformer/multi_token_prediction.py index ffc19e0cc60..57450beb3e3 100755 --- a/megatron/core/transformer/multi_token_prediction.py +++ b/megatron/core/transformer/multi_token_prediction.py @@ -371,8 +371,8 @@ def save_metrics_to_tracker( """Save normalized MTP loss and acceptance counts for logging. This compatibility path is used by tests and callers that already - computed a normalized per-layer loss. Dynamic-CP code should use - ``save_loss_to_tracker`` so loss is weighted by token counts. + computed a normalized per-layer loss. Dynamic-CP code uses + ``save_loss_to_tracker`` to normalize each local contribution safely. """ if layer_number is None: return @@ -402,11 +402,12 @@ def save_loss_to_tracker( reduce_group: Optional[torch.distributed.ProcessGroup] = None, avg_group: Optional[torch.distributed.ProcessGroup] = None, ): - """Save the mtp loss sum and token count for logging. + """Normalize and accumulate a local MTP loss for logging. - Stores raw sums so that the global per-token loss can be computed - correctly after all-reduce, even when token counts differ across - ranks (e.g. Dynamic CP) or microbatches. + MTP is normalized independently for each microbatch. The tracker + accumulates those normalized losses, then reduction combines the sums + across ranks. This intentionally preserves sequence-packing semantics + instead of changing the metric to global ``sum(loss) / sum(tokens)``. Args: loss_sum (torch.Tensor): Sum of per-element losses on this rank. @@ -415,8 +416,8 @@ def save_loss_to_tracker( num_layers (int): The number of total layers. correct (Optional[torch.Tensor]): Number of correct MTP predictions. total (Optional[torch.Tensor]): Total number of MTP predictions. - reduce_group (torch.distributed.ProcessGroup): The group for sum-reducing losses. - avg_group (torch.distributed.ProcessGroup): The group for sum-reducing before averaging. + reduce_group (torch.distributed.ProcessGroup): Group for summing losses. + avg_group (torch.distributed.ProcessGroup): Group for averaging losses. """ if layer_number is None: return @@ -424,9 +425,8 @@ def save_loss_to_tracker( tracker = MTPLossLoggingHelper.tracker if "loss_sums" not in tracker: tracker["loss_sums"] = torch.zeros(num_layers, device=torch.cuda.current_device()) - tracker["num_tokens"] = torch.zeros(num_layers, device=torch.cuda.current_device()) + loss_sum = (loss_sum * (num_tokens > 0).to(loss_sum.dtype)) / num_tokens.clamp(min=1) tracker["loss_sums"][layer_number] += loss_sum.detach() - tracker["num_tokens"][layer_number] += num_tokens.detach() if correct is not None and total is not None: if "correct_values" not in tracker: tracker["correct_values"] = torch.zeros( @@ -485,7 +485,6 @@ def clean_loss_in_tracker(): tracker = MTPLossLoggingHelper.tracker if "loss_sums" in tracker: tracker["loss_sums"].zero_() - tracker["num_tokens"].zero_() if "values" in tracker: tracker["values"].zero_() if "correct_values" in tracker: @@ -499,21 +498,21 @@ def clean_loss_in_tracker(): def reduce_loss_in_tracker(): """Collect and reduce the mtp losses across ranks. - Packs loss sums and token counts into a single tensor for one - all-reduce, then computes per-token loss. This produces correct - weighted-average results even when ranks hold different numbers - of tokens (e.g. Dynamic CP with variable CP sizes). + Each element is already a sum of normalized microbatch losses. Sum + reductions preserve additive groups, while the DP+CP average keeps the + legacy logging contract. """ tracker = MTPLossLoggingHelper.tracker if "loss_sums" not in tracker: return - packed = torch.cat([tracker["loss_sums"], tracker["num_tokens"]]) - for group_key in ('reduce_group', 'avg_group'): - group = tracker.get(group_key) - if group is not None: - torch.distributed.all_reduce(packed, group=group) - loss_sums, num_tokens = packed.chunk(2) - tracker["values"] = loss_sums / num_tokens.clamp(min=1) + values = tracker["loss_sums"] + if tracker.get('reduce_group') is not None: + torch.distributed.all_reduce(values, group=tracker['reduce_group']) + if tracker.get('avg_group') is not None: + torch.distributed.all_reduce( + values, group=tracker['avg_group'], op=torch.distributed.ReduceOp.AVG + ) + tracker["values"] = values @staticmethod def track_mtp_metrics(loss_scale, iteration, writer, wandb_writer=None, total_loss_dict=None): diff --git a/megatron/training/training.py b/megatron/training/training.py index e4075078f0e..09ecaeab5dc 100644 --- a/megatron/training/training.py +++ b/megatron/training/training.py @@ -3130,10 +3130,10 @@ def training_log( # Log MTP metrics. if args.mtp_num_layers is not None: - # MTP tracker stores raw loss sums and token counts, so after reduction - # tracker["values"] already equals the per-token loss (loss_sum / num_tokens) - # aggregated across all ranks and microbatches. No further scaling needed. - mtp_loss_scale = 1.0 + # The tracker stores a sum of normalized microbatch losses. + # Sequence-packing schedulers may change the number of microbatches for + # this step, so use the scheduled count passed to training_log. + mtp_loss_scale = 1 / (num_microbatches or get_num_microbatches()) MTPLossLoggingHelper.track_mtp_metrics( mtp_loss_scale, iteration, writer, wandb_writer, total_loss_dict ) diff --git a/tests/unit_tests/transformer/test_multi_token_prediction.py b/tests/unit_tests/transformer/test_multi_token_prediction.py index 849d8a4d42c..2a1b4b7bc39 100644 --- a/tests/unit_tests/transformer/test_multi_token_prediction.py +++ b/tests/unit_tests/transformer/test_multi_token_prediction.py @@ -685,8 +685,8 @@ def test_forward_backward(self, tmp_path_dist_ckpt, tp, cp, full_recompute): labels=labels, loss_mask=loss_mask, ) - # forward only fills raw loss_sums / num_tokens. Trigger the reduction - # so tracker["values"] (per-token loss across DP+CP) becomes available. + # Forward accumulates normalized losses. Trigger the DP+CP + # reduction so tracker["values"] becomes available. MTPLossLoggingHelper.reduce_loss_in_tracker() tracker = MTPLossLoggingHelper.tracker assert "values" in tracker @@ -741,9 +741,7 @@ def set_ckpt_path(ckpt_path): labels=labels, loss_mask=loss_mask, ) - # reduce_loss_in_tracker performs sum-reduce of loss_sums and - # num_tokens across DP+CP, then computes sum/sum -- already the - # correct global per-token loss, no extra CP averaging needed. + # Combine normalized loss contributions across DP+CP. MTPLossLoggingHelper.reduce_loss_in_tracker() tracker = MTPLossLoggingHelper.tracker assert "values" in tracker @@ -845,8 +843,7 @@ def test_packed_sequences(self, tp, cp): assert output.shape[0] == 1 # batch size assert output.shape[1] == total_seq_length - # Verify MTP loss was computed; reduce raw loss_sums/num_tokens into - # tracker["values"] (per-token loss) first. + # Verify MTP loss was computed; reduce local contributions first. MTPLossLoggingHelper.reduce_loss_in_tracker() tracker = MTPLossLoggingHelper.tracker assert "values" in tracker @@ -1234,7 +1231,7 @@ def test_save_metrics_to_tracker(self): assert tracker["avg_group"] is None def test_save_loss_to_tracker(self): - """Test saving loss sum and token count to tracker.""" + """Test saving a normalized loss to the tracker.""" loss_sum = torch.tensor(1.3) num_tokens = torch.tensor(5.0) layer_number = 2 @@ -1247,14 +1244,11 @@ def test_save_loss_to_tracker(self): num_layers=num_layers, ) - # Tracker now stores raw loss sums and token counts; per-token loss - # is computed in reduce_loss_in_tracker. assert "loss_sums" in MTPLossLoggingHelper.tracker - assert "num_tokens" in MTPLossLoggingHelper.tracker assert MTPLossLoggingHelper.tracker["loss_sums"].shape == (num_layers,) - assert MTPLossLoggingHelper.tracker["num_tokens"].shape == (num_layers,) - assert MTPLossLoggingHelper.tracker["loss_sums"][layer_number] == loss_sum - assert MTPLossLoggingHelper.tracker["num_tokens"][layer_number] == num_tokens + assert torch.isclose( + MTPLossLoggingHelper.tracker["loss_sums"][layer_number], loss_sum / num_tokens + ) assert MTPLossLoggingHelper.tracker["reduce_group"] is None assert MTPLossLoggingHelper.tracker["avg_group"] is None @@ -1271,7 +1265,7 @@ def __init__(self, gather_output): assert _mtp_logits_are_vocab_sharded(DummyOutputLayer(gather_output=True), False) is True def test_track_mtp_metrics(self): - """Test tracking MTP metrics including token-weighted loss and acceptance rate.""" + """Test tracking normalized MTP loss and acceptance rate.""" loss_sum = torch.tensor(2.3) num_tokens = torch.tensor(1.0) num_layers = self.num_layers @@ -1371,10 +1365,36 @@ def log(self, metrics, iteration): # Verify tracker is cleaned assert torch.all(MTPLossLoggingHelper.tracker["loss_sums"] == 0) - assert torch.all(MTPLossLoggingHelper.tracker["num_tokens"] == 0) assert MTPLossLoggingHelper.tracker["reduce_group"] is None assert MTPLossLoggingHelper.tracker["avg_group"] is None + def test_microbatch_means_are_not_globally_token_weighted(self): + """MTP logging preserves the pre-#4226 microbatch-normalized semantics.""" + MTPLossLoggingHelper.save_loss_to_tracker( + loss_sum=torch.tensor(8.0), num_tokens=torch.tensor(2.0), layer_number=0, num_layers=1 + ) + MTPLossLoggingHelper.save_loss_to_tracker( + loss_sum=torch.tensor(4.0), num_tokens=torch.tensor(4.0), layer_number=0, num_layers=1 + ) + + class DummyWriter: + def __init__(self): + self.scalars = {} + + def add_scalar(self, name, value, iteration): + self.scalars[name] = value + + writer = DummyWriter() + MTPLossLoggingHelper.track_mtp_metrics( + loss_scale=0.5, iteration=1, writer=writer, total_loss_dict={} + ) + + logged_loss = torch.as_tensor(writer.scalars["mtp_1 loss"]) + microbatch_mean_average = torch.tensor(((8.0 / 2.0) + (4.0 / 4.0)) / 2.0) + global_token_weighted = torch.tensor((8.0 + 4.0) / (2.0 + 4.0)) + assert torch.isclose(logged_loss, microbatch_mean_average) + assert not torch.isclose(logged_loss, global_token_weighted) + def test_track_mtp_loss_preserves_legacy_normalized_loss_semantics(self): """MTP loss logging should not become token-weighted when acceptance counters are added.""" first_loss = torch.tensor(10.0) @@ -1566,8 +1586,8 @@ def test_forward_backward_mamba(self, tmp_path_dist_ckpt, tp, cp): labels=labels, loss_mask=loss_mask, ) - # forward only fills raw loss_sums / num_tokens. Reduce them first so - # tracker["values"] (per-token loss across DP+CP) becomes available. + # Forward accumulates normalized losses. Reduce them first so + # tracker["values"] becomes available. MTPLossLoggingHelper.reduce_loss_in_tracker() tracker = MTPLossLoggingHelper.tracker assert "values" in tracker @@ -1617,8 +1637,7 @@ def set_ckpt_path(ckpt_path): labels=labels, loss_mask=loss_mask, ) - # reduce_loss_in_tracker already computes the cross-DP+CP per-token - # loss (sum/sum), no extra CP averaging needed. + # Combine normalized loss contributions across DP+CP. MTPLossLoggingHelper.reduce_loss_in_tracker() tracker = MTPLossLoggingHelper.tracker assert "values" in tracker @@ -1922,8 +1941,7 @@ def model_provider( ) assert torch.isfinite(output).all(), f"Non-finite output (TP={tp})" - # Reduce raw loss_sums/num_tokens into tracker["values"] (per-token - # loss across DP+CP) before reading. + # Reduce normalized loss contributions before reading. MTPLossLoggingHelper.reduce_loss_in_tracker() tracker = MTPLossLoggingHelper.tracker assert "values" in tracker, f"MTP loss not logged (TP={tp})" From 814701f3253cdf8ba25e25fe4d48ac89372ef5a5 Mon Sep 17 00:00:00 2001 From: xiaoyao0115 <1804647152@qq.com> Date: Thu, 9 Jul 2026 08:35:45 -0700 Subject: [PATCH 2/2] Update MTP logging golden values Signed-off-by: xiaoyao0115 <1804647152@qq.com> --- .../golden_values_dev_dgx_gb200.json | 18 +++++++++--------- .../golden_values_dev_dgx_gb200_2nd.json | 8 ++++---- .../golden_values_dev_dgx_h100.json | 12 ++++++------ .../golden_values_dev_dgx_h100_2nd.json | 4 ++-- .../golden_values_dev_dgx_gb200.json | 8 ++++---- 5 files changed, 25 insertions(+), 25 deletions(-) diff --git a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph/golden_values_dev_dgx_gb200.json b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph/golden_values_dev_dgx_gb200.json index 8e0b6544b90..ae215b3314a 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph/golden_values_dev_dgx_gb200.json +++ b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph/golden_values_dev_dgx_gb200.json @@ -433,7 +433,7 @@ "step_interval": 1, "values": { "1": 10.91979, - "2": 10.92265, + "2": 10.92264, "3": 10.92219, "4": 10.91957, "5": 10.90846, @@ -442,7 +442,7 @@ "8": 10.92446, "9": 10.91998, "10": 10.92341, - "11": 10.91589, + "11": 10.91588, "12": 10.90792, "13": 10.92213, "14": 10.91569, @@ -455,8 +455,8 @@ "21": 10.90603, "22": 10.90322, "23": 10.90169, - "24": 10.89089, - "25": 10.88251, + "24": 10.8909, + "25": 10.8825, "26": 10.89359, "27": 10.8887, "28": 10.87619, @@ -477,7 +477,7 @@ "43": 10.84328, "44": 10.82898, "45": 10.84347, - "46": 10.83297, + "46": 10.83298, "47": 10.83911, "48": 10.82542, "49": 10.83132, @@ -494,7 +494,7 @@ "60": 10.77895, "61": 10.77556, "62": 10.76277, - "63": 10.77595, + "63": 10.77594, "64": 10.76136, "65": 10.7585, "66": 10.75798, @@ -514,17 +514,17 @@ "80": 10.67858, "81": 10.67147, "82": 10.65165, - "83": 10.63056, + "83": 10.63057, "84": 10.61714, "85": 10.60392, "86": 10.63183, "87": 10.62791, - "88": 10.62832, + "88": 10.62833, "89": 10.59789, "90": 10.59506, "91": 10.60606, "92": 10.58205, - "93": 10.55313, + "93": 10.55314, "94": 10.58516, "95": 10.57313, "96": 10.56963, diff --git a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph/golden_values_dev_dgx_gb200_2nd.json b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph/golden_values_dev_dgx_gb200_2nd.json index d280e191db7..58d969d35d7 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph/golden_values_dev_dgx_gb200_2nd.json +++ b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph/golden_values_dev_dgx_gb200_2nd.json @@ -494,7 +494,7 @@ "60": 10.77895, "61": 10.77556, "62": 10.76277, - "63": 10.77595, + "63": 10.77594, "64": 10.76136, "65": 10.7585, "66": 10.75798, @@ -514,17 +514,17 @@ "80": 10.67858, "81": 10.67147, "82": 10.65165, - "83": 10.63056, + "83": 10.63057, "84": 10.61714, "85": 10.60392, "86": 10.63183, "87": 10.62791, - "88": 10.62832, + "88": 10.62833, "89": 10.59789, "90": 10.59506, "91": 10.60606, "92": 10.58205, - "93": 10.55313, + "93": 10.55314, "94": 10.58516, "95": 10.57313, "96": 10.56963, diff --git a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph/golden_values_dev_dgx_h100.json b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph/golden_values_dev_dgx_h100.json index 48b68ba5823..cc3963c29d9 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph/golden_values_dev_dgx_h100.json +++ b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph/golden_values_dev_dgx_h100.json @@ -437,7 +437,7 @@ "3": 10.93384, "4": 10.92739, "5": 10.90724, - "6": 10.91816, + "6": 10.91817, "7": 10.92486, "8": 10.92528, "9": 10.93457, @@ -445,17 +445,17 @@ "11": 10.91896, "12": 10.91863, "13": 10.92814, - "14": 10.91204, + "14": 10.91203, "15": 10.92041, "16": 10.92467, "17": 10.92235, - "18": 10.9072, + "18": 10.90719, "19": 10.91438, "20": 10.90506, "21": 10.91161, "22": 10.89778, "23": 10.90483, - "24": 10.88965, + "24": 10.88964, "25": 10.89765, "26": 10.88453, "27": 10.89849, @@ -491,7 +491,7 @@ "57": 10.78961, "58": 10.79824, "59": 10.78095, - "60": 10.77504, + "60": 10.77503, "61": 10.77627, "62": 10.7614, "63": 10.78392, @@ -511,7 +511,7 @@ "77": 10.69055, "78": 10.68188, "79": 10.66968, - "80": 10.67687, + "80": 10.67688, "81": 10.66904, "82": 10.65016, "83": 10.6267, diff --git a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph/golden_values_dev_dgx_h100_2nd.json b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph/golden_values_dev_dgx_h100_2nd.json index 786d23d265f..357c399d4b7 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph/golden_values_dev_dgx_h100_2nd.json +++ b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph/golden_values_dev_dgx_h100_2nd.json @@ -491,7 +491,7 @@ "57": 10.78961, "58": 10.79824, "59": 10.78095, - "60": 10.77504, + "60": 10.77503, "61": 10.77627, "62": 10.7614, "63": 10.78392, @@ -511,7 +511,7 @@ "77": 10.69055, "78": 10.68188, "79": 10.66968, - "80": 10.67687, + "80": 10.67688, "81": 10.66904, "82": 10.65016, "83": 10.6267, diff --git a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph_1node/golden_values_dev_dgx_gb200.json b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph_1node/golden_values_dev_dgx_gb200.json index 8dd231f1acc..d82c4eb4512 100644 --- a/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph_1node/golden_values_dev_dgx_gb200.json +++ b/tests/functional_tests/test_cases/moe/gpt3_moe_mcore_te_tp4_ep2_etp2_pp2_scoped_cudagraph_1node/golden_values_dev_dgx_gb200.json @@ -464,7 +464,7 @@ "30": 10.93961, "31": 10.91806, "32": 10.92899, - "33": 10.92158, + "33": 10.92159, "34": 10.92688, "35": 10.91862, "36": 10.917, @@ -475,13 +475,13 @@ "41": 10.90546, "42": 10.88722, "43": 10.89763, - "44": 10.87485, + "44": 10.87484, "45": 10.88603, "46": 10.87926, "47": 10.87569, "48": 10.86052, "49": 10.86102, - "50": 10.84774, + "50": 10.84773, "51": 10.86037, "52": 10.84549, "53": 10.85137, @@ -497,7 +497,7 @@ "63": 10.80803, "64": 10.80094, "65": 10.78782, - "66": 10.78838, + "66": 10.78839, "67": 10.78222, "68": 10.76003, "69": 10.78043,