#!/usr/bin/env python3 """Auto-upload each new training checkpoint to HF so the maintainer can bench the latest weights during training. Watches the trainer log for `ckpt @ ` lines (emitted right after the save completes), copies ckpt.pt to a step-tagged name, and uploads it to SlayerLab/gollem-v5-ckpts//. Usage: python ckpt_autopush.py TRAIN_LOG CKPT_PATH HF_SUBDIR e.g. python ckpt_autopush.py train_64m_v2.log run_64m_v2/ckpt.pt v2_muon """ import os import shutil import sys import time import torch from huggingface_hub import HfApi logf = sys.argv[1] ckpt_path = sys.argv[2] subdir = sys.argv[3] repo = "SlayerLab/gollem-v5-ckpts" api = HfApi() seen = set() # Seed already-uploaded tags so a restart does not re-push existing ckpts. try: existing = {f.split("/")[-1] for f in api.list_repo_files(repo) if f.startswith(subdir + "/")} except Exception: existing = set() print(f"[ckpt_autopush] watching {logf} -> {repo}/{subdir}/ (existing={len(existing)})", flush=True) while True: try: steps = [ln.split("@", 1)[1].strip() for ln in open(logf, encoding="utf-8", errors="ignore") if "ckpt @" in ln] except FileNotFoundError: time.sleep(30) continue for s in steps: if not s.isdigit() or s in seen: continue seen.add(s) tag = f"ckpt_{int(s) // 1000}k.pt" if tag in existing: continue if not os.path.exists(ckpt_path): continue local = os.path.join(os.path.dirname(ckpt_path) or ".", tag) try: shutil.copyfile(ckpt_path, local) # GUARD (migration-race fix): ckpt.pt may have advanced past the log-step before # this copy ran (root cause of v1_muon 120k/200k/240k/280k all = step-320000). # Trust the copied file's actual step, not the log line -> label ALWAYS = content. try: actual = int(torch.load(local, map_location="cpu")["step"]) ctag = f"ckpt_{actual // 1000}k.pt" if ctag != tag: print(f"[ckpt_autopush] WARN log-step {s} != ckpt.step {actual}; " f"relabel {tag} -> {ctag}", flush=True) clocal = os.path.join(os.path.dirname(ckpt_path) or ".", ctag) os.replace(local, clocal) local, tag = clocal, ctag if tag in existing: continue seen.add(str(actual)) except Exception as e: print(f"[ckpt_autopush] step-verify failed ({e!r}); using log tag {tag}", flush=True) api.upload_file(path_or_fileobj=local, path_in_repo=f"{subdir}/{tag}", repo_id=repo, commit_message=f"auto ckpt {tag}") existing.add(tag) print(f"[ckpt_autopush] uploaded {subdir}/{tag}", flush=True) except Exception as e: print(f"[ckpt_autopush] upload {tag} failed: {e!r}", flush=True) seen.discard(s) # retry next cycle time.sleep(60)