Mizan-Rerank-v3 / benchmark /benchmark_rerankers.py
ALJIACHI's picture
Release Mizan-Rerank-v3
ea26409 verified
Raw History Blame Contribute Delete
12.9 kB
"""Benchmark Arabic rerankers on public Hugging Face reranking sets.
Compares Mizan-Rerank-v3 with Mizan-Rerank-V2, gte-multilingual-reranker-base and
bge-reranker-v2-m3 (any other cross-encoder can be added with --model name=repo_or_path).
Datasets (downloaded from the Hub at pinned revisions):
namaa_mrtydi mteb/NamaaMrTydiReranking (MTEB) 918 queries, 1 positive + ~4 hard negatives
arabic_hard_negatives Omartificial-Intelligence-Space/
Arabic-With-Ranked-Hard-Negatives 12,373 queries, 1 positive + 4 mined negatives
Both are also reported on their "unseen" subset: cases whose query or positive passage does
not occur in the Mizan-Rerank-v3 training pool (hashes in seen_in_training.json). Your own
listwise JSONL files can be added with --listwise name=path.jsonl (fields: query, positive,
optional partial_or_neutral, hard_negatives as strings or {"text": ...}).
Metrics per query (graded gains: positive 3, partial 1, negative 0):
nDCG@10, MRR@10 (rank of the positive), Hit@1 (positive ranked first).
"average" = unweighted mean over the --average datasets (default: every dataset, using the
_unseen subset instead of the full set wherever one exists, so no query is counted twice).
pip install torch transformers huggingface_hub pyarrow
python benchmark_rerankers.py # all four models, public sets
python benchmark_rerankers.py --models mizan-v3 --fp16
python benchmark_rerankers.py --model mine=./my_reranker --models mine bge-v2-m3
"""
import argparse
import gc
import hashlib
import json
import math
import statistics
import time
from dataclasses import dataclass
from pathlib import Path
import pyarrow.parquet as pq
import torch
from huggingface_hub import snapshot_download
from transformers import AutoModelForSequenceClassification, AutoTokenizer
SCRIPT_DIR = Path(__file__).resolve().parent
SEEN_FILE = SCRIPT_DIR / "seen_in_training.json"
MODELS = {
"mizan-v3": "ALJIACHI/Mizan-Rerank-v3",
"mizan-v2": "ALJIACHI/Mizan-Rerank-V2",
"gte-base": "Alibaba-NLP/gte-multilingual-reranker-base",
"bge-v2-m3": "BAAI/bge-reranker-v2-m3",
}
HUB_DATASETS = {
"namaa_mrtydi": ("mteb/NamaaMrTydiReranking", "ecf1ff8c18469251bfe6acf1219e63e287bbb037"),
"arabic_hard_negatives": ("Omartificial-Intelligence-Space/Arabic-With-Ranked-Hard-Negatives", "fec243e3767f502d393d330bf97bfb6dce868830"),
}
METRICS = ("ndcg@10", "mrr@10", "hit@1")
@dataclass(frozen=True)
class Group:
query: str
documents: tuple[str, ...]
gains: tuple[int, ...]
def normalize(text: str) -> str:
return " ".join(str(text).split())
def digest(text: str) -> str:
return hashlib.sha256(normalize(text).encode("utf-8")).hexdigest()
def read_parquet_dir(directory: Path) -> list[dict]:
files = sorted(directory.glob("*.parquet"))
if not files:
raise FileNotFoundError(f"No parquet files in {directory}")
return [row for path in files for row in pq.read_table(path).to_pylist()]
def download(repo: str, revision: str | None) -> Path:
return Path(snapshot_download(repo, repo_type="dataset", revision=revision))
def load_namaa_mrtydi(root: Path) -> list[Group]:
corpus = {row["_id"]: row["text"] or "" for row in read_parquet_dir(root / "corpus")}
queries = {row["_id"]: row["text"] for row in read_parquet_dir(root / "queries")}
relevant = {(row["query-id"], row["corpus-id"]) for row in read_parquet_dir(root / "data") if row["score"] > 0}
groups = []
for row in read_parquet_dir(root / "top_ranked"):
ids = row["corpus-ids"]
groups.append(Group(queries[row["query-id"]], tuple(corpus[i] for i in ids), tuple(3 if (row["query-id"], i) in relevant else 0 for i in ids)))
return groups
def load_arabic_hard_negatives(root: Path) -> list[Group]:
groups = []
for row in read_parquet_dir(root / "data"):
negatives = [row[f"negative{i}"] for i in range(1, 5) if row.get(f"negative{i}") and str(row[f"negative{i}"]).strip()]
groups.append(Group(row["query"], (row["positive"], *negatives), (3, *([0] * len(negatives)))))
return groups
def load_listwise(path: Path) -> list[Group]:
groups = []
for line in path.read_text(encoding="utf-8").splitlines():
if not line.strip():
continue
record = json.loads(line)
documents, gains = [record["positive"]], [3]
if record.get("partial_or_neutral"):
documents.append(record["partial_or_neutral"])
gains.append(1)
for negative in record["hard_negatives"]:
documents.append(negative["text"] if isinstance(negative, dict) else negative)
gains.append(0)
groups.append(Group(record["query"], tuple(documents), tuple(gains)))
return groups
def unseen_indices(groups: list[Group], seen: dict) -> list[int]:
queries, passages = set(seen["query_sha256"]), set(seen["passage_sha256"])
return [
index for index, group in enumerate(groups)
if digest(group.query) not in queries and digest(group.documents[group.gains.index(3)]) not in passages
]
def dcg(gains: list[int], k: int) -> float:
return sum((2**gain - 1) / math.log2(rank + 2) for rank, gain in enumerate(gains[:k]))
def group_metrics(group: Group, scores: list[float]) -> dict[str, float]:
order = sorted(range(len(scores)), key=lambda index: scores[index], reverse=True)
ranked = [group.gains[index] for index in order]
ideal = dcg(sorted(group.gains, reverse=True), 10)
rank = ranked.index(3) + 1
return {"ndcg@10": dcg(ranked, 10) / ideal, "mrr@10": 1.0 / rank if rank <= 10 else 0.0, "hit@1": float(rank == 1)}
def score(model_path: str, groups: list[Group], device: str, batch_size: int, max_length: int, fp16: bool) -> list[list[float]]:
tokenizer = AutoTokenizer.from_pretrained(model_path, trust_remote_code=True)
dtype = torch.float16 if fp16 else torch.float32
model = AutoModelForSequenceClassification.from_pretrained(model_path, trust_remote_code=True, torch_dtype=dtype).to(device).eval()
pairs = [(group.query, document) for group in groups for document in group.documents]
order = sorted(range(len(pairs)), key=lambda index: len(pairs[index][0]) + len(pairs[index][1]))
flat = [0.0] * len(pairs)
with torch.inference_mode():
for start in range(0, len(order), batch_size):
batch = order[start : start + batch_size]
features = tokenizer(
[pairs[index][0] for index in batch], [pairs[index][1] for index in batch],
padding=True, truncation=True, max_length=max_length, return_tensors="pt",
).to(device)
logits = model(**features).logits.float()
logits = logits.squeeze(-1) if logits.shape[-1] == 1 else logits[:, -1]
for index, value in zip(batch, logits.cpu().tolist()):
flat[index] = value
del model
gc.collect()
if torch.cuda.is_available():
torch.cuda.empty_cache()
grouped, offset = [], 0
for group in groups:
grouped.append(flat[offset : offset + len(group.documents)])
offset += len(group.documents)
return grouped
def summarize(rows: list[dict]) -> dict:
return {"queries": len(rows), **{metric: statistics.fmean(row[metric] for row in rows) for metric in METRICS}}
def print_tables(results: dict, datasets: list[str], average: list[str]) -> None:
for metric in METRICS:
header = "| Model | " + " | ".join(datasets) + " | **Average** |"
print(f"\n### {metric}\n\n{header}\n|" + "---|" * (len(datasets) + 2))
for name, row in sorted(results.items(), key=lambda item: -item[1]["average"][metric]):
cells = " | ".join(f"{row['datasets'][dataset][metric]:.4f}" for dataset in datasets)
print(f"| {name} | {cells} | **{row['average'][metric]:.4f}** |")
print(f"\nAverage = unweighted mean over: {', '.join(average)}")
def parse_named(values: list[str], flag: str) -> dict[str, str]:
named = {}
for value in values or []:
name, _, target = value.partition("=")
if not name or not target:
raise ValueError(f"{flag} expects name=value, got {value!r}")
named[name] = target
return named
def build_parser() -> argparse.ArgumentParser:
parser = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
parser.add_argument("--model", action="append", help="name=repo_or_path; adds or overrides a model")
parser.add_argument("--models", nargs="+", help="Model names to run (default: all)")
parser.add_argument("--datasets", nargs="+", choices=list(HUB_DATASETS), default=list(HUB_DATASETS))
parser.add_argument("--listwise", action="append", help="name=path.jsonl extra listwise dataset")
parser.add_argument("--no-unseen", action="store_true", help="Skip the *_unseen subsets")
parser.add_argument("--average", nargs="+", help="Datasets included in the average (default: all scored)")
parser.add_argument("--device", default="cuda" if torch.cuda.is_available() else "cpu")
parser.add_argument("--batch-size", type=int, default=32)
parser.add_argument("--max-length", type=int, default=2048)
parser.add_argument("--fp16", action="store_true", help="Score in float16 (faster; results differ slightly from float32)")
parser.add_argument("--limit", type=int, default=None, help="Use only the first N queries of each dataset (smoke tests)")
parser.add_argument("--output", type=Path, default=SCRIPT_DIR / "results.json")
return parser
def main() -> int:
args = build_parser().parse_args()
models = {**MODELS, **parse_named(args.model, "--model")}
selected = args.models or list(models)
unknown = [name for name in selected if name not in models]
if unknown:
raise ValueError(f"Unknown model(s) {unknown}; known: {sorted(models)}")
datasets: dict[str, list[Group]] = {}
subsets: dict[str, tuple[str, list[int]]] = {}
loaders = {"namaa_mrtydi": load_namaa_mrtydi, "arabic_hard_negatives": load_arabic_hard_negatives}
seen = json.loads(SEEN_FILE.read_text(encoding="utf-8")) if SEEN_FILE.is_file() and not args.no_unseen else None
for name in args.datasets:
datasets[name] = loaders[name](download(*HUB_DATASETS[name]))
for name, path in parse_named(args.listwise, "--listwise").items():
datasets[name] = load_listwise(Path(path))
if args.limit:
datasets = {name: groups[: args.limit] for name, groups in datasets.items()}
if seen:
subsets = {f"{name}_unseen": (name, unseen_indices(datasets[name], seen)) for name in args.datasets}
columns = [column for name in datasets for column in (name, f"{name}_unseen") if column in datasets or column in subsets]
for column in columns:
size = len(subsets[column][1]) if column in subsets else len(datasets[column])
print(f"{column:<32} {size:>6} queries")
average = args.average or [column for column in columns if f"{column}_unseen" not in subsets]
missing = [name for name in average if name not in columns]
if missing:
raise ValueError(f"--average names not scored: {missing}; scored: {columns}")
results = {}
for name in selected:
started = time.perf_counter()
print(f"\nScoring {name} ({models[name]})")
per_dataset = {}
for dataset, groups in datasets.items():
rows = [group_metrics(group, values) for group, values in zip(groups, score(models[name], groups, args.device, args.batch_size, args.max_length, args.fp16))]
per_dataset[dataset] = summarize(rows)
subset = subsets.get(f"{dataset}_unseen")
if subset:
per_dataset[f"{dataset}_unseen"] = summarize([rows[index] for index in subset[1]])
for column in columns:
print(f" {column:<30} " + " ".join(f"{metric} {per_dataset[column][metric]:.4f}" for metric in METRICS))
results[name] = {
"model": models[name],
"seconds": time.perf_counter() - started,
"datasets": per_dataset,
"average": {metric: statistics.fmean(per_dataset[dataset][metric] for dataset in average) for metric in METRICS},
}
print_tables(results, columns, average)
args.output.write_text(json.dumps({
"settings": {"max_length": args.max_length, "batch_size": args.batch_size, "fp16": args.fp16, "device": args.device, "limit": args.limit},
"datasets": {column: len(subsets[column][1]) if column in subsets else len(datasets[column]) for column in columns},
"average_over": average,
"results": results,
}, indent=2), encoding="utf-8")
print(f"\nResults written to {args.output}")
return 0
if __name__ == "__main__":
raise SystemExit(main())