svd-code / sdg /launcher.py
fzzhang's picture
Upload folder using huggingface_hub
58258b8 verified
Raw History Blame Contribute Delete
22.1 kB
"""
Orchestrator for distributed sharded SDG pipeline.
Usage:
# Launch all shards
python -m sdg.launcher --config sdg/configs/my_experiment.yaml launch
# Dry run (print commands without executing)
python -m sdg.launcher --config sdg/configs/my_experiment.yaml launch --dry-run
# Check status of all shards
python -m sdg.launcher --config sdg/configs/my_experiment.yaml check
# Rerun failed/not-started shards
python -m sdg.launcher --config sdg/configs/my_experiment.yaml rerun
"""
from __future__ import annotations
import argparse
import json
import os
import subprocess
import sys
from datetime import datetime, timezone
from pathlib import Path
from pymongo import MongoClient
from sdg.config import SDGConfig
WORKDIR = "/nlp/scr4/nlp/crfm/text2image/text2image-rlhf/reasoning/virtual-world-data"
# =============================================================================
# Helpers
# =============================================================================
def _verify_mongo(uri: str) -> None:
"""Fail fast if MongoDB is unreachable."""
try:
client = MongoClient(uri, serverSelectionTimeoutMS=5000)
client.admin.command("ping")
print(f"MongoDB reachable at {uri}")
except Exception as e:
print(f"ERROR: Cannot reach MongoDB at {uri}: {e}", file=sys.stderr)
sys.exit(1)
def _shard_status(run_dir: Path, shard_id: int, run_name: str, num_shards: int) -> str:
"""Determine status of a single shard: completed / running / failed / not_started."""
shard_dir = run_dir / "shards" / f"shard_{shard_id:03d}"
stats_file = shard_dir / "stats.json"
log_file = run_dir / "logs" / f"shard_{shard_id:03d}.log"
if stats_file.exists():
return "completed"
# Check if job is in the queue via squeue
job_name = f"{run_name}-shard-{shard_id:03d}"
try:
result = subprocess.run(
["squeue", "-u", os.environ.get("USER", ""), "--name", job_name,
"--noheader", "-o", "%.18i %.100j %.8T"],
capture_output=True, text=True, timeout=10,
)
line = result.stdout.strip()
if line:
state = line.split()[-1] # RUNNING, PENDING, etc.
return state.lower()
except (FileNotFoundError, subprocess.TimeoutExpired):
pass # squeue not available or timed out
if log_file.exists():
return "failed"
return "not_started"
def _resolve_run_dir(args: argparse.Namespace) -> Path:
"""Resolve run directory from --config (preferred) or --run-dir."""
if args.config:
config = SDGConfig.from_yaml(args.config)
return config.output_path
return Path(args.run_dir)
def _parse_log_progress(log_file: Path) -> str:
"""Parse a shard log file and return a one-line progress summary."""
if not log_file.exists():
return ""
lines = log_file.read_text().splitlines()
stage = "starting"
detail = ""
passed_total = 0
for line in lines:
stripped = line.strip()
if "Loading model:" in stripped:
stage = "loading model"
elif "Capturing CUDA graphs" in stripped:
stage = "loading model (CUDA graphs)"
elif "Model loaded successfully" in stripped:
stage = "model loaded"
elif "Loading dataset:" in stripped:
stage = "loading dataset"
elif "Processed prompts:" in stripped:
stage = "generating (vLLM)"
detail = stripped
elif "[Generation]" in stripped:
stage = "generation"
detail = stripped
elif stripped.startswith("Round ") and "pending rows" in stripped:
stage = "validation"
detail = stripped
elif "newly passed" in stripped:
try:
parts = stripped.split()
idx = parts.index("newly")
passed_total += int(parts[idx - 1])
except (ValueError, IndexError):
pass
elif "PIPELINE COMPLETE" in stripped:
stage = "done"
if stage == "generation":
return detail
elif stage == "generating (vLLM)":
return detail
elif stage == "validation":
return f"{detail} | {passed_total} passed so far"
elif stage == "done":
return "pipeline complete"
return stage
def _load_shard_stats(run_dir: Path, shard_id: int) -> dict | None:
"""Load stats.json for a completed shard."""
stats_file = run_dir / "shards" / f"shard_{shard_id:03d}" / "stats.json"
if stats_file.exists():
with open(stats_file) as f:
return json.load(f)
return None
# =============================================================================
# --launch
# =============================================================================
def cmd_launch(args: argparse.Namespace) -> None:
config = SDGConfig.from_yaml(args.config)
num_shards = args.num_shards if args.num_shards else config.num_shards
run_name = config.run_name
if not config.mongo_uri:
print("ERROR: mongo_uri not set in config. Set it in your YAML.", file=sys.stderr)
sys.exit(1)
if not args.dry_run:
_verify_mongo(config.mongo_uri)
run_dir = config.output_path
run_dir.mkdir(parents=True, exist_ok=True)
log_dir = run_dir / "logs"
log_dir.mkdir(parents=True, exist_ok=True)
shard_commands = []
nlprun_commands = []
exclude_flag = ""
if config.machines_to_exclude:
exclude_flag = f" --exclude {','.join(config.machines_to_exclude)}"
for i in range(num_shards):
cmd = f"python -m sdg.generate --config {args.config} --shard-id {i} --num-shards {num_shards}"
job_name = f"{run_name}-shard-{i:03d}"
nlprun_cmd = (
f"nlprun -a tonyreasoningtraces -q sphinx{exclude_flag} -c 4 -g 1 --memory 100g "
f"-w {WORKDIR} --job-name {job_name} \"{cmd}\""
)
shard_commands.append(cmd)
nlprun_commands.append(nlprun_cmd)
# Save launch metadata
meta = {
"run_name": run_name,
"config_path": str(args.config),
"num_shards": num_shards,
"mongo_uri": config.mongo_uri,
"timestamp": datetime.now(timezone.utc).isoformat(),
"shard_commands": shard_commands,
"nlprun_commands": nlprun_commands,
}
meta_path = run_dir / "launch_meta.json"
with open(meta_path, "w") as f:
json.dump(meta, f, indent=2)
print(f"Wrote {meta_path}")
# Copy config
config.to_yaml(run_dir / "config.yaml")
# Execute or print
for i, nlprun_cmd in enumerate(nlprun_commands):
if args.dry_run:
print(f"[DRY RUN] {nlprun_cmd}")
else:
print(f"Launching shard {i:03d}: {nlprun_cmd}")
subprocess.run(nlprun_cmd, shell=True, check=True)
print(f"\nLaunched {num_shards} shards for run: {run_name}")
# =============================================================================
# --check
# =============================================================================
def cmd_check(args: argparse.Namespace) -> None:
run_dir = _resolve_run_dir(args)
meta_path = run_dir / "launch_meta.json"
if not meta_path.exists():
print(f"ERROR: {meta_path} not found", file=sys.stderr)
sys.exit(1)
with open(meta_path) as f:
meta = json.load(f)
run_name = meta["run_name"]
num_shards = meta["num_shards"]
# Collect all statuses first
statuses = {}
for i in range(num_shards):
statuses[i] = _shard_status(run_dir, i, run_name, num_shards)
# Show tail of logs for running shards first
running_shards = [i for i, s in statuses.items() if s == "running"]
if running_shards:
print(f"{'='*60}")
print(f"Logs for {len(running_shards)} running shard(s)")
print(f"{'='*60}")
for i in running_shards:
log_file = run_dir / "logs" / f"shard_{i:03d}.log"
print(f"\n{'─'*60}")
print(f" shard {i:03d} ({log_file})")
print(f"{'─'*60}")
if log_file.exists():
lines = log_file.read_text().splitlines()
tail = lines[-20:] if len(lines) > 20 else lines
for line in tail:
print(f" {line}")
else:
print(" (no log file yet)")
print()
# Status table
print(f"Run: {run_name} ({num_shards} shards)")
print(f"{'Shard':<6} {'Status':<12} {'Seeds':>7} {'Passed':>8} {'Failed':>8} {'Progress'}")
print("-" * 90)
for i in range(num_shards):
status = statuses[i]
stats = _load_shard_stats(run_dir, i) if status == "completed" else None
seeds = str(stats["summary"]["total_examples"]) if stats else "-"
passed = str(stats["summary"]["total_passed"]) if stats else "-"
failed = str(stats["summary"]["total_failed"]) if stats else "-"
progress = ""
if status == "running":
log_file = run_dir / "logs" / f"shard_{i:03d}.log"
progress = _parse_log_progress(log_file)
print(f"{i:03d} {status:<12} {seeds:>7} {passed:>8} {failed:>8} {progress}")
# Summary
counts = {}
for s in statuses.values():
counts[s] = counts.get(s, 0) + 1
parts = []
for label in ["completed", "running", "pending", "failed", "not_started"]:
if label in counts:
parts.append(f"{counts[label]}/{num_shards} {label}")
print(f"\nOverall: {', '.join(parts)}")
# If all completed, combine outputs and prompt for upload
force_upload = getattr(args, "force_upload", False)
completed_count = counts.get("completed", 0)
failed_count = counts.get("failed", 0)
terminal_count = completed_count + failed_count
if completed_count == num_shards:
print(f"\n{'='*60}")
print("All shards completed!")
print(f"{'='*60}")
_combine_outputs(run_dir, num_shards, meta)
elif force_upload and terminal_count == num_shards and completed_count > 0:
print(f"\n{'='*60}")
print(f"--force-upload: {completed_count}/{num_shards} completed, "
f"{failed_count}/{num_shards} failed. Uploading completed shards only.")
print(f"{'='*60}")
_combine_outputs(run_dir, num_shards, meta)
elif force_upload:
print(f"\n--force-upload requested but {num_shards - terminal_count} shard(s) "
f"are still running/pending. Skipping upload.")
def _aggregate_stats(run_dir: Path, num_shards: int) -> dict:
"""Aggregate stats from all shards into a combined dict."""
combined_stats = {
"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 i in range(num_shards):
stats = _load_shard_stats(run_dir, i)
if stats is None:
continue
for key in ["total_examples", "total_passed", "total_passed_samples", "total_failed"]:
combined_stats["summary"][key] += stats["summary"].get(key, 0)
for check_name in combined_stats["checks"]:
if check_name in stats.get("checks", {}):
for k in ["passed", "failed"]:
combined_stats["checks"][check_name][k] += stats["checks"][check_name][k]
for key in combined_stats["failure_breakdown"]:
if key in stats.get("failure_breakdown", {}):
combined_stats["failure_breakdown"][key] += stats["failure_breakdown"][key]
for key in combined_stats["failure_breakdown_examples"]:
if key in stats.get("failure_breakdown_examples", {}):
combined_stats["failure_breakdown_examples"][key] += stats["failure_breakdown_examples"][key]
total_ex = combined_stats["summary"]["total_examples"]
total_pass = combined_stats["summary"]["total_passed"]
combined_stats["summary"]["pass_rate_pct"] = round(total_pass / total_ex * 100, 2) if total_ex else 0.0
return combined_stats
def _combine_outputs(run_dir: Path, num_shards: int, meta: dict) -> None:
"""Combine shard outputs into a single output.jsonl and stats.json, then prompt for HF upload."""
# Combine JSONL
combined_jsonl = run_dir / "output.jsonl"
total_lines = 0
with open(combined_jsonl, "w") as out_f:
for i in range(num_shards):
shard_jsonl = run_dir / "shards" / f"shard_{i:03d}" / "output.jsonl"
if shard_jsonl.exists():
with open(shard_jsonl) as in_f:
for line in in_f:
out_f.write(line)
total_lines += 1
print(f"Combined {total_lines} rows into {combined_jsonl}")
# Aggregate and write stats
combined_stats = _aggregate_stats(run_dir, num_shards)
# Sanity check: JSONL rows must match expected output count from stats.
# With all_valid selection policy, total_passed_samples (individual passing
# samples) is the correct count; total_passed counts seeds, not rows.
expected_rows = combined_stats["summary"].get("total_passed_samples") or combined_stats["summary"]["total_passed"]
if total_lines != expected_rows:
raise RuntimeError(
f"INTEGRITY ERROR: output.jsonl has {total_lines} rows but "
f"stats report {expected_rows} expected output rows. "
f"Shard outputs may be corrupted or out of sync with stats."
)
stats_path = run_dir / "stats.json"
with open(stats_path, "w") as f:
json.dump(combined_stats, f, indent=2)
print(f"Aggregated stats to {stats_path}")
# Print summary for review
s = combined_stats["summary"]
checks = combined_stats["checks"]
fb = combined_stats["failure_breakdown"]
fbe = combined_stats["failure_breakdown_examples"]
print(f"\n{'='*60}")
print("COMBINED STATISTICS")
print(f"{'='*60}")
print(f" Total examples: {s['total_examples']:,}")
print(f" Total passed: {s['total_passed']:,} (seeds)")
if s.get("total_passed_samples"):
print(f" Passed samples: {s['total_passed_samples']:,} (individual)")
print(f" Total failed: {s['total_failed']:,}")
print(f" Pass rate: {s['pass_rate_pct']}%")
print(f"\n --- Checks (sample-level, summed across all rounds) ---")
print(f" Static check: {checks['static_check']['passed']:,} passed, {checks['static_check']['failed']:,} failed")
print(f" Cycle consistency: {checks['cycle_consistency']['passed']:,} passed, {checks['cycle_consistency']['failed']:,} failed")
print(f" Factual accuracy: {checks['factual_accuracy']['passed']:,} passed, {checks['factual_accuracy']['failed']:,} failed")
print(f" Correctness: {checks['correctness']['passed']:,} passed, {checks['correctness']['failed']:,} failed")
print(f"\n --- Failure Breakdown (sample-level) ---")
print(f" Failed static: {fb['failed_static_check']:,}")
print(f" Failed cycle: {fb['failed_cycle_consistency']:,}")
print(f" Failed factual: {fb['failed_factual_accuracy']:,}")
print(f" Failed correctness:{fb['failed_correctness']:,}")
print(f" Exhausted samples: {fb['exhausted_samples']:,}")
print(f"\n --- Failure Breakdown (example-level, by deepest check reached) ---")
print(f" Failed static: {fbe['failed_static_check']:,}")
print(f" Failed cycle: {fbe['failed_cycle_consistency']:,}")
print(f" Failed factual: {fbe['failed_factual_accuracy']:,}")
print(f" Failed correctness:{fbe['failed_correctness']:,}")
print(f" Total failed: {fbe['total_failed']:,}")
print(f"\n --- Output ---")
print(f" output.jsonl: {total_lines:,} rows")
print(f" Path: {combined_jsonl}")
print(f"{'='*60}")
# Prompt for HF upload
repo_id = f"teetone/{meta['run_name']}"
print(f"\nUpload to HuggingFace as: https://huggingface.co/datasets/{repo_id} ?")
answer = input("[y/N] ").strip().lower()
if answer != "y":
print("Skipping upload.")
return
try:
import tempfile
from huggingface_hub import HfApi
api = HfApi()
api.create_repo(repo_id, repo_type="dataset", exist_ok=True)
# Upload a README with dataset card so HF only loads output.jsonl as data
readme_content = (
"---\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 tmp:
tmp.write(readme_content)
tmp_path = tmp.name
api.upload_file(
path_or_fileobj=tmp_path,
path_in_repo="README.md",
repo_id=repo_id,
repo_type="dataset",
)
os.unlink(tmp_path)
api.upload_file(
path_or_fileobj=str(combined_jsonl),
path_in_repo="output.jsonl",
repo_id=repo_id,
repo_type="dataset",
)
api.upload_file(
path_or_fileobj=str(stats_path),
path_in_repo="stats.json",
repo_id=repo_id,
repo_type="dataset",
)
print(f"Uploaded to HuggingFace: https://huggingface.co/datasets/{repo_id}")
except Exception as e:
print(f"WARNING: HuggingFace upload failed: {e}")
# =============================================================================
# --rerun
# =============================================================================
def cmd_rerun(args: argparse.Namespace) -> None:
config = SDGConfig.from_yaml(args.config)
run_dir = config.output_path
meta_path = run_dir / "launch_meta.json"
if not meta_path.exists():
print(f"ERROR: {meta_path} not found", file=sys.stderr)
sys.exit(1)
with open(meta_path) as f:
meta = json.load(f)
run_name = meta["run_name"]
num_shards = meta["num_shards"]
exclude_flag = ""
if config.machines_to_exclude:
exclude_flag = f" --exclude {','.join(config.machines_to_exclude)}"
relaunch = []
for i in range(num_shards):
status = _shard_status(run_dir, i, run_name, num_shards)
if status in ("failed", "not_started"):
relaunch.append(i)
if not relaunch:
print("No failed or not_started shards to rerun.")
return
if args.dry_run:
print(f"[DRY RUN] Would relaunch {len(relaunch)} failed shards: {relaunch}")
for i in relaunch:
cmd = f"python -m sdg.generate --config {args.config} --shard-id {i} --num-shards {num_shards}"
job_name = f"{run_name}-shard-{i:03d}"
nlprun_cmd = (
f"nlprun -a tonyreasoningtraces -q sphinx{exclude_flag} -c 4 -g 1 --memory 100g "
f"-w {WORKDIR} --job-name {job_name} \"{cmd}\""
)
print(f" shard {i:03d}: {nlprun_cmd}")
return
print(f"Relaunching {len(relaunch)} failed shards: {relaunch}")
for i in relaunch:
cmd = f"python -m sdg.generate --config {args.config} --shard-id {i} --num-shards {num_shards}"
job_name = f"{run_name}-shard-{i:03d}"
nlprun_cmd = (
f"nlprun -a tonyreasoningtraces -q sphinx{exclude_flag} -c 4 -g 1 --memory 100g "
f"-w {WORKDIR} --job-name {job_name} \"{cmd}\""
)
print(f"Relaunching shard {i:03d}: {nlprun_cmd}")
subprocess.run(nlprun_cmd, shell=True, check=True)
print(f"\nRelaunched {len(relaunch)} shards")
# =============================================================================
# CLI
# =============================================================================
def main() -> None:
parser = argparse.ArgumentParser(description="SDG distributed pipeline orchestrator")
parser.add_argument("--config", required=True, help="Path to YAML config")
parser.add_argument("--run-dir", help="Run directory (alternative to --config for check/rerun)")
subparsers = parser.add_subparsers(dest="command", required=True)
launch_parser = subparsers.add_parser("launch", help="Launch all shards")
launch_parser.add_argument("--num-shards", type=int, help="Override num_shards from config")
launch_parser.add_argument("--dry-run", action="store_true", help="Print commands without executing")
check_parser = subparsers.add_parser("check", help="Check shard status")
check_parser.add_argument(
"--force-upload",
action="store_true",
help="Upload to HuggingFace as long as every shard is either completed or failed "
"(no running/pending). Only completed shards' data is uploaded.",
)
rerun_parser = subparsers.add_parser("rerun", help="Rerun failed shards")
rerun_parser.add_argument("--dry-run", action="store_true", help="Print commands without executing")
args = parser.parse_args()
if args.command == "launch":
cmd_launch(args)
elif args.command == "check":
cmd_check(args)
elif args.command == "rerun":
cmd_rerun(args)
if __name__ == "__main__":
main()