ChartPipeline / scripts /score_chart_template_quality.py
Ray1ee01's picture
Upload folder using huggingface_hub
58e6885 verified
Raw
History Blame Contribute Delete
13.2 kB
import argparse
import base64
import csv
import json
import os
import sys
from concurrent.futures import ThreadPoolExecutor, as_completed
from pathlib import Path
from openai import OpenAI
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
os.chdir(ROOT)
from config import api_key, base_url
SCORE_FIELDS = [
"chart_readability",
"data_encoding_clarity",
"label_legibility",
"visual_aesthetics",
"layout_integrity",
"rendering_correctness",
"variation_quality",
"overall",
]
DEFAULT_PROMPT = """
You are evaluating chart template implementation quality.
You will receive one generated infographic image. The image may include an infographic title, subtitle, decorative background, theme illustration, icons, or non-chart imagery.
Important: ignore everything except the chart itself.
- Ignore the infographic title and subtitle.
- Ignore theme images, illustrations, decorative icons, and non-chart imagery.
- Ignore the overall page layout outside the chart.
- Ignore whether the underlying real-world data facts are true.
- First identify the chart region, then score only that chart region.
Score each field from 1 to 10. Use the full scale. A score of 10 means excellent, 1 means unusable.
Scoring fields:
- chart_readability: whether the chart body is easy to understand. Consider marks, axes, ticks, legends, series separation, and data-value relationships.
- data_encoding_clarity: whether the chart type clearly expresses the data relationship. Consider temporal order, category order, proportions, grouping, stacking, multi-series structure, and whether the chosen visual encoding is understandable.
- label_legibility: whether chart-internal text is readable. Only evaluate axis labels, tick labels, legends, and data labels inside the chart. Do not evaluate infographic title/subtitle.
- visual_aesthetics: whether the chart itself looks visually pleasing. Consider colors, shape/line style, spacing, density, and visual hierarchy inside the chart.
- layout_integrity: whether the chart-internal layout is intact. Penalize overlap, clipping, misalignment, cramped legends, label collisions, and chart elements overflowing the chart area.
- rendering_correctness: whether there are obvious implementation or rendering bugs. Penalize empty charts, fallback screenshots, missing marks, broken shapes, abnormal blocks, missing axes, or obviously wrong scaling.
- variation_quality: whether this chart template variation appears stable, reusable, and distinctive as an implementation, independent of the specific topic.
- overall: weighted overall chart-only quality. Suggested weighting: chart_readability 25%, data_encoding_clarity 20%, layout_integrity 20%, rendering_correctness 15%, label_legibility 10%, visual_aesthetics 10%. Use variation_quality as an auxiliary judgment.
Return only valid JSON with this schema:
{
"chart_readability": 1,
"data_encoding_clarity": 1,
"label_legibility": 1,
"visual_aesthetics": 1,
"layout_integrity": 1,
"rendering_correctness": 1,
"variation_quality": 1,
"overall": 1,
"chart_region_description": "short description of the chart region you evaluated",
"reason": "short reason for the scores"
}
""".strip()
def parse_args():
parser = argparse.ArgumentParser()
parser.add_argument("--input-root", default="output/chart_template_samples/20260527_083945")
parser.add_argument("--output-dir", default=None)
parser.add_argument("--model", default="gpt-4o")
parser.add_argument("--prompt", default=DEFAULT_PROMPT)
parser.add_argument("--limit", type=int, default=0)
parser.add_argument("--workers", type=int, default=1)
parser.add_argument("--templates", nargs="+", default=None)
parser.add_argument("--resume", action="store_true")
parser.add_argument("--aggregate-only", action="store_true")
parser.add_argument("--detail", default="high", choices=["low", "high", "auto"])
return parser.parse_args()
def read_jsonl(path: Path):
if not path.is_file():
return []
with open(path, "r", encoding="utf-8") as f:
return [json.loads(line) for line in f if line.strip()]
def append_jsonl(path: Path, record: dict):
path.parent.mkdir(parents=True, exist_ok=True)
with open(path, "a", encoding="utf-8") as f:
f.write(json.dumps(record, ensure_ascii=False) + "\n")
def sample_key(record: dict):
return f"{record.get('template_key')}::{record.get('sample_index')}::{record.get('sample_dir')}"
def encode_image(path: Path):
with open(path, "rb") as f:
return base64.b64encode(f.read()).decode("utf-8")
def mime_type(path: Path):
suffix = path.suffix.lower()
if suffix == ".jpg" or suffix == ".jpeg":
return "image/jpeg"
if suffix == ".webp":
return "image/webp"
return "image/png"
def collect_tasks(input_root: Path, limit: int, templates: set[str] | None):
manifest_path = input_root / "manifest.jsonl"
records = read_jsonl(manifest_path)
tasks = []
failures = []
for record in records:
if not record.get("success"):
continue
if templates and record.get("template_key") not in templates:
continue
png_path = record.get("final_png")
if not png_path:
failures.append({**record, "error": "missing final_png"})
continue
image_path = Path(png_path)
if not image_path.is_file():
image_path = ROOT / png_path
if not image_path.is_file():
failures.append({**record, "error": f"final_png not found: {png_path}"})
continue
tasks.append({**record, "image_path": str(image_path)})
if limit > 0 and len(tasks) >= limit:
break
return tasks, failures
def scored_keys(path: Path):
return {sample_key(record) for record in read_jsonl(path)}
def score_one(args, record: dict):
image_path = Path(record["image_path"])
client = OpenAI(api_key=api_key, base_url=base_url)
image_b64 = encode_image(image_path)
response = client.chat.completions.create(
model=args.model,
response_format={"type": "json_object"},
messages=[
{
"role": "user",
"content": [
{"type": "text", "text": args.prompt},
{
"type": "image_url",
"image_url": {
"url": f"data:{mime_type(image_path)};base64,{image_b64}",
"detail": args.detail,
},
},
],
}
],
)
payload = json.loads(response.choices[0].message.content)
scores = {field: float(payload[field]) for field in SCORE_FIELDS}
return {
"template_key": record["template_key"],
"sample_index": record["sample_index"],
"sample_dir": record["sample_dir"],
"input_image": str(image_path),
"scores": scores,
"chart_region_description": payload.get("chart_region_description", ""),
"reason": payload.get("reason", ""),
"model": args.model,
}
def aggregate(output_dir: Path):
rows = read_jsonl(output_dir / "sample_scores.jsonl")
failures = read_jsonl(output_dir / "failed_scores.jsonl")
groups = {}
failure_counts = {}
for record in failures:
template_key = record.get("template_key", "")
failure_counts[template_key] = failure_counts.get(template_key, 0) + 1
for record in rows:
template_key = record["template_key"]
group = groups.setdefault(template_key, [])
group.append(record)
summary_rows = []
for template_key, records in groups.items():
item = {
"template_key": template_key,
"n_scored": len(records),
"n_failed": failure_counts.get(template_key, 0),
}
for field in SCORE_FIELDS:
values = [record["scores"][field] for record in records]
item[f"{field}_mean"] = round(sum(values) / len(values), 4)
summary_rows.append(item)
for template_key, count in failure_counts.items():
if template_key not in groups:
summary_rows.append({
"template_key": template_key,
"n_scored": 0,
"n_failed": count,
**{f"{field}_mean": "" for field in SCORE_FIELDS},
})
summary_rows.sort(
key=lambda row: (
row["overall_mean"] == "",
row["overall_mean"] if row["overall_mean"] != "" else 999,
row["template_key"],
)
)
with open(output_dir / "template_scores.json", "w", encoding="utf-8") as f:
json.dump(summary_rows, f, indent=2, ensure_ascii=False)
fieldnames = [
"template_key",
"n_scored",
"n_failed",
*[f"{field}_mean" for field in SCORE_FIELDS],
]
with open(output_dir / "template_scores.csv", "w", newline="", encoding="utf-8") as f:
writer = csv.DictWriter(f, fieldnames=fieldnames)
writer.writeheader()
writer.writerows(summary_rows)
return summary_rows
def write_combined_json(output_dir: Path, summary: dict, template_scores: list[dict]):
samples = read_jsonl(output_dir / "sample_scores.jsonl")
failures = read_jsonl(output_dir / "failed_scores.jsonl")
payload = {
"summary": summary,
"score_fields": SCORE_FIELDS,
"samples": samples,
"templates": template_scores,
"failures": failures,
}
with open(output_dir / "all_scores.json", "w", encoding="utf-8") as f:
json.dump(payload, f, indent=2, ensure_ascii=False)
def main():
args = parse_args()
input_root = Path(args.input_root)
output_dir = Path(args.output_dir) if args.output_dir else (
ROOT / "output" / "chart_template_quality_scores" / input_root.name
)
output_dir.mkdir(parents=True, exist_ok=True)
templates = set(args.templates) if args.templates else None
tasks, missing_failures = collect_tasks(input_root, args.limit, templates)
sample_scores_path = output_dir / "sample_scores.jsonl"
failed_scores_path = output_dir / "failed_scores.jsonl"
if args.aggregate_only:
template_scores = aggregate(output_dir)
summary = {
"input_root": str(input_root),
"output_dir": str(output_dir),
"model": args.model,
"total_tasks": len(tasks),
"pending_scored": 0,
"missing_png": len(missing_failures),
"templates_scored": len([row for row in template_scores if row["n_scored"] > 0]),
"score_fields": SCORE_FIELDS,
}
with open(output_dir / "summary.json", "w", encoding="utf-8") as f:
json.dump(summary, f, indent=2, ensure_ascii=False)
write_combined_json(output_dir, summary, template_scores)
print(f"wrote {output_dir / 'all_scores.json'}", flush=True)
return
existing = scored_keys(sample_scores_path) if args.resume else set()
if args.resume:
existing.update(scored_keys(failed_scores_path))
pending = [record for record in tasks if sample_key(record) not in existing]
for record in missing_failures:
if not args.resume or sample_key(record) not in existing:
append_jsonl(failed_scores_path, record)
print(f"input_root={input_root}", flush=True)
print(f"output_dir={output_dir}", flush=True)
print(f"tasks={len(tasks)} pending={len(pending)} missing_png={len(missing_failures)} workers={args.workers}", flush=True)
if args.workers > 1:
with ThreadPoolExecutor(max_workers=args.workers) as executor:
futures = [executor.submit(score_one, args, record) for record in pending]
for index, future in enumerate(as_completed(futures), 1):
result = future.result()
append_jsonl(sample_scores_path, result)
print(f"[{index}/{len(pending)}] {result['template_key']} sample_{result['sample_index']} overall={result['scores']['overall']}", flush=True)
else:
for index, record in enumerate(pending, 1):
result = score_one(args, record)
append_jsonl(sample_scores_path, result)
print(f"[{index}/{len(pending)}] {result['template_key']} sample_{result['sample_index']} overall={result['scores']['overall']}", flush=True)
template_scores = aggregate(output_dir)
summary = {
"input_root": str(input_root),
"output_dir": str(output_dir),
"model": args.model,
"total_tasks": len(tasks),
"pending_scored": len(pending),
"missing_png": len(missing_failures),
"templates_scored": len([row for row in template_scores if row["n_scored"] > 0]),
"score_fields": SCORE_FIELDS,
}
with open(output_dir / "summary.json", "w", encoding="utf-8") as f:
json.dump(summary, f, indent=2, ensure_ascii=False)
write_combined_json(output_dir, summary, template_scores)
if __name__ == "__main__":
main()