diff --git a/megatron/core/tensor_parallel/layers.py b/megatron/core/tensor_parallel/layers.py index c072c52bd05..3318972baec 100644 --- a/megatron/core/tensor_parallel/layers.py +++ b/megatron/core/tensor_parallel/layers.py @@ -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 diff --git a/tests/unit_tests/tensor_parallel/test_layers.py b/tests/unit_tests/tensor_parallel/test_layers.py index dbc27f502c6..43fda3a7a41 100644 --- a/tests/unit_tests/tensor_parallel/test_layers.py +++ b/tests/unit_tests/tensor_parallel/test_layers.py @@ -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)