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