From 53d0f4b17a24457fb86716412b151905ffa79009 Mon Sep 17 00:00:00 2001 From: dela Date: Tue, 25 Aug 2026 20:59:29 +0800 Subject: [PATCH] Cut 1B-run I/O: rarer SwanLab, ckpt, and generate Every micro-step was hitting SwanLab, and every 100 steps wrote a 5GB ckpt plus greedy decode. 0.5b now logs every 20, eval/held-out every 500, saves _last every 1000, samples every 2000. --- train_k3.py | 106 ++++++++++++++++++++++++++++++++++++++-------------- 1 file changed, 78 insertions(+), 28 deletions(-) diff --git a/train_k3.py b/train_k3.py index 3eb3143..8129f52 100644 --- a/train_k3.py +++ b/train_k3.py @@ -35,6 +35,9 @@ _TOY_TRAIN = { "warmup": 50, "grad_acc": 1, "eval_every": 100, + "log_every": 10, + "ckpt_every": 100, + "gen_every": 200, } _B500M_TRAIN = { "tokenizer": "Qwen/Qwen3-8B", @@ -46,7 +49,10 @@ _B500M_TRAIN = { "lr": 3e-4, "warmup": 64, "grad_acc": 8, - "eval_every": 100, + "eval_every": 500, + "log_every": 20, + "ckpt_every": 1000, + "gen_every": 2000, } @@ -93,6 +99,10 @@ def _init_swanlab(cfg: K3Config, args: argparse.Namespace): "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: @@ -201,6 +211,24 @@ def main() -> None: ) 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", @@ -265,6 +293,10 @@ def main() -> None: 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) @@ -344,7 +376,11 @@ def main() -> None: 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") + 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, @@ -429,13 +465,17 @@ def main() -> None: 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) + 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 log_now: + + if eval_now: held = _heldout_loss(model, held_chunks, device, use_bf16) if held is not None: metrics["heldout/loss"] = held @@ -446,17 +486,36 @@ def main() -> None: + (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 + 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, - ) + tracker.log( + {f"gen/{prefix[:24]}": swanlab.Text(sample)}, + step=micro_step, + ) + if ckpt_now: payload = _payload( cfg, model, @@ -469,17 +528,8 @@ def main() -> None: 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: + if tracker is not None and (log_now or eval_now): tracker.log(metrics, step=micro_step) payload = _payload(