gollem-v5-ckpts / ckpt_autopush.py
Maggio33's picture
Upload ckpt_autopush.py with huggingface_hub
2146c52 verified
Raw History Blame
3.09 kB
#!/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 @ <step>`
lines (emitted right after the save completes), copies ckpt.pt to a step-tagged
name, and uploads it to SlayerLab/gollem-v5-ckpts/<subdir>/.
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)