ChartPipeline / scripts /match_templates_to_data.py
Ray1ee01's picture
Upload folder using huggingface_hub
58e6885 verified
Raw
History Blame Contribute Delete
5.31 kB
"""
Dry-run: for every template, count how many JSON files in one or more
data directories are compatible (per the same check used by
infographics_generator).
Usage:
# default (read pools from config.data_resource_dirs)
python scripts/match_templates_to_data.py --out scripts/_template_match.csv
# explicit pools
python scripts/match_templates_to_data.py \
--data /data/liduan/resources/claude_data_v2 /data/liduan/resources/claude_new \
--out scripts/_template_match.csv
"""
import argparse
import csv
import json
import os
import sys
from pathlib import Path
from collections import defaultdict
ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))
from config import data_resource_dirs as CFG_DATA_DIRS
from modules.chart_engine.template.template_registry import scan_templates
from modules.infographics_generator.template_utils import (
analyze_templates,
check_template_compatibility,
)
def parse_args():
p = argparse.ArgumentParser()
p.add_argument("--data", nargs="+", default=None,
help="One or more directories of input JSONs "
"(default: config.data_resource_dirs)")
p.add_argument("--out", default="scripts/_template_match.csv", help="CSV report path")
p.add_argument("--engines", nargs="+",
default=["d3-js", "echarts-js", "echarts_py"],
help="Which engines to consider")
p.add_argument("--sample", type=int, default=0,
help="If >0, randomly subsample this many files (per --seed)")
p.add_argument("--seed", type=int, default=42)
return p.parse_args()
def main():
args = parse_args()
data_dirs = [Path(p) for p in (args.data or CFG_DATA_DIRS)]
for d in data_dirs:
if not d.is_dir():
raise SystemExit(f"--data path is not a directory: {d}")
json_files = []
for d in data_dirs:
files = sorted(d.glob("*.json"))
json_files.extend(files)
print(f" {d}: {len(files)} json files")
print(f"Total: {len(json_files)} json files across {len(data_dirs)} dir(s)")
if args.sample and args.sample < len(json_files):
import random
random.Random(args.seed).shuffle(json_files)
json_files = json_files[:args.sample]
print(f"Sub-sampled to {len(json_files)} files (seed={args.seed})")
datas = []
for f in json_files:
with open(f, "r") as fh:
d = json.load(fh)
d["name"] = str(f)
# Disambiguate filenames when the same basename exists in both pools:
# use "<dir-name>/<filename>" so picker / driver can map back uniquely.
rel = f"{f.parent.name}/{f.name}"
datas.append((rel, d))
templates = scan_templates()
_, template_requirements = analyze_templates(templates)
# Enumerate ALL templates (engine/chart_type/chart_name) keys
all_template_keys = []
for engine, t_dict in templates.items():
if engine not in args.engines:
continue
for chart_type, c_dict in t_dict.items():
for chart_name in c_dict:
if "base" in chart_name:
continue
all_template_keys.append((engine, chart_type, chart_name))
print(f"Total templates (engines={args.engines}): {len(all_template_keys)}")
# For each data, compute compatible templates ONCE (it returns the full list)
# then we tally per template.
counter = defaultdict(int)
matched_files = defaultdict(list)
for name, data in datas:
compat = check_template_compatibility(data, templates, None)
for tpl_key, _ordered_fields in compat:
counter[tpl_key] += 1
matched_files[tpl_key].append(name)
os.makedirs(os.path.dirname(args.out) or ".", exist_ok=True)
with open(args.out, "w", newline="") as fh:
w = csv.writer(fh)
w.writerow(["engine", "chart_type", "chart_name", "num_compatible_data"])
for engine, chart_type, chart_name in all_template_keys:
tpl_key = f"{engine}/{chart_type}/{chart_name}"
w.writerow([engine, chart_type, chart_name, counter.get(tpl_key, 0)])
print(f"Wrote {args.out}")
# Distribution histogram
buckets = [0, 1, 5, 10, 20, 30, 50, 100, 10**9]
bucket_labels = ["0", "1-4", "5-9", "10-19", "20-29", "30-49", "50-99", "100+"]
bucket_counts = [0] * len(bucket_labels)
for engine, chart_type, chart_name in all_template_keys:
n = counter.get(f"{engine}/{chart_type}/{chart_name}", 0)
for i in range(len(bucket_labels)):
if buckets[i] <= n < buckets[i + 1]:
bucket_counts[i] += 1
break
print()
print("Templates by number of compatible data files in",
", ".join(str(d) for d in data_dirs))
for lab, c in zip(bucket_labels, bucket_counts):
print(f" {lab:>8s}: {c}")
n_geq_20 = sum(1 for k in all_template_keys
if counter.get("/".join(k), 0) >= 20)
n_geq_5 = sum(1 for k in all_template_keys
if counter.get("/".join(k), 0) >= 5)
print()
print(f"Templates with >= 20 compatible data: {n_geq_20}")
print(f"Templates with >= 5 compatible data: {n_geq_5}")
if __name__ == "__main__":
main()