"""load_ckpt must read checkpoints written before the ffn/config renames.""" from dataclasses import asdict import pytest import torch from kda.models.config import KDAConfig from kda.models.k3_config import K3Config from kda.training.toy import load_ckpt, save_ckpt def _legacy(state, new, old): """Undo the ffn rename, reproducing a pre-rename checkpoint.""" renamed = { k.replace(f".{new}.", f".{old}.").replace(f".{new}_norm.", f".{old}_norm."): v for k, v in state.items() } assert any(f".{old}." in k for k in renamed), "fixture renamed nothing" return renamed @pytest.mark.parametrize( ("config", "old"), [ (KDAConfig(num_hidden_layers=2), "mlp"), # dense: was named .mlp (K3Config(num_hidden_layers=2), "moe"), # K3: was named .moe ], ids=["kda-mlp", "k3-moe"], ) def test_legacy_ffn_names_still_load(tmp_path, config, old): from kda.models.causal_lm import CausalLM model = CausalLM(config) path = str(tmp_path / "legacy.pt") torch.save( { "model_state": _legacy(model.state_dict(), "ffn", old), "config": asdict(config), }, path, ) loaded, loaded_config = load_ckpt(path) assert type(loaded_config) is type(config) for name, want in model.state_dict().items(): torch.testing.assert_close(loaded.state_dict()[name], want) @pytest.mark.parametrize("config", [KDAConfig(num_hidden_layers=2), K3Config(num_hidden_layers=2)]) def test_roundtrip_picks_the_right_config_class(tmp_path, config): from kda.models.causal_lm import CausalLM path = str(tmp_path / "ckpt.pt") save_ckpt(CausalLM(config), config, path) _, loaded_config = load_ckpt(path) assert loaded_config == config