diff --git a/megatron/core/models/gpt/gpt_model.py b/megatron/core/models/gpt/gpt_model.py index b2350e1ecf9..ac2e3f8bab1 100644 --- a/megatron/core/models/gpt/gpt_model.py +++ b/megatron/core/models/gpt/gpt_model.py @@ -682,6 +682,7 @@ def _postprocess( compute_language_model_loss=self.compute_language_model_loss, config=self.config, cp_group=self.pg_collection.cp, + tp_group=self.tp_group, packed_seq_params=packed_seq_params, scale_logits_fn=self._scale_logits if self.config.use_mup else None, input_ids=input_ids, diff --git a/megatron/core/models/hybrid/hybrid_model.py b/megatron/core/models/hybrid/hybrid_model.py index fafd50d7b5f..01dd8332008 100644 --- a/megatron/core/models/hybrid/hybrid_model.py +++ b/megatron/core/models/hybrid/hybrid_model.py @@ -559,6 +559,7 @@ def forward( compute_language_model_loss=self.compute_language_model_loss, config=self.config, cp_group=self.pg_collection.cp, + tp_group=self.tp_group, packed_seq_params=packed_seq_params, scale_logits_fn=self._scale_logits if self.config.use_mup else None, input_ids=input_ids, diff --git a/megatron/core/transformer/multi_token_prediction.py b/megatron/core/transformer/multi_token_prediction.py index 348f50053b5..d284b608e1a 100755 --- a/megatron/core/transformer/multi_token_prediction.py +++ b/megatron/core/transformer/multi_token_prediction.py @@ -340,80 +340,216 @@ def _roll_tensor_packed_seq(tensor, shifts, dims, packed_seq_params, cp_group=No class MTPLossLoggingHelper: - """Helper class for logging MTP losses.""" + """Helper class for logging MTP losses and acceptance rates.""" tracker = {} @staticmethod - def save_loss_to_tracker( + def save_metrics_to_tracker( loss: torch.Tensor, + correct: torch.Tensor, + total: torch.Tensor, layer_number: int, num_layers: int, - reduce_group: Optional[torch.distributed.ProcessGroup] = None, - avg_group: Optional[torch.distributed.ProcessGroup] = None, + reduce_group: torch.distributed.ProcessGroup = None, + avg_group: torch.distributed.ProcessGroup = None, ): - """Save the mtp loss for logging. + """Save the mtp metrics (loss, correct, total) for logging. + Args: - loss (torch.Tensor): The loss tensor. + loss (torch.Tensor): The normalized loss value for this MTP layer. + correct (torch.Tensor): Number of correct predictions. + total (torch.Tensor): Total number of predictions. layer_number (int): Layer index of the loss. num_layers (int): The number of total layers. reduce_group (torch.distributed.ProcessGroup): The group for reducing the loss. - mean_group (torch.distributed.ProcessGroup): The group for averaging the loss. + avg_group (torch.distributed.ProcessGroup): The group for averaging the loss. """ # Skip mtp loss logging if layer_number is None. if layer_number is None: return tracker = MTPLossLoggingHelper.tracker - if "values" not in tracker: - tracker["values"] = torch.zeros(num_layers, device=torch.cuda.current_device()) - tracker["values"][layer_number] += loss.detach() + if "loss_values" not in tracker: + tracker["loss_values"] = torch.zeros(num_layers, device=torch.cuda.current_device()) + if "correct_values" not in tracker: + tracker["correct_values"] = torch.zeros(num_layers, device=torch.cuda.current_device()) + if "total_values" not in tracker: + tracker["total_values"] = torch.zeros(num_layers, device=torch.cuda.current_device()) + + tracker["loss_values"][layer_number] += loss.detach() + tracker["correct_values"][layer_number] += correct.detach() + tracker["total_values"][layer_number] += total.detach() tracker["reduce_group"] = reduce_group tracker["avg_group"] = avg_group - def clean_loss_in_tracker(): - """Clear the mtp losses.""" + @staticmethod + def clean_metrics_in_tracker(): + """Clear the mtp metrics.""" tracker = MTPLossLoggingHelper.tracker - tracker["values"].zero_() + if "loss_values" in tracker: + tracker["loss_values"].zero_() + if "correct_values" in tracker: + tracker["correct_values"].zero_() + if "total_values" in tracker: + tracker["total_values"].zero_() tracker["reduce_group"] = None tracker["avg_group"] = None - def reduce_loss_in_tracker(): - """Collect and reduce the mtp losses across ranks.""" + @staticmethod + def reduce_metrics_in_tracker(): + """Collect and reduce the mtp metrics across ranks.""" tracker = MTPLossLoggingHelper.tracker - if "values" not in tracker: + if "loss_values" not in tracker: return - values = tracker["values"] - # Reduce mtp losses across ranks. + + loss_values = tracker["loss_values"] if tracker.get('reduce_group') is not None: - torch.distributed.all_reduce(values, group=tracker.get('reduce_group')) + torch.distributed.all_reduce(loss_values, group=tracker.get('reduce_group')) if tracker.get('avg_group') is not None: torch.distributed.all_reduce( - values, group=tracker['avg_group'], op=torch.distributed.ReduceOp.AVG + loss_values, group=tracker['avg_group'], op=torch.distributed.ReduceOp.AVG ) + for key in ["correct_values", "total_values"]: + if key not in tracker: + continue + values = tracker[key] + if tracker.get('reduce_group') is not None: + torch.distributed.all_reduce(values, group=tracker.get('reduce_group')) + if tracker.get('avg_group') is not None: + torch.distributed.all_reduce( + values, group=tracker['avg_group'], op=torch.distributed.ReduceOp.SUM + ) + + @staticmethod def track_mtp_metrics(loss_scale, iteration, writer, wandb_writer=None, total_loss_dict=None): """Track the Multi-Token Prediction (MTP) metrics for logging.""" - MTPLossLoggingHelper.reduce_loss_in_tracker() + MTPLossLoggingHelper.reduce_metrics_in_tracker() tracker = MTPLossLoggingHelper.tracker - if "values" not in tracker: + if "loss_values" not in tracker: return - mtp_losses = tracker["values"] * loss_scale + + mtp_losses = tracker["loss_values"] * loss_scale + mtp_corrects = tracker.get("correct_values", torch.zeros_like(mtp_losses)) + mtp_totals = tracker.get("total_values", torch.ones_like(mtp_losses)) + + # Process-local logging state; cumulative rates intentionally reset after restart/resume. + if ( + "cumulative_correct_values" not in tracker + or tracker["cumulative_correct_values"].shape != mtp_corrects.shape + ): + tracker["cumulative_correct_values"] = torch.zeros_like(mtp_corrects) + if ( + "cumulative_total_values" not in tracker + or tracker["cumulative_total_values"].shape != mtp_totals.shape + ): + tracker["cumulative_total_values"] = torch.zeros_like(mtp_totals) + + tracker["cumulative_correct_values"] += mtp_corrects + tracker["cumulative_total_values"] += mtp_totals + mtp_cumulative_corrects = tracker["cumulative_correct_values"] + mtp_cumulative_totals = tracker["cumulative_total_values"] + mtp_num_layers = mtp_losses.shape[0] for i in range(mtp_num_layers): - name = f"mtp_{i + 1} loss" + loss_name = f"mtp_{i+1} loss" + step_acc_name = f"mtp_{i+1}_acceptance_rate" + cum_acc_name = f"mtp_{i+1}_cumulative_acceptance_rate" + loss = mtp_losses[i] + # Empty masks can leave no valid MTP positions, so clamp denominators to avoid NaNs. + step_rate = (mtp_corrects[i] / torch.clamp(mtp_totals[i], min=1)) * 100.0 + cum_rate = ( + mtp_cumulative_corrects[i] / torch.clamp(mtp_cumulative_totals[i], min=1) + ) * 100.0 + if total_loss_dict is not None: - if name in total_loss_dict: - total_loss_dict[name] += loss - else: - total_loss_dict[name] = loss + total_loss_dict[loss_name] = ( + total_loss_dict.get(loss_name, torch.zeros_like(loss)) + loss + ) + if writer is not None: - writer.add_scalar(name, loss, iteration) + writer.add_scalar(loss_name, loss, iteration) + writer.add_scalar(step_acc_name, step_rate, iteration) + writer.add_scalar(cum_acc_name, cum_rate, iteration) if wandb_writer is not None: - wandb_writer.log({f"{name}": loss}, iteration) + wandb_writer.log({f"{loss_name}": loss}, iteration) + wandb_writer.log({f"{step_acc_name}": step_rate}, iteration) + wandb_writer.log({f"{cum_acc_name}": cum_rate}, iteration) - MTPLossLoggingHelper.clean_loss_in_tracker() + MTPLossLoggingHelper.clean_metrics_in_tracker() + + +def _mtp_logits_are_vocab_sharded( + output_layer: Callable, runtime_gather_output: Optional[bool] +) -> bool: + """Return whether MTP logits are still vocab-sharded across tensor-parallel ranks.""" + if runtime_gather_output is not None: + return not runtime_gather_output + return not getattr(output_layer, "gather_output", False) + + +def _vocab_parallel_argmax( + vocab_parallel_logits: Tensor, tp_group: torch.distributed.ProcessGroup, tp_size: int +) -> Tensor: + """Return global argmax ids from logits sharded across the vocab dimension.""" + vocab_shard_size = vocab_parallel_logits.size(-1) + local_max_vals, local_argmax = vocab_parallel_logits.max(dim=-1) # [s, b], [s, b] + + gathered_max_vals = [torch.empty_like(local_max_vals) for _ in range(tp_size)] + gathered_argmax = [torch.empty_like(local_argmax) for _ in range(tp_size)] + torch.distributed.all_gather(gathered_max_vals, local_max_vals, group=tp_group) + torch.distributed.all_gather(gathered_argmax, local_argmax, group=tp_group) + + stacked_max_vals = torch.stack(gathered_max_vals, dim=0) + stacked_argmax = torch.stack(gathered_argmax, dim=0) + winning_rank = stacked_max_vals.argmax(dim=0) # [s, b] + winning_local_argmax = torch.gather(stacked_argmax, 0, winning_rank.unsqueeze(0)).squeeze( + 0 + ) # [s, b] + return winning_rank * vocab_shard_size + winning_local_argmax # [s, b] + + +def _compute_mtp_acceptance_counts( + mtp_logits: Tensor, + mtp_labels: Tensor, + loss_mask: Tensor, + output_layer: Callable, + runtime_gather_output: Optional[bool], + tp_group: Optional[torch.distributed.ProcessGroup] = None, +) -> tuple[Tensor, Tensor]: + """Compute MTP acceptance correct/total counts.""" + with torch.no_grad(): + logits_are_vocab_sharded = _mtp_logits_are_vocab_sharded( + output_layer, runtime_gather_output + ) + if ( + tp_group is None + and logits_are_vocab_sharded + and parallel_state.is_initialized() + and parallel_state.get_tensor_model_parallel_world_size() > 1 + ): + raise ValueError( + "tp_group must be provided when computing MTP acceptance counts " + "from vocab-sharded logits under tensor model parallelism." + ) + tp_size = torch.distributed.get_world_size(group=tp_group) if tp_group is not None else 1 + + # Apply TP rank offsets only when logits are vocab-sharded; gathered logits already + # contain global vocab ids in their last dimension. + if tp_group is not None and tp_size > 1 and logits_are_vocab_sharded: + preds = _vocab_parallel_argmax(mtp_logits, tp_group, tp_size) + else: + preds = torch.argmax(mtp_logits, dim=-1) # [s, b] + + labels_match = mtp_labels.transpose(0, 1).contiguous() # [b, s] => [s, b] + mask_match = loss_mask.transpose(0, 1).contiguous() # [b, s] => [s, b] + valid_positions = mask_match.bool() + correct = ((preds == labels_match) & valid_positions).sum().float() + total = valid_positions.sum().float() + + return correct, total @dataclass @@ -632,6 +768,7 @@ def process_mtp_loss( compute_language_model_loss: Callable, config: TransformerConfig, cp_group: Optional[torch.distributed.ProcessGroup] = None, + tp_group: Optional[torch.distributed.ProcessGroup] = None, packed_seq_params: Optional[PackedSeqParams] = None, scale_logits_fn: Optional[Callable[[Tensor], Tensor]] = None, input_ids: Optional[Tensor] = None, @@ -652,6 +789,7 @@ def process_mtp_loss( compute_language_model_loss (Callable): Method to compute language model loss. config (TransformerConfig): Model configuration containing mtp_num_layers etc. cp_group (Optional[ProcessGroup]): Context parallelism process group. + tp_group (Optional[ProcessGroup]): Tensor parallelism process group. packed_seq_params (Optional[PackedSeqParams]): Packed sequence parameters. scale_logits_fn (Optional[Callable[[Tensor], Tensor]]): Optional function to scale logits before loss computation (e.g., MuP output scaling). @@ -707,15 +845,23 @@ def process_mtp_loss( loss_mask, num_tokens = roll_tensor( loss_mask, shifts=-1, dims=-1, cp_group=cp_group, packed_seq_params=packed_seq_params ) + mtp_loss = compute_language_model_loss(mtp_labels, mtp_logits) + mtp_loss = loss_mask * mtp_loss + if is_training: - # Safe divide without sync: mask numerator when num_tokens==0, divide by clamp(min=1) mtp_loss_for_log = ( torch.sum(mtp_loss) * (num_tokens > 0).to(mtp_loss.dtype) ) / num_tokens.clamp(min=1) - MTPLossLoggingHelper.save_loss_to_tracker( + correct, total = _compute_mtp_acceptance_counts( + mtp_logits, mtp_labels, loss_mask, output_layer, runtime_gather_output, tp_group + ) + + MTPLossLoggingHelper.save_metrics_to_tracker( mtp_loss_for_log, + correct, + total, mtp_layer_number, config.mtp_num_layers, avg_group=parallel_state.get_data_parallel_group(with_context_parallel=True), diff --git a/tests/unit_tests/transformer/test_multi_token_prediction.py b/tests/unit_tests/transformer/test_multi_token_prediction.py index 136c0b4d436..948c08cf37f 100644 --- a/tests/unit_tests/transformer/test_multi_token_prediction.py +++ b/tests/unit_tests/transformer/test_multi_token_prediction.py @@ -25,6 +25,7 @@ from megatron.core.transformer.multi_token_prediction import ( MTPLossLoggingHelper, MultiTokenPredictionBlock, + _mtp_logits_are_vocab_sharded, process_mtp_loss, roll_tensor, ) @@ -456,9 +457,9 @@ def test_forward_backward(self, tmp_path_dist_ckpt, tp, cp, full_recompute): ) tracker = MTPLossLoggingHelper.tracker mtp_loss_ref = None - assert "values" in tracker - mtp_loss_ref = tracker['values'].clone() - MTPLossLoggingHelper.clean_loss_in_tracker() + assert "loss_values" in tracker + mtp_loss_ref = tracker['loss_values'].clone() + MTPLossLoggingHelper.clean_metrics_in_tracker() iteration = 123 num_floating_point_operations_so_far = 456 @@ -509,14 +510,14 @@ def set_ckpt_path(ckpt_path): loss_mask=loss_mask, ) tracker = MTPLossLoggingHelper.tracker - assert "values" in tracker - mtp_loss = tracker['values'].clone() + assert "loss_values" in tracker + mtp_loss = tracker['loss_values'].clone() # Average MTP loss across CP ranks for comparison with reference pg_collection = ProcessGroupCollection.use_mpu_process_groups(required_pgs=['cp']) torch.distributed.all_reduce( mtp_loss, group=pg_collection.cp, op=torch.distributed.ReduceOp.AVG ) - MTPLossLoggingHelper.clean_loss_in_tracker() + MTPLossLoggingHelper.clean_metrics_in_tracker() assert torch.allclose(output_ref, output, rtol=1e-03, atol=1e-03) assert torch.allclose(mtp_loss, mtp_loss_ref, rtol=1e-02, atol=1e-02) @@ -615,10 +616,10 @@ def test_packed_sequences(self, tp, cp): # Verify MTP loss was computed tracker = MTPLossLoggingHelper.tracker - assert "values" in tracker - mtp_loss = tracker['values'].clone() + assert "loss_values" in tracker + mtp_loss = tracker['loss_values'].clone() assert mtp_loss.shape[0] == args.mtp_num_layers - MTPLossLoggingHelper.clean_loss_in_tracker() + MTPLossLoggingHelper.clean_metrics_in_tracker() # Backward pass loss = output.mean() @@ -876,39 +877,61 @@ def teardown_method(self, method): # Clean up the tracker after each test MTPLossLoggingHelper.tracker = {} - def test_save_loss_to_tracker(self): - """Test saving loss to tracker.""" - # Create a dummy loss tensor + def test_save_metrics_to_tracker(self): + """Test saving metrics to tracker.""" loss = torch.tensor(1.3) + correct = torch.tensor(5.0) + total = torch.tensor(10.0) layer_number = 2 num_layers = self.num_layers - # Test saving loss - MTPLossLoggingHelper.save_loss_to_tracker( - loss=loss, layer_number=layer_number, num_layers=num_layers + MTPLossLoggingHelper.save_metrics_to_tracker( + loss=loss, + correct=correct, + total=total, + layer_number=layer_number, + num_layers=num_layers, ) - # Verify tracker state - assert "values" in MTPLossLoggingHelper.tracker - assert MTPLossLoggingHelper.tracker["values"].shape == (num_layers,) - assert MTPLossLoggingHelper.tracker["values"][layer_number] == loss - assert MTPLossLoggingHelper.tracker["reduce_group"] is None - assert MTPLossLoggingHelper.tracker["avg_group"] is None + tracker = MTPLossLoggingHelper.tracker + assert "loss_values" in tracker + assert tracker["loss_values"].shape == (num_layers,) + assert tracker["loss_values"][layer_number] == loss + assert tracker["correct_values"][layer_number] == correct + assert tracker["total_values"][layer_number] == total + assert tracker["reduce_group"] is None + assert tracker["avg_group"] is None + + def test_mtp_logits_are_vocab_sharded(self): + """Test detection for vocab-sharded versus gathered MTP logits.""" + + class DummyOutputLayer: + def __init__(self, gather_output): + self.gather_output = gather_output + + assert _mtp_logits_are_vocab_sharded(DummyOutputLayer(gather_output=True), None) is False + assert _mtp_logits_are_vocab_sharded(DummyOutputLayer(gather_output=False), None) is True + assert _mtp_logits_are_vocab_sharded(DummyOutputLayer(gather_output=True), True) is False + assert _mtp_logits_are_vocab_sharded(DummyOutputLayer(gather_output=True), False) is True def test_track_mtp_metrics(self): - """Test tracking MTP metrics.""" - # First save some losses - loss = torch.tensor(2.3) + """Test tracking MTP metrics including acceptance rate.""" num_layers = self.num_layers + loss = torch.tensor(2.3) + correct = torch.tensor(7.0) + total = torch.tensor(10.0) + for i in range(num_layers): - MTPLossLoggingHelper.save_loss_to_tracker( - loss=loss, layer_number=i, num_layers=num_layers + MTPLossLoggingHelper.save_metrics_to_tracker( + loss=loss, correct=correct, total=total, layer_number=i, num_layers=num_layers ) - # Create dummy writer and loss dict class DummyWriter: + def __init__(self): + self.scalars = {} + def add_scalar(self, name, value, iteration): - pass + self.scalars[name] = value class DummyWandBWriter: def log(self, metrics, iteration): @@ -920,7 +943,6 @@ def log(self, metrics, iteration): wandb_writer = DummyWandBWriter() total_loss_dict = {} - # Test tracking metrics MTPLossLoggingHelper.track_mtp_metrics( loss_scale=loss_scale, iteration=iteration, @@ -929,16 +951,99 @@ def log(self, metrics, iteration): total_loss_dict=total_loss_dict, ) - # Verify total_loss_dict is populated + # Verify loss uses the legacy normalized MTP loss scaled by loss_scale. + expected_loss = loss * loss_scale + for i in range(num_layers): + assert f"mtp_{i+1} loss" in writer.scalars + assert torch.isclose(torch.as_tensor(writer.scalars[f"mtp_{i+1} loss"]), expected_loss) + assert torch.isclose(total_loss_dict[f"mtp_{i+1} loss"], expected_loss) + + # Verify acceptance rate is computed as (correct / total) * 100 + expected_rate = (correct / total) * 100.0 for i in range(num_layers): - assert f"mtp_{i + 1} loss" in total_loss_dict - assert total_loss_dict[f"mtp_{i + 1} loss"] == loss * loss_scale + assert f"mtp_{i+1}_acceptance_rate" in writer.scalars + assert torch.isclose( + torch.as_tensor(writer.scalars[f"mtp_{i+1}_acceptance_rate"]), expected_rate + ) + assert f"mtp_{i+1}_cumulative_acceptance_rate" in writer.scalars + assert torch.isclose( + torch.as_tensor(writer.scalars[f"mtp_{i+1}_cumulative_acceptance_rate"]), + expected_rate, + ) + + raw_counter_suffixes = ("_sum", "_tokens", "_correct", "_total") + assert not any(key.endswith(raw_counter_suffixes) for key in total_loss_dict) + + second_correct = torch.tensor(3.0) + second_total = torch.tensor(10.0) + for i in range(num_layers): + MTPLossLoggingHelper.save_metrics_to_tracker( + loss=loss, + correct=second_correct, + total=second_total, + layer_number=i, + num_layers=num_layers, + ) + + MTPLossLoggingHelper.track_mtp_metrics( + loss_scale=loss_scale, + iteration=iteration + 1, + writer=writer, + wandb_writer=wandb_writer, + total_loss_dict=total_loss_dict, + ) + + expected_second_rate = (second_correct / second_total) * 100.0 + expected_cumulative_rate = ((correct + second_correct) / (total + second_total)) * 100.0 + for i in range(num_layers): + assert torch.isclose( + torch.as_tensor(writer.scalars[f"mtp_{i+1}_acceptance_rate"]), expected_second_rate + ) + assert torch.isclose( + torch.as_tensor(writer.scalars[f"mtp_{i+1}_cumulative_acceptance_rate"]), + expected_cumulative_rate, + ) + assert torch.isclose(total_loss_dict[f"mtp_{i+1} loss"], expected_loss * 2) # Verify tracker is cleaned - assert torch.all(MTPLossLoggingHelper.tracker["values"] == 0) + assert torch.all(MTPLossLoggingHelper.tracker["loss_values"] == 0) assert MTPLossLoggingHelper.tracker["reduce_group"] is None assert MTPLossLoggingHelper.tracker["avg_group"] is None + 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) + second_loss = torch.tensor(2.0) + correct = torch.tensor(0.0) + total = torch.tensor(1.0) + loss_scale = torch.tensor(0.5) + layer_number = 0 + + MTPLossLoggingHelper.save_metrics_to_tracker( + loss=first_loss, correct=correct, total=total, layer_number=layer_number, num_layers=1 + ) + MTPLossLoggingHelper.save_metrics_to_tracker( + loss=second_loss, correct=correct, total=total, layer_number=layer_number, 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=loss_scale, iteration=1, writer=writer, total_loss_dict={} + ) + + logged_loss = torch.as_tensor(writer.scalars["mtp_1 loss"]) + expected_legacy_loss = (first_loss + second_loss) * loss_scale + token_weighted_loss = torch.tensor(40.0 / 12.0) + assert torch.isclose(logged_loss, expected_legacy_loss) + assert not torch.isclose(logged_loss, token_weighted_loss) + class TestMultiTokenPredictionHybrid: """Test Multi-Token Prediction with Mamba hybrid models.""" @@ -1098,9 +1203,9 @@ def test_forward_backward_mamba(self, tmp_path_dist_ckpt, tp, cp): ) tracker = MTPLossLoggingHelper.tracker mtp_loss_ref = None - assert "values" in tracker - mtp_loss_ref = tracker['values'].clone() - MTPLossLoggingHelper.clean_loss_in_tracker() + assert "loss_values" in tracker + mtp_loss_ref = tracker['loss_values'].clone() + MTPLossLoggingHelper.clean_metrics_in_tracker() iteration = 123 num_floating_point_operations_so_far = 456 @@ -1146,13 +1251,13 @@ def set_ckpt_path(ckpt_path): loss_mask=loss_mask, ) tracker = MTPLossLoggingHelper.tracker - assert "values" in tracker - mtp_loss = tracker['values'].clone() + assert "loss_values" in tracker + mtp_loss = tracker['loss_values'].clone() pg_collection = ProcessGroupCollection.use_mpu_process_groups(required_pgs=['cp']) torch.distributed.all_reduce( mtp_loss, group=pg_collection.cp, op=torch.distributed.ReduceOp.AVG ) - MTPLossLoggingHelper.clean_loss_in_tracker() + MTPLossLoggingHelper.clean_metrics_in_tracker() assert torch.allclose(output_ref, output, rtol=1e-03, atol=1e-03) assert torch.allclose(mtp_loss, mtp_loss_ref, rtol=1e-02, atol=1e-02)