|
13 | 13 | # limitations under the License. |
14 | 14 |
|
15 | 15 |
|
| 16 | +import json |
16 | 17 | from types import SimpleNamespace |
17 | 18 | from unittest.mock import MagicMock, Mock, patch |
18 | 19 |
|
|
25 | 26 | from megatron.bridge.models.gpt_provider import GPTModelProvider |
26 | 27 | from megatron.bridge.models.hybrid.hybrid_builder import HybridModelConfig |
27 | 28 | from megatron.bridge.models.hybrid.hybrid_provider import HybridModelProvider |
| 29 | +from megatron.bridge.models.llama_nemotron.llama_nemotron_provider import LlamaNemotronHeterogeneousProvider |
28 | 30 | from megatron.bridge.models.transformer_config import TransformerConfig |
29 | 31 | from megatron.bridge.training.callbacks import CallbackManager |
30 | 32 | from megatron.bridge.training.checkpointing import load_checkpoint |
@@ -556,9 +558,42 @@ def test_build_with_provider(self): |
556 | 558 |
|
557 | 559 | result = _build_distributed_model(cfg, pg_collection=MagicMock()) |
558 | 560 |
|
| 561 | + mock_provider.finalize.assert_called_once_with() |
559 | 562 | mock_provider.provide_distributed_model.assert_called_once() |
560 | 563 | assert result == mock_dist_model |
561 | 564 |
|
| 565 | + def test_finalizes_heterogeneous_provider_before_entering_distributed_model_build(self): |
| 566 | + """Finalize deferred heterogeneous fields before entering provider model construction.""" |
| 567 | + block = { |
| 568 | + "attention": {"no_op": False, "replace_with_linear": False, "num_query_groups": 4}, |
| 569 | + "mlp": {"no_op": False, "replace_with_linear": False, "ffn_hidden_size": 256}, |
| 570 | + } |
| 571 | + provider = LlamaNemotronHeterogeneousProvider( |
| 572 | + num_layers=1, |
| 573 | + hidden_size=64, |
| 574 | + num_attention_heads=4, |
| 575 | + heterogeneous_layers_config_encoded_json=json.dumps({"block_configs": [block]}), |
| 576 | + ) |
| 577 | + expected_model = [MagicMock()] |
| 578 | + |
| 579 | + def provide_distributed_model(**_kwargs): |
| 580 | + assert len(provider.per_block_parameters) == 1 |
| 581 | + return expected_model |
| 582 | + |
| 583 | + provider.provide_distributed_model = Mock(side_effect=provide_distributed_model) |
| 584 | + cfg = SimpleNamespace( |
| 585 | + model=provider, |
| 586 | + ddp=MagicMock(), |
| 587 | + optimizer=SimpleNamespace(overlap_param_gather_with_optimizer_step=False), |
| 588 | + dist=SimpleNamespace(use_megatron_fsdp=False, use_torch_fsdp2=False), |
| 589 | + rng=SimpleNamespace(data_parallel_random_init=False), |
| 590 | + ) |
| 591 | + |
| 592 | + result = _build_distributed_model(cfg, pg_collection=MagicMock()) |
| 593 | + |
| 594 | + provider.provide_distributed_model.assert_called_once() |
| 595 | + assert result == expected_model |
| 596 | + |
562 | 597 |
|
563 | 598 | def test_restart_rebinds_overlap_callbacks_to_rebuilt_model(): |
564 | 599 | """A restart must replace callbacks bound to the discarded model.""" |
|
0 commit comments