Spaces:
Sleeping
Sleeping
File size: 17,017 Bytes
58e6885 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 | """Run whole-image GPT Image post-processing over variation_skill_runs outputs."""
from __future__ import annotations
import argparse
import csv
import html
import os
import re
import sys
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from modules.full_image_polisher.full_image_polisher import (
DEFAULT_PROMPT,
GLOBAL_COHERENCE_PROMPT,
polish_full_image,
polish_two_pass,
)
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description="Post-process rendered variation PNGs with an image edit backend.")
parser.add_argument("--runs-root", default="data/output/variation_skill_runs")
parser.add_argument("--variation", default="", help="Comma-separated variation filter.")
parser.add_argument("--run-dir", action="append", default=[], help="Specific run directory to process; can be repeated.")
parser.add_argument(
"--render-dir",
default="latest_all_samples",
help="Render directory name inside each run, or latest_all_samples for the highest render_all_samples_roundN.",
)
parser.add_argument("--sample", default="sample_00", help="Sample filter, comma-separated, or all.")
parser.add_argument("--output-subdir", default="postprocess_gpt_image2")
parser.add_argument("--limit", type=int, default=None)
parser.add_argument("--resume", action="store_true")
parser.add_argument("--include-failed", action="store_true", help="Also process rows whose render ok field is not True.")
parser.add_argument("--mode", choices=("global", "slot", "two_pass"), default="two_pass")
parser.add_argument("--image-backend", choices=("auto", "openai", "pinco"), default="auto")
parser.add_argument("--model", default="gpt-image-2")
parser.add_argument("--size", default="auto")
parser.add_argument("--quality", default="auto")
parser.add_argument("--input-fidelity", choices=("high", "low"), default="high")
parser.add_argument("--output-format", default="png")
parser.add_argument("--prompt", default=DEFAULT_PROMPT)
parser.add_argument("--data-context", choices=("none", "compact", "strict"), default="compact")
parser.add_argument("--data-context-max-rows", type=int, default=80)
parser.add_argument("--api-key-env", default="OPENAI_API_KEY")
parser.add_argument("--base-url", default=None)
parser.add_argument("--dry-run", action="store_true")
parser.add_argument("--resize-to-input", action="store_true")
parser.add_argument("--slot-guided", action="store_true", help="Use each final SVG to build slot mask/overlay guidance.")
parser.add_argument("--slot-overlay-input", action="store_true", help="Send labeled slot overlay as a second image input. Off by default to avoid label leakage.")
parser.add_argument("--include-title-slots", action="store_true")
parser.add_argument("--slot-padding-px", type=int, default=8)
parser.add_argument(
"--pinco-command",
default=None,
help=(
"Command template for local Pinco inference. Placeholders: {input}, {mask}, "
"{foreground}, {prompt_file}, {output}, {model}, {width}, {height}."
),
)
parser.add_argument("--pinco-url", default=None, help="HTTP endpoint for a Pinco inpainting service.")
parser.add_argument("--pinco-timeout", type=int, default=600)
return parser.parse_args()
def _variation_filter(spec: str) -> set[str]:
return {item.strip() for item in spec.split(",") if item.strip()}
def discover_run_dirs(args: argparse.Namespace) -> list[Path]:
if args.run_dir:
return [Path(path) for path in args.run_dir]
runs_root = Path(args.runs_root)
wanted = _variation_filter(args.variation)
run_dirs: list[Path] = []
for variation_dir in sorted(path for path in runs_root.iterdir() if path.is_dir()):
if wanted and variation_dir.name not in wanted:
continue
dated = sorted(
[path for path in variation_dir.iterdir() if path.is_dir()],
key=lambda path: path.stat().st_mtime,
reverse=True,
)
if dated:
run_dirs.append(dated[0])
return run_dirs
def _resolve_row_path(path_text: str) -> Path:
path = Path(path_text)
if path.is_absolute():
return path
root_path = (ROOT / path).resolve()
if root_path.exists():
return root_path
return (ROOT.parent / path).resolve()
def _round_number(path: Path) -> int:
match = re.search(r"round(\d+)$", path.name)
return int(match.group(1)) if match else -1
def resolve_render_dir(run_dir: Path, render_dir_name: str) -> Path | None:
if render_dir_name != "latest_all_samples":
candidate = run_dir / render_dir_name
return candidate if candidate.is_dir() else None
candidates = [path for path in run_dir.glob("render_all_samples_round*") if path.is_dir()]
if not candidates:
return None
return sorted(candidates, key=lambda path: (_round_number(path), path.stat().st_mtime), reverse=True)[0]
def sample_filter(spec: str) -> set[str] | None:
if spec == "all":
return None
return {item.strip() for item in spec.split(",") if item.strip()}
def _png_for_row(row: dict[str, str]) -> Path | None:
final_svg = row.get("final_svg", "")
if not final_svg:
return None
svg_path = _resolve_row_path(final_svg)
png_path = svg_path.with_suffix(".png")
if png_path.is_file():
return png_path
candidates = sorted(svg_path.parent.glob(f"{svg_path.stem}*.png"))
return candidates[-1] if candidates else None
def _svg_for_row(row: dict[str, str]) -> Path | None:
final_svg = row.get("final_svg", "")
if not final_svg:
return None
svg_path = _resolve_row_path(final_svg)
return svg_path if svg_path.is_file() else None
def _data_for_row(row: dict[str, str]) -> Path | None:
final_svg = row.get("final_svg", "")
if not final_svg:
return None
svg_path = _resolve_row_path(final_svg)
candidates = [
svg_path.with_suffix("") / "data.json",
svg_path.parent / "data.json",
]
chart_svg = row.get("chart_svg", "")
if chart_svg:
chart_svg_path = _resolve_row_path(chart_svg)
candidates.append(chart_svg_path.parent / "data.json")
for candidate in candidates:
if candidate.is_file():
return candidate
return None
def _reference_map_for_run(run_dir: Path) -> Path | None:
candidate = run_dir / "reference_element_map.json"
return candidate if candidate.is_file() else None
def load_jobs(run_dirs: list[Path], render_dir_name: str, sample_spec: str, include_failed: bool) -> list[dict[str, str]]:
jobs: list[dict[str, str]] = []
wanted_samples = sample_filter(sample_spec)
for run_dir in run_dirs:
render_dir = resolve_render_dir(run_dir, render_dir_name)
if render_dir is None:
continue
tasks_csv = render_dir / "_tasks.csv"
if not tasks_csv.is_file():
continue
with tasks_csv.open(newline="", encoding="utf-8") as fh:
for row in csv.DictReader(fh):
if wanted_samples is not None and row.get("sample") not in wanted_samples:
continue
if not include_failed and row.get("ok") != "True":
continue
png_path = _png_for_row(row)
if not png_path:
continue
svg_path = _svg_for_row(row)
data_path = _data_for_row(row)
jobs.append({
"run_dir": str(run_dir),
"render_dir": render_dir.name,
"chart_name": row.get("chart_name", ""),
"sample": row.get("sample", ""),
"input_png": str(png_path),
"input_svg": str(svg_path) if svg_path else "",
"input_data": str(data_path) if data_path else "",
"reference_map": str(_reference_map_for_run(run_dir) or ""),
})
return jobs
def read_done(tasks_csv: Path) -> set[str]:
if not tasks_csv.is_file():
return set()
with tasks_csv.open(newline="", encoding="utf-8") as fh:
return {row["input_png"] for row in csv.DictReader(fh) if row.get("ok") == "True"}
def output_for(job: dict[str, str], output_subdir: str, mode: str = "global") -> Path:
run_dir = Path(job["run_dir"])
input_png = Path(job["input_png"])
suffix_by_mode = {
"global": "coherence_polished",
"slot": "slot_guided_polished",
"two_pass": "two_pass_final",
}
suffix = suffix_by_mode[mode]
rel_parts = [job["render_dir"], job["chart_name"], f"{input_png.stem}.{suffix}.png"]
return run_dir / output_subdir / Path(*rel_parts)
def write_preview(preview_path: Path, rows: list[dict[str, str]]) -> None:
preview_path.parent.mkdir(parents=True, exist_ok=True)
parent = preview_path.parent.resolve()
def rel(path_text: str) -> str:
if not path_text:
return ""
path = Path(path_text)
if not path.is_absolute():
path = (ROOT / path).resolve()
return os.path.relpath(path, parent)
parts = [
"<!doctype html><meta charset='utf-8'>",
"<style>body{font-family:-apple-system,BlinkMacSystemFont,sans-serif;background:#f7f7f7;margin:0;padding:20px}"
".grid{display:grid;grid-template-columns:repeat(auto-fill,minmax(420px,1fr));gap:14px}"
".card{background:#fff;border:1px solid #ddd;border-radius:6px;padding:10px}"
".bad{border-color:#c00;background:#fff7f7}.pair{display:grid;grid-template-columns:1fr 1fr;gap:8px}"
"img{width:100%;background:#fff;border:1px solid #eee}.meta{font-size:12px;color:#555;word-break:break-word}</style>",
"<h1>GPT Image Whole-Image Postprocess</h1>",
f"<p>tasks={len(rows)} ok={sum(row['ok']=='True' for row in rows)} fail={sum(row['ok']!='True' for row in rows)}</p>",
"<div class='grid'>",
]
for row in rows:
cls = "card" if row["ok"] == "True" else "card bad"
source = html.escape(rel(row["input_png"]))
output = html.escape(rel(row["output_png"]))
output_img = f"<img src='{output}'>" if row["ok"] == "True" and output else "<div>no output</div>"
parts.append(
f"<div class='{cls}'><div class='pair'><img src='{source}'>{output_img}</div>"
f"<div class='meta'><b>{html.escape(row['chart_name'])}</b> {html.escape(row['sample'])}<br>"
f"ok={row['ok']} {html.escape(row.get('err', ''))}<br>{html.escape(row['output_png'])}</div></div>"
)
parts.append("</div>")
preview_path.write_text("\n".join(parts), encoding="utf-8")
def main() -> int:
args = parse_args()
try:
from config import api_key, base_url
except Exception:
api_key = None
base_url = None
run_dirs = discover_run_dirs(args)
jobs = load_jobs(run_dirs, args.render_dir, args.sample, args.include_failed)
if args.limit is not None:
jobs = jobs[: args.limit]
output_roots = {Path(job["run_dir"]) / args.output_subdir for job in jobs}
for output_root in output_roots:
output_root.mkdir(parents=True, exist_ok=True)
task_logs = {output_root: output_root / "postprocess_tasks.csv" for output_root in output_roots}
done: set[str] = set()
if args.resume:
for tasks_csv in task_logs.values():
done.update(read_done(tasks_csv))
jobs = [job for job in jobs if job["input_png"] not in done]
fieldnames = [
"run_dir", "render_dir", "chart_name", "sample", "input_png", "input_svg",
"input_data", "reference_map", "output_png", "manifest", "ok", "err",
]
rows_by_output_root: dict[Path, list[dict[str, str]]] = {root: [] for root in output_roots}
effective_mode = "slot" if args.slot_guided else args.mode
for index, job in enumerate(jobs, 1):
output_png = output_for(job, args.output_subdir, mode=effective_mode)
output_root = Path(job["run_dir"]) / args.output_subdir
row = {**job, "output_png": str(output_png), "manifest": "", "ok": "False", "err": ""}
try:
if effective_mode == "two_pass":
if not job.get("input_svg"):
raise ValueError("two_pass mode requires input_svg from _tasks.csv")
manifest = polish_two_pass(
input_png=Path(job["input_png"]),
output_png=output_png,
api_key=api_key,
base_url=args.base_url or base_url,
svg_path=Path(job["input_svg"]),
reference_map=Path(job["reference_map"]) if job.get("reference_map") else None,
data_json=Path(job["input_data"]) if job.get("input_data") else None,
data_context_mode=args.data_context,
data_context_max_rows=args.data_context_max_rows,
model=args.model,
size=args.size,
quality=args.quality,
input_fidelity=args.input_fidelity,
output_format=args.output_format,
api_key_env=args.api_key_env,
dry_run=args.dry_run,
resize_to_input=args.resize_to_input,
slot_overlay_input=args.slot_overlay_input,
include_title_slots=args.include_title_slots,
slot_padding_px=args.slot_padding_px,
image_backend=args.image_backend,
pinco_command=args.pinco_command,
pinco_url=args.pinco_url,
pinco_timeout=args.pinco_timeout,
)
else:
mode_prompt = args.prompt
if effective_mode == "global" and mode_prompt == DEFAULT_PROMPT:
mode_prompt = GLOBAL_COHERENCE_PROMPT
manifest = polish_full_image(
input_png=Path(job["input_png"]),
output_png=output_png,
api_key=api_key,
base_url=args.base_url or base_url,
svg_path=Path(job["input_svg"]) if effective_mode == "slot" and job.get("input_svg") else None,
reference_map=Path(job["reference_map"]) if job.get("reference_map") else None,
data_json=Path(job["input_data"]) if job.get("input_data") else None,
data_context_mode=args.data_context if effective_mode == "global" else "none",
data_context_max_rows=args.data_context_max_rows,
model=args.model,
size=args.size,
quality=args.quality,
input_fidelity=args.input_fidelity,
output_format=args.output_format,
prompt=mode_prompt,
api_key_env=args.api_key_env,
dry_run=args.dry_run,
resize_to_input=args.resize_to_input,
slot_guided=effective_mode == "slot",
slot_overlay_input=args.slot_overlay_input,
include_title_slots=args.include_title_slots,
slot_padding_px=args.slot_padding_px,
image_backend=args.image_backend,
pinco_command=args.pinco_command,
pinco_url=args.pinco_url,
pinco_timeout=args.pinco_timeout,
)
row["manifest"] = str(Path(manifest["output_png"]).with_name(f"{Path(manifest['output_png']).stem}.manifest.json"))
row["ok"] = "True"
except Exception as exc:
row["err"] = f"{type(exc).__name__}: {exc}"[:500]
rows_by_output_root.setdefault(output_root, []).append(row)
print(f"[{index}/{len(jobs)}] {job['chart_name']}/{job['sample']} ok={row['ok']} output={row['output_png']}")
for output_root, rows in rows_by_output_root.items():
if not rows:
continue
tasks_csv = output_root / "postprocess_tasks.csv"
mode = "a" if args.resume and tasks_csv.exists() else "w"
with tasks_csv.open(mode, newline="", encoding="utf-8") as fh:
writer = csv.DictWriter(fh, fieldnames=fieldnames)
if mode == "w":
writer.writeheader()
writer.writerows(rows)
all_rows = list(csv.DictReader(tasks_csv.open(newline="", encoding="utf-8")))
write_preview(output_root / "preview.html", all_rows)
print(f"Wrote {tasks_csv}")
print(f"Wrote {output_root / 'preview.html'}")
return 0 if all(row["ok"] == "True" for rows in rows_by_output_root.values() for row in rows) else 1
if __name__ == "__main__":
raise SystemExit(main())
|