"""Rerank query/document pairs with zerank-2-reranker on MLX. pip install mlx-lm transformers from rerank import Reranker reranker = Reranker("fcmeyer/zerank-2-reranker-mlx-bf16") for document, score in reranker.rank("What is 2+2?", ["4", "1 million"]): print(f"{score:+.3f} {document}") The model is a cross-encoder: it reads the query and one document together and returns a single relevance logit. Higher is more relevant. Scores are comparable within one query, not across queries; `sigmoid(score / 5)` maps them to a 0-1 probability, matching the upstream model card. """ from __future__ import annotations import math import mlx.core as mx from mlx_lm import load from transformers import AutoTokenizer # From 1_LogitScore/config.json: the score is this token's logit at the final # position. Token 9454 decodes to "Yes". TRUE_TOKEN_ID = 9454 # The documented context length from the upstream model card. MAX_LENGTH = 32768 # Cap on batch_size * padded_length, which bounds peak memory regardless of # whether you pass many short documents or one very long one. TOKEN_BUDGET = 32768 class Reranker: def __init__(self, model_path: str, max_length: int = MAX_LENGTH): self.model, _ = load(model_path) self.tokenizer = AutoTokenizer.from_pretrained(model_path) self.max_length = max_length def _render(self, query: str, document: str) -> str: return ( "<|im_start|>system\n" + query + "<|im_end|>\n" "<|im_start|>user\n" + document + "<|im_end|>\n" "<|im_start|>assistant\n" ) def _encode(self, query: str, document: str) -> list[int]: ids = self.tokenizer( self._render(query, document), add_special_tokens=False ).input_ids if len(ids) <= self.max_length: return ids # Truncate the document, never the query, and cut on a character # boundary so the trailing template tokens stay intact. lo, hi = 0, max(len(document) - (len(ids) - self.max_length), 0) while lo < hi: mid = (lo + hi + 1) // 2 n = len(self.tokenizer( self._render(query, document[:mid]), add_special_tokens=False ).input_ids) lo, hi = (mid, hi) if n <= self.max_length else (lo, mid - 1) return self.tokenizer( self._render(query, document[:lo]), add_special_tokens=False ).input_ids def _forward(self, batch: list[list[int]]) -> list[float]: lengths = [len(ids) for ids in batch] width = max(lengths) pad = self.tokenizer.pad_token_id padded = mx.array([ids + [pad] * (width - len(ids)) for ids in batch]) # Right padding is safe: attention is causal, so the last real token # never attends to the pads after it. hidden = self.model.model(padded) # Project only the final position of each row. Materialising logits for # the whole sequence would allocate batch * length * 151936 floats. last = hidden[mx.arange(len(batch)), mx.array(lengths) - 1] logits = self.model.model.embed_tokens.as_linear(last) scores = logits[:, TRUE_TOKEN_ID].astype(mx.float32) mx.eval(scores) return scores.tolist() def score(self, pairs: list[tuple[str, str]]) -> list[float]: """Score (query, document) pairs, returning one relevance logit each.""" encoded = [self._encode(q, d) for q, d in pairs] results: list[float | None] = [None] * len(encoded) batch: list[int] = [] for index in sorted(range(len(encoded)), key=lambda i: len(encoded[i])): width = max(len(encoded[i]) for i in batch + [index]) if batch and (len(batch) + 1) * width > TOKEN_BUDGET: for i, s in zip(batch, self._forward([encoded[i] for i in batch])): results[i] = s batch = [] batch.append(index) for i, s in zip(batch, self._forward([encoded[i] for i in batch])): results[i] = s return results # type: ignore[return-value] def rank( self, query: str, documents: list[str], top_k: int | None = None ) -> list[tuple[str, float]]: """Score documents against one query and return them best-first.""" scores = self.score([(query, d) for d in documents]) ranked = sorted(zip(documents, scores), key=lambda x: -x[1]) return ranked[:top_k] if top_k else ranked def probability(score: float) -> float: """Map a relevance logit to the 0-1 range used by the upstream model card.""" return 1.0 / (1.0 + math.exp(-score / 5)) if __name__ == "__main__": import sys reranker = Reranker(sys.argv[1] if len(sys.argv) > 1 else "out/bf16") query = "What is 2+2?" for document, score in reranker.rank( query, ["4", "The answer is definitely 1 million"] ): print(f"{score:+.4f} (p={probability(score):.3f}) {document!r}")