Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
6 changes: 6 additions & 0 deletions megatron/core/model_parallel_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
31 changes: 25 additions & 6 deletions megatron/core/pipeline_parallel/schedules.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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)
Expand Down
1 change: 1 addition & 0 deletions megatron/training/arguments.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
1 change: 1 addition & 0 deletions tests/unit_tests/models/test_hybrid_moe_model.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down
75 changes: 75 additions & 0 deletions tests/unit_tests/pipeline_parallel/test_mtp_loss_scale.py
Original file line number Diff line number Diff line change
@@ -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)
Loading