"""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