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
Original file line number Diff line number Diff line change
Expand Up @@ -180,7 +180,9 @@ def compute_language_model_loss(self, labels: Tensor, logits: Tensor) -> Tensor:
elif self.config.cross_entropy_fusion_impl == 'native':
loss = fused_vocab_parallel_cross_entropy(logits, labels, self.pg_collection.tp)
else:
loss = tensor_parallel.vocab_parallel_cross_entropy(logits, labels)
loss = tensor_parallel.vocab_parallel_cross_entropy(
logits, labels, tp_group=self.tp_group
)

# [s b] => [b, s]
loss = loss.transpose(0, 1).contiguous()
Expand Down
43 changes: 23 additions & 20 deletions megatron/core/tensor_parallel/cross_entropy.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,11 +4,8 @@

import torch

from megatron.core.parallel_state import (
get_tensor_model_parallel_group,
get_tensor_model_parallel_rank,
get_tensor_model_parallel_world_size,
)
from megatron.core.parallel_state import get_tensor_model_parallel_group
from megatron.core.utils import get_pg_rank, get_pg_size

from .utils import VocabUtility

Expand Down Expand Up @@ -121,21 +118,22 @@ def calculate_gradients(

class _VocabParallelCrossEntropy(torch.autograd.Function):
@staticmethod
def forward(ctx, vocab_parallel_logits, target, label_smoothing=0.0):
def forward(ctx, vocab_parallel_logits, target, label_smoothing=0.0, tp_group=None):
"""Vocab parallel cross entropy forward function."""

if tp_group is None:
tp_group = get_tensor_model_parallel_group()

vocab_parallel_logits, logits_max = VocabParallelCrossEntropy.calculate_logits_max(
vocab_parallel_logits
)
torch.distributed.all_reduce(
logits_max, op=torch.distributed.ReduceOp.MAX, group=get_tensor_model_parallel_group()
)
torch.distributed.all_reduce(logits_max, op=torch.distributed.ReduceOp.MAX, group=tp_group)

# Get the partition's vocab indices
get_vocab_range = VocabUtility.vocab_range_from_per_partition_vocab_size
partition_vocab_size = vocab_parallel_logits.size()[-1]
rank = get_tensor_model_parallel_rank()
world_size = get_tensor_model_parallel_world_size()
rank = get_pg_rank(tp_group)
world_size = get_pg_size(tp_group)
vocab_start_index, vocab_end_index = get_vocab_range(partition_vocab_size, rank, world_size)

(target_mask, masked_target_1d, predicted_logits, sum_exp_logits, exp_logits) = (
Expand All @@ -146,15 +144,11 @@ def forward(ctx, vocab_parallel_logits, target, label_smoothing=0.0):

# All reduce is needed to get the chunks from other GPUs.
torch.distributed.all_reduce(
predicted_logits,
op=torch.distributed.ReduceOp.SUM,
group=get_tensor_model_parallel_group(),
predicted_logits, op=torch.distributed.ReduceOp.SUM, group=tp_group
)

torch.distributed.all_reduce(
sum_exp_logits,
op=torch.distributed.ReduceOp.SUM,
group=get_tensor_model_parallel_group(),
sum_exp_logits, op=torch.distributed.ReduceOp.SUM, group=tp_group
)

exp_logits, loss = VocabParallelCrossEntropy.calculate_cross_entropy_loss(
Expand Down Expand Up @@ -213,10 +207,15 @@ def backward(ctx, grad_output):
grad_2d, arange_1d, masked_target_1d, softmax_update, grad_input, grad_output
)

return grad_input, None, None
return grad_input, None, None, None


def vocab_parallel_cross_entropy(vocab_parallel_logits, target, label_smoothing=0.0):
def vocab_parallel_cross_entropy(
vocab_parallel_logits: torch.Tensor,
target: torch.Tensor,
label_smoothing: float = 0.0,
tp_group: torch.distributed.ProcessGroup | None = None,
) -> torch.Tensor:
"""
Performs cross entropy loss when logits are split across tensor parallel ranks

Expand All @@ -228,5 +227,9 @@ def vocab_parallel_cross_entropy(vocab_parallel_logits, target, label_smoothing=

label_smoothing: smoothing factor, must be in range [0.0, 1.0)
default is no smoothing (=0.0)

tp_group: the tensor parallel group over which to all reduce
"""
return _VocabParallelCrossEntropy.apply(vocab_parallel_logits, target, label_smoothing)
return _VocabParallelCrossEntropy.apply(
vocab_parallel_logits, target, label_smoothing, tp_group
)
76 changes: 76 additions & 0 deletions tests/unit_tests/tensor_parallel/test_cross_entropy.py
Original file line number Diff line number Diff line change
@@ -1,10 +1,86 @@
# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.

from types import SimpleNamespace

import numpy as np
import torch

from megatron.core.models.common.language_module import language_module as language_module_module
from megatron.core.tensor_parallel import cross_entropy as cross_entropy_module
from megatron.core.tensor_parallel.cross_entropy import vocab_parallel_cross_entropy
from tests.unit_tests.test_utilities import Utils


class _FakeTPGroup:
def rank(self):
return 0

def size(self):
return 1


def test_vocab_parallel_cross_entropy_uses_explicit_tp_group(monkeypatch):
tp_group = _FakeTPGroup()
all_reduce_groups = []

def fake_all_reduce(tensor, op=None, group=None):
all_reduce_groups.append(group)
return tensor

def fail_parallel_state_call(*args, **kwargs):
raise AssertionError("explicit tp_group should avoid parallel_state")

monkeypatch.setattr(torch.distributed, "all_reduce", fake_all_reduce)
monkeypatch.setattr(
cross_entropy_module, "get_tensor_model_parallel_group", fail_parallel_state_call
)

vocab_parallel_logits = torch.tensor([[1.0, 2.0, 3.0], [0.5, -0.5, 1.0]])
target = torch.tensor([2, 0])
expected_output = torch.nn.functional.cross_entropy(
vocab_parallel_logits.clone(), target, reduction="none"
)

output = vocab_parallel_cross_entropy(vocab_parallel_logits, target, tp_group=tp_group)

torch.testing.assert_close(output, expected_output)
assert all_reduce_groups == [tp_group, tp_group, tp_group]


def test_language_module_unfused_loss_passes_tp_group(monkeypatch):
tp_group = _FakeTPGroup()
captured = {}

def fake_vocab_parallel_cross_entropy(logits, labels, label_smoothing=0.0, tp_group=None):
captured["logits"] = logits
captured["labels"] = labels
captured["label_smoothing"] = label_smoothing
captured["tp_group"] = tp_group
return torch.zeros_like(labels, dtype=logits.dtype)

monkeypatch.setattr(
language_module_module.tensor_parallel,
"vocab_parallel_cross_entropy",
fake_vocab_parallel_cross_entropy,
)

module = SimpleNamespace(
config=SimpleNamespace(cross_entropy_loss_fusion=False), tp_group=tp_group
)
labels = torch.tensor([[0, 1, 2], [2, 1, 0]])
logits = torch.randn(3, 2, 4)

loss = language_module_module.LanguageModule.compute_language_model_loss(
module, labels=labels, logits=logits
)

assert captured["logits"] is logits
assert captured["tp_group"] is tp_group
assert captured["label_smoothing"] == 0.0
torch.testing.assert_close(captured["labels"], labels.transpose(0, 1).contiguous())
assert loss.shape == labels.shape


def test_vocab_parallel_cross_entropy():
Utils.initialize_model_parallel(4, 2)
vocab_parallel_logits = torch.range(0, 7).repeat(16, 4).cuda()
Expand Down
Loading