Fit 0.5b training on 32GB: SDPA MLA, block checkpoint, chunked CE

Whole-mixer checkpoint plus T×T MLA scores OOM'd a 31GB GPU on backward.
Checkpoint each AttnRes block, run absorbed MLA through SDPA, and compute
CE in vocab chunks so [B,T,V] logits are never materialized.

--max-tokens is now the training budget; default --steps 2000 no longer
caps a 1B-token run at 250 optimizer steps.
This commit is contained in:
dela
2026-08-25 20:09:27 +08:00
parent 7a12f61de1
commit 49aede9cb2
7 changed files with 117 additions and 36 deletions
+9 -4
View File
@@ -247,6 +247,7 @@ def main() -> None:
help="router z-loss weight (default 0.001; 0 disables)",
)
args = p.parse_args()
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
if args.gen_prefix is None:
args.gen_prefix = ["人工智能的发展", "The history of computing"]
@@ -387,9 +388,10 @@ def main() -> None:
t0 = time.perf_counter()
tokens_at_t0 = tokens
for chunk_index, x, y in iter_indexed(train_chunks, start=chunk_index):
if args.max_tokens is not None and tokens >= args.max_tokens:
break
if micro_step >= args.steps:
if args.max_tokens is not None:
if tokens >= args.max_tokens:
break
elif micro_step >= args.steps:
break
x, y = x.to(device), y.to(device)
scale = lr_scale(opt_step, args.warmup, horizon)
@@ -431,7 +433,7 @@ def main() -> None:
micro_step % args.eval_every == 0
or micro_step == 1
or (args.max_tokens is not None and tokens >= args.max_tokens)
or micro_step >= args.steps
or (args.max_tokens is None and micro_step >= args.steps)
)
if log_now:
held = _heldout_loss(model, held_chunks, device, use_bf16)
@@ -474,6 +476,9 @@ def main() -> None:
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:
tracker.log(metrics, step=micro_step)