"""Render the model-card charts from results_full.json and internal_results.json into ../assets/.""" import argparse import json from pathlib import Path import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt # noqa: E402 SCRIPT_DIR = Path(__file__).resolve().parent ASSETS = SCRIPT_DIR.parent / "assets" ORDER = ("mizan-v3", "mizan-v2", "gte-base", "bge-v2-m3") LABELS = { "mizan-v3": "Mizan-Rerank-v3 (306M)", "mizan-v2": "Mizan-Rerank-V2 (306M)", "gte-base": "gte-multilingual-reranker-base (306M)", "bge-v2-m3": "bge-reranker-v2-m3 (568M)", "retriever": "Retriever only (no reranker)", } COLORS = {"mizan-v3": "#0F766E", "mizan-v2": "#94A3B8", "gte-base": "#CBD5E1", "bge-v2-m3": "#F59E0B", "retriever": "#E2E8F0"} DATASET_LABELS = { "namaa_mrtydi": "NamaaMrTydi\n(MTEB, full)", "namaa_mrtydi_unseen": "NamaaMrTydi\n(unseen subset)", "arabic_hard_negatives": "Arabic Hard Negatives\n(full)", "arabic_hard_negatives_unseen": "Arabic Hard Negatives\n(unseen subset)", "gemini_test": "Adversarial short\n(Gemini test)", "longctx_test": "Adversarial long-context\n(test)", "multi_llm": "Multi-LLM\n(6 generators)", } HELD_OUT = ("namaa_mrtydi_unseen", "arabic_hard_negatives_unseen", "gemini_test", "longctx_test", "multi_llm") def style(ax, title: str, ylabel: str) -> None: ax.set_title(title, fontsize=13, fontweight="bold", loc="left") ax.set_ylabel(ylabel) ax.spines[["top", "right"]].set_visible(False) ax.grid(axis="y", alpha=0.3) ax.set_axisbelow(True) def efficiency_chart(results: dict, speed: dict, path: Path) -> None: quality = {key: sum(results[key]["datasets"][dataset]["ndcg@10"] for dataset in HELD_OUT) / len(HELD_OUT) for key in ORDER} fig, ax = plt.subplots(figsize=(8.5, 4.6), dpi=150) for key in ORDER: row = speed["results"][key] ax.scatter(row["pairs_per_second"], quality[key], s=row["parameters"] / 1e6 * 2.2, color=COLORS[key], edgecolor="#334155", linewidth=0.8, zorder=3) offset = (-12, 0.006) if key == "gte-base" else (12, -0.004) ax.annotate(f"{LABELS[key]}\n{quality[key]:.3f} nDCG@10 ยท {row['pairs_per_second']:.0f} pairs/s", (row["pairs_per_second"], quality[key]), xytext=offset, textcoords="offset points", fontsize=8.5, ha="right" if key == "gte-base" else "left", va="center", fontweight="bold" if key == "mizan-v3" else None) rates = [speed["results"][key]["pairs_per_second"] for key in ORDER] ax.set_xlim(min(rates) * 0.6, max(rates) * 1.5) ax.set_ylim(min(quality.values()) - 0.03, max(quality.values()) + 0.03) settings = speed["settings"] ax.set_xlabel(f"throughput, query-passage pairs/s ({settings['device']}, {'fp16' if settings['fp16'] else 'fp32'}, max_length {settings['max_length']})") style(ax, "Accuracy vs. speed (bubble size = parameters; up and right is better)", "held-out mean nDCG@10") ax.grid(axis="x", alpha=0.3) fig.tight_layout() fig.savefig(path) plt.close(fig) def grouped_bars(results: dict, datasets: tuple[str, ...], metric: str, title: str, path: Path, floor: float) -> None: fig, ax = plt.subplots(figsize=(11, 5.2), dpi=150) width = 0.8 / len(ORDER) for position, key in enumerate(ORDER): values = [results[key]["datasets"][dataset][metric] for dataset in datasets] bars = ax.bar([index + (position - (len(ORDER) - 1) / 2) * width for index in range(len(datasets))], values, width, label=LABELS[key], color=COLORS[key], edgecolor="white") for bar, value in zip(bars, values): ax.text(bar.get_x() + bar.get_width() / 2, value + 0.004, f"{value:.3f}", ha="center", va="bottom", fontsize=7.5, rotation=90) ax.set_xticks(range(len(datasets)), [DATASET_LABELS[dataset] for dataset in datasets], fontsize=9) ax.set_ylim(floor, 1.0) style(ax, title, metric.replace("ndcg", "nDCG").replace("hit", "Hit").replace("mrr", "MRR")) ax.legend(frameon=False, fontsize=8.5, ncol=2, loc="upper right") fig.tight_layout() fig.savefig(path) plt.close(fig) def summary_chart(results: dict, fiqh: dict, path: Path) -> None: held_out = {key: sum(results[key]["datasets"][dataset]["ndcg@10"] for dataset in HELD_OUT) / len(HELD_OUT) for key in ORDER} panels = ( ("Held-out benchmarks\n(mean nDCG@10 over 5 sets)", held_out), ("Real fiqh search queries\n(nDCG@10, 224 queries)", {key: fiqh[key]["ndcg@10"] for key in ORDER}), ) fig, axes = plt.subplots(1, 2, figsize=(12, 3.8), dpi=150) for ax, (title, values) in zip(axes, panels): keys = sorted(ORDER, key=lambda key: values[key]) bars = ax.barh([LABELS[key] for key in keys], [values[key] for key in keys], color=[COLORS[key] for key in keys]) for bar, key in zip(bars, keys): ax.text(bar.get_width() + 0.003, bar.get_y() + bar.get_height() / 2, f"{values[key]:.3f}", va="center", fontsize=9, fontweight="bold" if key == "mizan-v3" else None) ax.set_xlim(max(0.0, min(values.values()) - 0.08), max(values.values()) + 0.04) ax.set_title(title, fontsize=11, fontweight="bold", loc="left") ax.spines[["top", "right"]].set_visible(False) ax.grid(axis="x", alpha=0.3) ax.set_axisbelow(True) fig.tight_layout() fig.savefig(path) plt.close(fig) def training_chart(curve: dict, path: Path) -> None: fig, ax = plt.subplots(figsize=(8, 4.2), dpi=150) ax.plot(curve["epoch"], curve["longctx"], marker="o", color="#0F766E", label="Long-context adversarial (validation)") ax.plot(curve["epoch"], curve["gemini"], marker="s", color="#F59E0B", label="Short adversarial (validation)") best = max(range(len(curve["epoch"])), key=lambda index: curve["longctx"][index] + curve["gemini"][index]) ax.axvline(curve["epoch"][best], color="#64748B", linestyle="--", linewidth=1) ax.text(curve["epoch"][best] - 0.1, curve["longctx"][best] - 0.012, f"released checkpoint\n(epoch {curve['epoch'][best]})", fontsize=8.5, color="#334155", ha="right", va="top") ax.set_xticks(curve["epoch"], ["gte base"] + [str(epoch) for epoch in curve["epoch"][1:]]) ax.set_xlabel("epoch") style(ax, "Validation nDCG@10 during training", "nDCG@10") ax.legend(frameon=False, fontsize=9, loc="lower right") fig.tight_layout() fig.savefig(path) plt.close(fig) def main() -> int: parser = argparse.ArgumentParser(description=__doc__) parser.add_argument("--results", type=Path, default=SCRIPT_DIR / "results_full.json") parser.add_argument("--internal", type=Path, default=SCRIPT_DIR / "internal_results.json") parser.add_argument("--output-dir", type=Path, default=ASSETS) args = parser.parse_args() results = json.loads(args.results.read_text(encoding="utf-8"))["results"] internal = json.loads(args.internal.read_text(encoding="utf-8")) args.output_dir.mkdir(parents=True, exist_ok=True) summary_chart(results, internal["fiqh_eval"]["models"], args.output_dir / "summary.png") speed_file = SCRIPT_DIR / "speed_results.json" if speed_file.is_file(): efficiency_chart(results, json.loads(speed_file.read_text(encoding="utf-8")), args.output_dir / "efficiency.png") grouped_bars(results, ("namaa_mrtydi", "namaa_mrtydi_unseen", "arabic_hard_negatives_unseen"), "ndcg@10", "Public Arabic reranking benchmarks", args.output_dir / "public_benchmarks.png", 0.6) grouped_bars(results, ("gemini_test", "longctx_test", "multi_llm"), "ndcg@10", "Adversarial Arabic reranking (held-out test sets)", args.output_dir / "adversarial_benchmarks.png", 0.5) grouped_bars(results, HELD_OUT, "hit@1", "Hit@1: correct passage ranked first", args.output_dir / "hit_at_1.png", 0.0) training_chart(internal["training_curve"], args.output_dir / "training_curve.png") for path in sorted(args.output_dir.glob("*.png")): print(path) return 0 if __name__ == "__main__": raise SystemExit(main())