"""Cross-encoder reranker over ONNX Runtime (local, key-free).""" import json import os import numpy as np import onnxruntime as ort from huggingface_hub import hf_hub_download from tokenizers import Tokenizer RERANK_REPO = "Xenova/bge-reranker-base" RERANK_ONNX = "onnx/model_quantized.onnx" # int8: ~3x faster on CPU, negligible quality loss MAX_TOKENS = 512 _MAX_DOC_CHARS = int(os.environ.get("CANLEX_RERANK_DOC_CHARS", "2000")) # 1000 -> 2000 adopted 2026-07-23: the doubled reading # window held every 203-Q eval metric and lifted Hit@5 # +0.01; the cross-encoder now judges a full-size piece # rather than 60% of one. # cap doc text before tokenizing. Cross-encoder cost is # ~linear in tokens per doc, so halving this ~halves the # dominant query cost; env-sweepable, eval-gated (a # tighter cap scores long sections on less of their text # -- fine in practice: the median section is ~520 chars, # and sub-chunked pieces cap at 1800). Swept 2026-06: # 3000 -> 1000 (with pool 50 -> 16) held the eval exactly. def make_ort_session(model_path): """A CPU InferenceSession tuned for shared, small-vCPU hosts. Two deviations from onnxruntime defaults: (1) spin-waiting between requests is disabled -- by default ORT worker threads busy-spin after an inference for a microsecond-scale latency win, which on the Space's shared 2 vCPUs burns quota between queries; (2) the intra-op thread count is env-configurable (CANLEX_ORT_THREADS) so concurrent sessions can be capped below core count if contention is ever observed. Unset, ORT's default (all cores) is kept, so single-query numerics and speed are unchanged.""" opts = ort.SessionOptions() threads = int(os.environ.get("CANLEX_ORT_THREADS", "0")) if threads: opts.intra_op_num_threads = threads opts.add_session_config_entry("session.intra_op.allow_spinning", "0") opts.add_session_config_entry("session.inter_op.allow_spinning", "0") return ort.InferenceSession(model_path, sess_options=opts, providers=["CPUExecutionProvider"]) class Reranker: """Cross-encoder that scores (query, section) pairs for relevance. Loads a BGE reranker as ONNX and runs it on CPU -- no API key; the model is downloaded once and cached. A cross-encoder reads the query and section jointly, so it judges true relevance far better than the BM25/embedding similarities used to build the candidate pool. """ def __init__(self): model_path = hf_hub_download(RERANK_REPO, RERANK_ONNX) tok_path = hf_hub_download(RERANK_REPO, "tokenizer.json") cfg_path = hf_hub_download(RERANK_REPO, "config.json") with open(cfg_path, encoding="utf-8") as fh: self.pad_id = json.load(fh).get("pad_token_id", 0) self.session = make_ort_session(model_path) self.input_names = {i.name for i in self.session.get_inputs()} self.tokenizer = Tokenizer.from_file(tok_path) # only_second keeps the query intact and truncates the (longer) section. self.tokenizer.enable_truncation(max_length=MAX_TOKENS, strategy="only_second") def score(self, query, documents): """Return a relevance logit for each document paired with the query. Higher means more relevant. The returned list is aligned with `documents`. """ if not documents: return [] encs = self.tokenizer.encode_batch( [(query, doc[:_MAX_DOC_CHARS]) for doc in documents]) width = max(len(e.ids) for e in encs) input_ids = np.full((len(encs), width), self.pad_id, dtype=np.int64) attention = np.zeros((len(encs), width), dtype=np.int64) type_ids = np.zeros((len(encs), width), dtype=np.int64) for row, enc in enumerate(encs): n = len(enc.ids) input_ids[row, :n] = enc.ids attention[row, :n] = enc.attention_mask type_ids[row, :n] = enc.type_ids feed = {"input_ids": input_ids, "attention_mask": attention} if "token_type_ids" in self.input_names: feed["token_type_ids"] = type_ids logits = self.session.run(None, feed)[0] return np.asarray(logits, dtype=np.float32).reshape(-1).tolist() def main(): import time print(f"Loading {RERANK_REPO} ({RERANK_ONNX}) ...") reranker = Reranker() query = "powers of arrest without warrant" docs = [ "Arrest without warrant. A peace officer may arrest without warrant a " "person who has committed a criminal offence.", "Definitions. In this Act, fish means any fish and includes shellfish, " "crustaceans and marine animals.", "Importation. It is prohibited to import cannabis except as authorized " "under this Act.", ] start = time.perf_counter() scores = reranker.score(query, docs) elapsed = (time.perf_counter() - start) * 1000 print(f"\nQuery: {query!r}") for doc, score in sorted(zip(docs, scores), key=lambda x: x[1], reverse=True): print(f" {score:8.3f} {doc[:62]}") print(f"\n{elapsed:.0f} ms for {len(docs)} documents") if __name__ == "__main__": main()