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.
This commit is contained in:
dela
2026-08-25 20:59:29 +08:00
parent 8442f92c58
commit 53d0f4b17a
+78 -28
View File
@@ -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(