File size: 2,082 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 | """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 `<HF_USER>/<RUN_NAME>-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)
|