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:
dela
2026-08-26 14:23:29 +08:00
parent 5cc0555563
commit 24c9d56b72
5 changed files with 130 additions and 30 deletions
+1
View File
@@ -78,6 +78,7 @@ def test_preset_0_5b_schedule():
assert cfg.hidden_size == 768
assert cfg.num_heads * cfg.head_dim == cfg.hidden_size
assert cfg.num_hidden_layers == 24
assert cfg.vocab_size == 64000
assert cfg.tie_word_embeddings
assert cfg.chunk_size == 64
assert cfg.gradient_checkpointing is True
+81
View File
@@ -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