"""Wait for this worker's 8 shards to complete, then upload to the staging HF repo. Env vars: HF_USER — HuggingFace username (default: fzzhang) RUN_NAME — experiment name (default: math53K run name) Polls every 2 minutes until all 8 stats.json files appear, then uploads each shard folder to `/-staging` and writes a `workers/worker_N.done` marker that the combiner watches for. """ import os import re import time from pathlib import Path from huggingface_hub import HfApi HF_USER = os.environ.get("HF_USER", "fzzhang") RUN_NAME = os.environ.get( "RUN_NAME", "qwen3_4b_openthoughts3_math53K_instill_n8_valredundancy5_round1" ) GPUS_PER_WORKER = 8 SHARDS_DIR = Path(f"output/{RUN_NAME}/shards") STAGING_REPO = f"{HF_USER}/{RUN_NAME}-staging" m = re.search(r"worker-(\d+)$", os.environ.get("MY_POD_NAME", "")) if not m: raise SystemExit("cannot infer worker idx from MY_POD_NAME") widx = int(m.group(1)) my_shards = list(range(widx * GPUS_PER_WORKER, (widx + 1) * GPUS_PER_WORKER)) print(f"worker {widx} owns shards {my_shards}", flush=True) def done(sid: int) -> bool: return (SHARDS_DIR / f"shard_{sid:03d}" / "stats.json").exists() while True: d = [s for s in my_shards if done(s)] print(f"[wait] {len(d)}/{len(my_shards)} shards complete", flush=True) if len(d) == len(my_shards): break time.sleep(120) api = HfApi() api.create_repo(STAGING_REPO, repo_type="dataset", exist_ok=True, private=True) print(f"uploading to {STAGING_REPO}", flush=True) for sid in my_shards: sp = SHARDS_DIR / f"shard_{sid:03d}" api.upload_folder( folder_path=str(sp), path_in_repo=f"shards/shard_{sid:03d}", repo_id=STAGING_REPO, repo_type="dataset", commit_message=f"upload shard {sid:03d}", ) print(f" shard_{sid:03d} done", flush=True) api.upload_file( path_or_fileobj=f"worker {widx} done\n".encode(), path_in_repo=f"workers/worker_{widx}.done", repo_id=STAGING_REPO, repo_type="dataset", ) print(f"worker {widx} fully uploaded", flush=True)