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()
|