File size: 3,094 Bytes
1a298ca
 
 
 
 
 
 
 
 
 
 
 
 
 
2146c52
1a298ca
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
2146c52
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1a298ca
2146c52
1a298ca
2146c52
1a298ca
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
#!/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)