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
12 changes: 9 additions & 3 deletions megatron/core/tensor_parallel/layers.py
Original file line number Diff line number Diff line change
Expand Up @@ -89,11 +89,17 @@
dist_reduce_scatter_func = torch.distributed._reduce_scatter_base


def param_is_not_tensor_parallel_duplicate(param, tp_group=None):
"""Returns true if the passed-in parameter is not a duplicate parameter
on another TP rank."""
def param_is_not_tensor_parallel_duplicate(param, tp_group=None, expert_tp_group=None):
"""Return whether a parameter contributes to a unique model-parallel shard.

Parameters reduced over expert data parallel groups use the expert tensor-parallel
group for duplicate filtering. Other parameters use the regular tensor-parallel group.
"""
if hasattr(param, "tensor_model_parallel") and param.tensor_model_parallel:
return True
# allreduce=False marks parameters reduced over expert DP, so filter their duplicates over ETP.
if not getattr(param, "allreduce", True) and expert_tp_group is not None:
tp_group = expert_tp_group
# Prefer provided tp_group when available (new explicit path).
if tp_group is not None:
return tp_group.rank() == 0
Expand Down
33 changes: 32 additions & 1 deletion tests/unit_tests/tensor_parallel/test_layers.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,11 +2,42 @@
import pytest
import torch

from megatron.core.tensor_parallel.layers import linear_with_frozen_weight
from megatron.core.tensor_parallel.layers import (
linear_with_frozen_weight,
param_is_not_tensor_parallel_duplicate,
)
from megatron.core.tensor_parallel.mappings import gather_from_tensor_model_parallel_region
from tests.unit_tests.test_utilities import Utils


class _RankGroup:
"""Process-group stub that reports a fixed local rank."""

def __init__(self, rank):
self._rank = rank

def rank(self):
return self._rank


@pytest.mark.parametrize(
("allreduce", "regular_tp_rank", "expert_tp_rank", "expected"),
[(True, 0, 1, True), (True, 1, 0, False), (False, 0, 1, False), (False, 1, 0, True)],
)
def test_param_is_not_tensor_parallel_duplicate_uses_parameter_parallel_group(
allreduce, regular_tp_rank, expert_tp_rank, expected
):
"""Use expert TP only for parameters reduced over expert data parallel groups."""
param = torch.nn.Parameter(torch.ones(1))
param.allreduce = allreduce

actual = param_is_not_tensor_parallel_duplicate(
param, tp_group=_RankGroup(regular_tp_rank), expert_tp_group=_RankGroup(expert_tp_rank)
)

assert actual is expected


@pytest.mark.parametrize("tensor_parallel,allreduce_dgrad", [(1, False), (8, True)])
def test_LinearWithFrozenWeight(tensor_parallel, allreduce_dgrad):
Utils.initialize_model_parallel(tensor_parallel, 1)
Expand Down
Loading