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
+69 -19
View File
@@ -35,6 +35,9 @@ _TOY_TRAIN = {
"warmup": 50, "warmup": 50,
"grad_acc": 1, "grad_acc": 1,
"eval_every": 100, "eval_every": 100,
"log_every": 10,
"ckpt_every": 100,
"gen_every": 200,
} }
_B500M_TRAIN = { _B500M_TRAIN = {
"tokenizer": "Qwen/Qwen3-8B", "tokenizer": "Qwen/Qwen3-8B",
@@ -46,7 +49,10 @@ _B500M_TRAIN = {
"lr": 3e-4, "lr": 3e-4,
"warmup": 64, "warmup": 64,
"grad_acc": 8, "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, "gradient_checkpointing": cfg.gradient_checkpointing,
"moe_aux_loss_coef": cfg.moe_aux_loss_coef, "moe_aux_loss_coef": cfg.moe_aux_loss_coef,
"moe_z_loss_coef": cfg.moe_z_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: 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("--grad-acc", type=int, default=train_defaults["grad_acc"])
p.add_argument("--eval-every", type=int, default=train_defaults["eval_every"]) 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( p.add_argument(
"--langs", "--langs",
default="zh,en", default="zh,en",
@@ -265,6 +293,10 @@ def main() -> None:
raise SystemExit( raise SystemExit(
"KDA training needs bf16; this GPU does not support it (avoid V100 fp16)" "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} ...") print(f"loading tokenizer {args.tokenizer} ...")
tok = load_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" f"({args.steps * tpm:,} tokens); pass --max-tokens for a real run"
) )
else: 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( train_chunks, held_chunks, n_ids = load_pretrain_chunks(
tok, tok,
@@ -429,13 +465,17 @@ def main() -> None:
if elapsed > 0: if elapsed > 0:
metrics["train/tok_s"] = (tokens - tokens_at_t0) / elapsed metrics["train/tok_s"] = (tokens - tokens_at_t0) / elapsed
log_now = ( ended = (
micro_step % args.eval_every == 0 args.max_tokens is not None and tokens >= args.max_tokens
or micro_step == 1 ) or (args.max_tokens is None and micro_step >= args.steps)
or (args.max_tokens is not None and tokens >= args.max_tokens) log_now = micro_step % args.log_every == 0 or micro_step == 1 or ended
or (args.max_tokens is None and micro_step >= args.steps) 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) held = _heldout_loss(model, held_chunks, device, use_bf16)
if held is not None: if held is not None:
metrics["heldout/loss"] = held metrics["heldout/loss"] = held
@@ -446,7 +486,25 @@ def main() -> None:
+ (f" held {held:.4f}" if held is not None else "") + (f" held {held:.4f}" if held is not None else "")
+ f" aux {metrics['moe/aux']:.4f} z {metrics['moe/z']:.4f}" + 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: 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: for prefix in args.gen_prefix:
sample = gen_sample(prefix) sample = gen_sample(prefix)
print(f" gen[{prefix[:16]}]: {sample}") print(f" gen[{prefix[:16]}]: {sample}")
@@ -457,6 +515,7 @@ def main() -> None:
{f"gen/{prefix[:24]}": swanlab.Text(sample)}, {f"gen/{prefix[:24]}": swanlab.Text(sample)},
step=micro_step, step=micro_step,
) )
if ckpt_now:
payload = _payload( payload = _payload(
cfg, cfg,
model, model,
@@ -469,17 +528,8 @@ def main() -> None:
best_heldout=best_heldout, best_heldout=best_heldout,
) )
_save(_sibling(args.out, "_last"), payload) _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 del payload
if device == "cuda": if tracker is not None and (log_now or eval_now):
torch.cuda.empty_cache()
if tracker is not None:
tracker.log(metrics, step=micro_step) tracker.log(metrics, step=micro_step)
payload = _payload( payload = _payload(