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
+17 -1
View File
@@ -133,6 +133,9 @@ def _payload(
"tokens": tokens,
"chunk_index": chunk_index,
"best_heldout": best_heldout,
"batch": args.batch,
"seq_len": args.seq_len,
"grad_acc": args.grad_acc,
}
@@ -392,10 +395,23 @@ def main() -> None:
)
print(
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:
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(
model.parameters(),