use anyhow::{Context, Result};
use ort::session::Session;
use ort::value::Tensor;
use parking_lot::Mutex;
use std::path::{Path, PathBuf};
use std::sync::OnceLock;
use tokenizers::Tokenizer;
use tokenizers::tokenizer::TruncationParams;
#[derive(Debug, Clone)]
pub struct PromptTemplate {
pub prefix: String,
pub middle: String,
pub suffix: String,
}
impl PromptTemplate {
#[must_use]
pub fn render(&self, query: &str, document: &str) -> String {
format!(
"{}{}{}{}{}",
self.prefix, query, self.middle, document, self.suffix
)
}
}
#[derive(Debug, Clone)]
pub enum ScoreStrategy {
ClassifierLogit,
YesNoLogit {
yes_id: usize,
no_id: usize,
prompt: PromptTemplate,
},
}
pub(crate) fn classifier_score_from_logits(logits: &[f32]) -> Result<f32> {
logits
.last()
.copied()
.context("classifier reranker produced an empty logits tensor")
}
pub(crate) fn yes_no_score_from_logits(
final_pos_logits: &[f32],
yes_id: usize,
no_id: usize,
) -> Result<f32> {
let l_yes = *final_pos_logits.get(yes_id).with_context(|| {
format!(
"yes_id {yes_id} out of range (vocab {})",
final_pos_logits.len()
)
})?;
let l_no = *final_pos_logits.get(no_id).with_context(|| {
format!(
"no_id {no_id} out of range (vocab {})",
final_pos_logits.len()
)
})?;
let m = l_yes.max(l_no);
let e_yes = (l_yes - m).exp();
let e_no = (l_no - m).exp();
Ok(e_yes / (e_yes + e_no))
}
pub struct OnnxReranker {
model_dir: PathBuf,
session_file: String,
tokenizer: Tokenizer,
strategy: ScoreStrategy,
max_context: usize,
threads: usize,
use_coreml: bool,
session: OnceLock<Mutex<Session>>,
build_lock: std::sync::Mutex<()>,
}
impl OnnxReranker {
pub fn new(
model_dir: &Path,
session_file: &str,
strategy: ScoreStrategy,
max_context: usize,
threads: usize,
use_coreml: bool,
) -> Result<Self> {
let tok_path = model_dir.join("tokenizer.json");
let mut tokenizer = Tokenizer::from_file(&tok_path)
.map_err(|e| anyhow::anyhow!("failed to load tokenizer {}: {e}", tok_path.display()))?;
let max_length = max_context.max(1);
tokenizer
.with_truncation(Some(TruncationParams {
max_length,
..Default::default()
}))
.map_err(|e| anyhow::anyhow!("failed to set tokenizer truncation: {e}"))?;
Ok(Self {
model_dir: model_dir.to_path_buf(),
session_file: session_file.to_string(),
tokenizer,
strategy,
max_context: max_length,
threads: threads.max(1),
use_coreml,
session: OnceLock::new(),
build_lock: std::sync::Mutex::new(()),
})
}
#[allow(clippy::vec_init_then_push)] fn execution_providers(&self) -> Vec<ort::ep::ExecutionProviderDispatch> {
let mut providers = Vec::new();
#[cfg(target_os = "macos")]
if self.use_coreml {
providers.push(ort::ep::CoreML::default().build());
}
#[cfg(not(target_os = "macos"))]
let _ = self.use_coreml;
#[cfg(feature = "cuda")]
{
providers.push(ort::ep::CUDA::default().build());
}
providers.push(ort::ep::CPU::default().build());
providers
}
fn build_session(&self) -> Result<Session> {
let model_path = self.model_dir.join(&self.session_file);
let session = Session::builder()
.context("ort Session::builder failed")?
.with_execution_providers(self.execution_providers())
.context("failed to set execution providers")?
.with_intra_threads(self.threads)
.context("failed to set intra-op threads")?
.commit_from_file(&model_path)
.with_context(|| format!("failed to load ONNX model {}", model_path.display()))?;
Ok(session)
}
fn session(&self) -> Result<&Mutex<Session>> {
if let Some(s) = self.session.get() {
return Ok(s);
}
let _guard = self
.build_lock
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(s) = self.session.get() {
return Ok(s);
}
let built = self.build_session()?;
let _ = self.session.set(Mutex::new(built));
Ok(self.session.get().expect("session set above"))
}
fn encode_pair(&self, query: &str, document: &str) -> Result<(Vec<i64>, Vec<i64>)> {
let (mut ids, mut mask) = match &self.strategy {
ScoreStrategy::ClassifierLogit => {
let enc = self
.tokenizer
.encode((query, document), true)
.map_err(|e| anyhow::anyhow!("tokenizer.encode(pair) failed: {e}"))?;
to_i64_pair(&enc)
}
ScoreStrategy::YesNoLogit { prompt, .. } => {
let text = prompt.render(query, document);
let enc = self
.tokenizer
.encode(text, true)
.map_err(|e| anyhow::anyhow!("tokenizer.encode failed: {e}"))?;
to_i64_pair(&enc)
}
};
if ids.len() > self.max_context {
ids.truncate(self.max_context);
mask.truncate(self.max_context);
}
Ok((ids, mask))
}
pub fn score_pair(&self, query: &str, document: &str) -> Result<f32> {
let (ids, mask) = self.encode_pair(query, document)?;
let seq = ids.len();
let shape = vec![1_i64, seq as i64];
let id_tensor =
Tensor::from_array((shape.clone(), ids)).context("failed to build input_ids tensor")?;
let mask_tensor = Tensor::from_array((shape.clone(), mask))
.context("failed to build attention_mask tensor")?;
let mut inputs = ort::inputs![
"input_ids" => id_tensor,
"attention_mask" => mask_tensor,
];
if matches!(self.strategy, ScoreStrategy::YesNoLogit { .. }) {
let position_ids: Vec<i64> = (0..seq as i64).collect();
let pos_tensor = Tensor::from_array((shape, position_ids))
.context("failed to build position_ids tensor")?;
inputs.push((
std::borrow::Cow::Borrowed("position_ids"),
pos_tensor.into(),
));
}
let session = self.session()?;
let mut guard = session.lock();
let outputs = guard
.run(inputs)
.context("ONNX reranker forward pass failed")?;
let (out_shape, logits) = extract_logits_f32(&outputs[0])?;
match &self.strategy {
ScoreStrategy::ClassifierLogit => classifier_score_from_logits(&logits),
ScoreStrategy::YesNoLogit { yes_id, no_id, .. } => {
let vocab = out_shape[out_shape.len() - 1] as usize;
anyhow::ensure!(vocab > 0, "reranker output has zero-width vocab dim");
let n = logits.len();
anyhow::ensure!(
n >= vocab,
"reranker logits ({n}) shorter than vocab ({vocab})"
);
let final_pos = &logits[n - vocab..];
yes_no_score_from_logits(final_pos, *yes_id, *no_id)
}
}
}
pub fn rerank(
&self,
query: &str,
documents: &[&str],
top_k: usize,
) -> Result<Vec<(usize, f32)>> {
if documents.is_empty() {
return Ok(Vec::new());
}
if !crate::search::fastembed_reranker::reranker_enabled() {
tracing::debug!(
"Reranker disabled (SEMANTEX_RERANKER!=on); returning identity ordering"
);
let n = documents.len().min(top_k);
return Ok((0..n).map(|i| (i, 0.0_f32)).collect());
}
let mut scored: Vec<(usize, f32)> = Vec::with_capacity(documents.len());
for (i, doc) in documents.iter().enumerate() {
let s = self.score_pair(query, doc)?;
scored.push((i, s));
}
scored.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
scored.truncate(top_k);
Ok(scored)
}
}
fn to_i64_pair(enc: &tokenizers::Encoding) -> (Vec<i64>, Vec<i64>) {
let ids = enc.get_ids().iter().map(|&x| i64::from(x)).collect();
let mask = enc
.get_attention_mask()
.iter()
.map(|&x| i64::from(x))
.collect();
(ids, mask)
}
fn extract_logits_f32(value: &ort::value::DynValue) -> Result<(Vec<i64>, Vec<f32>)> {
use ort::tensor::TensorElementType;
match value.data_type() {
TensorElementType::Float32 => {
let (shape, data) = value
.try_extract_tensor::<f32>()
.context("failed to extract f32 logits from reranker output")?;
Ok((shape.iter().copied().collect(), data.to_vec()))
}
TensorElementType::Float16 => {
let (shape, data) = value
.try_extract_tensor::<half::f16>()
.context("failed to extract f16 logits from reranker output")?;
let floats: Vec<f32> = data.iter().map(|h| h.to_f32()).collect();
Ok((shape.iter().copied().collect(), floats))
}
other => anyhow::bail!(
"reranker logits output has unsupported dtype {other:?} (expected f32 or f16)"
),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn classifier_score_reads_last_logit() {
assert_eq!(classifier_score_from_logits(&[0.42]).unwrap(), 0.42);
assert_eq!(classifier_score_from_logits(&[-1.0, 3.5]).unwrap(), 3.5);
}
#[test]
fn classifier_score_errors_on_empty() {
assert!(classifier_score_from_logits(&[]).is_err());
}
#[test]
fn yes_no_score_is_monotonic_and_bounded() {
let logits = [0.0, 0.0, 2.0, 1.0];
let s = yes_no_score_from_logits(&logits, 2, 3).unwrap();
assert!((s - 0.731_058_6).abs() < 1e-5, "got {s}");
assert!(yes_no_score_from_logits(&[0.0, 0.0, 10.0, 0.0], 2, 3).unwrap() > 0.99);
assert!(yes_no_score_from_logits(&[0.0, 0.0, 0.0, 10.0], 2, 3).unwrap() < 0.01);
}
#[test]
fn yes_no_score_errors_on_oob_token() {
assert!(yes_no_score_from_logits(&[0.1, 0.2], 5, 0).is_err());
}
#[test]
fn prompt_template_renders_in_order() {
let t = PromptTemplate {
prefix: "<I>".into(),
middle: "<M>".into(),
suffix: "<S>".into(),
};
assert_eq!(t.render("Q", "D"), "<I>Q<M>D<S>");
}
#[test]
fn new_rejects_missing_dir_but_is_lazy_for_files() {
let tmp = tempfile::TempDir::new().unwrap();
let r = OnnxReranker::new(
tmp.path(),
"model_int8.onnx",
ScoreStrategy::ClassifierLogit,
512,
2,
false,
);
assert!(
r.is_err(),
"missing tokenizer.json should fail at construction"
);
let missing = std::path::Path::new("/no/such/reranker/dir");
assert!(
OnnxReranker::new(
missing,
"model_int8.onnx",
ScoreStrategy::ClassifierLogit,
512,
2,
false,
)
.is_err()
);
}
#[test]
fn rerank_is_identity_when_disabled() {
use crate::search::fastembed_reranker::ENV_ENABLE;
with_env(&[(ENV_ENABLE, None)], || {
let tmp = tempfile::TempDir::new().unwrap();
std::fs::write(tmp.path().join("tokenizer.json"), MINIMAL_TOKENIZER_JSON).unwrap();
let r = OnnxReranker::new(
tmp.path(),
"model_int8.onnx",
ScoreStrategy::ClassifierLogit,
512,
2,
false,
)
.expect("construct with tokenizer present");
let docs = ["a", "b", "c"];
let docs_ref: Vec<&str> = docs.to_vec();
let out = r.rerank("q", &docs_ref, 2).expect("identity rerank");
assert_eq!(out, vec![(0, 0.0_f32), (1, 0.0_f32)]);
});
}
#[test]
fn encode_pair_truncates_to_max_context() {
const MAX: usize = 8;
let tmp = tempfile::TempDir::new().unwrap();
std::fs::write(tmp.path().join("tokenizer.json"), MINIMAL_TOKENIZER_JSON).unwrap();
let r = OnnxReranker::new(
tmp.path(),
"model.onnx",
ScoreStrategy::ClassifierLogit,
MAX,
2,
false,
)
.expect("construct");
let long_doc = vec!["word"; 600].join(" ");
let (ids, mask) = r.encode_pair("a query", &long_doc).expect("encode");
assert!(
ids.len() <= MAX,
"input_ids must be truncated to <= max_context ({MAX}); got {}",
ids.len()
);
assert_eq!(ids.len(), mask.len(), "ids/mask lengths must match");
}
const MINIMAL_TOKENIZER_JSON: &str = r#"{
"version": "1.0",
"truncation": null,
"padding": null,
"added_tokens": [],
"normalizer": null,
"pre_tokenizer": {"type": "Whitespace"},
"post_processor": null,
"decoder": null,
"model": {"type": "WordLevel", "vocab": {"[UNK]": 0}, "unk_token": "[UNK]"}
}"#;
#[test]
#[ignore = "requires downloading the bge ONNX reranker model"]
fn onnx_classifier_ranks_on_topic() {
use crate::config::SemantexConfig;
use crate::search::fastembed_reranker::ENV_ENABLE;
use crate::search::reranker_download::ensure_reranker_model;
use crate::search::reranker_model::{RerankerChoice, select_reranker_choice_from_env};
with_env(&[(ENV_ENABLE, Some("on"))], || {
unsafe { std::env::set_var("SEMANTEX_RERANKER_MODEL", "bge-onnx") };
let spec = match select_reranker_choice_from_env() {
RerankerChoice::Onnx(s) => s,
other @ RerankerChoice::Fastembed(_) => {
panic!("expected ONNX choice for bge-onnx, got {other:?}")
}
};
let config = SemantexConfig::default();
let dir = ensure_reranker_model(&config.models_dir(), &spec.files)
.expect("download (offline?)");
let r = OnnxReranker::new(
&dir,
&spec.session_file,
spec.strategy.clone(),
spec.max_context,
4,
false,
)
.expect("construct");
let docs = [
"Pizza is a popular Italian dish.",
"fn binary_search(a: &[i32], t: i32) -> Option<usize> { /* ... */ }",
];
let docs_ref: Vec<&str> = docs.to_vec();
let out = r
.rerank("how does binary search work", &docs_ref, 2)
.expect("rerank");
assert_eq!(out[0].0, 1, "on-topic code doc should rank first");
});
}
#[test]
#[ignore = "requires downloading the ~1.1 GB qwen3 ONNX reranker model"]
fn onnx_qwen3_yesno_ranks_on_topic() {
use crate::config::SemantexConfig;
use crate::search::fastembed_reranker::ENV_ENABLE;
use crate::search::reranker_download::ensure_reranker_model;
use crate::search::reranker_model::{RerankerChoice, select_reranker_choice_from_env};
with_env(&[(ENV_ENABLE, Some("on"))], || {
unsafe { std::env::set_var("SEMANTEX_RERANKER_MODEL", "qwen3") };
let spec = match select_reranker_choice_from_env() {
RerankerChoice::Onnx(s) => s,
other @ RerankerChoice::Fastembed(_) => {
panic!("expected ONNX choice for qwen3, got {other:?}")
}
};
assert!(
matches!(spec.strategy, ScoreStrategy::YesNoLogit { .. }),
"qwen3 must resolve to YesNoLogit"
);
let config = SemantexConfig::default();
let dir = ensure_reranker_model(&config.models_dir(), &spec.files)
.expect("download (offline?)");
let r = OnnxReranker::new(
&dir,
&spec.session_file,
spec.strategy.clone(),
spec.max_context,
4,
false,
)
.expect("construct");
let docs = [
"Pizza is a popular Italian dish.",
"fn binary_search(a: &[i32], t: i32) -> Option<usize> { /* ... */ }",
];
let docs_ref: Vec<&str> = docs.to_vec();
let out = r
.rerank("how does binary search work", &docs_ref, 2)
.expect("rerank (fp16 extraction must succeed)");
assert_eq!(out[0].0, 1, "on-topic code doc should rank first");
assert!(out[0].1 > out[1].1, "scores must separate on-topic vs off");
});
}
fn with_env<F: FnOnce()>(vars: &[(&str, Option<&str>)], f: F) {
let _guard = crate::search::RERANKER_TEST_ENV_LOCK
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
let prior: Vec<(String, Option<String>)> = vars
.iter()
.map(|(k, _)| ((*k).to_string(), std::env::var(*k).ok()))
.collect();
unsafe {
for (k, v) in vars {
match v {
Some(val) => std::env::set_var(k, val),
None => std::env::remove_var(k),
}
}
}
let result = std::panic::catch_unwind(std::panic::AssertUnwindSafe(f));
unsafe {
for (k, v) in &prior {
match v {
Some(val) => std::env::set_var(k, val),
None => std::env::remove_var(k),
}
}
}
if let Err(e) = result {
std::panic::resume_unwind(e);
}
}
}