@@ -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