Keep attnres on resume, fix final chunk_index, default 0.5b to Yi-6B
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.
This commit is contained in:
@@ -0,0 +1,81 @@
|
||||
"""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
|
||||
Reference in New Issue
Block a user