svd-code / scripts /combine_and_publish.py
fzzhang's picture
Upload folder using huggingface_hub
58258b8 verified
Raw History Blame Contribute Delete
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)