Text Ranking
MLX
Safetensors
English
qwen3
reranker
cross-encoder
apple-silicon
finance
legal
code
stem
medical
8-bit precision
Instructions to use fcmeyer/zerank-2-reranker-MLX-8bit with libraries, inference providers, notebooks, and local apps. Follow these links to get started.
- Libraries
- MLX
How to use fcmeyer/zerank-2-reranker-MLX-8bit with MLX:
# Download the model from the Hub pip install huggingface_hub[hf_xet] hf download fcmeyer/zerank-2-reranker-MLX-8bit --local-dir zerank-2-reranker-MLX-8bit
- Notebooks
- Google Colab
- Kaggle
- Local Apps Settings
- LM Studio
- Atomic Chat
Download rerank.py from fcmeyer/zerank-2-reranker-MLX-8bit: direct link, hf CLI and curl.
- Browser
- Download file 5.01 kB
-
https://huggingface.co/fcmeyer/zerank-2-reranker-MLX-8bit/resolve/main/rerank.py
- Command line
-
hf download hf://fcmeyer/zerank-2-reranker-MLX-8bit/rerank.py
-
curl -L -o rerank.py https://huggingface.co/fcmeyer/zerank-2-reranker-MLX-8bit/resolve/main/rerank.py
5.01 kB
| """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}") | |