"""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)