@@ -635,3 +635,95 @@ def fake_spec_unsupported(
635635
636636 assert result == "te_spec_unsupported"
637637 assert "use_grouped_gemm_for_dense_mlp" not in captured
638+
639+
640+ class TestCallableSpecPpRank :
641+ """Tests for forwarding pp_rank to callable transformer layer specs."""
642+
643+ @staticmethod
644+ def _provider (** kwargs ) -> GPTModelProvider :
645+ return GPTModelProvider (
646+ num_layers = 2 ,
647+ hidden_size = 128 ,
648+ num_attention_heads = 4 ,
649+ vocab_size = 1000 ,
650+ ** kwargs ,
651+ )
652+
653+ def test_provide_forwards_pp_rank_resolved_from_the_pg_collection (self ):
654+ """A spec callable that accepts pp_rank receives the rank of _pg_collection.pp."""
655+ captured : dict = {}
656+
657+ def spec_fn (config , vp_stage = None , pp_rank = None ):
658+ captured ["vp_stage" ] = vp_stage
659+ captured ["pp_rank" ] = pp_rank
660+ return "block_spec"
661+
662+ pp_group = object ()
663+ provider = self ._provider (transformer_layer_spec = spec_fn )
664+ provider ._pg_collection = type ("PG" , (), {"pp" : pp_group , "tp" : object (), "cp" : object ()})()
665+
666+ with patch ("megatron.bridge.models.gpt_provider.get_pg_rank" , return_value = 3 ) as mock_rank :
667+ with patch ("megatron.bridge.models.gpt_provider.MCoreGPTModel" ):
668+ provider .provide (pre_process = True , post_process = True , vp_stage = 1 )
669+
670+ mock_rank .assert_called_once_with (pp_group )
671+ assert captured ["pp_rank" ] == 3 , "pp_rank must come from the provider's pipeline group"
672+ assert captured ["vp_stage" ] == 1
673+
674+ def test_provide_omits_pp_rank_for_specs_that_do_not_accept_it (self ):
675+ """Regression: a spec callable without a pp_rank parameter is called unchanged."""
676+ captured : dict = {}
677+
678+ def spec_fn (config , vp_stage = None ):
679+ captured ["vp_stage" ] = vp_stage
680+ return "block_spec"
681+
682+ provider = self ._provider (transformer_layer_spec = spec_fn )
683+ provider ._pg_collection = type ("PG" , (), {"pp" : object (), "tp" : object (), "cp" : object ()})()
684+
685+ with patch ("megatron.bridge.models.gpt_provider.MCoreGPTModel" ) as mock_model :
686+ provider .provide (pre_process = True , post_process = True , vp_stage = 2 )
687+
688+ assert captured ["vp_stage" ] == 2
689+ assert mock_model .call_args .kwargs ["transformer_layer_spec" ] == "block_spec"
690+
691+ def test_provide_omits_pp_rank_when_no_pg_collection_is_set (self ):
692+ """Without a process group collection the spec keeps its own pp_rank default."""
693+ captured : dict = {}
694+
695+ def spec_fn (config , vp_stage = None , pp_rank = "unset" ):
696+ captured ["pp_rank" ] = pp_rank
697+ return "block_spec"
698+
699+ provider = self ._provider (transformer_layer_spec = spec_fn )
700+
701+ with patch ("megatron.bridge.models.gpt_provider.MCoreGPTModel" ):
702+ provider .provide (pre_process = True , post_process = True )
703+
704+ assert captured ["pp_rank" ] == "unset"
705+
706+ @patch ("megatron.core.models.gpt.gpt_layer_specs.get_gpt_mtp_block_spec" )
707+ def test_mtp_block_spec_forwards_pp_rank_to_callable_spec (self , mock_get_mtp ):
708+ """mtp_block_spec resolves pp_rank the same way when re-invoking the spec callable."""
709+ from megatron .bridge .models .gpt_provider import mtp_block_spec
710+
711+ captured : dict = {}
712+ block_spec = Mock ()
713+ block_spec .layer_specs = ["layer_a" ]
714+
715+ def spec_fn (config , vp_stage = None , pp_rank = None ):
716+ captured ["pp_rank" ] = pp_rank
717+ return block_spec
718+
719+ pp_group = object ()
720+ provider = self ._provider (mtp_num_layers = 1 , transformer_layer_spec = spec_fn )
721+ provider ._pg_collection = type ("PG" , (), {"pp" : pp_group , "tp" : object (), "cp" : object ()})()
722+ mock_get_mtp .return_value = "mtp_spec"
723+
724+ with patch ("megatron.bridge.models.gpt_provider.get_pg_rank" , return_value = 5 ) as mock_rank :
725+ result = mtp_block_spec (provider , vp_stage = 0 )
726+
727+ mock_rank .assert_called_once_with (pp_group )
728+ assert captured ["pp_rank" ] == 5
729+ assert result == "mtp_spec"
0 commit comments