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:
+16
@@ -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,
|
||||
}
|
||||
|
||||
|
||||
@@ -393,9 +396,22 @@ 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"{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(),
|
||||
|
||||
Reference in New Issue
Block a user