ChartPipeline / scripts /run_variation_sample_check.py
Ray1ee01's picture
Upload folder using huggingface_hub
58e6885 verified
Raw
History Blame Contribute Delete
14.6 kB
"""
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())