fcmeyer's picture
Upload folder using huggingface_hub
313c369 verified
Raw History Blame Contribute Delete
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}")