OpenTransformer commited on
Commit
6b17c0a
Β·
verified Β·
1 Parent(s): bdcf028

Restore original n.py (was overwritten by delta patch upload)

Browse files
Files changed (1) hide show
  1. n.py +2 -116
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, threading, hashlib
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.resume_delta and not args.fresh:
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")