Spaces:
Sleeping
Sleeping
File size: 14,632 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 | """
Run infographics_generator against saved variation sample inputs.
The generated template samples live as:
<samples-root>/<chart_name>/sample_00/<chart_name>_sample_00
This script runs one or more samples per chart_name, records a CSV report, and
builds a compact preview.html for quick visual inspection.
"""
from __future__ import annotations
import argparse
import csv
import contextlib
import html
import logging
import os
import re
import sys
import time
import traceback
from concurrent.futures import ProcessPoolExecutor, as_completed
from pathlib import Path
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
FALLBACK_MARKER = "This is a fallback SVG using a PNG screenshot"
SHAPE_TAGS = ("path", "rect", "circle", "line", "polygon", "polyline", "ellipse", "use", "image")
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser()
parser.add_argument(
"--samples-root",
default="../data/all_variations_check/template_variations_20260527",
help="Directory containing <chart_name>/sample_XX sample folders.",
)
parser.add_argument("--output-dir", default="output/full_pipeline_check_20260602_all_variations")
parser.add_argument("--samples", default="sample_00", help="'all' or comma-separated sample names.")
parser.add_argument("--charts", default="", help="Optional comma-separated chart_name filter.")
parser.add_argument("--threads", type=int, default=4)
parser.add_argument("--limit", type=int, default=None)
parser.add_argument("--resume", action="store_true")
parser.add_argument("--output-png", action="store_true")
parser.add_argument("--chrome-path", default="/opt/google/chrome/chrome")
parser.add_argument(
"--reuse-workers",
action="store_true",
help="Reuse worker processes across samples. Faster, but can leak renderer state between outputs.",
)
return parser.parse_args()
def worker_init(chrome_path: str) -> None:
os.environ.setdefault("MPLCONFIGDIR", "/tmp/matplotlib-chartgalaxy")
os.environ.setdefault("CHROMIUM_TMP", f"/tmp/chartpipeline_chromium_{os.environ.get('USER', 'user')}")
os.environ.setdefault("ALLOWED_CHART_TYPES_FILE", "/tmp/nonexistent_allowed_chart_types.json")
os.environ.setdefault("PUPPETEER_EXECUTABLE_PATH", chrome_path)
os.environ.setdefault("RENDER_LONGEST_SIDE", "1600")
os.environ.setdefault("CHARTPIPELINE_SKIP_VARIATION_STATS_UPDATE", "1")
os.environ["NO_PROXY"] = "localhost,127.0.0.1"
os.environ["no_proxy"] = "localhost,127.0.0.1"
os.environ.setdefault("OMP_NUM_THREADS", "1")
os.environ.setdefault("MKL_NUM_THREADS", "1")
Path(os.environ["CHROMIUM_TMP"]).mkdir(parents=True, exist_ok=True)
def count_svg_elements(svg_path: Path) -> tuple[int, int, bool]:
if not svg_path.is_file():
return 0, 0, False
content = svg_path.read_text(encoding="utf-8", errors="ignore")
n_shapes = sum(len(re.findall(rf"<{tag}\b", content)) for tag in SHAPE_TAGS)
n_text = len(re.findall(r"<text\b", content))
return n_shapes, n_text, FALLBACK_MARKER in content
def configure_worker_logging(log_fh) -> None:
"""Point generator module logging at the per-sample log file.
ProcessPool workers are reused across samples. The generator configures a
module logger at import time, so without resetting it here later samples can
keep a handler whose stream belongs to a previous, already-closed log file.
"""
formatter = logging.Formatter("%(asctime)s - %(levelname)s - %(message)s")
for logger_name in ("modules.infographics_generator.infographics_generator",):
logger = logging.getLogger(logger_name)
for handler in logger.handlers[:]:
logger.removeHandler(handler)
try:
handler.close()
except Exception:
pass
handler = logging.StreamHandler(log_fh)
handler.setLevel(logging.INFO)
handler.setFormatter(formatter)
logger.addHandler(handler)
logger.setLevel(logging.INFO)
logger.propagate = False
def sample_input_for(chart_dir: Path, sample_name: str) -> Path | None:
chart_name = chart_dir.name
sample_dir = chart_dir / sample_name
preferred = sample_dir / f"{chart_name}_{sample_name}"
if preferred.is_file():
return preferred
if not sample_dir.is_dir():
return None
candidates = [
p for p in sample_dir.iterdir()
if p.is_file()
and not p.suffix
and not p.name.startswith(".")
and p.name != "render.log"
]
return sorted(candidates)[0] if candidates else None
def discover_jobs(samples_root: Path, samples_spec: str, charts_spec: str = "") -> list[tuple[str, str, str, str]]:
jobs: list[tuple[str, str, str, str]] = []
chart_filter = {c.strip() for c in charts_spec.split(",") if c.strip()}
for chart_dir in sorted(p for p in samples_root.iterdir() if p.is_dir()):
if chart_filter and chart_dir.name not in chart_filter:
continue
if samples_spec == "all":
sample_names = sorted(p.name for p in chart_dir.iterdir() if p.is_dir() and p.name.startswith("sample_"))
else:
sample_names = [s.strip() for s in samples_spec.split(",") if s.strip()]
for sample_name in sample_names:
input_path = sample_input_for(chart_dir, sample_name)
if input_path:
jobs.append((chart_dir.name, sample_name, str(input_path), f"{chart_dir.name}/{sample_name}"))
return jobs
def latest_matching_svg(out_dir: Path, chart_name: str, sample_name: str) -> Path | None:
candidates = sorted(out_dir.glob(f"*_{chart_name}_{sample_name}.svg"))
return candidates[-1] if candidates else None
def latest_subfolder(out_dir: Path, chart_name: str, sample_name: str) -> Path | None:
candidates = sorted(
p for p in out_dir.iterdir()
if p.is_dir() and p.name.endswith(f"_{chart_name}_{sample_name}")
) if out_dir.is_dir() else []
return candidates[-1] if candidates else None
def run_one(job: tuple[str, str, str, str], output_root: str, output_png: bool) -> dict[str, str]:
chart_name, sample_name, input_path, data_key = job
template_out_dir = Path(output_root) / chart_name
template_out_dir.mkdir(parents=True, exist_ok=True)
output_path = template_out_dir / sample_name
log_path = template_out_dir / f"{sample_name}.log"
t0 = time.time()
ok = False
err = ""
with log_path.open("w", encoding="utf-8", errors="ignore") as log_fh:
with contextlib.redirect_stdout(log_fh), contextlib.redirect_stderr(log_fh):
try:
from modules.infographics_generator.infographics_generator import process
configure_worker_logging(log_fh)
ok = bool(process(
input=input_path,
output=str(output_path),
base_url="",
api_key="",
chart_name=chart_name,
output_png=output_png,
))
except BaseException as exc:
tb = traceback.format_exc().splitlines()
tail = tb[-1] if tb else str(exc)
err = f"{type(exc).__name__}: {exc} | {tail}"[:500]
elapsed = time.time() - t0
final_svg = latest_matching_svg(template_out_dir, chart_name, sample_name)
subfolder = latest_subfolder(template_out_dir, chart_name, sample_name)
chart_svg = subfolder / "chart.svg" if subfolder else None
n_shapes, n_text, final_fallback = count_svg_elements(final_svg) if final_svg else (0, 0, False)
_, _, chart_fallback = count_svg_elements(chart_svg) if chart_svg else (0, 0, False)
final_size = final_svg.stat().st_size if final_svg else 0
if ok and (not final_svg or final_fallback or chart_fallback):
ok = False
if ok and n_shapes + n_text < 8:
ok = False
return {
"chart_name": chart_name,
"sample": sample_name,
"input": data_key,
"ok": str(bool(ok)),
"elapsed_s": f"{elapsed:.2f}",
"final_svg": str(final_svg or ""),
"final_svg_bytes": str(final_size),
"chart_svg": str(chart_svg or ""),
"final_svg_fallback_png": str(final_fallback),
"chart_svg_fallback_png": str(chart_fallback),
"n_shapes": str(n_shapes),
"n_text": str(n_text),
"log": str(log_path),
"err": err,
}
def read_done(tasks_csv: Path) -> set[tuple[str, str]]:
if not tasks_csv.is_file():
return set()
with tasks_csv.open(newline="") as fh:
return {(r["chart_name"], r["sample"]) for r in csv.DictReader(fh)}
def write_summary(out_dir: Path, rows: list[dict[str, str]]) -> None:
grouped: dict[str, list[dict[str, str]]] = {}
for row in rows:
grouped.setdefault(row["chart_name"], []).append(row)
summary_rows = []
for chart_name, items in grouped.items():
done = len(items)
n_success = sum(r["ok"] == "True" for r in items)
n_fail = done - n_success
n_fallback = sum(
r["final_svg_fallback_png"] == "True" or r["chart_svg_fallback_png"] == "True"
for r in items
)
n_empty = sum(int(r["n_shapes"] or 0) + int(r["n_text"] or 0) < 8 for r in items)
mean_shapes = sum(int(r["n_shapes"] or 0) for r in items) / done
mean_text = sum(int(r["n_text"] or 0) for r in items) / done
mean_time = sum(float(r["elapsed_s"] or 0) for r in items) / done
summary_rows.append({
"chart_name": chart_name,
"done": done,
"n_success": n_success,
"n_fail": n_fail,
"n_fallback_png": n_fallback,
"n_empty": n_empty,
"mean_shapes": f"{mean_shapes:.1f}",
"mean_text": f"{mean_text:.1f}",
"mean_elapsed_s": f"{mean_time:.2f}",
})
summary_rows.sort(key=lambda r: (-r["n_fail"], -r["n_fallback_png"], -r["n_empty"], r["chart_name"]))
fieldnames = [
"chart_name", "done", "n_success", "n_fail", "n_fallback_png",
"n_empty", "mean_shapes", "mean_text", "mean_elapsed_s",
]
with (out_dir / "_summary.csv").open("w", newline="") as fh:
writer = csv.DictWriter(fh, fieldnames=fieldnames)
writer.writeheader()
writer.writerows(summary_rows)
def write_preview(out_dir: Path, rows: list[dict[str, str]]) -> None:
html_path = out_dir / "preview.html"
parent = html_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:#f6f6f6;margin:0;padding:20px}"
".grid{display:grid;grid-template-columns:repeat(auto-fill,minmax(220px,1fr));gap:12px}"
".card{background:#fff;border:1px solid #ddd;border-radius:6px;padding:8px}.bad{border-color:#d00;background:#fff6f6}"
"object{width:100%;height:190px;background:white}.meta{font-size:11px;color:#555;word-break:break-word}</style>",
"<h1>Variation Pipeline Check</h1>",
f"<p>tasks={len(rows)} ok={sum(r['ok']=='True' for r in rows)} fail={sum(r['ok']!='True' for r in rows)}</p>",
"<div class='grid'>",
]
for row in rows:
cls = "card" if row["ok"] == "True" else "card bad"
svg = rel(row["final_svg"])
obj = f"<object data='{html.escape(svg)}' type='image/svg+xml'></object>" if svg else "<div>no svg</div>"
parts.append(
f"<div class='{cls}'>{obj}<div class='meta'>"
f"<b>{html.escape(row['chart_name'])}</b> {html.escape(row['sample'])}<br>"
f"ok={row['ok']} shapes={row['n_shapes']} text={row['n_text']} time={row['elapsed_s']}s<br>"
f"{html.escape(row['err'])}</div></div>"
)
parts.append("</div>")
html_path.write_text("\n".join(parts), encoding="utf-8")
def main() -> int:
args = parse_args()
samples_root = Path(args.samples_root)
out_dir = Path(args.output_dir)
out_dir.mkdir(parents=True, exist_ok=True)
tasks_csv = out_dir / "_tasks.csv"
fieldnames = [
"chart_name", "sample", "input", "ok", "elapsed_s", "final_svg",
"final_svg_bytes", "chart_svg", "final_svg_fallback_png",
"chart_svg_fallback_png", "n_shapes", "n_text", "err",
"log",
]
jobs = discover_jobs(samples_root, args.samples, args.charts)
if args.limit:
jobs = jobs[: args.limit]
done_keys = read_done(tasks_csv) if args.resume else set()
jobs = [j for j in jobs if (j[0], j[1]) not in done_keys]
print(f"Discovered jobs={len(jobs) + len(done_keys)} resume_done={len(done_keys)} to_run={len(jobs)}")
mode = "a" if args.resume and tasks_csv.exists() else "w"
with tasks_csv.open(mode, newline="") as fh:
writer = csv.DictWriter(fh, fieldnames=fieldnames)
if mode == "w":
writer.writeheader()
executor_kwargs = {
"max_workers": args.threads,
"initializer": worker_init,
"initargs": (args.chrome_path,),
}
if not args.reuse_workers:
executor_kwargs["max_tasks_per_child"] = 1
with ProcessPoolExecutor(**executor_kwargs) as pool:
futures = [pool.submit(run_one, job, str(out_dir), args.output_png) for job in jobs]
t0 = time.time()
for index, future in enumerate(as_completed(futures), 1):
row = future.result()
writer.writerow(row)
fh.flush()
print(
f"[{index}/{len(futures)}] {row['chart_name']}/{row['sample']} "
f"ok={row['ok']} shapes={row['n_shapes']} text={row['n_text']} "
f"time={row['elapsed_s']}s elapsed={time.time() - t0:.1f}s"
)
rows = list(csv.DictReader(tasks_csv.open(newline="")))
write_summary(out_dir, rows)
write_preview(out_dir, rows)
n_fail = sum(r["ok"] != "True" for r in rows)
print(f"Wrote {tasks_csv}")
print(f"Wrote {out_dir / '_summary.csv'}")
print(f"Wrote {out_dir / 'preview.html'}")
print(f"Total rows={len(rows)} failures={n_fail}")
return 1 if n_fail else 0
if __name__ == "__main__":
raise SystemExit(main())
|