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:
+9
-4
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user