Warn and reset chunk_index when resume changes batch or seq_len

Trying a larger micro-batch on an existing 0.5b run re-packs wiki
chunks; keep tokens/opt_step and restart the data cursor.
This commit is contained in:
dela
2026-08-25 22:00:28 +08:00
parent e7185cbf49
commit 47c72e5bb8
+16
View File
@@ -133,6 +133,9 @@ def _payload(
"tokens": tokens, "tokens": tokens,
"chunk_index": chunk_index, "chunk_index": chunk_index,
"best_heldout": best_heldout, "best_heldout": best_heldout,
"batch": args.batch,
"seq_len": args.seq_len,
"grad_acc": args.grad_acc,
} }
@@ -393,9 +396,22 @@ def main() -> None:
print( print(
f"packed tokens {n_ids:,} -> {train_chunks.size(0)} train / " f"packed tokens {n_ids:,} -> {train_chunks.size(0)} train / "
f"{held_chunks.size(0)} held-out chunks of [{args.batch}, {args.seq_len}] " f"{held_chunks.size(0)} held-out chunks of [{args.batch}, {args.seq_len}] "
f"{tpm} tok/micro"
) )
if train_chunks.size(0) == 0: if train_chunks.size(0) == 0:
raise SystemExit("no training chunks; raise --limit or lower --batch/--seq-len") raise SystemExit("no training chunks; raise --limit or lower --batch/--seq-len")
if args.resume:
old_batch = payload.get("batch")
old_seq = payload.get("seq_len")
if old_batch is not None and (
int(old_batch) != args.batch or int(old_seq or args.seq_len) != args.seq_len
):
print(
f"warning: resume pack [{old_batch}, {old_seq}] -> "
f"[{args.batch}, {args.seq_len}]; reset chunk_index 0 "
f"(tokens/opt_step kept)"
)
chunk_index = 0
optim = torch.optim.AdamW( optim = torch.optim.AdamW(
model.parameters(), model.parameters(),