swanlab.init always opened a new experiment on --resume. Save the run id in the checkpoint and pass resume=True, id=... on the next start. --swanlab-id overrides; --swanlab-new forces a fresh experiment.
615 lines
21 KiB
Python
615 lines
21 KiB
Python
"""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 _swanlab_run_id(run) -> str | None:
|
||
for attr in ("id", "run_id"):
|
||
val = getattr(run, attr, None)
|
||
if isinstance(val, str) and val:
|
||
return val
|
||
public = getattr(run, "public", None)
|
||
if public is not None:
|
||
for attr in ("cloud_run_id", "run_id", "id"):
|
||
val = getattr(public, attr, None)
|
||
if isinstance(val, str) and val:
|
||
return val
|
||
return None
|
||
|
||
|
||
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
|
||
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"
|
||
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 _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="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))
|
||
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
|
||
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()
|