microduck_rl/scripts/hf/uploader.py
Upstream Snapshot 47372443ff Import upstream snapshot d424a0c899f6b33cbd3daeb279913134349c0b63
Upstream: https://github.com/pollen-robotics/microduck_rl
Upstream-Commit: d424a0c899f6b33cbd3daeb279913134349c0b63
Upstream-Branch: develop
2026-08-28 15:41:56 +08:00

76 lines
2.4 KiB
Python

"""Checkpoint uploader run inside an HF Job.
Watches `logs/rsl_rl/**/model_*.pt` and uploads new/updated files to the
target HF Model repo. Designed to be `nohup uv run`-launched from the job
bootstrap, with auth coming from the HF_TOKEN secret injected by `hf jobs run`.
"""
from __future__ import annotations
import os
import sys
import time
from pathlib import Path
from huggingface_hub import HfApi, CommitOperationAdd
def main() -> int:
repo_id = os.environ.get("CKPT_REPO")
if not repo_id:
print("[uploader] CKPT_REPO not set, exiting", flush=True)
return 1
poll_interval = float(os.environ.get("CKPT_POLL_INTERVAL", "60"))
root = Path(os.environ.get("CKPT_ROOT", "logs/rsl_rl"))
one_shot = os.environ.get("CKPT_ONE_SHOT") == "1"
api = HfApi()
api.create_repo(repo_id, repo_type="model", private=True, exist_ok=True)
mode = "one-shot" if one_shot else f"every {poll_interval}s"
print(f"[uploader] watching {root} -> {repo_id} ({mode})", flush=True)
sent: dict[Path, float] = {}
while True:
try:
files = list(root.glob("**/model_*.pt"))
# also pick up the dumped configs once
files += [p for p in root.glob("**/params/*.yaml")]
files += [p for p in root.glob("**/params/*.json")]
to_upload: list[CommitOperationAdd] = []
for f in files:
try:
mtime = f.stat().st_mtime
except FileNotFoundError:
continue
if sent.get(f) == mtime:
continue
# use path-in-repo relative to logs/rsl_rl so the repo mirrors run dirs
rel = f.relative_to(root)
to_upload.append(
CommitOperationAdd(path_in_repo=str(rel), path_or_fileobj=str(f))
)
sent[f] = mtime
if to_upload:
msg = f"upload {len(to_upload)} file(s)"
api.create_commit(
repo_id=repo_id,
repo_type="model",
operations=to_upload,
commit_message=msg,
)
print(f"[uploader] pushed {len(to_upload)} file(s)", flush=True)
except Exception as e:
print(f"[uploader] error: {e}", flush=True)
if one_shot:
return 0
time.sleep(poll_interval)
if __name__ == "__main__":
sys.exit(main())