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