use ndarray::{ArrayView, ArrayView3, Dim, Dimension, IxDynImpl, s};
use ort::session::Session;
use ort::value::Value;
use thiserror::Error;
use tokenizers::{EncodeInput, InputSequence, Tokenizer};
use crate::core::config::reranker::RerankerHead;
#[cfg_attr(alef, alef(skip))]
#[derive(Debug, Error)]
pub enum RerankError {
#[error("Tokenizer error: {0}")]
Tokenizer(String),
#[error("ONNX Runtime error: {0}")]
Ort(#[from] ort::Error),
#[error("Tensor shape error: {0}")]
Shape(String),
#[error("Model produced no output tensors")]
NoOutput,
}
const QWEN3_QUERY_PREFIX: &str = "<|im_start|>system\n\
Judge whether the Document meets the requirements based on the Query and the Instruct provided. \
Note that the answer can only be \"yes\" or \"no\".<|im_end|>\n\
<|im_start|>user\n\
<Instruct>: Given a web search query, retrieve relevant passages that answer the query\n\
<Query>: ";
const QWEN3_DOCUMENT_SUFFIX: &str = "<|im_end|>\n<|im_start|>assistant\n<think>\n\n</think>\n\n";
#[cfg_attr(alef, alef(skip))]
pub struct RerankerEngine {
tokenizer: Tokenizer,
session: Session,
need_token_type_ids: bool,
head: RerankerHead,
true_token_id: Option<u32>,
false_token_id: Option<u32>,
}
impl RerankerEngine {
pub(crate) fn new(
tokenizer: Tokenizer,
session: Session,
head: RerankerHead,
true_token_id: Option<u32>,
false_token_id: Option<u32>,
) -> Self {
let need_token_type_ids = session.inputs().iter().any(|input| input.name() == "token_type_ids");
Self {
tokenizer,
session,
need_token_type_ids,
head,
true_token_id,
false_token_id,
}
}
pub(crate) fn rerank(&self, query: &str, documents: &[&str], batch_size: usize) -> Result<Vec<f32>, RerankError> {
if documents.is_empty() {
return Ok(Vec::new());
}
let batch_size = if batch_size == 0 { 32 } else { batch_size };
let mut all_scores = Vec::with_capacity(documents.len());
for batch in documents.chunks(batch_size) {
let batch_scores = self.rerank_batch(query, batch)?;
all_scores.extend(batch_scores);
}
Ok(all_scores)
}
fn rerank_batch(&self, query: &str, documents: &[&str]) -> Result<Vec<f32>, RerankError> {
let owned_prompts: Vec<String>;
let encodings = if self.head == RerankerHead::Qwen3Generative {
owned_prompts = documents
.iter()
.map(|doc| format!("{QWEN3_QUERY_PREFIX}{query}\n\n<Document>: {doc}{QWEN3_DOCUMENT_SUFFIX}"))
.collect();
let inputs: Vec<EncodeInput<'_>> = owned_prompts
.iter()
.map(|prompt| EncodeInput::Single(InputSequence::Raw(std::borrow::Cow::Borrowed(prompt.as_str()))))
.collect();
self.tokenizer
.encode_batch(inputs, true)
.map_err(|e| RerankError::Tokenizer(e.to_string()))?
} else {
let pairs: Vec<EncodeInput<'_>> = documents
.iter()
.map(|doc| {
EncodeInput::Dual(
InputSequence::Raw(std::borrow::Cow::Borrowed(query)),
InputSequence::Raw(std::borrow::Cow::Borrowed(doc)),
)
})
.collect();
self.tokenizer
.encode_batch(pairs, true)
.map_err(|e| RerankError::Tokenizer(e.to_string()))?
};
let encoding_length = encodings
.first()
.ok_or_else(|| RerankError::Tokenizer("Empty encodings".to_string()))?
.len();
let batch_size = documents.len();
let max_size = encoding_length * batch_size;
let mut ids_array = Vec::with_capacity(max_size);
let mut mask_array = Vec::with_capacity(max_size);
let mut type_ids_array = if self.need_token_type_ids {
Vec::with_capacity(max_size)
} else {
Vec::new()
};
for encoding in &encodings {
ids_array.extend(encoding.get_ids().iter().map(|&x| x as i64));
mask_array.extend(encoding.get_attention_mask().iter().map(|&x| x as i64));
if self.need_token_type_ids {
type_ids_array.extend(encoding.get_type_ids().iter().map(|&x| x as i64));
}
}
let ids_tensor = ndarray::Array::from_shape_vec((batch_size, encoding_length), ids_array)
.map_err(|e| RerankError::Shape(e.to_string()))?;
let mask_tensor = ndarray::Array::from_shape_vec((batch_size, encoding_length), mask_array)
.map_err(|e| RerankError::Shape(e.to_string()))?;
let last_token_indices: Vec<usize> = mask_tensor
.outer_iter()
.map(|row| {
row.iter()
.rposition(|&m| m != 0)
.unwrap_or(encoding_length.saturating_sub(1))
})
.collect();
let mut session_inputs = ort::inputs![
"input_ids" => Value::from_array(ids_tensor)?,
"attention_mask" => Value::from_array(mask_tensor)?,
];
if self.need_token_type_ids {
let type_ids_tensor = ndarray::Array::from_shape_vec((batch_size, encoding_length), type_ids_array)
.map_err(|e| RerankError::Shape(e.to_string()))?;
session_inputs.push(("token_type_ids".into(), Value::from_array(type_ids_tensor)?.into()));
}
#[allow(unsafe_code)]
let outputs = unsafe {
let session_ptr = &self.session as *const Session as *mut Session;
(*session_ptr).run(session_inputs)
}
.map_err(RerankError::Ort)?;
let (_, output_value) = outputs.iter().next().ok_or(RerankError::NoOutput)?;
let tensor: ArrayView<f32, Dim<IxDynImpl>> = output_value.try_extract_array().map_err(RerankError::Ort)?;
let scores = match self.head {
RerankerHead::CrossEncoder => match tensor.dim().ndim() {
1 => tensor.slice(s![..]).iter().copied().collect(),
2 => tensor.slice(s![.., 0]).iter().copied().collect(),
n => return Err(RerankError::Shape(format!("Expected 1D or 2D output tensor, got {n}D"))),
},
RerankerHead::Qwen3Generative => {
let true_id = self
.true_token_id
.ok_or_else(|| RerankError::Shape("Qwen3 head requires a resolved true_token_id".to_string()))?;
let false_id = self
.false_token_id
.ok_or_else(|| RerankError::Shape("Qwen3 head requires a resolved false_token_id".to_string()))?;
if tensor.dim().ndim() != 3 {
return Err(RerankError::Shape(format!(
"Qwen3 generative head expects a 3D [batch, seq, vocab] output tensor, got {}D",
tensor.dim().ndim()
)));
}
let logits: ArrayView3<f32> = tensor
.view()
.into_dimensionality::<ndarray::Ix3>()
.map_err(|e| RerankError::Shape(format!("Failed to reshape Qwen3 output to 3D: {e}")))?;
qwen3_scores(&logits, true_id, false_id, &last_token_indices)?
}
};
Ok(scores)
}
}
fn qwen3_scores(
logits: &ArrayView3<f32>,
true_id: u32,
false_id: u32,
last_token_indices: &[usize],
) -> Result<Vec<f32>, RerankError> {
let (batch, seq_len, vocab) = logits.dim();
if seq_len == 0 {
return Err(RerankError::Shape(
"Qwen3 generative head received a zero-length sequence".to_string(),
));
}
let true_id = true_id as usize;
let false_id = false_id as usize;
if true_id >= vocab || false_id >= vocab {
return Err(RerankError::Shape(format!(
"Qwen3 true/false token id out of vocab range: true={true_id}, false={false_id}, vocab={vocab}"
)));
}
let mut scores = Vec::with_capacity(batch);
for b in 0..batch {
let idx = last_token_indices
.get(b)
.copied()
.unwrap_or(seq_len - 1)
.min(seq_len - 1);
let false_logit = logits[[b, idx, false_id]];
let true_logit = logits[[b, idx, true_id]];
let max_logit = false_logit.max(true_logit);
let false_exp = (false_logit - max_logit).exp();
let true_exp = (true_logit - max_logit).exp();
let denom = false_exp + true_exp;
scores.push(true_exp / denom);
}
Ok(scores)
}
#[allow(unsafe_code)]
unsafe impl Send for RerankerEngine {}
#[allow(unsafe_code)]
unsafe impl Sync for RerankerEngine {}
#[cfg(test)]
mod tests {
use super::super::sigmoid_f32 as sigmoid;
use super::*;
#[test]
fn sigmoid_zero_gives_half() {
let s = sigmoid(0.0);
assert!((s - 0.5).abs() < 1e-6, "sigmoid(0) should be 0.5, got {s}");
}
#[test]
fn sigmoid_large_positive_approaches_one() {
let s = sigmoid(100.0);
assert!(s > 0.99, "sigmoid(100) should be close to 1.0, got {s}");
}
#[test]
fn sigmoid_large_negative_approaches_zero() {
let s = sigmoid(-100.0);
assert!(s < 0.01, "sigmoid(-100) should be close to 0.0, got {s}");
}
#[test]
fn rerank_error_display_does_not_panic() {
let err = RerankError::Tokenizer("test".to_string());
assert!(format!("{err}").contains("Tokenizer"));
let err = RerankError::Shape("bad shape".to_string());
assert!(format!("{err}").contains("shape"));
let err = RerankError::NoOutput;
assert!(format!("{err}").contains("no output"));
}
#[test]
fn rerank_error_implements_error_trait() {
let err = RerankError::Shape("test".to_string());
let _: &dyn std::error::Error = &err;
}
fn make_logits(rows: Vec<Vec<Vec<f32>>>) -> ndarray::Array3<f32> {
let batch = rows.len();
let seq = rows[0].len();
let vocab = rows[0][0].len();
let flat: Vec<f32> = rows.into_iter().flatten().flatten().collect();
ndarray::Array3::from_shape_vec((batch, seq, vocab), flat).expect("valid shape")
}
#[test]
fn qwen3_scores_hand_computed_softmax() {
let logits = make_logits(vec![
vec![vec![9.0, 9.0, 9.0, 9.0], vec![0.0, 0.0, 1.0, 2.0]],
vec![vec![9.0, 9.0, 9.0, 9.0], vec![0.0, 0.0, 2.0, 0.0]],
]);
let view = logits.view();
let scores = qwen3_scores(&view, 3, 2, &[1, 1]).expect("scores must compute");
assert_eq!(scores.len(), 2);
assert!(
(scores[0] - 0.7310586).abs() < 1e-5,
"expected P(yes) ~= 0.7310586 for item 0, got {}",
scores[0]
);
assert!(
(scores[1] - 0.1192029).abs() < 1e-5,
"expected P(yes) ~= 0.1192029 for item 1, got {}",
scores[1]
);
for &s in &scores {
assert!((0.0..=1.0).contains(&s), "score {s} must be in [0, 1]");
}
}
#[test]
fn qwen3_scores_equal_logits_gives_half() {
let logits = make_logits(vec![vec![vec![0.0, 0.0, 3.0, 3.0]]]);
let view = logits.view();
let scores = qwen3_scores(&view, 3, 2, &[0]).expect("scores must compute");
assert_eq!(scores.len(), 1);
assert!(
(scores[0] - 0.5).abs() < 1e-6,
"equal logits must give P(yes)=0.5, got {}",
scores[0]
);
}
#[test]
fn qwen3_scores_only_reads_last_token() {
let logits = make_logits(vec![vec![vec![0.0, 0.0, 100.0, -100.0], vec![0.0, 0.0, -100.0, 100.0]]]);
let view = logits.view();
let scores = qwen3_scores(&view, 3, 2, &[1]).expect("scores must compute");
assert!(
scores[0] > 0.99,
"must score based on last token only, got {}",
scores[0]
);
}
#[test]
fn qwen3_scores_reads_last_unmasked_token_under_right_padding() {
let logits = make_logits(vec![vec![vec![0.0, 0.0, -100.0, 100.0], vec![0.0, 0.0, 100.0, -100.0]]]);
let view = logits.view();
let scores = qwen3_scores(&view, 3, 2, &[0]).expect("scores must compute");
assert!(
scores[0] > 0.99,
"must read the last unmasked token (index 0), got {}",
scores[0]
);
}
#[test]
fn qwen3_scores_rejects_zero_length_sequence() {
let logits = ndarray::Array3::<f32>::zeros((1, 0, 4));
let view = logits.view();
let err = qwen3_scores(&view, 3, 2, &[0]).expect_err("zero-length sequence must error");
assert!(matches!(err, RerankError::Shape(_)));
}
#[test]
fn qwen3_scores_rejects_out_of_range_token_ids() {
let logits = make_logits(vec![vec![vec![0.0, 0.0, 1.0, 2.0]]]);
let view = logits.view();
let err = qwen3_scores(&view, 99, 2, &[0]).expect_err("out-of-range true_id must error");
assert!(matches!(err, RerankError::Shape(_)));
}
}