"""L7: toy training loop — overfit 起步. target: 端到端验证模型 + 数据流 + optimizer + ckpt + generate. toy data: 建一份 256-token vocab 的小数据集: e.g. 1000 个长度 32 随机 token 序列 起步只取 batch=4, 看能否在 ~320 steps 内把 loss 压到 < 0.1 (overfit 单 batch). step: optimizer = AdamW(lr=1e-3, wd=0.01) loss.backward(); optimizer.step(); optimizer.zero_grad() every N steps: 打印 loss end: 保存 ckpt to ckpts/kda_toy.pt ckpt: save: torch.save({"model_state": model.state_dict(), "config": asdict(config)}, path) load: torch.load -> model.load_state_dict """ from __future__ import annotations import os from dataclasses import asdict, fields import torch from ..models.causal_lm import CausalLM from ..models.config import KDAConfig from ..models.k3_config import K3Config def make_toy_data(batch: int = 4, seq_len: int = 32, vocab: int = 256, seed: int = 42): """单 batch overfit 数据: 同一组序列循环.""" torch.manual_seed(seed) seq = torch.randint(0, vocab, (batch, seq_len), dtype=torch.long) return seq # 用作 input_ids 和 labels (shift one inside forward) def train_one_batch(model, optimizer, input_ids, labels): optimizer.zero_grad(set_to_none=True) loss = model(input_ids, labels=labels) loss.backward() optimizer.step() return loss.detach() def save_ckpt(model, config, path: str): os.makedirs(os.path.dirname(path) or ".", exist_ok=True) torch.save({"model_state": model.state_dict(), "config": asdict(config)}, path) #: The feed-forward submodule was named after its contents (``mlp`` in the #: dense config, ``moe`` in K3) before both were unified under ``ffn``. #: Checkpoints saved before that rename still carry the old prefixes. _LEGACY_PREFIXES = { ".mlp.": ".ffn.", ".mlp_norm.": ".ffn_norm.", ".moe.": ".ffn.", ".moe_norm.": ".ffn_norm.", } def _rename_legacy_keys(state: dict) -> dict: def fix(key: str) -> str: for old, new in _LEGACY_PREFIXES.items(): if old in key: return key.replace(old, new) return key return {fix(k): v for k, v in state.items()} def _config_from(payload_config: dict) -> K3Config | KDAConfig: """Pick the config class the checkpoint was written with. ``moe_latent_size`` is a K3-only field, so its presence identifies the hybrid K3 architecture; anything else is the dense KDA config. """ cls = K3Config if "moe_latent_size" in payload_config else KDAConfig known = {item.name for item in fields(cls)} return cls(**{k: v for k, v in payload_config.items() if k in known}) def load_ckpt(path: str, model: CausalLM | None = None) -> tuple[CausalLM, K3Config | KDAConfig]: payload = torch.load(path, map_location="cpu", weights_only=False) config = _config_from(payload["config"]) if model is None: model = CausalLM(config) model.load_state_dict(_rename_legacy_keys(payload["model_state"])) return model, config def main(): """主入口: overfit 起步. 320 steps 期望 loss < 0.1.""" device = "cuda" if torch.cuda.is_available() else "cpu" config = KDAConfig() model = CausalLM(config).to(device) tokens = make_toy_data(seq_len=32, vocab=config.vocab_size).to(device) optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3, weight_decay=0.01) for step in range(320): loss = train_one_batch(model, optimizer, tokens, tokens) if step % 64 == 0 or step == 319: print(f"step {step:3d} loss {loss.item():.4f}") save_ckpt(model, config, "ckpts/kda_toy.pt") if __name__ == "__main__": main()