File size: 5,651 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 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 | """Wait for all workers to upload, gather, combine, and publish the final dataset.
Aggregates stats with the FULL launcher.py schema (`summary`, `checks`,
`failure_breakdown`, `failure_breakdown_examples`, `num_shards`) so the published
`stats.json` matches the original `sdg/launcher.py` output structure.
Env vars:
HF_USER — HuggingFace username (default: fzzhang)
RUN_NAME — experiment name (default: math53K run name)
NUM_WORKERS — total worker count (default: 4)
NUM_SHARDS — total shards (default: 32)
"""
import json
import os
import tempfile
import time
from pathlib import Path
from huggingface_hub import HfApi, snapshot_download
HF_USER = os.environ.get("HF_USER", "fzzhang")
RUN_NAME = os.environ.get(
"RUN_NAME", "qwen3_4b_openthoughts3_math53K_instill_n8_valredundancy5_round1"
)
NUM_WORKERS = int(os.environ.get("NUM_WORKERS", "4"))
NUM_SHARDS = int(os.environ.get("NUM_SHARDS", "32"))
STAGING_REPO = f"{HF_USER}/{RUN_NAME}-staging"
FINAL_REPO = f"{HF_USER}/{RUN_NAME}"
api = HfApi()
print(f"waiting for {NUM_WORKERS} workers to upload...", flush=True)
while True:
try:
files = api.list_repo_files(STAGING_REPO, repo_type="dataset")
except Exception as e:
print(f" list error (staging repo may not exist yet): {e}", flush=True)
files = []
done = [
f for f in files
if f.startswith("workers/worker_") and f.endswith(".done")
]
print(
f" [{time.strftime('%H:%M:%S')}] {len(done)}/{NUM_WORKERS} workers done",
flush=True,
)
if len(done) >= NUM_WORKERS:
break
time.sleep(120)
print("downloading staging repo...", flush=True)
local = snapshot_download(
repo_id=STAGING_REPO, repo_type="dataset", local_dir="./gathered"
)
shards_dir = Path(local) / "shards"
present = sorted(shards_dir.glob("shard_*"))
print(f" {len(present)} shard dirs present", flush=True)
# Concatenate output.jsonl files in shard order
combined = Path("combined_output.jsonl")
total = 0
with open(combined, "w") as out:
for sid in range(NUM_SHARDS):
f = shards_dir / f"shard_{sid:03d}" / "output.jsonl"
if not f.exists():
print(f" WARN: missing {f}")
continue
with open(f) as src:
for line in src:
out.write(line)
total += 1
print(f"combined {total} rows", flush=True)
# Aggregate stats with the FULL launcher.py schema
agg = {
"summary": {
"total_examples": 0,
"total_passed": 0,
"total_passed_samples": 0,
"total_failed": 0,
"pass_rate_pct": 0.0,
},
"checks": {
"static_check": {"passed": 0, "failed": 0},
"cycle_consistency": {"passed": 0, "failed": 0},
"factual_accuracy": {"passed": 0, "failed": 0},
"correctness": {"passed": 0, "failed": 0},
},
"failure_breakdown": {
"failed_static_check": 0,
"failed_cycle_consistency": 0,
"failed_factual_accuracy": 0,
"failed_correctness": 0,
"exhausted_samples": 0,
"total_failed": 0,
},
"failure_breakdown_examples": {
"failed_static_check": 0,
"failed_cycle_consistency": 0,
"failed_factual_accuracy": 0,
"failed_correctness": 0,
"total_failed": 0,
},
"num_shards": NUM_SHARDS,
}
for sid in range(NUM_SHARDS):
sf = shards_dir / f"shard_{sid:03d}" / "stats.json"
if not sf.exists():
continue
s = json.load(open(sf))
for k in ["total_examples", "total_passed", "total_passed_samples", "total_failed"]:
agg["summary"][k] += s["summary"].get(k, 0)
for cn in agg["checks"]:
if cn in s.get("checks", {}):
for k in ["passed", "failed"]:
agg["checks"][cn][k] += s["checks"][cn][k]
for k in agg["failure_breakdown"]:
if k in s.get("failure_breakdown", {}):
agg["failure_breakdown"][k] += s["failure_breakdown"][k]
for k in agg["failure_breakdown_examples"]:
if k in s.get("failure_breakdown_examples", {}):
agg["failure_breakdown_examples"][k] += s["failure_breakdown_examples"][k]
te = agg["summary"]["total_examples"]
tp = agg["summary"]["total_passed"]
agg["summary"]["pass_rate_pct"] = round(tp / te * 100, 2) if te else 0.0
stats_path = Path("combined_stats.json")
json.dump(agg, open(stats_path, "w"), indent=2)
print(f"stats: {agg['summary']}", flush=True)
# Sanity: row count should match total_passed (first_valid policy)
expected = agg["summary"].get("total_passed_samples") or agg["summary"]["total_passed"]
if total != expected:
raise RuntimeError(
f"INTEGRITY: combined output.jsonl has {total} rows but stats expect "
f"{expected}. Shard outputs may be corrupted."
)
print(f"publishing to {FINAL_REPO}", flush=True)
api.create_repo(FINAL_REPO, repo_type="dataset", exist_ok=True)
readme = (
"---\n"
"configs:\n"
"- config_name: default\n"
" data_files:\n"
" - split: train\n"
" path: output.jsonl\n"
"---\n"
)
with tempfile.NamedTemporaryFile(mode="w", suffix=".md", delete=False) as f:
f.write(readme)
rp = f.name
api.upload_file(
path_or_fileobj=rp,
path_in_repo="README.md",
repo_id=FINAL_REPO,
repo_type="dataset",
)
api.upload_file(
path_or_fileobj=str(combined),
path_in_repo="output.jsonl",
repo_id=FINAL_REPO,
repo_type="dataset",
)
api.upload_file(
path_or_fileobj=str(stats_path),
path_in_repo="stats.json",
repo_id=FINAL_REPO,
repo_type="dataset",
)
print(f"DONE: https://huggingface.co/datasets/{FINAL_REPO}", flush=True)
|