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