Mizan-Rerank-v3 / benchmark /plot_results.py
ALJIACHI's picture
Model card: focus on accuracy and speed for its size
742ee3c verified
Raw History Blame Contribute Delete
7.99 kB
"""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())