Upstream: https://github.com/pollen-robotics/microduck_rl Upstream-Commit: d424a0c899f6b33cbd3daeb279913134349c0b63 Upstream-Branch: develop
76 lines
2.4 KiB
Python
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())
|