mirror of
https://github.com/radixark/miles.git
synced 2026-10-02 07:14:53 +08:00
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.
262 lines
9.9 KiB
Python
262 lines
9.9 KiB
Python
"""LoRA detection, adapter parameters, and training checkpoint state."""
|
|
|
|
from argparse import Namespace
|
|
from types import SimpleNamespace
|
|
|
|
import pytest
|
|
import torch
|
|
from tests.fast.fixtures.args_fixtures import make_trainer_args
|
|
|
|
import miles.backends.megatron_utils.lora.utils as lora_utils
|
|
from miles.backends.megatron_utils.lora.utils import (
|
|
_is_adapter_param_name,
|
|
is_lora_enabled,
|
|
load_lora_adapter,
|
|
save_lora_checkpoint,
|
|
)
|
|
from miles.utils.lora.utils import LORA_ADAPTER_NAME, is_lora_weight_name
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# is_lora_enabled
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestIsLoraEnabled:
|
|
def test_enabled_by_rank(self):
|
|
args = Namespace(lora_rank=32, lora_adapter_path=None)
|
|
assert is_lora_enabled(args) is True
|
|
|
|
def test_enabled_by_adapter_path(self):
|
|
args = Namespace(lora_rank=0, lora_adapter_path="/some/path")
|
|
assert is_lora_enabled(args) is True
|
|
|
|
def test_enabled_by_both(self):
|
|
args = Namespace(lora_rank=16, lora_adapter_path="/some/path")
|
|
assert is_lora_enabled(args) is True
|
|
|
|
def test_disabled(self):
|
|
args = Namespace(lora_rank=0, lora_adapter_path=None)
|
|
assert is_lora_enabled(args) is False
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# is_lora_weight_name / _is_adapter_param_name
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
class TestIsLoraWeightName:
|
|
@pytest.mark.parametrize(
|
|
"name",
|
|
[
|
|
"model.layers.0.self_attn.q_proj.lora_A.weight",
|
|
"model.layers.0.self_attn.q_proj.lora_B.weight",
|
|
"base_model.model.layers.5.mlp.gate_proj.lora_A.default.weight",
|
|
"base_model.model.layers.5.mlp.gate_proj.lora_B.default.weight",
|
|
],
|
|
)
|
|
def test_positive(self, name):
|
|
assert is_lora_weight_name(name) is True
|
|
|
|
@pytest.mark.parametrize(
|
|
"name",
|
|
[
|
|
"model.layers.0.self_attn.q_proj.weight",
|
|
"model.embed_tokens.weight",
|
|
"lm_head.weight",
|
|
"model.layers.0.mlp.gate_proj.weight",
|
|
],
|
|
)
|
|
def test_negative(self, name):
|
|
assert is_lora_weight_name(name) is False
|
|
|
|
|
|
class TestIsAdapterParamName:
|
|
@pytest.mark.parametrize(
|
|
"name",
|
|
[
|
|
"module.decoder.layers.0.self_attention.linear_qkv.lora_A.weight",
|
|
"module.decoder.layers.0.self_attention.linear_qkv.adapter.linear_in.weight",
|
|
"module.decoder.layers.0.self_attention.linear_qkv.adapter.linear_out.weight",
|
|
],
|
|
)
|
|
def test_positive(self, name):
|
|
assert _is_adapter_param_name(name) is True
|
|
|
|
@pytest.mark.parametrize(
|
|
"name",
|
|
[
|
|
"module.decoder.layers.0.self_attention.linear_qkv.weight",
|
|
"module.decoder.layers.0.mlp.linear_fc1.weight",
|
|
"module.embedding.word_embeddings.weight",
|
|
],
|
|
)
|
|
def test_negative(self, name):
|
|
assert _is_adapter_param_name(name) is False
|
|
|
|
|
|
# ---------------------------------------------------------------------------
|
|
# LORA_ADAPTER_NAME constant
|
|
# ---------------------------------------------------------------------------
|
|
|
|
|
|
def test_lora_adapter_name_constant():
|
|
assert LORA_ADAPTER_NAME == "miles_lora"
|
|
|
|
|
|
class _AdapterModel(torch.nn.Module):
|
|
def __init__(self):
|
|
super().__init__()
|
|
self.base = torch.nn.Linear(2, 2)
|
|
self.lora_A = torch.nn.Parameter(torch.zeros(1, 2))
|
|
self.lora_B = torch.nn.Parameter(torch.zeros(2, 1))
|
|
|
|
|
|
def _single_rank(monkeypatch):
|
|
monkeypatch.setattr(
|
|
lora_utils,
|
|
"get_parallel_state",
|
|
lambda: SimpleNamespace(tp=SimpleNamespace(rank=0), pp=SimpleNamespace(rank=0)),
|
|
)
|
|
|
|
|
|
def test_load_lora_adapter_rejects_a_shard_that_does_not_match_the_adapter(tmp_path, monkeypatch):
|
|
"""A silently partial load resumes training on a half-initialized adapter."""
|
|
_single_rank(monkeypatch)
|
|
torch.save({"lora_A": torch.ones(1, 2), "stale_lora_C": torch.ones(1)}, tmp_path / "adapter_megatron_rank0.pt")
|
|
|
|
with pytest.raises(RuntimeError, match=r"missing=\['lora_B'\], unexpected=\['stale_lora_C'\]"):
|
|
lora_utils.load_lora_adapter([_AdapterModel()], str(tmp_path))
|
|
|
|
|
|
def test_load_lora_adapter_rejects_shards_saved_under_another_layout(tmp_path, monkeypatch):
|
|
"""Falling through to fresh adapter weights would hide a resharding mistake."""
|
|
_single_rank(monkeypatch)
|
|
torch.save({"lora_A": torch.ones(1, 2)}, tmp_path / "adapter_megatron_rank1.pt")
|
|
|
|
with pytest.raises(FileNotFoundError, match="none for global rank 0"):
|
|
lora_utils.load_lora_adapter([_AdapterModel()], str(tmp_path))
|
|
|
|
|
|
class TestSaveLoraCheckpointTrainingState:
|
|
def _save(self, tmp_path, *, no_save_optim, scheduler=None):
|
|
publisher = SimpleNamespace(write_adapter=lambda *_: None)
|
|
|
|
adapter = torch.nn.Parameter(torch.ones(2))
|
|
model = [SimpleNamespace(named_parameters=lambda: [("layers.0.self_attention.lora_A.weight", adapter)])]
|
|
args = make_trainer_args(megatron_to_hf_mode="bridge", no_save_optim=no_save_optim)
|
|
optimizer = SimpleNamespace(state_dict=lambda: {"step": 7})
|
|
save_lora_checkpoint(
|
|
model,
|
|
args,
|
|
str(tmp_path / "checkpoint"),
|
|
publisher=publisher,
|
|
optimizer=optimizer,
|
|
opt_param_scheduler=scheduler,
|
|
iteration=3,
|
|
)
|
|
return sorted(path.name for path in (tmp_path / "checkpoint").iterdir())
|
|
|
|
@staticmethod
|
|
def _state(tmp_path):
|
|
return torch.load(tmp_path / "checkpoint" / "training_state_rank0.pt", weights_only=False)
|
|
|
|
def test_training_state_is_written_by_default(self, tmp_path):
|
|
scheduler = SimpleNamespace(state_dict=lambda: {"lr": 0.5})
|
|
files = self._save(tmp_path, no_save_optim=False, scheduler=scheduler)
|
|
|
|
assert files == ["adapter_megatron_rank0.pt", "training_state_rank0.pt"]
|
|
state = self._state(tmp_path)
|
|
assert state["optimizer"] == {"step": 7}
|
|
assert state["opt_param_scheduler"] == {"lr": 0.5}
|
|
assert state["iteration"] == 3
|
|
|
|
def test_no_save_optim_drops_the_optimizer_and_keeps_the_resume_metadata(self, tmp_path):
|
|
"""--no-save-optim is about optimizer state; losing the step and the LR schedule with it
|
|
would silently restart a resumed run from iteration 0."""
|
|
scheduler = SimpleNamespace(state_dict=lambda: {"lr": 0.5})
|
|
files = self._save(tmp_path, no_save_optim=True, scheduler=scheduler)
|
|
|
|
assert files == ["adapter_megatron_rank0.pt", "training_state_rank0.pt"]
|
|
state = self._state(tmp_path)
|
|
assert state["optimizer"] is None
|
|
assert state["opt_param_scheduler"] == {"lr": 0.5}
|
|
assert state["iteration"] == 3
|
|
|
|
|
|
class TestLoadTrainingState:
|
|
@staticmethod
|
|
def _recorder():
|
|
loaded = []
|
|
return loaded, SimpleNamespace(load_state_dict=loaded.append)
|
|
|
|
def _write(self, tmp_path, optimizer_state):
|
|
torch.save(
|
|
{"iteration": 3, "optimizer": optimizer_state, "opt_param_scheduler": {"lr": 0.5}},
|
|
tmp_path / "training_state_rank0.pt",
|
|
)
|
|
|
|
def test_an_optimizer_free_checkpoint_still_restores_the_step_and_the_schedule(self, tmp_path):
|
|
self._write(tmp_path, None)
|
|
optimizer_loads, optimizer = self._recorder()
|
|
scheduler_loads, scheduler = self._recorder()
|
|
|
|
assert lora_utils._load_training_state(tmp_path, optimizer, scheduler) == (3, False)
|
|
assert optimizer_loads == []
|
|
assert scheduler_loads == [{"lr": 0.5}]
|
|
|
|
def test_a_full_checkpoint_restores_the_optimizer(self, tmp_path):
|
|
self._write(tmp_path, {"step": 7})
|
|
optimizer_loads, optimizer = self._recorder()
|
|
scheduler_loads, scheduler = self._recorder()
|
|
|
|
assert lora_utils._load_training_state(tmp_path, optimizer, scheduler) == (3, True)
|
|
assert optimizer_loads == [{"step": 7}]
|
|
assert scheduler_loads == [{"lr": 0.5}]
|
|
|
|
|
|
class TestLoadTrainingStateOptimizerGate:
|
|
"""--no-load-optim must keep the fresh optimizer without losing the step or the LR schedule."""
|
|
|
|
@staticmethod
|
|
def _recorder():
|
|
loaded = []
|
|
return loaded, SimpleNamespace(load_state_dict=loaded.append)
|
|
|
|
@staticmethod
|
|
def _write_training_state(tmp_path):
|
|
torch.save(
|
|
{"iteration": 11, "optimizer": {"step": 7}, "opt_param_scheduler": {"lr": 0.5}},
|
|
tmp_path / "training_state_rank0.pt",
|
|
)
|
|
|
|
def test_no_load_optim_skips_the_optimizer_and_keeps_the_rest(self, tmp_path):
|
|
self._write_training_state(tmp_path)
|
|
optimizer_loads, optimizer = self._recorder()
|
|
scheduler_loads, scheduler = self._recorder()
|
|
|
|
assert lora_utils._load_training_state(tmp_path, optimizer, scheduler, load_optimizer=False) == (11, False)
|
|
assert optimizer_loads == []
|
|
assert scheduler_loads == [{"lr": 0.5}]
|
|
|
|
def test_load_lora_adapter_forwards_the_flag(self, tmp_path, monkeypatch):
|
|
rank0 = SimpleNamespace(rank=0)
|
|
monkeypatch.setattr(lora_utils, "get_parallel_state", lambda: SimpleNamespace(tp=rank0, pp=rank0))
|
|
name = "layers.0.self_attention.lora_A.weight"
|
|
torch.save({name: torch.ones(2)}, tmp_path / "adapter_megatron_rank0.pt")
|
|
self._write_training_state(tmp_path)
|
|
model = [SimpleNamespace(named_parameters=lambda: [(name, torch.nn.Parameter(torch.zeros(2)))])]
|
|
optimizer_loads, optimizer = self._recorder()
|
|
scheduler_loads, scheduler = self._recorder()
|
|
|
|
loaded, iteration, optimizer_restored = load_lora_adapter(
|
|
model,
|
|
str(tmp_path),
|
|
optimizer=optimizer,
|
|
opt_param_scheduler=scheduler,
|
|
load_optimizer=False,
|
|
)
|
|
|
|
assert (loaded, iteration, optimizer_restored) == (True, 11, False)
|
|
assert optimizer_loads == []
|
|
assert scheduler_loads == [{"lr": 0.5}]
|