Files
miles/tests/fast/backends/megatron_utils/test_model_initialize.py
T
Tom 6501f53c71 Migrate unit tests and preserve Verifiers contract coverage
Materialize the final GU operation with unit-test migrations and regression coverage. Retire the AllConfig parallel-property test case alongside the explicitly deleted op48-14; preserve TrainerConfig product, no-duplicate-field, and immutability assertions.
2026-10-01 15:43:36 +08:00

321 lines
12 KiB
Python

import sys
import types
from contextlib import ExitStack
from pathlib import Path
from typing import TYPE_CHECKING
from unittest.mock import MagicMock, patch
import pytest
from tests.fast.fixtures.args_fixtures import make_trainer_args, make_trainer_config
if TYPE_CHECKING:
from miles.backends.megatron_utils.model import LoadCheckpointOutput
def _stub_module(name: str, attrs: dict[str, object] | None = None, is_package: bool = False) -> types.ModuleType:
module = types.ModuleType(name)
if is_package:
module.__path__ = []
if attrs is not None:
for attr_name, value in attrs.items():
setattr(module, attr_name, value)
sys.modules[name] = module
return module
class _DummyDDP:
pass
class _DummyModel:
pass
class _DummyOptimizer:
pass
class _DummyChainedOptimizer:
pass
class _DummyDistributedOptimizer:
pass
class _DummyScheduler:
pass
class _DummyOptimizerConfig:
def __init__(self, **kwargs):
self.__dict__.update(kwargs)
class _FakeModelChunk:
role: str | None = None
@pytest.fixture(scope="module", autouse=True)
def _mock_megatron_environment():
original_modules = dict(sys.modules)
try:
_stub_module("megatron", is_package=True)
core_module = _stub_module("megatron.core", is_package=True)
core_module.mpu = types.SimpleNamespace()
core_module.tensor_parallel = _stub_module(
"megatron.core.tensor_parallel",
{"model_parallel_cuda_manual_seed": MagicMock()},
is_package=True,
)
_stub_module(
"megatron.core.tensor_parallel.random",
{"_get_all_rng_states": MagicMock(), "_set_all_rng_states": MagicMock()},
)
_stub_module(
"megatron.core.distributed",
{
"DistributedDataParallel": _DummyDDP,
"finalize_model_grads": MagicMock(),
},
)
_stub_module(
"megatron.core.enums",
{"ModelType": types.SimpleNamespace(encoder_or_decoder="encoder_or_decoder")},
)
_stub_module("megatron.core.models", is_package=True)
_stub_module("megatron.core.models.gpt", {"GPTModel": _DummyModel})
_stub_module(
"megatron.core.optimizer",
{
"OptimizerConfig": _DummyOptimizerConfig,
"get_megatron_optimizer": MagicMock(),
"Adam": _DummyOptimizer,
"CPUAdam": _DummyOptimizer,
},
is_package=True,
)
_stub_module("megatron.core.optimizer.emerging_optimizers", {"TensorParallelMuon": _DummyOptimizer})
_stub_module("megatron.core.optimizer.muon", {"get_megatron_muon_optimizer": MagicMock()})
_stub_module("megatron.core.optimizer.distrib_optimizer", {"DistributedOptimizer": _DummyDistributedOptimizer})
_stub_module(
"megatron.core.optimizer.optimizer",
{
"ChainedOptimizer": _DummyChainedOptimizer,
"MegatronOptimizer": _DummyOptimizer,
},
)
_stub_module("megatron.core.optimizer_param_scheduler", {"OptimizerParamScheduler": _DummyScheduler})
_stub_module("megatron.core.packed_seq_params", {"PackedSeqParams": MagicMock()})
_stub_module("megatron.core.pipeline_parallel", {"get_forward_backward_func": MagicMock()})
_stub_module("megatron.core.transformer", is_package=True)
_stub_module("megatron.core.transformer.utils", {"sharded_state_dict_default": MagicMock()})
_stub_module("megatron.core.utils", {"get_model_config": MagicMock(), "unwrap_model": MagicMock()})
_stub_module("megatron.core.config", {"set_experimental_flag": MagicMock()})
_stub_module("megatron.core.num_microbatches_calculator", {"init_num_microbatches_calculator": MagicMock()})
_stub_module("megatron.training", is_package=True)
_stub_module(
"megatron.training.global_vars",
{
"get_args": MagicMock(),
"_build_tokenizer": MagicMock(),
"set_args": MagicMock(),
},
)
_stub_module("megatron.training.training", {"get_model": MagicMock()})
_stub_module(
"megatron.training.checkpointing",
{
"load_checkpoint": MagicMock(),
"save_checkpoint": MagicMock(),
},
)
_stub_module("sglang.srt.debug_utils", is_package=True)
_stub_module(
"sglang.srt.debug_utils.dumper",
{
"DumperConfig": MagicMock(),
"_get_rank": MagicMock(return_value=0),
"dumper": MagicMock(),
},
)
_stub_module(
"miles.backends.megatron_utils.lora.bridge",
{
"_ensure_model_list": MagicMock(),
"_setup_lora_model_via_bridge": MagicMock(),
},
)
_stub_module(
"miles.backends.megatron_utils.model_provider",
{
"get_model_provider_func": MagicMock(),
"LinearForLastLayer": _DummyModel,
},
)
yield
finally:
sys.modules.clear()
sys.modules.update(original_modules)
def _patch_initialize_side_effects(stack: ExitStack) -> None:
stack.enter_context(patch("miles.backends.megatron_utils.model.clear_memory"))
stack.enter_context(patch("miles.backends.megatron_utils.model.check_peak_gpu_memory_after_load"))
stack.enter_context(patch("miles.backends.megatron_utils.model.check_model_hashes"))
def test_initialize_does_not_step_scheduler_restored_from_checkpoint():
from miles.backends.megatron_utils.model import LoadCheckpointOutput, initialize_model_and_optimizer
args = make_trainer_config(use_checkpoint_opt_param_scheduler=True, global_batch_size=8, finetune=False)
model = [_FakeModelChunk()]
optimizer = object()
opt_param_scheduler = MagicMock()
with ExitStack() as stack:
stack.enter_context(
patch(
"miles.backends.megatron_utils.model.setup_model_and_optimizer",
return_value=(model, optimizer, opt_param_scheduler),
)
)
stack.enter_context(
patch("miles.backends.megatron_utils.model.load_checkpoint", return_value=(100, True, False))
)
_patch_initialize_side_effects(stack)
result = initialize_model_and_optimizer(args)
assert result == (
model,
optimizer,
opt_param_scheduler,
LoadCheckpointOutput(loaded_rollout_id=100, start_rollout_id=101),
)
opt_param_scheduler.step.assert_not_called()
def test_initialize_steps_scheduler_when_checkpoint_did_not_restore_it():
from miles.backends.megatron_utils.model import LoadCheckpointOutput, initialize_model_and_optimizer
args = make_trainer_config(use_checkpoint_opt_param_scheduler=False, global_batch_size=8, finetune=False)
model = [_FakeModelChunk()]
optimizer = object()
opt_param_scheduler = MagicMock()
with ExitStack() as stack:
stack.enter_context(
patch(
"miles.backends.megatron_utils.model.setup_model_and_optimizer",
return_value=(model, optimizer, opt_param_scheduler),
)
)
stack.enter_context(
patch("miles.backends.megatron_utils.model.load_checkpoint", return_value=(100, True, False))
)
_patch_initialize_side_effects(stack)
result = initialize_model_and_optimizer(args)
assert result == (
model,
optimizer,
opt_param_scheduler,
LoadCheckpointOutput(loaded_rollout_id=100, start_rollout_id=101),
)
opt_param_scheduler.step.assert_called_once_with(increment=800)
def _load_model_state_with(
*,
tmp_path: Path,
finetune: bool,
iteration: int,
lora_rank: int = 0,
restored_trained_iteration: bool | None = None,
) -> "LoadCheckpointOutput":
from miles.backends.megatron_utils.model import load_model_state
load_dir = tmp_path / "ckpt"
load_dir.mkdir()
(load_dir / "latest_checkpointed_iteration.txt").write_text(str(iteration))
if restored_trained_iteration is None:
restored_trained_iteration = not finetune or iteration > 0
with ExitStack() as stack:
stack.enter_context(
patch(
"miles.backends.megatron_utils.model.load_checkpoint",
return_value=(iteration, restored_trained_iteration, False),
)
)
_patch_initialize_side_effects(stack)
return load_model_state(
make_trainer_args(
use_checkpoint_opt_param_scheduler=True,
global_batch_size=8,
finetune=finetune,
lora_rank=lora_rank,
megatron_to_hf_mode="core",
lora_adapter_path=None,
load=str(load_dir),
),
model=[_FakeModelChunk()],
optimizer=None,
opt_param_scheduler=None,
role="actor",
checkpointing_context=None,
)
class TestWhereALoadSaysTheRunStarts:
def test_a_finetune_load_starts_the_run_at_rollout_zero(self, tmp_path: Path):
"""--finetune means there is no run to continue, so rollout 0 is still ahead rather than behind."""
assert _load_model_state_with(tmp_path=tmp_path, finetune=True, iteration=0).start_rollout_id == 0
def test_a_resumed_load_starts_the_run_after_the_checkpoint_it_read(self, tmp_path: Path):
"""The checkpoint's own rollout is done, so the run continues at the next one."""
assert _load_model_state_with(tmp_path=tmp_path, finetune=False, iteration=100).start_rollout_id == 101
def test_a_run_that_restored_the_iteration_zero_checkpoint_it_wrote_starts_at_one(self, tmp_path: Path):
"""A real resume from the very first checkpoint must not be read as a finetune that starts over."""
output = _load_model_state_with(tmp_path=tmp_path, finetune=False, iteration=0)
assert output.start_rollout_id == 1
def test_weight_initialization_with_a_trained_iteration_is_refused(self, tmp_path: Path):
"""Weight initialization cannot claim a nonzero trained iteration."""
with pytest.raises(AssertionError, match="Weight initialization returned a trained iteration"):
_load_model_state_with(tmp_path=tmp_path, finetune=True, iteration=100, restored_trained_iteration=False)
class TestALoraAdapterThatCarriesItsOwnIteration:
def test_a_lora_resume_at_iteration_zero_preserves_the_restored_training_state(self, tmp_path: Path) -> None:
"""An adapter's saved iteration zero must restore the rollout state just like later iterations."""
output = _load_model_state_with(
tmp_path=tmp_path, finetune=True, iteration=0, lora_rank=8, restored_trained_iteration=True
)
assert output.start_rollout_id == 1
def test_a_lora_resume_under_finetune_continues_after_the_iteration_the_adapter_names(self, tmp_path: Path):
"""LoRA saves write no tracker, so a lora resume always arrives here with --finetune set."""
output = _load_model_state_with(tmp_path=tmp_path, finetune=True, iteration=100, lora_rank=8)
assert output.start_rollout_id == 101
def test_a_lora_run_that_really_starts_from_scratch_starts_at_rollout_zero(self, tmp_path: Path):
"""An adapter with no training state starts a new run at rollout zero."""
assert _load_model_state_with(tmp_path=tmp_path, finetune=True, iteration=0, lora_rank=8).start_rollout_id == 0
def test_a_lora_run_that_really_starts_from_scratch_restored_no_trained_iteration(self, tmp_path: Path):
"""Nothing was trained, so no rollout state was ever saved for the rollout side to restore."""
output = _load_model_state_with(tmp_path=tmp_path, finetune=True, iteration=0, lora_rank=8)
assert output.start_rollout_id == 0
def test_a_lora_resume_restored_a_trained_iteration(self, tmp_path: Path):
"""The adapter carries a trained iteration, so the rollout state saved beside it must be restored."""
output = _load_model_state_with(tmp_path=tmp_path, finetune=True, iteration=100, lora_rank=8)
assert output.start_rollout_id == output.loaded_rollout_id + 1