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:
+78
-28
@@ -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(
|
||||
|
||||
Reference in New Issue
Block a user