#!/usr/bin/env python3 """Run model configurations concurrently without fail-fast behavior.""" from __future__ import annotations import argparse import json import subprocess import sys from concurrent.futures import ThreadPoolExecutor, as_completed from pathlib import Path from typing import Any from pipeline_common import REPO_ROOT, atomic_write_json, run_id, shell_join, utc_now def parse_args() -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("configs", nargs="+", type=Path) parser.add_argument("--workers", type=int, default=2) parser.add_argument("--stage", action="append", dest="stages") parser.add_argument("--dry-run", action="store_true") parser.add_argument("--force", action="store_true") parser.add_argument("--resume-failed", action="store_true") return parser.parse_args() def run_one(config: Path, batch_dir: Path, args: argparse.Namespace) -> dict[str, Any]: name = config.stem stdout_path = batch_dir / f"{name}.stdout.log" stderr_path = batch_dir / f"{name}.stderr.log" argv = [sys.executable, str(REPO_ROOT / "scripts" / "run_model.py"), "--config", str(config.resolve())] for stage in args.stages or []: argv.extend(["--stage", stage]) for enabled, flag in ( (args.dry_run, "--dry-run"), (args.force, "--force"), (args.resume_failed, "--resume-failed") ): if enabled: argv.append(flag) started_at = utc_now() with stdout_path.open("wb") as stdout_handle, stderr_path.open("wb") as stderr_handle: completed = subprocess.run(argv, cwd=REPO_ROOT, stdout=stdout_handle, stderr=stderr_handle, check=False) return { "config": str(config.resolve()), "command": shell_join(argv), "started_at": started_at, "ended_at": utc_now(), "exit_code": completed.returncode, "status": "PASS" if completed.returncode == 0 else "FAIL", "stdout_log": str(stdout_path), "stderr_log": str(stderr_path), } def main() -> int: args = parse_args() if args.workers < 1: raise SystemExit("--workers must be at least 1") current_run_id = run_id() batch_dir = REPO_ROOT / "logs" / "batch" / current_run_id batch_dir.mkdir(parents=True, exist_ok=False) configs = [path.resolve() for path in args.configs] results: list[dict[str, Any]] = [] with ThreadPoolExecutor(max_workers=args.workers) as executor: futures = {executor.submit(run_one, config, batch_dir, args): config for config in configs} for future in as_completed(futures): config = futures[future] try: results.append(future.result()) except Exception as error: results.append({ "config": str(config), "command": "", "started_at": utc_now(), "ended_at": utc_now(), "exit_code": 127, "status": "FAIL", "stdout_log": "", "stderr_log": "", "error": str(error), }) results.sort(key=lambda item: item["config"]) summary = { "schema_version": "1.0", "run_id": current_run_id, "worker_count": args.workers, "total": len(results), "passed": sum(item["status"] == "PASS" for item in results), "failed": sum(item["status"] == "FAIL" for item in results), "results": results, } atomic_write_json(batch_dir / "batch_summary.json", summary) print(json.dumps(summary, sort_keys=True)) return 0 if summary["failed"] == 0 else 1 if __name__ == "__main__": raise SystemExit(main())