diff --git a/megatron/core/model_parallel_config.py b/megatron/core/model_parallel_config.py index dabe0d0aced..a593457a7fe 100644 --- a/megatron/core/model_parallel_config.py +++ b/megatron/core/model_parallel_config.py @@ -121,6 +121,12 @@ class ModelParallelConfig: None, no function is called on the loss. """ + mtp_grad_scale_func: Optional[Callable] = None + """If using loss scaling for MTP (Multi-Token Prediction), this function should return the + scalar or size-1 scale value for MTP loss. The value is converted to the output tensor + device. If None, falls back to grad_scale_func with torch.ones(1). + """ + no_sync_func: Optional[Callable] = None """Function that creates a context that suppresses asynchronous data-parallel communication. If the model is an instance of core.distributed.DistributedDataParallel, the default is to use diff --git a/megatron/core/pipeline_parallel/schedules.py b/megatron/core/pipeline_parallel/schedules.py index f07a1915ed7..d85b267ed7e 100644 --- a/megatron/core/pipeline_parallel/schedules.py +++ b/megatron/core/pipeline_parallel/schedules.py @@ -226,6 +226,28 @@ def get_tensor_device(tensor: Union[torch.Tensor, Dict[str, torch.Tensor]]): return tensor.device +def _get_mtp_loss_scale(config, device: torch.device) -> torch.Tensor: + """Get the MTP loss scale on the output tensor device.""" + + def _normalize_loss_scale(loss_scale, scale_func_name: str) -> torch.Tensor: + loss_scale = torch.as_tensor(loss_scale, device=device) + if loss_scale.numel() != 1: + raise ValueError( + f"{scale_func_name} must return a scalar or size-1 tensor for MTP loss scaling, " + f"but returned a tensor with {loss_scale.numel()} elements." + ) + return loss_scale + + mtp_grad_scale_func = getattr(config, 'mtp_grad_scale_func', None) + if mtp_grad_scale_func is not None: + return _normalize_loss_scale(mtp_grad_scale_func(), "mtp_grad_scale_func") + if config.grad_scale_func is not None: + return _normalize_loss_scale( + config.grad_scale_func(torch.ones(1, device=device)), "grad_scale_func" + ) + return torch.ones(1, device=device) + + def forward_step_calc_loss( model, output_tensor, @@ -306,13 +328,10 @@ def forward_step_calc_loss( # Set the loss scale for Multi-Token Prediction (MTP) loss. if hasattr(config, 'mtp_num_layers') and config.mtp_num_layers is not None: - # Calculate the loss scale based on the grad_scale_func if available, else default to 1. + # Calculate the loss scale based on mtp_grad_scale_func if available, + # else fall back to grad_scale_func, else default to 1. device = get_tensor_device(output_tensor) - loss_scale = ( - config.grad_scale_func(torch.ones(1, device=device)) - if config.grad_scale_func is not None - else torch.ones(1, device=device) - ) + loss_scale = _get_mtp_loss_scale(config, device) # Set the loss scale if config.calculate_per_token_loss: MTPLossAutoScaler.set_loss_scale(loss_scale) diff --git a/megatron/training/arguments.py b/megatron/training/arguments.py index 18cb5ef8a15..dcafa9e8139 100644 --- a/megatron/training/arguments.py +++ b/megatron/training/arguments.py @@ -2007,6 +2007,7 @@ def _add_network_size_args(parser): "timers", "finalize_model_grads_func", "grad_scale_func", + "mtp_grad_scale_func", "no_sync_func", "grad_sync_func", "param_sync_func", diff --git a/tests/unit_tests/models/test_hybrid_moe_model.py b/tests/unit_tests/models/test_hybrid_moe_model.py index 5aef44eccde..00cb57e3164 100644 --- a/tests/unit_tests/models/test_hybrid_moe_model.py +++ b/tests/unit_tests/models/test_hybrid_moe_model.py @@ -123,6 +123,7 @@ "gated_linear_unit": False, "glu_linear_offset": 0.0, "grad_scale_func": None, + "mtp_grad_scale_func": None, "grad_sync_func": None, "gradient_accumulation_fusion": True, "hetereogenous_dist_checkpoint": False, diff --git a/tests/unit_tests/pipeline_parallel/test_mtp_loss_scale.py b/tests/unit_tests/pipeline_parallel/test_mtp_loss_scale.py new file mode 100644 index 00000000000..81b8682f2e6 --- /dev/null +++ b/tests/unit_tests/pipeline_parallel/test_mtp_loss_scale.py @@ -0,0 +1,75 @@ +# Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + +import pytest +import torch + +from megatron.core import ModelParallelConfig +from megatron.core.pipeline_parallel.schedules import _get_mtp_loss_scale + + +def test_mtp_grad_scale_func_config(): + """Test that mtp_grad_scale_func config defaults to None and can be set.""" + config = ModelParallelConfig() + assert config.mtp_grad_scale_func is None + + scale_fn = lambda: torch.tensor(0.5) + config = ModelParallelConfig(mtp_grad_scale_func=scale_fn) + assert config.mtp_grad_scale_func is scale_fn + assert config.mtp_grad_scale_func().item() == 0.5 + + +def test_mtp_loss_scale_selection(): + """Test MTP loss scale selection and device normalization.""" + + device = ( + torch.device('cuda', torch.cuda.current_device()) + if torch.cuda.is_available() + else torch.device('cpu') + ) + + # Case 1: mtp_grad_scale_func takes priority + config = ModelParallelConfig( + mtp_grad_scale_func=lambda: torch.tensor([0.25], device='cpu'), + grad_scale_func=lambda x: x * 2.0, + ) + loss_scale = _get_mtp_loss_scale(config, device) + assert loss_scale.item() == 0.25 + assert loss_scale.device == device + + # Case 2: Falls back to grad_scale_func + config = ModelParallelConfig(grad_scale_func=lambda x: x * 3.0) + loss_scale = _get_mtp_loss_scale(config, device) + assert loss_scale.item() == 3.0 + assert loss_scale.device == device + + # Case 3: Falls back to ones + config = ModelParallelConfig() + loss_scale = _get_mtp_loss_scale(config, device) + assert loss_scale.item() == 1.0 + assert loss_scale.device == device + + +@pytest.mark.parametrize( + "config,scale_func_name", + [ + ( + ModelParallelConfig(mtp_grad_scale_func=lambda: torch.tensor([0.25, 0.5])), + "mtp_grad_scale_func", + ), + ( + ModelParallelConfig(grad_scale_func=lambda _: torch.tensor([0.25, 0.5])), + "grad_scale_func", + ), + ], +) +def test_mtp_loss_scale_rejects_non_scalar_scale(config, scale_func_name): + """Test that MTP loss scaling rejects per-token or per-sample scale values.""" + + device = ( + torch.device('cuda', torch.cuda.current_device()) + if torch.cuda.is_available() + else torch.device('cpu') + ) + + with pytest.raises(ValueError, match=f"{scale_func_name} must return a scalar"): + _get_mtp_loss_scale(config, device)