| |
| """Sidecar: copy logs to scratch and push heartbeat/checkpoints to Hugging Face. |
| |
| Does not train. Safe to run in a loop next to the job. Large Orbax dirs are |
| tarred, then uploaded as a single file so `latest.tar` can be overwritten. |
| """ |
| import argparse |
| import json |
| import os |
| import socket |
| import subprocess |
| import sys |
| import time |
| import traceback |
|
|
|
|
| def _load_json(path): |
| try: |
| with open(path) as f: |
| return json.load(f) |
| except Exception: |
| return {} |
|
|
|
|
| def _write_json(path, payload): |
| tmp = path + ".tmp" |
| with open(tmp, "w") as f: |
| json.dump(payload, f, indent=2, sort_keys=True) |
| f.write("\n") |
| os.replace(tmp, path) |
|
|
|
|
| def _latest_ckpt_dir(workdir): |
| best = None |
| best_step = -1 |
| try: |
| names = os.listdir(workdir) |
| except FileNotFoundError: |
| return None, -1 |
| for name in names: |
| if not name.startswith("checkpoint_"): |
| continue |
| if name.endswith(".orbax-checkpoint-tmp") or ".orbax-checkpoint-tmp" in name: |
| continue |
| path = os.path.join(workdir, name) |
| if not os.path.isdir(path): |
| continue |
| try: |
| step = int(name.split("_", 1)[1]) |
| except ValueError: |
| continue |
| if step > best_step: |
| best_step = step |
| best = path |
| return best, best_step |
|
|
|
|
| def _stage_ckpt_dirs(stage_dir): |
| out = [] |
| if not os.path.isdir(stage_dir): |
| return out |
| for name in sorted(os.listdir(stage_dir)): |
| path = os.path.join(stage_dir, name) |
| if os.path.isdir(path) and not name.endswith(".tmp"): |
| out.append(path) |
| return out |
|
|
|
|
| def _tar(src_dir, dest_tar): |
| parent = os.path.dirname(os.path.abspath(src_dir)) |
| base = os.path.basename(src_dir.rstrip("/")) |
| |
| |
| os.makedirs(os.path.dirname(os.path.abspath(dest_tar)), exist_ok=True) |
| tmp = dest_tar + ".tmp" |
| proc = subprocess.run( |
| ["tar", "-C", parent, "-cf", tmp, base], |
| stdout=subprocess.DEVNULL, |
| stderr=subprocess.PIPE, |
| ) |
| if proc.returncode != 0: |
| err = proc.stderr.decode("utf-8", "replace").strip() |
| raise RuntimeError(f"tar {base} failed ({proc.returncode}): {err}") |
| os.replace(tmp, dest_tar) |
| return dest_tar |
|
|
|
|
| def _hf_api(token): |
| try: |
| from huggingface_hub import HfApi |
| except Exception as exc: |
| print(f"[hf_sync] huggingface_hub unavailable ({exc}); " |
| f"will stage files for the login-node pusher", flush=True) |
| return None |
| return HfApi(token=token) |
|
|
|
|
| def _rsync_to_login(local, remote_dir): |
| """Copy a file to gandalf so the login-node pusher can upload it.""" |
| if not remote_dir or not os.path.isfile(local): |
| return False |
| dest = remote_dir.rstrip("/") + "/" |
| cmd = ["rsync", "-a", "-e", |
| "ssh -o IdentitiesOnly=yes -o StrictHostKeyChecking=accept-new", |
| local, dest] |
| try: |
| subprocess.check_call(cmd, stdout=subprocess.DEVNULL, |
| stderr=subprocess.DEVNULL) |
| print(f"[hf_sync] staged {os.path.basename(local)} -> {dest}", flush=True) |
| return True |
| except Exception as exc: |
| print(f"[hf_sync] rsync to login failed: {exc}", flush=True) |
| return False |
|
|
|
|
| def _upload_file(api, local, repo, dest, token): |
| if api is None: |
| print(f"[hf_sync] skip upload {dest}: huggingface_hub missing", flush=True) |
| return False |
| api.upload_file( |
| path_or_fileobj=local, |
| path_in_repo=dest, |
| repo_id=repo, |
| repo_type="model", |
| token=token, |
| commit_message=f"sync {dest}", |
| ) |
| print(f"[hf_sync] uploaded {dest} ({os.path.getsize(local)} bytes)", flush=True) |
| return True |
|
|
|
|
| def sync_once(args, state): |
| workdir = os.path.abspath(args.workdir) |
| scratch_log = args.scratch_log |
| if args.log and os.path.isfile(args.log) and scratch_log: |
| os.makedirs(os.path.dirname(scratch_log), exist_ok=True) |
| subprocess.call(["rsync", "-a", args.log, scratch_log]) |
|
|
| hb_src = os.path.join(workdir, "heartbeat.json") |
| if os.path.isfile(hb_src) and scratch_log: |
| hb_scratch = os.path.join(os.path.dirname(scratch_log), "w12_inst_latent_heartbeat.json") |
| subprocess.call(["rsync", "-a", hb_src, hb_scratch]) |
|
|
| token = (os.environ.get("HF_TOKEN") |
| or os.environ.get("HUGGING_FACE_HUB_TOKEN") |
| or "") |
| if args.token_file and os.path.isfile(args.token_file): |
| with open(args.token_file) as f: |
| token = f.read().strip() or token |
| if not token: |
| print("[hf_sync] no HF token; scratch copy only", flush=True) |
| return state |
|
|
| api = _hf_api(token) |
| prefix = args.prefix.strip("/") |
| repo = args.repo |
|
|
| if os.path.isfile(hb_src): |
| mtime = os.path.getmtime(hb_src) |
| if mtime > float(state.get("heartbeat_mtime", 0)): |
| if _upload_file(api, hb_src, repo, f"{prefix}/heartbeat.json", token): |
| state["heartbeat_mtime"] = mtime |
| _rsync_to_login(hb_src, args.login_stage) |
|
|
| if args.log and os.path.isfile(args.log): |
| mtime = os.path.getmtime(args.log) |
| size = os.path.getsize(args.log) |
| last = state.get("log", {}) |
| due = (time.time() - float(last.get("uploaded_at", 0))) >= 300 |
| if due and (mtime > float(last.get("mtime", 0)) or size != last.get("size", -1)): |
| if _upload_file(api, args.log, repo, f"{prefix}/train.log", token): |
| state["log"] = {"mtime": mtime, "size": size, |
| "uploaded_at": time.time()} |
|
|
| ready = _load_json(os.path.join(workdir, "ckpt_ready.json")) |
| ckpt_dir, ckpt_step = _latest_ckpt_dir(workdir) |
| if ckpt_dir and ckpt_step > int(state.get("latest_step", -1)): |
| tar_path = os.path.join(args.tar_dir, "latest.tar") |
| print(f"[hf_sync] tarring {ckpt_dir} -> {tar_path}", flush=True) |
| _tar(ckpt_dir, tar_path) |
| uploaded = _upload_file(api, tar_path, repo, f"{prefix}/latest.tar", token) |
| staged = _rsync_to_login(tar_path, args.login_stage) |
| if uploaded or staged: |
| state["latest_step"] = ckpt_step |
| meta = { |
| "step": ckpt_step, |
| "event": ready.get("event", "periodic"), |
| "stage": ready.get("stage"), |
| "host": socket.gethostname(), |
| "job_id": os.environ.get("SLURM_JOB_ID", ""), |
| "src": ckpt_dir, |
| } |
| meta_path = os.path.join(args.tar_dir, "latest.json") |
| _write_json(meta_path, meta) |
| _upload_file(api, meta_path, repo, f"{prefix}/latest.json", token) |
| _rsync_to_login(meta_path, args.login_stage) |
| if uploaded: |
| state["latest_step"] = ckpt_step |
| keep = bool(ready.get("keep_snapshot")) or (ckpt_step > 0 and ckpt_step % 50000 == 0) |
| if keep and ckpt_step not in set(state.get("snapshots", [])): |
| dest = f"{prefix}/steps/step_{ckpt_step}.tar" |
| snap = os.path.join(args.tar_dir, f"step_{ckpt_step}.tar") |
| if snap != tar_path: |
| subprocess.call(["/bin/cp", "-f", tar_path, snap]) |
| if _upload_file(api, tar_path, repo, dest, token): |
| state.setdefault("snapshots", []).append(ckpt_step) |
| _rsync_to_login(snap if os.path.isfile(snap) else tar_path, args.login_stage) |
|
|
| stage_dir = os.path.join(workdir, "stage_ckpts") |
| uploaded_stages = set(state.get("stages", [])) |
| for path in _stage_ckpt_dirs(stage_dir): |
| name = os.path.basename(path) |
| if name in uploaded_stages: |
| continue |
| tar_path = os.path.join(args.tar_dir, f"{name}.tar") |
| print(f"[hf_sync] tarring stage {path}", flush=True) |
| _tar(path, tar_path) |
| if _upload_file(api, tar_path, repo, f"{prefix}/stages/{name}.tar", token): |
| uploaded_stages.add(name) |
| _rsync_to_login(tar_path, args.login_stage) |
| state["stages"] = sorted(uploaded_stages) |
| return state |
|
|
|
|
| def main(): |
| ap = argparse.ArgumentParser() |
| ap.add_argument("--workdir", required=True) |
| ap.add_argument("--log", default="") |
| ap.add_argument("--scratch-log", default="") |
| ap.add_argument("--repo", default="Avra98/Sudoku_superposition") |
| ap.add_argument("--prefix", default="runs/w12_inst_latent") |
| ap.add_argument("--token-file", default="") |
| ap.add_argument("--state", default="") |
| ap.add_argument("--tar-dir", default="") |
| ap.add_argument("--interval", type=int, default=60) |
| ap.add_argument("--login-stage", default="gandalf.berkeley.edu:/tmp/sudoku_hf_uploads") |
| ap.add_argument("--once", action="store_true") |
| args = ap.parse_args() |
| args.tar_dir = args.tar_dir or os.path.join(os.path.abspath(args.workdir), "_hf_tars") |
| os.makedirs(args.tar_dir, exist_ok=True) |
| state_path = args.state or os.path.join(args.tar_dir, "hf_sync_state.json") |
| state = _load_json(state_path) |
| print(f"[hf_sync] host={socket.gethostname()} workdir={args.workdir}", |
| flush=True) |
| while True: |
| try: |
| state = sync_once(args, state) |
| _write_json(state_path, state) |
| except Exception: |
| traceback.print_exc() |
| print("[hf_sync] cycle failed; will retry", flush=True) |
| if args.once: |
| break |
| time.sleep(max(5, args.interval)) |
| return 0 |
|
|
|
|
| if __name__ == "__main__": |
| sys.exit(main()) |
|
|