scripts/run_batch.py
scripts/run_batch.pyBrowse 33 files
2,244 tokens
9,100 bytes
Token encoding: o200k_base
Snapshot 24fd22b
← Back to SKILL.md
1#!/usr/bin/env python32"""3run_batch.py — Run a workflow many times, varying parameters per run.4 5Two modes:6 1. --count N --randomize-seed7 Submit N runs, each with a fresh random seed. Use for quick variations.8 2. --sweep '{"seed": [1,2,3], "steps": [20,30]}'9 Cartesian product of values. With cloud subscription, runs in parallel10 up to your tier's concurrent-job limit.11 12Both modes write each run's outputs into output-dir/run_NNN/.13 14Examples:15 python3 run_batch.py --workflow flux_dev.json \16 --args '{"prompt": "a cat"}' \17 --count 8 --randomize-seed \18 --output-dir ./outputs/cat-batch19 20 python3 run_batch.py --workflow sdxl.json \21 --args '{"prompt": "abstract"}' \22 --sweep '{"seed": [1,2,3], "steps": [20, 40]}' \23 --output-dir ./outputs/sweep24"""25 26from __future__ import annotations27 28import argparse29import itertools30import json31import sys32from concurrent.futures import ThreadPoolExecutor, as_completed33from pathlib import Path34 35sys.path.insert(0, str(Path(__file__).resolve().parent))36from _common import ( # noqa: E40237 DEFAULT_LOCAL_HOST, ENV_API_KEY, coerce_seed, emit_json, log,38 looks_like_video_workflow, resolve_api_key, unwrap_workflow,39)40from run_workflow import ( # noqa: E40241 ComfyRunner, download_outputs, inject_params,42)43from extract_schema import extract_schema # noqa: E40244 45 46def expand_sweep(sweep: dict, base_args: dict, count: int, randomize_seed: bool) -> list[dict]:47 """Generate a list of args dicts for each run."""48 if sweep:49 # Cartesian product50 keys = list(sweep.keys())51 values = [sweep[k] if isinstance(sweep[k], list) else [sweep[k]] for k in keys]52 runs = []53 for combo in itertools.product(*values):54 ar = dict(base_args)55 for k, v in zip(keys, combo):56 ar[k] = v57 runs.append(ar)58 return runs59 # Count mode60 runs = []61 for _ in range(count):62 ar = dict(base_args)63 if randomize_seed:64 ar["seed"] = coerce_seed(None)65 runs.append(ar)66 return runs67 68 69def execute_one(70 runner: ComfyRunner, workflow: dict, schema: dict, args: dict,71 *, output_dir: Path, timeout: int, ws: bool,72) -> dict:73 wf, warnings = inject_params(workflow, schema, args)74 sub = runner.submit(wf)75 if "_http_error" in sub:76 return {"status": "error", "error": "submission HTTP error",77 "details": sub.get("body"), "args": args}78 pid = sub.get("prompt_id")79 if not pid:80 return {"status": "error", "error": "no prompt_id", "response": sub, "args": args}81 if sub.get("node_errors"):82 return {"status": "error", "error": "validation failed",83 "node_errors": sub["node_errors"], "args": args}84 85 if ws:86 result = runner.monitor_ws(pid, timeout=timeout)87 else:88 result = runner.poll_status(pid, timeout=timeout)89 90 if result["status"] != "success":91 return {92 "status": result["status"],93 "prompt_id": pid,94 "details": result.get("data"),95 "args": args,96 }97 98 outputs = result.get("outputs") or runner.get_outputs(pid)99 downloaded = download_outputs(runner, outputs, output_dir, preserve_subfolder=False)100 return {101 "status": "success",102 "prompt_id": pid,103 "args": args,104 "outputs": downloaded,105 "warnings": warnings,106 }107 108 109def main(argv: list[str] | None = None) -> int:110 p = argparse.ArgumentParser(111 description="Submit a workflow many times with varying parameters.",112 )113 p.add_argument("--workflow", required=True)114 p.add_argument("--args", default="{}", help="Base parameters JSON")115 p.add_argument("--count", type=int, default=0,116 help="Number of runs (use with --randomize-seed)")117 p.add_argument("--sweep", default="",118 help='JSON dict of param→list of values. Cartesian product. '119 'e.g. \'{"seed":[1,2,3],"cfg":[5,8]}\'')120 p.add_argument("--randomize-seed", action="store_true",121 help="In --count mode, vary seed per run")122 p.add_argument("--host", default=DEFAULT_LOCAL_HOST)123 p.add_argument("--api-key", help=f"or set ${ENV_API_KEY}")124 p.add_argument("--partner-key")125 p.add_argument("--parallel", type=int, default=1,126 help="Concurrent submissions (cloud: up to your tier limit). "127 "Default 1 (sequential)")128 p.add_argument("--output-dir", default="./outputs/batch")129 p.add_argument("--timeout", type=int, default=0)130 p.add_argument("--ws", action="store_true")131 p.add_argument("--continue-on-error", action="store_true",132 help="Don't stop the batch when a run fails")133 args = p.parse_args(argv)134 135 if args.count <= 0 and not args.sweep:136 emit_json({"error": "Specify --count N or --sweep '{...}'"})137 return 1138 139 base_args = json.loads(args.args) if args.args.strip() else {}140 sweep = json.loads(args.sweep) if args.sweep.strip() else {}141 142 # Validate sweep shape143 if sweep:144 if not isinstance(sweep, dict):145 emit_json({"error": "--sweep must be a JSON object {param: [values]}"})146 return 1147 empty = [k for k, v in sweep.items() if isinstance(v, list) and len(v) == 0]148 if empty:149 emit_json({"error": f"--sweep parameters have empty value lists: {empty}"})150 return 1151 # If user passed BOTH --sweep and --count/--randomize-seed, --sweep wins152 if args.count or args.randomize_seed:153 log("--sweep set; ignoring --count / --randomize-seed (sweep defines the runs)")154 155 wf_path = Path(args.workflow).expanduser()156 if not wf_path.exists():157 emit_json({"error": f"Workflow not found: {args.workflow}"})158 return 1159 try:160 with wf_path.open(encoding="utf-8-sig") as f:161 workflow = unwrap_workflow(json.load(f))162 except (ValueError, json.JSONDecodeError) as e:163 emit_json({"error": str(e)})164 return 1165 166 schema = extract_schema(workflow)167 runs = expand_sweep(sweep, base_args, args.count, args.randomize_seed)168 log(f"Planned {len(runs)} run(s)")169 170 api_key = resolve_api_key(args.api_key)171 runner = ComfyRunner(host=args.host, api_key=api_key, partner_key=args.partner_key)172 173 ok, info = runner.check_server()174 if not ok:175 emit_json({"error": "Cannot reach server", "details": info, "host": args.host})176 return 1177 178 timeout = args.timeout179 if timeout <= 0:180 timeout = 900 if looks_like_video_workflow(workflow) else 300181 182 base_dir = Path(args.output_dir).expanduser()183 base_dir.mkdir(parents=True, exist_ok=True)184 185 results: list[dict] = []186 failures = 0187 188 if args.parallel > 1:189 with ThreadPoolExecutor(max_workers=args.parallel) as ex:190 future_to_idx = {}191 for i, ar in enumerate(runs):192 run_dir = base_dir / f"run_{i:04d}"193 fut = ex.submit(194 execute_one, runner, workflow, schema, ar,195 output_dir=run_dir, timeout=timeout, ws=args.ws,196 )197 future_to_idx[fut] = i198 for fut in as_completed(future_to_idx):199 i = future_to_idx[fut]200 try:201 r = fut.result()202 except Exception as e:203 r = {"status": "error", "error": str(e), "args": runs[i]}204 r["index"] = i205 results.append(r)206 if r["status"] != "success":207 failures += 1208 log(f" run {i} → {r['status']}: {r.get('error','?')}")209 if not args.continue_on_error:210 log(" --continue-on-error not set; aborting batch")211 break212 else:213 log(f" run {i} → success: {len(r.get('outputs', []))} files")214 else:215 for i, ar in enumerate(runs):216 run_dir = base_dir / f"run_{i:04d}"217 r = execute_one(runner, workflow, schema, ar,218 output_dir=run_dir, timeout=timeout, ws=args.ws)219 r["index"] = i220 results.append(r)221 if r["status"] != "success":222 failures += 1223 log(f" run {i} → {r['status']}: {r.get('error','?')}")224 if not args.continue_on_error:225 log(" --continue-on-error not set; aborting batch")226 break227 else:228 log(f" run {i} → success: {len(r.get('outputs', []))} files")229 230 results.sort(key=lambda x: x.get("index", 0))231 emit_json({232 "status": "success" if failures == 0 else "partial",233 "total": len(runs),234 "completed": sum(1 for r in results if r["status"] == "success"),235 "failed": failures,236 "output_dir": str(base_dir),237 "results": results,238 })239 return 0 if failures == 0 else 1240 241 242if __name__ == "__main__":243 sys.exit(main())244