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.
This commit is contained in:
@@ -0,0 +1,56 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user