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