Skip to content

Commit c57cd8f

Browse files
committed
test(training): cover expert parameter norm topology
Signed-off-by: yaoyu-33 <yaoyu.094@gmail.com>
1 parent dd114f4 commit c57cd8f

1 file changed

Lines changed: 49 additions & 0 deletions

File tree

tests/unit_tests/training/utils/test_train_utils.py

Lines changed: 49 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -2870,6 +2870,55 @@ def test_duplicate_filter_receives_tp_and_expert_tp_groups(
28702870
expert_tp_group=_patch_pg_collection.expt_tp,
28712871
)
28722872

2873+
def test_moe_param_norm_counts_logical_parameter_when_tp_ranks_differ(
2874+
self,
2875+
monkeypatch,
2876+
mock_model_config_fp32,
2877+
):
2878+
"""Expert parameters use the expert-TP rank when it differs from regular TP."""
2879+
2880+
class _RankGroup:
2881+
def __init__(self, rank):
2882+
self._rank = rank
2883+
2884+
def rank(self):
2885+
return self._rank
2886+
2887+
class _MixedModel(torch.nn.Module):
2888+
def __init__(self):
2889+
super().__init__()
2890+
self.dense = torch.nn.Parameter(torch.ones(4, device="cuda"))
2891+
self.expert = torch.nn.Parameter(torch.ones(4, device="cuda"))
2892+
self.expert.allreduce = False
2893+
2894+
regular_tp_group = _RankGroup(rank=1)
2895+
expert_tp_group = _RankGroup(rank=0)
2896+
reduce_group = _RankGroup(rank=0)
2897+
pg_collection = SimpleNamespace(
2898+
tp=regular_tp_group,
2899+
expt_tp=expert_tp_group,
2900+
dp_cp=reduce_group,
2901+
expt_dp=reduce_group,
2902+
mp=reduce_group,
2903+
tp_ep_pp=reduce_group,
2904+
)
2905+
monkeypatch.setattr(
2906+
"megatron.bridge.training.utils.train_utils.get_pg_collection",
2907+
lambda model: pg_collection,
2908+
)
2909+
2910+
with (
2911+
mock.patch(
2912+
"megatron.core.tensor_parallel.layers.get_tensor_model_parallel_rank",
2913+
return_value=regular_tp_group.rank(),
2914+
),
2915+
mock.patch("torch.distributed.get_process_group_ranks", return_value=[0]),
2916+
mock.patch("torch.distributed.all_reduce"),
2917+
):
2918+
actual_norm = calc_params_l2_norm(_MixedModel(), mock_model_config_fp32)
2919+
2920+
assert actual_norm == pytest.approx(2.0)
2921+
28732922
@mock.patch("megatron.bridge.training.utils.train_utils.calc_dtensor_params_l2_norm")
28742923
def test_megatron_fsdp_path(self, mock_calc_dtensor_norm, mock_model_config_fp32):
28752924
"""Test calc_params_l2_norm with use_megatron_fsdp=True."""

0 commit comments

Comments
 (0)