"""Train a Kimi-K3-like model (KDA + Gated MLA + LatentMoE) on bilingual wiki. 用法: uv run python train_k3.py # 8M toy, zh+en wiki uv run python train_k3.py --preset 0.5b # ~0.5B smoke (8M tokens) # 32–40GB Ampere, 1B-token 翻译前置预训练: uv run python train_k3.py --preset 0.5b --attnres block \\ --max-tokens 1000000000 --warmup 2000 """ from __future__ import annotations import argparse import os import time from dataclasses import asdict import torch from kda.layers.latent_moe import LatentMoE, moe_route_frac, moe_router_losses from kda.models.causal_lm import CausalLM from kda.models.k3_config import K3Config from kda.training.data import iter_indexed, load_pretrain_chunks, load_tokenizer from kda.training.schedule import lr_scale, tokens_per_micro, total_opt_steps from kda.training.toy import load_ckpt _TOY_TRAIN = { "tokenizer": "data/spm_4k.model", "out": "ckpts/k3_wiki.pt", "limit": 8000, "batch": 4, "seq_len": 256, "steps": 500, "lr": 2e-3, "warmup": 50, "grad_acc": 1, "eval_every": 100, } _B500M_TRAIN = { "tokenizer": "Qwen/Qwen3-8B", "out": "ckpts/k3_0.5b.pt", "limit": 20000, "batch": 2, "seq_len": 2048, "steps": 2000, "lr": 3e-4, "warmup": 64, "grad_acc": 8, "eval_every": 100, } def _sibling(path: str, suffix: str) -> str: root, ext = os.path.splitext(path) return f"{root}{suffix}{ext or '.pt'}" def _set_lr(optim: torch.optim.Optimizer, lr: float) -> None: for group in optim.param_groups: group["lr"] = lr def _init_swanlab(cfg: K3Config, args: argparse.Namespace): """Cloud monitor if SWANLAB_API_KEY is set; otherwise no-op.""" key = os.environ.get("SWANLAB_API_KEY") if not key: return None try: import swanlab except ImportError: print("SWANLAB_API_KEY set but swanlab is not installed") return None try: swanlab.login(api_key=key, save=False) # swanlab 0.9 Settings.project is nested; a string SWANLAB_PROJECT env crashes init. project = os.environ.pop("SWANLAB_PROJECT", None) or "kda" return swanlab.init( project=project, name=f"{args.preset}-{cfg.attnres}", config={ "preset": args.preset, "attnres": cfg.attnres, "attnres_block_size": cfg.attnres_block_size, "lr": args.lr, "batch": args.batch, "seq_len": args.seq_len, "steps": args.steps, "max_tokens": args.max_tokens, "grad_acc": args.grad_acc, "warmup": args.warmup, "langs": args.langs, "kda_backend": cfg.kda_backend, "gradient_checkpointing": cfg.gradient_checkpointing, "moe_aux_loss_coef": cfg.moe_aux_loss_coef, "moe_z_loss_coef": cfg.moe_z_loss_coef, }, ) except Exception as exc: print(f"swanlab init failed ({exc}); continuing without cloud monitor") return None def _payload( cfg: K3Config, model: CausalLM, optim: torch.optim.Optimizer, args: argparse.Namespace, *, micro_step: int, opt_step: int, tokens: int, chunk_index: int, best_heldout: float, ): return { "config": asdict(cfg), "model_state": model.state_dict(), "optimizer_state": optim.state_dict(), "tokenizer": args.tokenizer, "preset": args.preset, "micro_step": micro_step, "opt_step": opt_step, "tokens": tokens, "chunk_index": chunk_index, "best_heldout": best_heldout, } def _save(path: str, payload: dict) -> None: os.makedirs(os.path.dirname(path) or ".", exist_ok=True) tmp = path + ".tmp" torch.save(payload, tmp) os.replace(tmp, path) @torch.no_grad() def _heldout_loss( model, chunks, device, use_bf16, max_batches: int = 4 ) -> float | None: if chunks is None or chunks.numel() == 0: return None model.eval() losses = [] for x, y in zip(chunks[:max_batches], chunks[:max_batches]): x = x.to(device) y = y.to(device) with torch.autocast("cuda", dtype=torch.bfloat16, enabled=use_bf16): losses.append(float(model(x, labels=y).item())) model.train() return sum(losses) / max(len(losses), 1) def _apply_moe_coefs(model, cfg: K3Config) -> None: for module in model.modules(): if isinstance(module, LatentMoE): module.aux_loss_coef = cfg.moe_aux_loss_coef module.z_loss_coef = cfg.moe_z_loss_coef def _moe_log(model) -> dict: frac = moe_route_frac(model) if frac is None: return {} n = frac.numel() dead = float((frac < 1.0 / (2 * n)).sum().item()) return { "moe/max_frac": float(frac.max()), "moe/min_frac": float(frac.min()), "moe/n_dead": dead, } def main() -> None: pre = argparse.ArgumentParser(add_help=False) pre.add_argument("--preset", default="toy", choices=["toy", "0.5b"]) pre_args, _ = pre.parse_known_args() train_defaults = _B500M_TRAIN if pre_args.preset == "0.5b" else _TOY_TRAIN p = argparse.ArgumentParser(description=__doc__) p.add_argument("--preset", default="toy", choices=["toy", "0.5b"]) p.add_argument( "--tokenizer", "--sp", dest="tokenizer", default=train_defaults["tokenizer"] ) p.add_argument("--out", default=train_defaults["out"]) p.add_argument("--limit", type=int, default=train_defaults["limit"]) p.add_argument("--batch", type=int, default=train_defaults["batch"]) p.add_argument("--seq-len", type=int, default=train_defaults["seq_len"]) p.add_argument("--steps", type=int, default=train_defaults["steps"]) p.add_argument( "--max-tokens", type=int, default=None, help="stop after this many tokens (primary budget). --steps is a micro-step cap", ) p.add_argument("--lr", type=float, default=train_defaults["lr"]) p.add_argument( "--warmup", type=int, default=train_defaults["warmup"], help="linear warmup in optimizer steps, then cosine to 0.1× lr", ) p.add_argument("--grad-acc", type=int, default=train_defaults["grad_acc"]) p.add_argument("--eval-every", type=int, default=train_defaults["eval_every"]) p.add_argument( "--langs", default="zh,en", help="comma-separated wiki languages mixed 1:1 by token (default zh,en)", ) p.add_argument("--heldout-frac", type=float, default=0.01) p.add_argument("--resume", default=None, help="checkpoint to continue from") p.add_argument("--gen-prefix", action="append", default=None) p.add_argument("--device", default="auto") p.add_argument( "--attnres", default="off", choices=["off", "full", "block"], help="depth mixer: off=standard residual, block=K3 AttnRes, full=per-layer AttnRes", ) p.add_argument( "--attnres-block-size", type=int, default=None, help="DecoderBlocks per AttnRes block (block mode). Default ≈ L/8", ) p.add_argument( "--grad-checkpoint", dest="grad_checkpoint", action="store_true", default=None, help="activation checkpointing (0.5b preset default on)", ) p.add_argument( "--no-grad-checkpoint", dest="grad_checkpoint", action="store_false", ) p.add_argument( "--moe-aux-coef", type=float, default=None, help="Switch/GShard aux loss weight (default 0.01; 0 disables)", ) p.add_argument( "--moe-z-coef", type=float, default=None, help="router z-loss weight (default 0.001; 0 disables)", ) args = p.parse_args() os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True") if args.gen_prefix is None: args.gen_prefix = ["人工智能的发展", "The history of computing"] if args.seq_len % K3Config.preset(args.preset).chunk_size: raise SystemExit( f"seq-len {args.seq_len} must be divisible by chunk_size " f"{K3Config.preset(args.preset).chunk_size}" ) device = args.device if device == "auto": device = "cuda" if torch.cuda.is_available() else "cpu" use_bf16 = device == "cuda" and torch.cuda.is_bf16_supported() if device == "cuda" and not use_bf16: raise SystemExit( "KDA training needs bf16; this GPU does not support it (avoid V100 fp16)" ) print(f"loading tokenizer {args.tokenizer} ...") 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 langs = [part.strip() for part in args.langs.split(",") if part.strip()] tpm = tokens_per_micro(args.batch, args.seq_len) horizon = total_opt_steps( max_tokens=args.max_tokens, max_micro=args.steps, batch=args.batch, seq_len=args.seq_len, grad_acc=args.grad_acc, ) micro_step = 0 opt_step = 0 tokens = 0 chunk_index = 0 best_heldout = float("inf") best_train = float("inf") if args.resume: print(f"resume {args.resume}") model, loaded_cfg = load_ckpt(args.resume) if not isinstance(loaded_cfg, K3Config): raise SystemExit( f"train_k3.py requires a K3 checkpoint; {args.resume} has " 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 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}") micro_step = int(payload.get("micro_step", 0)) opt_step = int(payload.get("opt_step", 0)) tokens = int(payload.get("tokens", 0)) chunk_index = int(payload.get("chunk_index", 0)) best_heldout = float(payload.get("best_heldout", best_heldout)) else: model = CausalLM(cfg).to(device) _apply_moe_coefs(model, cfg) tracker = _init_swanlab(cfg, args) n = sum(p.numel() for p in model.parameters()) print( f"preset={args.preset} model={n:,} params ({n / 1e6:.1f}M) on {device} " f"bf16={use_bf16} checkpoint={cfg.gradient_checkpointing}" ) print( f"vocab={cfg.vocab_size} tied={cfg.tie_word_embeddings} " f"layers={cfg.layer_types()} attnres={cfg.attnres} langs={langs} " f"moe_aux={cfg.moe_aux_loss_coef:g} moe_z={cfg.moe_z_loss_coef:g}" ) if args.max_tokens is None: print( f"token budget: --steps {args.steps} micro " f"({args.steps * tpm:,} tokens); pass --max-tokens for a real run" ) else: print(f"token budget: {args.max_tokens:,} cosine horizon {horizon} opt steps") train_chunks, held_chunks, n_ids = load_pretrain_chunks( tok, langs=langs, limit=args.limit, batch=args.batch, seq_len=args.seq_len, heldout_frac=args.heldout_frac, ) print( f"packed tokens {n_ids:,} -> {train_chunks.size(0)} train / " f"{held_chunks.size(0)} held-out chunks of [{args.batch}, {args.seq_len}]" ) if train_chunks.size(0) == 0: raise SystemExit("no training chunks; raise --limit or lower --batch/--seq-len") optim = torch.optim.AdamW( model.parameters(), lr=args.lr, weight_decay=0.1 if args.preset == "0.5b" else 0.01, ) if args.resume: payload = torch.load(args.resume, map_location="cpu", weights_only=False) if payload.get("optimizer_state"): optim.load_state_dict(payload["optimizer_state"]) def gen_sample(prefix: str, max_new: int = 24) -> str: ids_ = tok.encode(prefix) if not ids_: return "" inp = torch.tensor([ids_], dtype=torch.long, device=device) was_training = model.training model.eval() with torch.inference_mode(): out = model.generate(inp, max_new) if was_training: model.train() return tok.decode(out[0].tolist()) model.train() t0 = time.perf_counter() tokens_at_t0 = tokens 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: break elif micro_step >= args.steps: break x, y = x.to(device), y.to(device) scale = lr_scale(opt_step, args.warmup, horizon) _set_lr(optim, args.lr * scale) with torch.autocast("cuda", dtype=torch.bfloat16, enabled=use_bf16): task = model(x, labels=y) aux, z_loss = moe_router_losses(model) loss = (task + aux + z_loss) / args.grad_acc loss.backward() do_step = (micro_step + 1) % args.grad_acc == 0 grad_norm = None if do_step: grad_norm = float(torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)) optim.step() optim.zero_grad(set_to_none=True) opt_step += 1 raw_loss = float(task.detach()) tokens += tpm micro_step += 1 lr_now = optim.param_groups[0]["lr"] if raw_loss < best_train: best_train = raw_loss metrics = { "train/loss": raw_loss, "train/lr": lr_now, "train/tokens": tokens, "moe/aux": float(aux.detach()), "moe/z": float(z_loss.detach()), } if grad_norm is not None: metrics["train/grad_norm"] = grad_norm elapsed = time.perf_counter() - t0 if elapsed > 0: metrics["train/tok_s"] = (tokens - tokens_at_t0) / elapsed log_now = ( micro_step % args.eval_every == 0 or micro_step == 1 or (args.max_tokens is not None and tokens >= args.max_tokens) or (args.max_tokens is None and micro_step >= args.steps) ) if log_now: held = _heldout_loss(model, held_chunks, device, use_bf16) if held is not None: metrics["heldout/loss"] = held metrics.update(_moe_log(model)) print( f"micro {micro_step:6d} opt {opt_step:6d} tok {tokens:,} " f"loss {raw_loss:.4f} lr {lr_now:.2e}" + (f" held {held:.4f}" if held is not None else "") + f" aux {metrics['moe/aux']:.4f} z {metrics['moe/z']:.4f}" ) if micro_step % (args.eval_every * 2) == 0 or micro_step <= args.eval_every: for prefix in args.gen_prefix: sample = gen_sample(prefix) print(f" gen[{prefix[:16]}]: {sample}") if tracker is not None: import swanlab tracker.log( {f"gen/{prefix[:24]}": swanlab.Text(sample)}, step=micro_step, ) payload = _payload( cfg, model, optim, args, micro_step=micro_step, opt_step=opt_step, tokens=tokens, chunk_index=chunk_index + 1, best_heldout=best_heldout, ) _save(_sibling(args.out, "_last"), payload) if held is not None and held < best_heldout: best_heldout = held payload["best_heldout"] = best_heldout _save(_sibling(args.out, "_best"), payload) print( f" best held-out {best_heldout:.4f} -> {_sibling(args.out, '_best')}" ) del payload if device == "cuda": torch.cuda.empty_cache() if tracker is not None: tracker.log(metrics, step=micro_step) payload = _payload( cfg, model, optim, args, micro_step=micro_step, opt_step=opt_step, tokens=tokens, chunk_index=chunk_index + 1, best_heldout=best_heldout, ) _save(args.out, payload) print( f"best train {best_train:.4f} best held-out {best_heldout:.4f}; " f"tokens {tokens:,} -> {args.out}" ) if tracker is not None: tracker.finish() if __name__ == "__main__": main()