Skip to content

Commit b9f37fc

Browse files
Merge commit 'c931bed154d09469a745bc89b8f57931f35d3444' into chore/bump-mcore-main-260904
2 parents 8517220 + c931bed commit b9f37fc

1 file changed

Lines changed: 20 additions & 4 deletions

File tree

tests/unit_tests/training/utils/test_train_utils.py

Lines changed: 20 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1060,12 +1060,22 @@ def test_memory_reporting_checkpoint_resume(
10601060
mock_report_memory.assert_called_once()
10611061

10621062
@pytest.mark.parametrize(
1063-
("model_num_layers", "mtp_num_layers", "hybrid_pattern", "moe_layer_freq", "expected_num_moe_layers"),
1063+
(
1064+
"model_num_layers",
1065+
"mtp_num_layers",
1066+
"hybrid_pattern",
1067+
"moe_layer_freq",
1068+
"expected_num_moe_layers",
1069+
"supports_num_moe_layers",
1070+
),
10641071
[
1065-
pytest.param(12, None, None, 2, 6, id="non_hybrid"),
1066-
pytest.param(4, 2, "MMME/*E/*E", None, 3, id="hybrid"),
1072+
pytest.param(12, None, None, 2, 6, True, id="non_hybrid-supported"),
1073+
pytest.param(12, None, None, 2, 6, False, id="non_hybrid-unsupported"),
1074+
pytest.param(4, 2, "MMME/*E/*E", None, 3, True, id="hybrid-supported"),
1075+
pytest.param(4, 2, "MMME/*E/*E", None, 3, False, id="hybrid-unsupported"),
10671076
],
10681077
)
1078+
@mock.patch("megatron.bridge.training.utils.train_utils._track_moe_metrics_supports_num_moe_layers")
10691079
@mock.patch("megatron.bridge.training.utils.train_utils.get_num_microbatches")
10701080
@mock.patch("megatron.bridge.training.utils.train_utils.reduce_max_stat_across_model_parallel_group")
10711081
@mock.patch("megatron.bridge.training.utils.train_utils.get_world_size_safe")
@@ -1084,6 +1094,7 @@ def test_moe_logging(
10841094
mock_get_world_size,
10851095
mock_reduce_lr,
10861096
mock_get_microbatches,
1097+
mock_supports_num_moe_layers,
10871098
mock_config,
10881099
mock_global_state,
10891100
loss_dict,
@@ -1092,12 +1103,14 @@ def test_moe_logging(
10921103
hybrid_pattern,
10931104
moe_layer_freq,
10941105
expected_num_moe_layers,
1106+
supports_num_moe_layers,
10951107
):
10961108
"""Test MoE (Mixture of Experts) logging when enabled."""
10971109
# Get fresh total_loss_dict for this test
10981110
total_loss_dict = self.get_fresh_total_loss_dict()
10991111

11001112
# Setup mocks
1113+
mock_supports_num_moe_layers.return_value = supports_num_moe_layers
11011114
mock_report_l2_norm_grad.return_value = {}
11021115
mock_report_throughput.return_value = {}
11031116
mock_report_runtime.return_value = {}
@@ -1139,7 +1152,10 @@ def test_moe_logging(
11391152
assert "load_balancing_loss" in call_args.kwargs["track_names"]
11401153
assert "z_loss" in call_args.kwargs["track_names"]
11411154
assert call_args.kwargs["num_layers"] == model_num_layers
1142-
assert call_args.kwargs["num_moe_layers"] == expected_num_moe_layers
1155+
if supports_num_moe_layers:
1156+
assert call_args.kwargs["num_moe_layers"] == expected_num_moe_layers
1157+
else:
1158+
assert "num_moe_layers" not in call_args.kwargs
11431159
assert call_args.kwargs["mtp_num_layers"] == mtp_num_layers
11441160

11451161
@mock.patch("megatron.bridge.training.utils.train_utils.get_num_microbatches")

0 commit comments

Comments
 (0)