Text Ranking
sentence-transformers
Safetensors
Arabic
new
cross-encoder
reranker
arabic
long-context
rag
islamic
custom_code
text-embeddings-inference
Instructions to use ALJIACHI/Mizan-Rerank-v3 with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- sentence-transformers
How to use ALJIACHI/Mizan-Rerank-v3 with sentence-transformers:
from sentence_transformers import CrossEncoder model = CrossEncoder("ALJIACHI/Mizan-Rerank-v3", trust_remote_code=True) query = "Which planet is known as the Red Planet?" passages = [ "Venus is often called Earth's twin because of its similar size and proximity.", "Mars, known for its reddish appearance, is often referred to as the Red Planet.", "Jupiter, the largest planet in our solar system, has a prominent red spot.", "Saturn, famous for its rings, is sometimes mistaken for the Red Planet." ] scores = model.predict([(query, passage) for passage in passages]) print(scores) - Notebooks
- Google Colab
- Kaggle
Download benchmark/plot_results.py from ALJIACHI/Mizan-Rerank-v3: direct link, hf CLI and curl.
- Browser
- Download file 7.99 kB
-
https://huggingface.co/ALJIACHI/Mizan-Rerank-v3/resolve/main/benchmark/plot_results.py
- Command line
-
hf download hf://ALJIACHI/Mizan-Rerank-v3/benchmark/plot_results.py
-
curl -L -o plot_results.py https://huggingface.co/ALJIACHI/Mizan-Rerank-v3/resolve/main/benchmark/plot_results.py
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()) | |