diff --git a/README.md b/README.md index b431165..14662a4 100644 --- a/README.md +++ b/README.md @@ -138,7 +138,7 @@ PYTHONPATH=. python train.py ### 复现训练 -目标是 **~0.5B zh↔en 指令翻译模型**(`K3Config.preset("0.5b")` = 482M,tied Qwen3 词表)。本机 RTX 3060 6GB 只跑 8M 全流程孪生;0.5B 预训练需要 **32–40GB Ampere bf16**。 +目标是 **~0.5B zh↔en 指令翻译模型**(`K3Config.preset("0.5b")` ≈ 415M,tied Yi-6B 64k 词表)。本机 RTX 3060 6GB 只跑 8M 全流程孪生;0.5B 预训练需要 **32–40GB Ampere bf16**。 成功标准是冻结集上的 `translation_success()`,**不是** wiki train loss。wiki 预训练没见过 `Translate to English:\n...`,预训练阶段 `eval_mt` 的 success_rate 预期 ≈0。 @@ -160,7 +160,7 @@ uv run python train_k3.py --preset 0.5b --attnres block \ --max-tokens 1000000000 --warmup 2000 ``` -`0.5b` 预设:`d=768`,`L=24`(6×3 KDA + 1 MLA),`H=12`,`head=64`,`chunk=64`,LatentMoE `ℓ=384` / 16 routed / Top-2 / shared 2,tied Qwen3 embedding,activation checkpoint 默认开。训练默认 seq 2048、micro-batch 2、grad-acc 8、lr 3e-4、warmup **64 optimizer steps**。checkpoint:`ckpts/k3_0.5b.pt`,另写 `_last` / `_best`(best 按 held-out CE)。 +`0.5b` 预设:`d=768`,`L=24`(6×3 KDA + 1 MLA),`H=12`,`head=64`,`chunk=64`,LatentMoE `ℓ=384` / 16 routed / Top-2 / shared 2,tied Yi-6B embedding(64k,有 EOS),activation checkpoint 默认开。训练默认 seq 2048、micro-batch 2、grad-acc 8、lr 3e-4、warmup **64 optimizer steps**。checkpoint:`ckpts/k3_0.5b.pt`,另写 `_last` / `_best`(best 按 held-out CE)。**不能**从 Qwen3 词表的旧 ckpt `--resume`。 Token 会计:`step` 仍是 micro-batch;有效 token = `batch × seq_len × micro_steps`。默认 0.5b 冒烟是 **8.2M token ≈ 0.017 tok/param**。翻译前置 LM 的最低有意义预算是 **1B token**(`--max-tokens`),不是 2000 step。 diff --git a/kda/models/k3_config.py b/kda/models/k3_config.py index 233d2c8..93fb055 100644 --- a/kda/models/k3_config.py +++ b/kda/models/k3_config.py @@ -9,13 +9,15 @@ Hybrid Attention (K3): 每 4 层 1 次 Gated MLA, 末层强制 MLA. Presets: toy — ~8M, 自训 8k SP, 本地过拟合 - 0.5b — ~482M, Qwen3 词表, 32–40GB bf16;默认 step 是冒烟,翻译前置用 --max-tokens + 0.5b — ~415M, Yi-6B 词表 (64k), 32–40GB bf16;默认 step 是冒烟,翻译前置用 --max-tokens """ from __future__ import annotations from dataclasses import dataclass -# Qwen3 config.json; train_k3 overrides with len(tokenizer). +# 01-ai/Yi-6B config.json; train_k3 overrides with len(tokenizer). +YI6B_VOCAB_SIZE = 64000 +# Kept for old Qwen3 checkpoints / docs. QWEN3_VOCAB_SIZE = 151936 @@ -24,7 +26,7 @@ class K3Config: # 主干 hidden_size: int = 256 num_hidden_layers: int = 4 - vocab_size: int = 8192 # toy: data/spm_4k; 0.5b: Qwen3 + vocab_size: int = 8192 # toy: data/spm_4k; 0.5b: Yi-6B 64k initializer_range: float = 0.02 norm_eps: float = 1e-6 tie_word_embeddings: bool = False @@ -76,11 +78,11 @@ class K3Config: return cls() if name in {"0.5b", "500m"}: # H * head_dim == hidden. Routed 16 Top-2; LatentMoE padded bmm. - # ~482M with tied Qwen3 embeddings. 6×(3 KDA + 1 MLA). + # ~415M with tied Yi-6B embeddings. 6×(3 KDA + 1 MLA). return cls( hidden_size=768, num_hidden_layers=24, - vocab_size=QWEN3_VOCAB_SIZE, + vocab_size=YI6B_VOCAB_SIZE, tie_word_embeddings=True, max_position_embeddings=2048, num_heads=12, diff --git a/tests/integration/test_k3_arch.py b/tests/integration/test_k3_arch.py index 91649ab..eb2acdf 100644 --- a/tests/integration/test_k3_arch.py +++ b/tests/integration/test_k3_arch.py @@ -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 diff --git a/tests/integration/test_train_k3_resume.py b/tests/integration/test_train_k3_resume.py new file mode 100644 index 0000000..6275785 --- /dev/null +++ b/tests/integration/test_train_k3_resume.py @@ -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 diff --git a/train_k3.py b/train_k3.py index 8530289..ba8b922 100644 --- a/train_k3.py +++ b/train_k3.py @@ -41,7 +41,7 @@ _TOY_TRAIN = { "gen_every": 200, } _B500M_TRAIN = { - "tokenizer": "Qwen/Qwen3-8B", + "tokenizer": "01-ai/Yi-6B", "out": "ckpts/k3_0.5b.pt", "limit": 20000, "batch": 2, @@ -184,6 +184,20 @@ def _apply_moe_coefs(model, cfg: K3Config) -> None: module.z_loss_coef = cfg.moe_z_loss_coef +def _apply_cli_overrides(cfg: K3Config, args: argparse.Namespace) -> None: + """Copy only flags the user actually passed. CLI defaults must not clobber a resume.""" + if args.attnres is not None: + cfg.attnres = args.attnres + if args.attnres_block_size is not None: + cfg.attnres_block_size = args.attnres_block_size + if args.grad_checkpoint is not None: + cfg.gradient_checkpointing = args.grad_checkpoint + if args.moe_aux_coef is not None: + cfg.moe_aux_loss_coef = args.moe_aux_coef + if args.moe_z_coef is not None: + cfg.moe_z_loss_coef = args.moe_z_coef + + def _moe_log(model) -> dict: frac = moe_route_frac(model) if frac is None: @@ -267,9 +281,10 @@ def main() -> None: p.add_argument("--device", default="auto") p.add_argument( "--attnres", - default="off", + default=None, choices=["off", "full", "block"], - help="depth mixer: off=standard residual, block=K3 AttnRes, full=per-layer AttnRes", + help="depth mixer: off=standard residual (preset default), block=K3 AttnRes, " + "full=per-layer AttnRes. Omit on --resume to keep the checkpoint value", ) p.add_argument( "--attnres-block-size", @@ -329,14 +344,7 @@ def main() -> None: tok = load_tokenizer(args.tokenizer) cfg = K3Config.preset(args.preset) cfg.vocab_size = tok.vocab_size - cfg.attnres = args.attnres - cfg.attnres_block_size = args.attnres_block_size - if args.grad_checkpoint is not None: - cfg.gradient_checkpointing = args.grad_checkpoint - if args.moe_aux_coef is not None: - cfg.moe_aux_loss_coef = args.moe_aux_coef - if args.moe_z_coef is not None: - cfg.moe_z_loss_coef = args.moe_z_coef + _apply_cli_overrides(cfg, args) langs = [part.strip() for part in args.langs.split(",") if part.strip()] tpm = tokens_per_micro(args.batch, args.seq_len) @@ -364,19 +372,16 @@ def main() -> None: f"{type(loaded_cfg).__name__}" ) cfg = loaded_cfg - cfg.attnres = args.attnres - cfg.attnres_block_size = args.attnres_block_size - if args.grad_checkpoint is not None: - cfg.gradient_checkpointing = args.grad_checkpoint - if args.moe_aux_coef is not None: - cfg.moe_aux_loss_coef = args.moe_aux_coef - if args.moe_z_coef is not None: - cfg.moe_z_loss_coef = args.moe_z_coef + _apply_cli_overrides(cfg, args) model.gradient_checkpointing = cfg.gradient_checkpointing model.to(device) payload = torch.load(args.resume, map_location="cpu", weights_only=False) if payload.get("tokenizer") and payload["tokenizer"] != args.tokenizer: - print(f"warning: ckpt tokenizer {payload['tokenizer']} != {args.tokenizer}") + raise SystemExit( + f"tokenizer mismatch: ckpt {payload['tokenizer']!r} vs " + f"CLI {args.tokenizer!r}; embeddings are not interchangeable " + f"(do not resume a Qwen ckpt with Yi)" + ) micro_step = int(payload.get("micro_step", 0)) opt_step = int(payload.get("opt_step", 0)) tokens = int(payload.get("tokens", 0)) @@ -469,6 +474,10 @@ def main() -> None: model.train() t0 = time.perf_counter() tokens_at_t0 = tokens + # Index of the next untrained chunk. Mid-loop saves use last_trained+1. + # The final save must NOT +1 again: the loop may break on a yielded chunk + # that was never trained (budget check is at the top). + next_chunk = chunk_index for chunk_index, x, y in iter_indexed(train_chunks, start=chunk_index): if args.max_tokens is not None: if tokens >= args.max_tokens: @@ -497,6 +506,7 @@ def main() -> None: lr_now = optim.param_groups[0]["lr"] if raw_loss < best_train: best_train = raw_loss + next_chunk = chunk_index + 1 metrics = { "train/loss": raw_loss, @@ -542,7 +552,7 @@ def main() -> None: micro_step=micro_step, opt_step=opt_step, tokens=tokens, - chunk_index=chunk_index + 1, + chunk_index=next_chunk, best_heldout=best_heldout, ) _save(_sibling(args.out, "_best"), payload) @@ -570,7 +580,7 @@ def main() -> None: micro_step=micro_step, opt_step=opt_step, tokens=tokens, - chunk_index=chunk_index + 1, + chunk_index=next_chunk, best_heldout=best_heldout, ) _save(_sibling(args.out, "_last"), payload) @@ -586,7 +596,7 @@ def main() -> None: micro_step=micro_step, opt_step=opt_step, tokens=tokens, - chunk_index=chunk_index + 1, + chunk_index=next_chunk, best_heldout=best_heldout, ) _save(args.out, payload) @@ -596,6 +606,12 @@ def main() -> None: ) if tracker is not None: tracker.finish() + if device == "cuda": + try: + torch.cuda.synchronize() + torch.cuda.empty_cache() + except Exception: + pass if __name__ == "__main__":