Files
K3/tests/integration/test_ckpt_compat.py
dela 584f7e9e73 Initial K3 snapshot: 0.5B KDA/MLA/MoE train path
Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias
backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
2026-08-25 14:43:17 +08:00

57 lines
1.7 KiB
Python

"""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