Files
K3/train_k3.py
T
dela 24c9d56b72 Keep attnres on resume, fix final chunk_index, default 0.5b to Yi-6B
CLI default attnres=off was overwriting block checkpoints on resume
so later loads hit Unexpected key(s). Only apply flags the user
passed. Track next_chunk so a budget-exit save does not skip the
untrained yield. 0.5b now uses 01-ai/Yi-6B (64k); refuse resume
when the ckpt tokenizer does not match.
2026-08-26 14:23:29 +08:00

619 lines
21 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""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.swanlab_env import prepare_swanlab_env, swanlab_run_id
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,
"log_every": 10,
"ckpt_every": 100,
"gen_every": 200,
}
_B500M_TRAIN = {
"tokenizer": "01-ai/Yi-6B",
"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": 500,
"log_every": 20,
"ckpt_every": 1000,
"gen_every": 2000,
}
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, resume_id: str | None = None
):
"""Cloud monitor if SWANLAB_API_KEY is set; otherwise no-op."""
key = os.environ.get("SWANLAB_API_KEY")
if not key:
return None
project = prepare_swanlab_env()
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)
run_id = resume_id or os.environ.get("SWANLAB_RUN_ID")
init_kw = dict(
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,
"eval_every": args.eval_every,
"log_every": args.log_every,
"ckpt_every": args.ckpt_every,
"gen_every": args.gen_every,
},
)
if run_id:
init_kw["id"] = run_id
init_kw["resume"] = True
print(f"swanlab resume id={run_id}")
run = swanlab.init(**init_kw)
got = swanlab_run_id(run)
if got:
args.swanlab_id = got
print(f"swanlab run id {got}")
return run
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,
"batch": args.batch,
"seq_len": args.seq_len,
"grad_acc": args.grad_acc,
"swanlab_id": getattr(args, "swanlab_id", None),
}
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 _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:
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(
"--log-every",
type=int,
default=train_defaults["log_every"],
help="swanlab scalar period in micro-steps",
)
p.add_argument(
"--ckpt-every",
type=int,
default=train_defaults["ckpt_every"],
help="write _last/_best this many micro-steps (1B default 1000)",
)
p.add_argument(
"--gen-every",
type=int,
default=train_defaults["gen_every"],
help="sample prefixes this often; 0 disables",
)
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(
"--swanlab-id",
default=None,
help="resume this SwanLab run (URL /runs/<id>); default: id stored in ckpt",
)
p.add_argument(
"--swanlab-new",
action="store_true",
help="start a new SwanLab run even when --resume",
)
p.add_argument("--gen-prefix", action="append", default=None)
p.add_argument("--device", default="auto")
p.add_argument(
"--attnres",
default=None,
choices=["off", "full", "block"],
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",
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)"
)
if device == "cuda":
torch.backends.cuda.matmul.allow_tf32 = True
torch.backends.cudnn.allow_tf32 = True
torch.set_float32_matmul_precision("high")
print(f"loading tokenizer {args.tokenizer} ...")
tok = load_tokenizer(args.tokenizer)
cfg = K3Config.preset(args.preset)
cfg.vocab_size = tok.vocab_size
_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)
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
_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:
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))
chunk_index = int(payload.get("chunk_index", 0))
best_heldout = float(payload.get("best_heldout", best_heldout))
if not args.swanlab_new and not args.swanlab_id:
args.swanlab_id = payload.get("swanlab_id") or args.swanlab_id
else:
model = CausalLM(cfg).to(device)
_apply_moe_coefs(model, cfg)
tracker = _init_swanlab(
cfg,
args,
resume_id=None if args.swanlab_new else args.swanlab_id,
)
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 "
f"log/{args.log_every} eval/{args.eval_every} ckpt/{args.ckpt_every} "
f"gen/{args.gen_every}"
)
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}] "
f"{tpm} tok/micro"
)
if train_chunks.size(0) == 0:
raise SystemExit("no training chunks; raise --limit or lower --batch/--seq-len")
if args.resume:
old_batch = payload.get("batch")
old_seq = payload.get("seq_len")
if old_batch is not None and (
int(old_batch) != args.batch or int(old_seq or args.seq_len) != args.seq_len
):
print(
f"warning: resume pack [{old_batch}, {old_seq}] -> "
f"[{args.batch}, {args.seq_len}]; reset chunk_index 0 "
f"(tokens/opt_step kept)"
)
chunk_index = 0
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
# 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:
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
next_chunk = chunk_index + 1
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
ended = (args.max_tokens is not None and tokens >= args.max_tokens) or (
args.max_tokens is None and micro_step >= args.steps
)
log_now = micro_step % args.log_every == 0 or micro_step == 1 or ended
eval_now = micro_step % args.eval_every == 0 or micro_step == 1 or ended
ckpt_now = micro_step % args.ckpt_every == 0 or ended
gen_now = args.gen_every > 0 and (
micro_step % args.gen_every == 0 or micro_step == 1 or ended
)
if eval_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 held is not None and held < best_heldout:
best_heldout = held
payload = _payload(
cfg,
model,
optim,
args,
micro_step=micro_step,
opt_step=opt_step,
tokens=tokens,
chunk_index=next_chunk,
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 gen_now:
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,
)
if ckpt_now:
payload = _payload(
cfg,
model,
optim,
args,
micro_step=micro_step,
opt_step=opt_step,
tokens=tokens,
chunk_index=next_chunk,
best_heldout=best_heldout,
)
_save(_sibling(args.out, "_last"), payload)
del payload
if tracker is not None and (log_now or eval_now):
tracker.log(metrics, step=micro_step)
payload = _payload(
cfg,
model,
optim,
args,
micro_step=micro_step,
opt_step=opt_step,
tokens=tokens,
chunk_index=next_chunk,
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 device == "cuda":
try:
torch.cuda.synchronize()
torch.cuda.empty_cache()
except Exception:
pass
if __name__ == "__main__":
main()