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)