File size: 4,424 Bytes
58258b8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
#!/usr/bin/env python
"""Watch a training run's hf/ dir and upload each new step-N export to a private HF
repo, so ephemeral-disk checkpoints survive a machine kill and are eval-able anywhere.

train_sft_qwen3.py writes a bf16 hf/step-N/ every steps_per_hf_export. This uploads each
one (once it has finished writing) to <repo> under path_in_repo=step-N/. One repo per
experiment; steps are subfolders. Idempotent via a local .pushed_steps.json marker, so a
restart of this watcher (or of the whole run) re-uploads nothing.

Run in the background alongside training (run_sft.sh does this for you):
  python push_exports.py --out <rundir> --repo fzzhang/<exp>_sft --until-step 2000

HF auth: uses your `huggingface-cli login` token (~/.cache/huggingface/token).
"""
from __future__ import annotations

import argparse
import json
import time
from pathlib import Path


def step_dirs(hf_dir: Path) -> dict[int, Path]:
    out: dict[int, Path] = {}
    if not hf_dir.is_dir():
        return out
    for d in hf_dir.glob("step-*"):
        if d.is_dir():
            try:
                out[int(d.name.split("-", 1)[1])] = d
            except (IndexError, ValueError):
                pass
    return out


def is_stable(d: Path, stable_sec: float) -> bool:
    """True if no file in the dir was modified within the last stable_sec (write finished)."""
    files = [p for p in d.rglob("*") if p.is_file()]
    if not files:
        return False
    return (time.time() - max(p.stat().st_mtime for p in files)) > stable_sec


def main() -> None:
    ap = argparse.ArgumentParser(description="Upload hf/step-N exports to HF as they land")
    ap.add_argument("--out", required=True, help="training run dir (contains hf/)")
    ap.add_argument("--repo", required=True, help="target HF repo id, e.g. fzzhang/<exp>_sft")
    ap.add_argument("--poll", type=int, default=120, help="seconds between scans")
    ap.add_argument("--stable-sec", type=float, default=60.0,
                    help="a step dir must be idle this long before upload (guards partial writes)")
    ap.add_argument("--every", type=int, default=1,
                    help="upload every Nth export (exports are every 100 steps; 1=all). "
                         "The --until-step export is always uploaded regardless.")
    ap.add_argument("--until-step", type=int, default=0,
                    help="exit after this step is pushed (0 = run forever / until killed)")
    ap.add_argument("--once", action="store_true",
                    help="scan once, push all unpushed steps (ignore --stable-sec), then exit")
    args = ap.parse_args()

    from huggingface_hub import HfApi

    api = HfApi()
    api.create_repo(args.repo, repo_type="model", private=True, exist_ok=True)

    hf_dir = Path(args.out) / "hf"
    marker = Path(args.out) / ".pushed_steps.json"
    pushed = set(json.loads(marker.read_text())) if marker.exists() else set()
    print(f"[push] {hf_dir} -> {args.repo} (private); already pushed: {sorted(pushed)}", flush=True)

    while True:
        for step in sorted(step_dirs(hf_dir)):
            if step in pushed:
                continue
            if args.every > 1 and (step // 100) % args.every != 0 and step != args.until_step:
                continue
            d = step_dirs(hf_dir)[step]
            if not args.once and not is_stable(d, args.stable_sec):
                continue  # still being written; try next poll
            print(f"[push] uploading step-{step} ...", flush=True)
            try:
                api.upload_folder(folder_path=str(d), repo_id=args.repo, repo_type="model",
                                  path_in_repo=f"step-{step}", commit_message=f"add step-{step}")
                pushed.add(step)
                marker.write_text(json.dumps(sorted(pushed)))
                print(f"[push] DONE step-{step} -> {args.repo}/step-{step}", flush=True)
            except Exception as e:  # network / rate-limit / partial — retry next poll
                print(f"[push] FAILED step-{step}: {type(e).__name__}: {e} (retry next poll)", flush=True)
        if args.once:
            print("[push] --once sweep complete; exiting.", flush=True)
            return
        if args.until_step and args.until_step in pushed:
            print(f"[push] final step-{args.until_step} pushed; exiting.", flush=True)
            return
        time.sleep(args.poll)


if __name__ == "__main__":
    main()