Skip to content

Commit 616a34b

Browse files
committed
fix(model): pass pp_rank to callable transformer layer specs
GPTModelProvider.provide() forwarded only vp_stage to a callable transformer_layer_spec, so Megatron-Core block-spec builders fell back to parallel_state.get_pipeline_model_parallel_rank(). Under use_decentralized_pg=True the MPU globals are never initialized, so building any block-spec model asserted with "pipeline_model parallel group is not initialized" before the first forward. mtp_block_spec() had the same omission. Adds unit tests covering the change (red-green verified). Detected by: megatron-bridge QA (test_eval_cp_gdn_metadata_e2e) Signed-off-by: Pruthviraj Prakash <pruprakash@nvidia.com>
1 parent 5c4092f commit 616a34b

2 files changed

Lines changed: 114 additions & 9 deletions

File tree

src/megatron/bridge/models/gpt_provider.py

Lines changed: 22 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -36,6 +36,7 @@
3636
from megatron.core.transformer import ModuleSpec
3737
from megatron.core.transformer.dot_product_attention import DotProductAttention as MCoreDotProductAttention
3838
from megatron.core.transformer.enums import AttnBackend
39+
from megatron.core.utils import get_pg_rank
3940

4041
from megatron.bridge.models.logit_dtype import logit_dtype_kwarg
4142
from megatron.bridge.models.model_provider import ModelProviderMixin
@@ -256,11 +257,9 @@ def provide(self, pre_process=None, post_process=None, vp_stage=None) -> MCoreGP
256257

257258
transformer_layer_spec = self.transformer_layer_spec
258259
if not isinstance(transformer_layer_spec, ModuleSpec):
259-
# Check if the transformer_layer_spec function accepts vp_stage parameter
260-
if "vp_stage" in inspect.signature(transformer_layer_spec).parameters:
261-
transformer_layer_spec = transformer_layer_spec(self, vp_stage=vp_stage)
262-
else:
263-
transformer_layer_spec = transformer_layer_spec(self)
260+
transformer_layer_spec = transformer_layer_spec(
261+
self, **_callable_spec_kwargs(transformer_layer_spec, self, vp_stage)
262+
)
264263

265264
assert self.vocab_size is not None, "vocab_size must be configured before calling provide()"
266265
if self.should_pad_vocab:
@@ -342,6 +341,21 @@ def provide(self, pre_process=None, post_process=None, vp_stage=None) -> MCoreGP
342341
return model
343342

344343

344+
def _callable_spec_kwargs(spec_fn: Callable, config: "GPTModelProvider", vp_stage: Optional[int]) -> dict[str, Any]:
345+
"""Build the keyword arguments a callable transformer layer spec declares."""
346+
params = inspect.signature(spec_fn).parameters
347+
kwargs: dict[str, Any] = {}
348+
if "vp_stage" in params:
349+
kwargs["vp_stage"] = vp_stage
350+
if "pp_rank" in params:
351+
# Resolve the pipeline rank from the provider's own groups so that block spec
352+
# builders do not fall back to the MPU globals, which decentralized runs never set.
353+
pg_collection = getattr(config, "_pg_collection", None)
354+
if pg_collection is not None:
355+
kwargs["pp_rank"] = get_pg_rank(pg_collection.pp)
356+
return kwargs
357+
358+
345359
def mtp_block_spec(config: "GPTModelProvider", vp_stage: Optional[int] = None) -> Optional[ModuleSpec]:
346360
"""Pass in the MTP block spec if model has MTP layers.
347361
@@ -355,10 +369,9 @@ def mtp_block_spec(config: "GPTModelProvider", vp_stage: Optional[int] = None) -
355369
from megatron.core.models.gpt.gpt_layer_specs import get_gpt_mtp_block_spec
356370

357371
if isinstance(config.transformer_layer_spec, Callable):
358-
if "vp_stage" in inspect.signature(config.transformer_layer_spec).parameters:
359-
spec = config.transformer_layer_spec(config, vp_stage=vp_stage)
360-
else:
361-
spec = config.transformer_layer_spec(config)
372+
spec = config.transformer_layer_spec(
373+
config, **_callable_spec_kwargs(config.transformer_layer_spec, config, vp_stage)
374+
)
362375
else:
363376
spec = config.transformer_layer_spec
364377
if hasattr(spec, "layer_specs") and len(spec.layer_specs) == 0:

tests/unit_tests/models/test_gpt_provider.py

Lines changed: 92 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)