Initial K3 snapshot: 0.5B KDA/MLA/MoE train path

Standalone tree split from LLMRL/projects/kda. Includes Triton dt_bias
backward fix, train_k3 --preset 0.5b, SFT, Docker runtime, and tests.
This commit is contained in:
dela
2026-08-25 14:43:17 +08:00
commit 584f7e9e73
140 changed files with 21592 additions and 0 deletions
+461
View File
@@ -0,0 +1,461 @@
"""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 moe_route_frac
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,
},
)
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 _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",
)
args = p.parse_args()
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
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
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)
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}"
)
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 and tokens >= args.max_tokens:
break
if 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):
loss = model(x, labels=y) / 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 = loss.item() * args.grad_acc
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}
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 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 "")
)
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')}"
)
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()