Files
K3/train_k3.py
T
dela 47c72e5bb8 Warn and reset chunk_index when resume changes batch or seq_len
Trying a larger micro-batch on an existing 0.5b run re-packs wiki
chunks; keep tokens/opt_step and restart the data cursor.
2026-08-25 22:00:28 +08:00

573 lines
19 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.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": "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": 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):
"""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,
"eval_every": args.eval_every,
"log_every": args.log_every,
"ckpt_every": args.ckpt_every,
"gen_every": args.gen_every,
},
)
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,
}
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(
"--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("--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)"
)
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
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 "
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
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
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=chunk_index + 1,
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=chunk_index + 1,
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=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()