Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
57 lines
1.7 KiB
Python
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
|