CLI default attnres=off was overwriting block checkpoints on resume so later loads hit Unexpected key(s). Only apply flags the user passed. Track next_chunk so a budget-exit save does not skip the untrained yield. 0.5b now uses 01-ai/Yi-6B (64k); refuse resume when the ckpt tokenizer does not match.
82 lines
2.0 KiB
Python
82 lines
2.0 KiB
Python
"""Resume must not clobber attnres or skip a chunk at budget exit."""
|
|
from argparse import Namespace
|
|
from dataclasses import asdict
|
|
|
|
import pytest
|
|
import torch
|
|
|
|
from kda.models.causal_lm import CausalLM
|
|
from kda.models.k3_config import K3Config
|
|
from kda.training.toy import load_ckpt
|
|
from train_k3 import _apply_cli_overrides
|
|
|
|
|
|
def _tiny_block():
|
|
return K3Config(
|
|
hidden_size=32,
|
|
num_hidden_layers=4,
|
|
num_heads=4,
|
|
head_dim=8,
|
|
chunk_size=4,
|
|
vocab_size=64,
|
|
moe_latent_size=16,
|
|
moe_d_ff=16,
|
|
n_routed=4,
|
|
top_k=2,
|
|
n_shared=1,
|
|
kv_lora_rank=8,
|
|
q_lora_rank=16,
|
|
qk_nope_head_dim=8,
|
|
v_head_dim=8,
|
|
attnres="block",
|
|
attnres_block_size=2,
|
|
)
|
|
|
|
|
|
def _cli(**kwargs):
|
|
base = dict(
|
|
attnres=None,
|
|
attnres_block_size=None,
|
|
grad_checkpoint=None,
|
|
moe_aux_coef=None,
|
|
moe_z_coef=None,
|
|
)
|
|
base.update(kwargs)
|
|
return Namespace(**base)
|
|
|
|
|
|
def test_resume_without_attnres_flag_keeps_block():
|
|
cfg = _tiny_block()
|
|
_apply_cli_overrides(cfg, _cli())
|
|
assert cfg.attnres == "block"
|
|
assert cfg.attnres_block_size == 2
|
|
|
|
|
|
def test_explicit_attnres_overrides_resume():
|
|
cfg = _tiny_block()
|
|
_apply_cli_overrides(cfg, _cli(attnres="full", attnres_block_size=1))
|
|
assert cfg.attnres == "full"
|
|
assert cfg.attnres_block_size == 1
|
|
|
|
|
|
def test_block_state_with_off_config_cannot_load(tmp_path):
|
|
cfg = _tiny_block()
|
|
model = CausalLM(cfg)
|
|
payload = {"config": asdict(cfg), "model_state": model.state_dict()}
|
|
payload["config"]["attnres"] = "off"
|
|
payload["config"]["attnres_block_size"] = None
|
|
path = str(tmp_path / "polluted.pt")
|
|
torch.save(payload, path)
|
|
with pytest.raises(RuntimeError, match="Unexpected key"):
|
|
load_ckpt(path)
|
|
|
|
|
|
def test_budget_break_does_not_skip_yielded_chunk():
|
|
next_chunk = 10
|
|
for chunk_index in (10, 11, 12):
|
|
tokens = 100
|
|
if tokens >= 100:
|
|
break
|
|
next_chunk = chunk_index + 1
|
|
assert next_chunk == 10
|