Restore original n.py (was overwritten by delta patch upload)
Browse files
n.py
CHANGED
|
@@ -4,8 +4,7 @@
|
|
| 4 |
# Enhanced inference: checkpoint name, tok/s, UK time
|
| 5 |
|
| 6 |
from __future__ import annotations
|
| 7 |
-
import argparse, json, math, pathlib, random, time, os, sys
|
| 8 |
-
from pathlib import Path
|
| 9 |
from contextlib import nullcontext
|
| 10 |
from typing import Dict, Any, List, Optional, Tuple
|
| 11 |
from datetime import datetime, timezone
|
|
@@ -136,8 +135,6 @@ SAT_BLOCK = 2
|
|
| 136 |
LR_CORE, LR_HEAD = 5e-5, 2e-4
|
| 137 |
EMIT_LAMBDA = 0.1
|
| 138 |
DEFAULT_SAVE_SEC = 24 * 3600
|
| 139 |
-
DEFAULT_DELTA_STEPS = 500 # lightweight weight-only save every N steps
|
| 140 |
-
DEFAULT_MAX_DELTAS = 5 # keep last N deltas (older pruned after full save)
|
| 141 |
CKDIR = pathlib.Path("ckpts_expansion")
|
| 142 |
|
| 143 |
DEFAULT_PRETRAIN_SOURCES = "OpenTransformer/goddess-crawl,OpenTransformer/agillm-crawl-data,OpenTransformer/web-crawl-2026,OpenTransformer/web-crawl-clean-v2,OpenTransformer/scraped-web-data,OpenTransformer/turbo-crawl,OpenTransformer/sft-data-clean,OpenTransformer/web-crawl-v1"
|
|
@@ -486,99 +483,6 @@ def sat_mask_cached(new_len: int, cached_len: int, block=SAT_BLOCK):
|
|
| 486 |
|
| 487 |
|
| 488 |
# βββββββββββββββββββββββββ Checkpoint helpers βββββββββββββββββββββββββ
|
| 489 |
-
|
| 490 |
-
# βββββββββββββββββββββββββ Delta Checkpoints (weight-only, async) βββββββββββββββββββββββββ
|
| 491 |
-
_delta_lock = threading.Lock()
|
| 492 |
-
_delta_thread: Optional[threading.Thread] = None
|
| 493 |
-
|
| 494 |
-
def _sha256_file(path: pathlib.Path) -> str:
|
| 495 |
-
"""Compute SHA256 of a file for integrity verification."""
|
| 496 |
-
h = hashlib.sha256()
|
| 497 |
-
with open(path, "rb") as f:
|
| 498 |
-
for chunk in iter(lambda: f.read(1 << 20), b""):
|
| 499 |
-
h.update(chunk)
|
| 500 |
-
return h.hexdigest()
|
| 501 |
-
|
| 502 |
-
def _do_delta_save(tensors: dict, path: pathlib.Path, meta: dict):
|
| 503 |
-
"""Background worker: write weight-only checkpoint + checksum."""
|
| 504 |
-
try:
|
| 505 |
-
path.parent.mkdir(exist_ok=True, parents=True)
|
| 506 |
-
tmp = path.with_suffix(path.suffix + ".dtmp")
|
| 507 |
-
torch.save({"weights": tensors, **meta}, tmp, _use_new_zipfile_serialization=False)
|
| 508 |
-
digest = _sha256_file(tmp)
|
| 509 |
-
tmp.replace(path)
|
| 510 |
-
# Write sidecar checksum
|
| 511 |
-
path.with_suffix(".sha256").write_text(f"{digest} {path.name}\n")
|
| 512 |
-
print(f" [delta] saved {path.name} ({digest[:12]}...)")
|
| 513 |
-
except Exception as e:
|
| 514 |
-
print(f" [delta] FAILED {path.name}: {e}")
|
| 515 |
-
|
| 516 |
-
def save_delta(core, ar_h, sat_h, step: int, seen_tok: int, save_dir: pathlib.Path, phase_name: str):
|
| 517 |
-
"""Save weight-only delta in background thread. Non-blocking."""
|
| 518 |
-
global _delta_thread
|
| 519 |
-
# Wait for any previous delta write to finish
|
| 520 |
-
if _delta_thread is not None and _delta_thread.is_alive():
|
| 521 |
-
_delta_thread.join(timeout=60)
|
| 522 |
-
# Snapshot weights to CPU (detach from GPU graph)
|
| 523 |
-
with _delta_lock:
|
| 524 |
-
tensors = {
|
| 525 |
-
"core": {k: v.detach().cpu() for k, v in core.state_dict().items()},
|
| 526 |
-
"ar": {k: v.detach().cpu() for k, v in ar_h.state_dict().items()},
|
| 527 |
-
"sat": {k: v.detach().cpu() for k, v in sat_h.state_dict().items()},
|
| 528 |
-
}
|
| 529 |
-
meta = {"step": step, "seen_tok": seen_tok, "wall_time": time.time(), "delta": True}
|
| 530 |
-
path = save_dir / f"{phase_name}_delta_step{step:08d}.pt"
|
| 531 |
-
_delta_thread = threading.Thread(target=_do_delta_save, args=(tensors, path, meta), daemon=True)
|
| 532 |
-
_delta_thread.start()
|
| 533 |
-
|
| 534 |
-
def _prune_deltas(save_dir: pathlib.Path, phase_name: str, max_deltas: int):
|
| 535 |
-
"""Keep only the most recent max_deltas delta files."""
|
| 536 |
-
if max_deltas is None or max_deltas <= 0:
|
| 537 |
-
return
|
| 538 |
-
try:
|
| 539 |
-
pattern = f"{phase_name}_delta_step*.pt"
|
| 540 |
-
deltas = sorted(
|
| 541 |
-
[p for p in save_dir.glob(pattern) if p.stat().st_size > 0],
|
| 542 |
-
key=lambda p: p.stat().st_mtime
|
| 543 |
-
)
|
| 544 |
-
excess = len(deltas) - max_deltas
|
| 545 |
-
if excess > 0:
|
| 546 |
-
for p in deltas[:excess]:
|
| 547 |
-
try:
|
| 548 |
-
p.unlink()
|
| 549 |
-
sha = p.with_suffix(".sha256")
|
| 550 |
-
if sha.exists(): sha.unlink()
|
| 551 |
-
print(f" [delta-prune] deleted {p.name}")
|
| 552 |
-
except Exception:
|
| 553 |
-
pass
|
| 554 |
-
except Exception as e:
|
| 555 |
-
print(f" [delta-prune] error: {e}")
|
| 556 |
-
|
| 557 |
-
def load_delta(path: pathlib.Path, core, ar_h, sat_h):
|
| 558 |
-
"""Load weight-only delta. Returns (step, seen_tok) or raises."""
|
| 559 |
-
# Verify checksum if sidecar exists
|
| 560 |
-
sha_path = path.with_suffix(".sha256")
|
| 561 |
-
if sha_path.exists():
|
| 562 |
-
expected = sha_path.read_text().split()[0]
|
| 563 |
-
actual = _sha256_file(path)
|
| 564 |
-
if expected != actual:
|
| 565 |
-
raise ValueError(f"Checksum mismatch for {path.name}: expected {expected[:12]}... got {actual[:12]}...")
|
| 566 |
-
print(f" [delta] checksum OK for {path.name}")
|
| 567 |
-
ck = torch.load(path, map_location="cpu", weights_only=False)
|
| 568 |
-
if not ck.get("delta"):
|
| 569 |
-
raise ValueError(f"{path.name} is not a delta checkpoint")
|
| 570 |
-
core.load_state_dict(ck["weights"]["core"])
|
| 571 |
-
ar_h.load_state_dict(ck["weights"]["ar"])
|
| 572 |
-
sat_h.load_state_dict(ck["weights"]["sat"])
|
| 573 |
-
return ck.get("step", 0), ck.get("seen_tok", 0)
|
| 574 |
-
|
| 575 |
-
def _flush_delta():
|
| 576 |
-
"""Wait for any in-flight delta save to complete."""
|
| 577 |
-
global _delta_thread
|
| 578 |
-
if _delta_thread is not None and _delta_thread.is_alive():
|
| 579 |
-
print(" [delta] flushing in-flight write...")
|
| 580 |
-
_delta_thread.join(timeout=120)
|
| 581 |
-
|
| 582 |
def save_ckpt(path: pathlib.Path, core, ar_h, sat_h, opt, scaler, meta):
|
| 583 |
path.parent.mkdir(exist_ok=True, parents=True)
|
| 584 |
tmp = path.with_suffix(path.suffix + ".tmp")
|
|
@@ -700,7 +604,6 @@ def _train_phase(
|
|
| 700 |
MAX_OOM_RETRIES = 2
|
| 701 |
now_wall = time.time()
|
| 702 |
last_save_mono = time.monotonic() - (now_wall - (resume_wall_time or now_wall))
|
| 703 |
-
last_delta_step = start_step
|
| 704 |
print(f"[{phase_name}] Starting. Goal: {total_tokens_needed:,} tokens. Batch={BATCH}, Block={BLOCK}")
|
| 705 |
print(f"[{phase_name}] AR_ONLY={args.ar_only}, TIE_WEIGHTS={tie_weights}, STREAMING={streaming}")
|
| 706 |
while seen_tok < total_tokens_needed:
|
|
@@ -774,19 +677,10 @@ def _train_phase(
|
|
| 774 |
now_mono = time.monotonic()
|
| 775 |
if now_mono - last_save_mono >= args.save_every_sec:
|
| 776 |
ck_name = f"{phase_name}_step{step:08d}.pt"
|
| 777 |
-
_flush_delta() # wait for any in-flight delta before full save
|
| 778 |
_prune_checkpoints(pathlib.Path(args.save_dir), phase_name, max_ckpts)
|
| 779 |
save_ckpt(pathlib.Path(args.save_dir) / ck_name, core, ar_h, sat_h, opt, scaler,
|
| 780 |
meta={"cfg": cfg, "step": step, "seen_tok": seen_tok, "wall_time": time.time(), "tie_weights": tie_weights})
|
| 781 |
last_save_mono = now_mono
|
| 782 |
-
# Prune old deltas after a full save (they're superseded)
|
| 783 |
-
_prune_deltas(pathlib.Path(args.save_dir), phase_name, args.delta_max_keep)
|
| 784 |
-
last_delta_step = step # reset delta counter after full save
|
| 785 |
-
# ββ Delta checkpoint (step-based, weight-only, async) ββ
|
| 786 |
-
if args.delta_every_steps > 0 and (step - last_delta_step) >= args.delta_every_steps:
|
| 787 |
-
_prune_deltas(pathlib.Path(args.save_dir), phase_name, args.delta_max_keep)
|
| 788 |
-
save_delta(core, ar_h, sat_h, step, seen_tok, pathlib.Path(args.save_dir), phase_name)
|
| 789 |
-
last_delta_step = step
|
| 790 |
if args.auto_grow:
|
| 791 |
steps_since_last_grow += 1
|
| 792 |
if steps_since_last_grow >= args.grow_every_steps:
|
|
@@ -800,7 +694,6 @@ def _train_phase(
|
|
| 800 |
except ValueError:
|
| 801 |
grow_plan = sorted(set(grow_plan + [BLOCK]))
|
| 802 |
pbar.close()
|
| 803 |
-
_flush_delta() # ensure any in-flight delta completes before final save
|
| 804 |
save_ckpt(pathlib.Path(args.save_dir) / f"{phase_name}_final.pt", core, ar_h, sat_h, opt, scaler,
|
| 805 |
meta={"cfg": cfg, "step": step, "seen_tok": seen_tok, "wall_time": time.time(), "tie_weights": tie_weights})
|
| 806 |
return step, seen_tok, time.time()
|
|
@@ -844,11 +737,7 @@ def train(args):
|
|
| 844 |
])
|
| 845 |
scaler = GradScaler(enabled=(args.amp and DEV.type == "cuda"))
|
| 846 |
start_step, seen_tok, last_wall = 0, 0, None
|
| 847 |
-
if args.
|
| 848 |
-
delta_step, delta_tok = load_delta(pathlib.Path(args.resume_delta), core, ar_h, sat_h)
|
| 849 |
-
start_step, seen_tok, last_wall = delta_step, delta_tok, None
|
| 850 |
-
print(f"Resumed from DELTA at step {start_step} (optimizer state reset β momentum rebuilds in ~100 steps)")
|
| 851 |
-
elif args.resume and not args.fresh:
|
| 852 |
start_step, seen_tok, last_wall = load_ckpt(pathlib.Path(args.resume), core, ar_h, sat_h, opt, scaler)
|
| 853 |
print(f"Resumed from step {start_step}")
|
| 854 |
# torch.compile AFTER loading checkpoint (key names differ)
|
|
@@ -1046,9 +935,6 @@ def main():
|
|
| 1046 |
tr.add_argument("--amp", action="store_true")
|
| 1047 |
tr.add_argument("--compile", action="store_true", help="Use torch.compile for speedup")
|
| 1048 |
tr.add_argument("--save_every_sec", type=int, default=DEFAULT_SAVE_SEC)
|
| 1049 |
-
tr.add_argument("--delta_every_steps", type=int, default=DEFAULT_DELTA_STEPS, help="Weight-only delta save every N steps (0=off)")
|
| 1050 |
-
tr.add_argument("--delta_max_keep", type=int, default=DEFAULT_MAX_DELTAS, help="Max delta checkpoints to keep")
|
| 1051 |
-
tr.add_argument("--resume_delta", type=str, help="Resume from a delta (weight-only, no optimizer state)")
|
| 1052 |
tr.add_argument("--save_dir", default=str(CKDIR))
|
| 1053 |
tr.add_argument("--resume", type=str)
|
| 1054 |
tr.add_argument("--x2", action="store_true")
|
|
|
|
| 4 |
# Enhanced inference: checkpoint name, tok/s, UK time
|
| 5 |
|
| 6 |
from __future__ import annotations
|
| 7 |
+
import argparse, json, math, pathlib, random, time, os, sys
|
|
|
|
| 8 |
from contextlib import nullcontext
|
| 9 |
from typing import Dict, Any, List, Optional, Tuple
|
| 10 |
from datetime import datetime, timezone
|
|
|
|
| 135 |
LR_CORE, LR_HEAD = 5e-5, 2e-4
|
| 136 |
EMIT_LAMBDA = 0.1
|
| 137 |
DEFAULT_SAVE_SEC = 24 * 3600
|
|
|
|
|
|
|
| 138 |
CKDIR = pathlib.Path("ckpts_expansion")
|
| 139 |
|
| 140 |
DEFAULT_PRETRAIN_SOURCES = "OpenTransformer/goddess-crawl,OpenTransformer/agillm-crawl-data,OpenTransformer/web-crawl-2026,OpenTransformer/web-crawl-clean-v2,OpenTransformer/scraped-web-data,OpenTransformer/turbo-crawl,OpenTransformer/sft-data-clean,OpenTransformer/web-crawl-v1"
|
|
|
|
| 483 |
|
| 484 |
|
| 485 |
# βββββββββββββββββββββββββ Checkpoint helpers βββββββββββββββββββββββββ
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 486 |
def save_ckpt(path: pathlib.Path, core, ar_h, sat_h, opt, scaler, meta):
|
| 487 |
path.parent.mkdir(exist_ok=True, parents=True)
|
| 488 |
tmp = path.with_suffix(path.suffix + ".tmp")
|
|
|
|
| 604 |
MAX_OOM_RETRIES = 2
|
| 605 |
now_wall = time.time()
|
| 606 |
last_save_mono = time.monotonic() - (now_wall - (resume_wall_time or now_wall))
|
|
|
|
| 607 |
print(f"[{phase_name}] Starting. Goal: {total_tokens_needed:,} tokens. Batch={BATCH}, Block={BLOCK}")
|
| 608 |
print(f"[{phase_name}] AR_ONLY={args.ar_only}, TIE_WEIGHTS={tie_weights}, STREAMING={streaming}")
|
| 609 |
while seen_tok < total_tokens_needed:
|
|
|
|
| 677 |
now_mono = time.monotonic()
|
| 678 |
if now_mono - last_save_mono >= args.save_every_sec:
|
| 679 |
ck_name = f"{phase_name}_step{step:08d}.pt"
|
|
|
|
| 680 |
_prune_checkpoints(pathlib.Path(args.save_dir), phase_name, max_ckpts)
|
| 681 |
save_ckpt(pathlib.Path(args.save_dir) / ck_name, core, ar_h, sat_h, opt, scaler,
|
| 682 |
meta={"cfg": cfg, "step": step, "seen_tok": seen_tok, "wall_time": time.time(), "tie_weights": tie_weights})
|
| 683 |
last_save_mono = now_mono
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 684 |
if args.auto_grow:
|
| 685 |
steps_since_last_grow += 1
|
| 686 |
if steps_since_last_grow >= args.grow_every_steps:
|
|
|
|
| 694 |
except ValueError:
|
| 695 |
grow_plan = sorted(set(grow_plan + [BLOCK]))
|
| 696 |
pbar.close()
|
|
|
|
| 697 |
save_ckpt(pathlib.Path(args.save_dir) / f"{phase_name}_final.pt", core, ar_h, sat_h, opt, scaler,
|
| 698 |
meta={"cfg": cfg, "step": step, "seen_tok": seen_tok, "wall_time": time.time(), "tie_weights": tie_weights})
|
| 699 |
return step, seen_tok, time.time()
|
|
|
|
| 737 |
])
|
| 738 |
scaler = GradScaler(enabled=(args.amp and DEV.type == "cuda"))
|
| 739 |
start_step, seen_tok, last_wall = 0, 0, None
|
| 740 |
+
if args.resume and not args.fresh:
|
|
|
|
|
|
|
|
|
|
|
|
|
| 741 |
start_step, seen_tok, last_wall = load_ckpt(pathlib.Path(args.resume), core, ar_h, sat_h, opt, scaler)
|
| 742 |
print(f"Resumed from step {start_step}")
|
| 743 |
# torch.compile AFTER loading checkpoint (key names differ)
|
|
|
|
| 935 |
tr.add_argument("--amp", action="store_true")
|
| 936 |
tr.add_argument("--compile", action="store_true", help="Use torch.compile for speedup")
|
| 937 |
tr.add_argument("--save_every_sec", type=int, default=DEFAULT_SAVE_SEC)
|
|
|
|
|
|
|
|
|
|
| 938 |
tr.add_argument("--save_dir", default=str(CKDIR))
|
| 939 |
tr.add_argument("--resume", type=str)
|
| 940 |
tr.add_argument("--x2", action="store_true")
|