Download scripts/combine_and_publish.py from fzzhang/svd-code: direct link, hf CLI and curl.
- Browser
- Download file 5.65 kB
-
https://huggingface.co/fzzhang/svd-code/resolve/main/scripts/combine_and_publish.py
- Command line
-
hf download hf://fzzhang/svd-code/scripts/combine_and_publish.py
-
curl -L -o combine_and_publish.py https://huggingface.co/fzzhang/svd-code/resolve/main/scripts/combine_and_publish.py
5.65 kB
| """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) | |