use std::collections::HashMap;
use std::fs;
use std::path::Path;
use std::sync::{Arc, Mutex};
use ft_api::{FrankenTorchSession, quantize_per_output_channel_i8};
use ft_autograd::TensorNodeId;
use ft_core::{DType, Device, ExecutionMode, TensorMeta};
use tokenizers::Tokenizer;
use wide::f32x8;
use frankensearch_core::error::{SearchError, SearchResult};
use frankensearch_core::traits::{RerankDocument, RerankScore, SyncRerank};
const H: usize = 384;
const L: usize = 6;
const NH: usize = 12;
const HD: usize = H / NH; const INTER: usize = 4 * H; const EPS: f64 = 1.0e-12;
const EPS_F32: f32 = 1.0e-12;
const ATTN_SCALE_F32: f32 = 0.176_776_69;
const CLS_Q_CACHE_MIN_SEQ: usize = 256;
pub(crate) const DEFAULT_MAX_LENGTH: usize = 512;
const MAX_BATCH_TOKENS: usize = 2048;
const FUSED_ATTN_MAX_SEQ: usize = DEFAULT_MAX_LENGTH;
const MODEL_NAME: &str = "ms-marco-minilm-l-6-v2";
pub(crate) const SAFETENSORS_PRIMARY: &str = "model_f32.safetensors";
pub(crate) const SAFETENSORS_FALLBACK: &str = "model.safetensors";
pub(crate) const TOKENIZER_JSON: &str = "tokenizer.json";
fn rerank_err(ctx: &str, e: impl std::fmt::Display) -> SearchError {
SearchError::RerankFailed {
model: MODEL_NAME.to_owned(),
source: format!("{ctx}: {e}").into(),
}
}
fn index_to_i64(index: usize, ctx: &str) -> SearchResult<i64> {
i64::try_from(index).map_err(|_| rerank_err(ctx, format!("index {index} exceeds i64::MAX")))
}
fn softmax_row_fused(row: &mut [f32], scale: f32) {
let n = row.len();
let mut max_raw = f32::NEG_INFINITY;
for &x in row.iter() {
max_raw = max_raw.max(x);
}
let max_v = f32x8::splat(max_raw);
let scale_v = f32x8::splat(scale);
let mut sum_v = f32x8::splat(0.0);
let mut i = 0;
while i + 32 <= n {
let e0 = ((f32x8_from_slice(&row[i..i + 8]) - max_v) * scale_v).exp();
let e1 = ((f32x8_from_slice(&row[i + 8..i + 16]) - max_v) * scale_v).exp();
let e2 = ((f32x8_from_slice(&row[i + 16..i + 24]) - max_v) * scale_v).exp();
let e3 = ((f32x8_from_slice(&row[i + 24..i + 32]) - max_v) * scale_v).exp();
row[i..i + 8].copy_from_slice(&e0.to_array());
row[i + 8..i + 16].copy_from_slice(&e1.to_array());
row[i + 16..i + 24].copy_from_slice(&e2.to_array());
row[i + 24..i + 32].copy_from_slice(&e3.to_array());
sum_v += e0;
sum_v += e1;
sum_v += e2;
sum_v += e3;
i += 32;
}
while i + 8 <= n {
let e = ((f32x8_from_slice(&row[i..i + 8]) - max_v) * scale_v).exp();
row[i..i + 8].copy_from_slice(&e.to_array());
sum_v += e;
i += 8;
}
let mut sum: f32 = sum_v.to_array().iter().sum();
while i < n {
let e = ((row[i] - max_raw) * scale).exp();
row[i] = e;
sum += e;
i += 1;
}
let inv = 1.0 / sum;
for x in row.iter_mut() {
*x *= inv;
}
}
fn fast_softmax_inplace(data: &mut [f32], rows: usize, n: usize, scale: f32) {
debug_assert_eq!(data.len(), rows * n);
if n == 0 {
return;
}
if rows >= 8 && rows * n >= 8192 && rayon::current_num_threads() > 1 {
use rayon::prelude::*;
data.par_chunks_exact_mut(n)
.for_each(|row| softmax_row_fused(row, scale));
} else {
data.chunks_exact_mut(n)
.for_each(|row| softmax_row_fused(row, scale));
}
}
#[inline]
fn gelu_vec8(x: f32x8) -> f32x8 {
const C: f32 = std::f32::consts::FRAC_1_SQRT_2;
let one = f32x8::splat(1.0);
let z = x * f32x8::splat(C);
let az = z.abs();
let t = one / (one + f32x8::splat(0.327_591_1) * az);
let a1 = f32x8::splat(0.254_829_6);
let a2 = f32x8::splat(-0.284_496_73);
let a3 = f32x8::splat(1.421_413_7);
let a4 = f32x8::splat(-1.453_152);
let a5 = f32x8::splat(1.061_405_4);
let poly = t * (a1 + t * (a2 + t * (a3 + t * (a4 + t * a5))));
let erf_abs = one - poly * (-(z * z)).exp();
let erf = erf_abs.copysign(z);
f32x8::splat(0.5) * x * (one + erf)
}
#[inline]
fn gelu_scalar(x: f32) -> f32 {
const C: f32 = std::f32::consts::FRAC_1_SQRT_2;
let z = x * C;
let az = z.abs();
let t = 1.0 / (1.0 + 0.327_591_1 * az);
let poly = t
* (0.254_829_6
+ t * (-0.284_496_73 + t * (1.421_413_7 + t * (-1.453_152 + t * 1.061_405_4))));
let erf = (1.0 - poly * (-(z * z)).exp()).copysign(z);
0.5 * x * (1.0 + erf)
}
fn fast_gelu_inplace(data: &mut [f32]) {
let process = |chunk: &mut [f32]| {
let n = chunk.len();
let mut i = 0;
while i + 32 <= n {
let g0 = gelu_vec8(f32x8_from_slice(&chunk[i..i + 8]));
let g1 = gelu_vec8(f32x8_from_slice(&chunk[i + 8..i + 16]));
let g2 = gelu_vec8(f32x8_from_slice(&chunk[i + 16..i + 24]));
let g3 = gelu_vec8(f32x8_from_slice(&chunk[i + 24..i + 32]));
chunk[i..i + 8].copy_from_slice(&g0.to_array());
chunk[i + 8..i + 16].copy_from_slice(&g1.to_array());
chunk[i + 16..i + 24].copy_from_slice(&g2.to_array());
chunk[i + 24..i + 32].copy_from_slice(&g3.to_array());
i += 32;
}
while i + 8 <= n {
let g = gelu_vec8(f32x8_from_slice(&chunk[i..i + 8]));
chunk[i..i + 8].copy_from_slice(&g.to_array());
i += 8;
}
while i < n {
chunk[i] = gelu_scalar(chunk[i]);
i += 1;
}
};
if data.len() >= 8192 && rayon::current_num_threads() > 1 {
use rayon::prelude::*;
data.par_chunks_mut(2048).for_each(process);
} else {
process(data);
}
}
#[inline]
fn f32x8_from_slice(slice: &[f32]) -> f32x8 {
let mut buf = [0.0f32; 8];
buf.copy_from_slice(&slice[..8]);
f32x8::new(buf)
}
#[inline]
fn dot_hd(q: &[f32], k: &[f32]) -> f32 {
debug_assert_eq!(q.len(), HD);
debug_assert_eq!(k.len(), HD);
(f32x8_from_slice(&q[0..]) * f32x8_from_slice(&k[0..])
+ f32x8_from_slice(&q[8..]) * f32x8_from_slice(&k[8..])
+ f32x8_from_slice(&q[16..]) * f32x8_from_slice(&k[16..])
+ f32x8_from_slice(&q[24..]) * f32x8_from_slice(&k[24..]))
.reduce_add()
}
#[inline]
fn q_lanes(q: &[f32]) -> [f32x8; 4] {
debug_assert_eq!(q.len(), HD);
[
f32x8_from_slice(&q[0..]),
f32x8_from_slice(&q[8..]),
f32x8_from_slice(&q[16..]),
f32x8_from_slice(&q[24..]),
]
}
#[inline]
fn dot_hd_q_lanes(q: [f32x8; 4], k: &[f32]) -> f32 {
debug_assert_eq!(k.len(), HD);
(q[0] * f32x8_from_slice(&k[0..])
+ q[1] * f32x8_from_slice(&k[8..])
+ q[2] * f32x8_from_slice(&k[16..])
+ q[3] * f32x8_from_slice(&k[24..]))
.reduce_add()
}
fn weighted_value_sum_hd(qkv: &[f32], probs: &[f32], s_len: usize, head: usize, out: &mut [f32]) {
debug_assert_eq!(probs.len(), s_len);
debug_assert_eq!(out.len(), HD);
const STRIDE: usize = 3 * H;
let mut acc0 = f32x8::splat(0.0);
let mut acc1 = f32x8::splat(0.0);
let mut acc2 = f32x8::splat(0.0);
let mut acc3 = f32x8::splat(0.0);
for (j, &prob) in probs.iter().enumerate() {
let p = f32x8::splat(prob);
let base = j * STRIDE + 2 * H + head * HD;
acc0 += p * f32x8_from_slice(&qkv[base..base + 8]);
acc1 += p * f32x8_from_slice(&qkv[base + 8..base + 16]);
acc2 += p * f32x8_from_slice(&qkv[base + 16..base + 24]);
acc3 += p * f32x8_from_slice(&qkv[base + 24..base + 32]);
}
out[0..8].copy_from_slice(&acc0.to_array());
out[8..16].copy_from_slice(&acc1.to_array());
out[16..24].copy_from_slice(&acc2.to_array());
out[24..32].copy_from_slice(&acc3.to_array());
}
#[derive(Default)]
struct AttnScratch {
q_hm: Vec<f32>,
kt: Vec<f32>,
v_hm: Vec<f32>,
scores: Vec<f32>,
ctx_hm: Vec<f32>,
}
impl AttnScratch {
fn ensure(&mut self, s_len: usize) {
let hm = s_len * H; let sc = NH * s_len * s_len;
if self.q_hm.len() < hm {
self.q_hm.resize(hm, 0.0);
}
if self.kt.len() < hm {
self.kt.resize(hm, 0.0);
}
if self.v_hm.len() < hm {
self.v_hm.resize(hm, 0.0);
}
if self.scores.len() < sc {
self.scores.resize(sc, 0.0);
}
if self.ctx_hm.len() < hm {
self.ctx_hm.resize(hm, 0.0);
}
}
}
fn fused_attention(
scratch: &mut AttnScratch,
qkv: &[f32],
s_len: usize,
scale: f32,
out: &mut [f32],
) {
debug_assert_eq!(qkv.len(), s_len * 3 * H);
debug_assert_eq!(out.len(), s_len * H);
if s_len == 0 {
return;
}
const STRIDE: usize = 3 * H;
scratch.ensure(s_len);
let hm = s_len * H;
let sc = NH * s_len * s_len;
let AttnScratch {
q_hm,
kt,
v_hm,
scores,
ctx_hm,
} = scratch;
let (q_hm, kt, v_hm) = (&mut q_hm[..hm], &mut kt[..hm], &mut v_hm[..hm]);
let scores = &mut scores[..sc];
let ctx_hm = &mut ctx_hm[..hm];
for j in 0..s_len {
let base = j * STRIDE;
for h in 0..NH {
let hmj = (h * s_len + j) * HD;
q_hm[hmj..hmj + HD].copy_from_slice(&qkv[base + h * HD..base + h * HD + HD]);
v_hm[hmj..hmj + HD]
.copy_from_slice(&qkv[base + 2 * H + h * HD..base + 2 * H + h * HD + HD]);
}
}
for h in 0..NH {
for d in 0..HD {
let col = H + h * HD + d;
let row = &mut kt[h * HD * s_len + d * s_len..h * HD * s_len + d * s_len + s_len];
for (j, slot) in row.iter_mut().enumerate() {
*slot = qkv[j * STRIDE + col];
}
}
}
let qm = TensorMeta::from_shape(vec![NH, s_len, HD], DType::F32, Device::Cpu);
let km = TensorMeta::from_shape(vec![NH, HD, s_len], DType::F32, Device::Cpu);
ft_api::bmm_tensor_contiguous_f32_into(q_hm, kt, &qm, &km, scores)
.expect("attn QKᵀ bmm: shapes are internally consistent");
fast_softmax_inplace(scores, NH * s_len, s_len, scale);
let sm = TensorMeta::from_shape(vec![NH, s_len, s_len], DType::F32, Device::Cpu);
let vm = TensorMeta::from_shape(vec![NH, s_len, HD], DType::F32, Device::Cpu);
ft_api::bmm_tensor_contiguous_f32_into(scores, v_hm, &sm, &vm, ctx_hm)
.expect("attn ·V bmm: shapes are internally consistent");
for h in 0..NH {
for i in 0..s_len {
let src = (h * s_len + i) * HD;
out[i * H + h * HD..i * H + h * HD + HD].copy_from_slice(&ctx_hm[src..src + HD]);
}
}
}
fn fused_attention_cls(
scratch: &mut AttnScratch,
qkv: &[f32],
s_len: usize,
scale: f32,
out: &mut [f32],
) {
debug_assert_eq!(qkv.len(), s_len * 3 * H);
debug_assert_eq!(out.len(), H);
if s_len == 0 {
out.fill(0.0);
return;
}
const STRIDE: usize = 3 * H;
scratch.ensure(s_len);
let sc = NH * s_len;
for h in 0..NH {
let row = &mut scratch.scores[h * s_len..(h + 1) * s_len];
if s_len >= CLS_Q_CACHE_MIN_SEQ {
let q = q_lanes(&qkv[h * HD..h * HD + HD]);
for (j, slot) in row.iter_mut().enumerate() {
let k_base = j * STRIDE + H + h * HD;
*slot = dot_hd_q_lanes(q, &qkv[k_base..k_base + HD]);
}
} else {
let q = &qkv[h * HD..h * HD + HD];
for (j, slot) in row.iter_mut().enumerate() {
let k_base = j * STRIDE + H + h * HD;
*slot = dot_hd(q, &qkv[k_base..k_base + HD]);
}
}
}
let scores = &mut scratch.scores[..sc];
fast_softmax_inplace(scores, NH, s_len, scale);
for h in 0..NH {
let row = &scores[h * s_len..(h + 1) * s_len];
weighted_value_sum_hd(qkv, row, s_len, h, &mut out[h * HD..h * HD + HD]);
}
}
#[derive(Clone)]
struct QLinear {
w_i8: Arc<Vec<i8>>,
w_scales: Arc<Vec<f32>>,
bias: Arc<Vec<f32>>,
out: usize,
in_: usize,
packed: bool,
}
pub(crate) struct Model {
s: FrankenTorchSession,
w: HashMap<String, TensorNodeId>,
qw: HashMap<String, QLinear>,
raw_params: HashMap<String, Arc<Vec<f32>>>,
weights_boundary: usize,
}
impl Model {
fn g(&self, name: &str) -> SearchResult<TensorNodeId> {
self.w
.get(name)
.copied()
.ok_or_else(|| rerank_err("weights", format!("missing weight tensor {name}")))
}
fn linear_raw(&self, x: &[f32], m: usize, prefix: &str) -> SearchResult<Vec<f32>> {
let q = self
.qw
.get(prefix)
.ok_or_else(|| rerank_err("linear_raw", format!("missing linear weights {prefix}")))?;
debug_assert_eq!(x.len(), m * q.in_);
let y = if q.packed {
ft_api::linear_int8_dynamic_prepacked_f32(
x,
m,
q.in_,
&q.w_i8,
&q.w_scales,
q.out,
Some(&q.bias),
)
} else {
ft_api::linear_int8_dynamic_f32(x, m, q.in_, &q.w_i8, &q.w_scales, q.out, Some(&q.bias))
};
Ok(y)
}
fn add_ln_raw(&self, a: &[f32], b: &[f32], m: usize, prefix: &str) -> SearchResult<Vec<f32>> {
let w = self
.raw_params
.get(&format!("{prefix}.weight"))
.ok_or_else(|| rerank_err("add_ln_raw", format!("missing {prefix}.weight")))?;
let bias = self
.raw_params
.get(&format!("{prefix}.bias"))
.ok_or_else(|| rerank_err("add_ln_raw", format!("missing {prefix}.bias")))?;
Ok(ft_api::add_layer_norm_forward_f32(
a,
b,
Some(w),
Some(bias),
m,
H,
EPS_F32,
))
}
fn encoder_layer_raw(
&self,
emb: &[f32],
total: usize,
offsets: &[usize],
lens: &[usize],
p: &str,
scale: f32,
scratch: &mut AttnScratch,
) -> SearchResult<Vec<f32>> {
let qkv = self.linear_raw(emb, total, &format!("{p}.attention.self.qkv"))?;
let mut ctx = vec![0.0f32; total * H];
for (&off, &len) in offsets.iter().zip(lens) {
let qkv_doc = &qkv[off * 3 * H..(off + len) * 3 * H];
fused_attention(
scratch,
qkv_doc,
len,
scale,
&mut ctx[off * H..(off + len) * H],
);
}
let attn = self.linear_raw(&ctx, total, &format!("{p}.attention.output.dense"))?;
let emb = self.add_ln_raw(
emb,
&attn,
total,
&format!("{p}.attention.output.LayerNorm"),
)?;
let mut inter = self.linear_raw(&emb, total, &format!("{p}.intermediate.dense"))?;
debug_assert_eq!(inter.len(), total * INTER);
fast_gelu_inplace(&mut inter);
let ffn = self.linear_raw(&inter, total, &format!("{p}.output.dense"))?;
self.add_ln_raw(&emb, &ffn, total, &format!("{p}.output.LayerNorm"))
}
fn encoder_layer_cls(
&self,
emb: &[f32],
offsets: &[usize],
lens: &[usize],
total: usize,
p: &str,
scale: f32,
scratch: &mut AttnScratch,
) -> SearchResult<Vec<f32>> {
let n_docs = lens.len();
let qkv = self.linear_raw(emb, total, &format!("{p}.attention.self.qkv"))?;
let mut ctx = vec![0.0f32; n_docs * H];
for (n, (&off, &len)) in offsets.iter().zip(lens).enumerate() {
let qkv_doc = &qkv[off * 3 * H..(off + len) * 3 * H];
fused_attention_cls(scratch, qkv_doc, len, scale, &mut ctx[n * H..(n + 1) * H]);
}
let attn = self.linear_raw(&ctx, n_docs, &format!("{p}.attention.output.dense"))?;
let mut emb_cls = vec![0.0f32; n_docs * H];
for (n, &off) in offsets.iter().enumerate() {
emb_cls[n * H..(n + 1) * H].copy_from_slice(&emb[off * H..off * H + H]);
}
let emb_cls = self.add_ln_raw(
&emb_cls,
&attn,
n_docs,
&format!("{p}.attention.output.LayerNorm"),
)?;
let mut inter = self.linear_raw(&emb_cls, n_docs, &format!("{p}.intermediate.dense"))?;
debug_assert_eq!(inter.len(), n_docs * INTER);
fast_gelu_inplace(&mut inter);
let ffn = self.linear_raw(&inter, n_docs, &format!("{p}.output.dense"))?;
self.add_ln_raw(&emb_cls, &ffn, n_docs, &format!("{p}.output.LayerNorm"))
}
fn linear(&mut self, x: TensorNodeId, prefix: &str) -> SearchResult<TensorNodeId> {
let q = self
.qw
.get(prefix)
.ok_or_else(|| rerank_err("linear", format!("missing int8 linear weights {prefix}")))?;
let w_i8 = Arc::clone(&q.w_i8);
let w_scales = Arc::clone(&q.w_scales);
let bias = Arc::clone(&q.bias);
let (out, in_, packed) = (q.out, q.in_, q.packed);
if packed {
self.s
.tensor_linear_int8_dynamic_prepacked(x, &w_i8, &w_scales, out, in_, Some(&bias))
.map_err(|e| rerank_err("linear.int8.packed", e))
} else {
self.s
.tensor_linear_int8_dynamic(x, &w_i8, &w_scales, out, in_, Some(&bias))
.map_err(|e| rerank_err("linear.int8", e))
}
}
fn idx(&mut self, vals: &[i64]) -> SearchResult<TensorNodeId> {
let f: Vec<f64> = vals.iter().map(|&v| v as f64).collect();
self.s
.tensor_variable(f, vec![vals.len()], false)
.map_err(|e| rerank_err("index_tensor", e))
}
fn add_ln(
&mut self,
a: TensorNodeId,
b: TensorNodeId,
prefix: &str,
) -> SearchResult<TensorNodeId> {
let w = self.g(&format!("{prefix}.weight"))?;
let bias = self.g(&format!("{prefix}.bias"))?;
self.s
.tensor_add_layer_norm(a, b, H, Some(w), Some(bias), EPS)
.map_err(|e| rerank_err("add_layer_norm", e))
}
fn gelu(&mut self, inter: TensorNodeId) -> SearchResult<TensorNodeId> {
let slice = self
.s
.tensor_values_f32_mut(inter)
.map_err(|e| rerank_err("ffn.gelu_mut", e))?;
fast_gelu_inplace(slice);
Ok(inter)
}
fn heads(&mut self, x: TensorNodeId, s_len: usize) -> SearchResult<TensorNodeId> {
let r = self
.s
.tensor_reshape(x, vec![s_len, NH, HD])
.map_err(|e| rerank_err("heads.reshape", e))?;
self.s
.tensor_transpose(r, 0, 1)
.map_err(|e| rerank_err("heads.transpose", e))
}
#[cfg(test)]
fn attn_fused(
&mut self,
qkv: TensorNodeId,
s_len: usize,
scale: f32,
) -> SearchResult<TensorNodeId> {
let ctx_vals = {
let qkv_v = self
.s
.tensor_values_f32_borrowed(qkv)
.map_err(|e| rerank_err("attn.qkv_vals", e))?;
let mut scratch = AttnScratch::default();
let mut ctx = vec![0.0f32; s_len * H];
fused_attention(&mut scratch, qkv_v, s_len, scale, &mut ctx);
ctx
};
self.s
.tensor_variable_f32(ctx_vals, vec![s_len, H], false)
.map_err(|e| rerank_err("attn.ctx", e))
}
fn attn_bmm(
&mut self,
q: TensorNodeId,
k: TensorNodeId,
v: TensorNodeId,
s_len: usize,
scale: f32,
) -> SearchResult<TensorNodeId> {
let q = self.heads(q, s_len)?;
let k = self.heads(k, s_len)?;
let v = self.heads(v, s_len)?;
let kt = self
.s
.tensor_transpose(k, 1, 2)
.map_err(|e| rerank_err("attn.kt", e))?; let scores = self
.s
.tensor_bmm(q, kt)
.map_err(|e| rerank_err("attn.qk", e))?;
{
let slice = self
.s
.tensor_values_f32_mut(scores)
.map_err(|e| rerank_err("attn.softmax_mut", e))?;
fast_softmax_inplace(slice, NH * s_len, s_len, scale);
}
let ctx = self
.s
.tensor_bmm(scores, v)
.map_err(|e| rerank_err("attn.ctx", e))?;
let ctx = self
.s
.tensor_transpose(ctx, 0, 1)
.map_err(|e| rerank_err("attn.ctx_t", e))?;
self.s
.tensor_reshape(ctx, vec![s_len, H])
.map_err(|e| rerank_err("attn.ctx_reshape", e))
}
#[cfg(test)]
fn attention(
&mut self,
emb: TensorNodeId,
p: &str,
s_len: usize,
scale: f32,
) -> SearchResult<TensorNodeId> {
if s_len <= FUSED_ATTN_MAX_SEQ {
let qkv = self.linear(emb, &format!("{p}.attention.self.qkv"))?;
self.attn_fused(qkv, s_len, scale)
} else {
let q = self.linear(emb, &format!("{p}.attention.self.query"))?;
let k = self.linear(emb, &format!("{p}.attention.self.key"))?;
let v = self.linear(emb, &format!("{p}.attention.self.value"))?;
self.attn_bmm(q, k, v, s_len, scale)
}
}
#[cfg(test)]
fn forward(&mut self, ids: &[i64], typ: &[i64]) -> SearchResult<f32> {
let s_len = ids.len();
let id_t = self.idx(ids)?;
let pos: Vec<i64> = (0..s_len)
.map(|i| index_to_i64(i, "forward.position"))
.collect::<SearchResult<_>>()?;
let pos_t = self.idx(&pos)?;
let typ_t = self.idx(typ)?;
let we = self.g("bert.embeddings.word_embeddings.weight")?;
let pe = self.g("bert.embeddings.position_embeddings.weight")?;
let te = self.g("bert.embeddings.token_type_embeddings.weight")?;
let e_word = self
.s
.tensor_index_select(we, 0, id_t)
.map_err(|e| rerank_err("embed.word", e))?;
let e_pos = self
.s
.tensor_index_select(pe, 0, pos_t)
.map_err(|e| rerank_err("embed.pos", e))?;
let e_typ = self
.s
.tensor_index_select(te, 0, typ_t)
.map_err(|e| rerank_err("embed.type", e))?;
let emb_wp = self
.s
.tensor_add(e_word, e_pos)
.map_err(|e| rerank_err("embed.add", e))?;
let mut emb = self.add_ln(emb_wp, e_typ, "bert.embeddings.LayerNorm")?;
let scale = ATTN_SCALE_F32;
for i in 0..L {
let p = format!("bert.encoder.layer.{i}");
let ctx = self.attention(emb, &p, s_len, scale)?;
let attn = self.linear(ctx, &format!("{p}.attention.output.dense"))?;
emb = self.add_ln(emb, attn, &format!("{p}.attention.output.LayerNorm"))?;
let inter = self.linear(emb, &format!("{p}.intermediate.dense"))?;
let inter = self.gelu(inter)?;
let ffn = self.linear(inter, &format!("{p}.output.dense"))?;
emb = self.add_ln(emb, ffn, &format!("{p}.output.LayerNorm"))?;
}
let cls = self
.s
.tensor_narrow(emb, 0, 0, 1)
.map_err(|e| rerank_err("pooler.narrow", e))?; let pooled = self.linear(cls, "bert.pooler.dense")?;
let pooled = self
.s
.tensor_tanh(pooled)
.map_err(|e| rerank_err("pooler.tanh", e))?;
let logit_t = self.linear(pooled, "classifier")?; let vals = self
.s
.tensor_values_f32(logit_t)
.map_err(|e| rerank_err("classifier.values", e))?;
let logit = vals
.first()
.copied()
.ok_or_else(|| rerank_err("classifier", "empty logit output"))?;
self.s.truncate_autograd_graph(self.weights_boundary);
Ok(logit)
}
fn forward_batch(&mut self, batch: &[(Vec<i64>, Vec<i64>)]) -> SearchResult<Vec<f32>> {
let n_docs = batch.len();
let lens: Vec<usize> = batch.iter().map(|(ids, _)| ids.len()).collect();
let total: usize = lens.iter().sum();
if total == 0 {
return Ok(vec![0.0; n_docs]);
}
let mut offsets = Vec::with_capacity(n_docs);
{
let mut o = 0usize;
for &l in &lens {
offsets.push(o);
o += l;
}
}
let mut ids_flat = Vec::with_capacity(total);
let mut pos_flat = Vec::with_capacity(total);
let mut typ_flat = Vec::with_capacity(total);
for (ids, typ) in batch {
for (i, (&id, &t)) in ids.iter().zip(typ.iter()).enumerate() {
ids_flat.push(id);
pos_flat.push(index_to_i64(i, "forward_batch.position")?);
typ_flat.push(t);
}
}
let id_t = self.idx(&ids_flat)?;
let pos_t = self.idx(&pos_flat)?;
let typ_t = self.idx(&typ_flat)?;
let we = self.g("bert.embeddings.word_embeddings.weight")?;
let pe = self.g("bert.embeddings.position_embeddings.weight")?;
let te = self.g("bert.embeddings.token_type_embeddings.weight")?;
let e_word = self
.s
.tensor_index_select(we, 0, id_t)
.map_err(|e| rerank_err("embed.word", e))?;
let e_pos = self
.s
.tensor_index_select(pe, 0, pos_t)
.map_err(|e| rerank_err("embed.pos", e))?;
let e_typ = self
.s
.tensor_index_select(te, 0, typ_t)
.map_err(|e| rerank_err("embed.type", e))?;
let emb_wp = self
.s
.tensor_add(e_word, e_pos)
.map_err(|e| rerank_err("embed.add", e))?;
let mut emb = self.add_ln(emb_wp, e_typ, "bert.embeddings.LayerNorm")?;
let scale = ATTN_SCALE_F32;
let cls_prepacked = if lens.iter().all(|&l| l <= FUSED_ATTN_MAX_SEQ) {
let mut scratch = AttnScratch::default();
let mut emb_vals = self
.s
.tensor_values_f32(emb)
.map_err(|e| rerank_err("batch.emb_extract", e))?;
for i in 0..L - 1 {
let p = format!("bert.encoder.layer.{i}");
emb_vals = self.encoder_layer_raw(
&emb_vals,
total,
&offsets,
&lens,
&p,
scale,
&mut scratch,
)?;
}
let p_last = format!("bert.encoder.layer.{}", L - 1);
let cls_vals = self.encoder_layer_cls(
&emb_vals,
&offsets,
&lens,
total,
&p_last,
scale,
&mut scratch,
)?;
emb = self
.s
.tensor_variable_f32(cls_vals, vec![n_docs, H], false)
.map_err(|e| rerank_err("batch.emb_reinsert", e))?;
true
} else {
for i in 0..L {
let p = format!("bert.encoder.layer.{i}");
let q = self.linear(emb, &format!("{p}.attention.self.query"))?;
let k = self.linear(emb, &format!("{p}.attention.self.key"))?;
let v = self.linear(emb, &format!("{p}.attention.self.value"))?;
let mut ctx_parts = Vec::with_capacity(n_docs);
for n in 0..n_docs {
let (off, len) = (offsets[n], lens[n]);
let qn = self
.s
.tensor_narrow(q, 0, off, len)
.map_err(|e| rerank_err("batch.q_narrow", e))?;
let kn = self
.s
.tensor_narrow(k, 0, off, len)
.map_err(|e| rerank_err("batch.k_narrow", e))?;
let vn = self
.s
.tensor_narrow(v, 0, off, len)
.map_err(|e| rerank_err("batch.v_narrow", e))?;
ctx_parts.push(self.attn_bmm(qn, kn, vn, len, scale)?);
}
let ctx = self
.s
.tensor_cat(&ctx_parts, 0)
.map_err(|e| rerank_err("batch.ctx_cat", e))?; let attn = self.linear(ctx, &format!("{p}.attention.output.dense"))?;
emb = self.add_ln(emb, attn, &format!("{p}.attention.output.LayerNorm"))?;
let inter = self.linear(emb, &format!("{p}.intermediate.dense"))?;
let inter = self.gelu(inter)?;
let ffn = self.linear(inter, &format!("{p}.output.dense"))?;
emb = self.add_ln(emb, ffn, &format!("{p}.output.LayerNorm"))?;
}
false
};
let cls_idx: Vec<i64> = if cls_prepacked {
(0..n_docs)
.map(|i| index_to_i64(i, "forward_batch.cls_idx"))
.collect::<SearchResult<_>>()?
} else {
offsets
.iter()
.map(|&o| index_to_i64(o, "forward_batch.cls_offset"))
.collect::<SearchResult<_>>()?
};
let cls_t = self.idx(&cls_idx)?;
let cls = self
.s
.tensor_index_select(emb, 0, cls_t)
.map_err(|e| rerank_err("batch.cls_gather", e))?; let pooled = self.linear(cls, "bert.pooler.dense")?;
let pooled = self
.s
.tensor_tanh(pooled)
.map_err(|e| rerank_err("pooler.tanh", e))?;
let logit_t = self.linear(pooled, "classifier")?; let vals = self
.s
.tensor_values_f32(logit_t)
.map_err(|e| rerank_err("classifier.values", e))?;
self.s.truncate_autograd_graph(self.weights_boundary);
if vals.len() != n_docs {
return Err(rerank_err(
"classifier",
format!("expected {n_docs} logits, got {}", vals.len()),
));
}
Ok(vals)
}
pub(crate) fn embed_forward(&mut self, batch: &[Vec<i64>]) -> SearchResult<Vec<Vec<f32>>> {
let n_docs = batch.len();
let lens: Vec<usize> = batch.iter().map(Vec::len).collect();
let total: usize = lens.iter().sum();
if total == 0 {
return Ok(vec![vec![0.0; H]; n_docs]);
}
let mut offsets = Vec::with_capacity(n_docs);
{
let mut o = 0usize;
for &l in &lens {
offsets.push(o);
o += l;
}
}
let mut ids_flat = Vec::with_capacity(total);
let mut pos_flat = Vec::with_capacity(total);
let mut typ_flat = Vec::with_capacity(total);
for ids in batch {
for (i, &id) in ids.iter().enumerate() {
ids_flat.push(id);
pos_flat.push(index_to_i64(i, "embed_forward.position")?);
typ_flat.push(0i64);
}
}
let id_t = self.idx(&ids_flat)?;
let pos_t = self.idx(&pos_flat)?;
let typ_t = self.idx(&typ_flat)?;
let we = self.g("bert.embeddings.word_embeddings.weight")?;
let pe = self.g("bert.embeddings.position_embeddings.weight")?;
let te = self.g("bert.embeddings.token_type_embeddings.weight")?;
let e_word = self
.s
.tensor_index_select(we, 0, id_t)
.map_err(|e| rerank_err("embed.word", e))?;
let e_pos = self
.s
.tensor_index_select(pe, 0, pos_t)
.map_err(|e| rerank_err("embed.pos", e))?;
let e_typ = self
.s
.tensor_index_select(te, 0, typ_t)
.map_err(|e| rerank_err("embed.type", e))?;
let emb_wp = self
.s
.tensor_add(e_word, e_pos)
.map_err(|e| rerank_err("embed.add", e))?;
let emb = self.add_ln(emb_wp, e_typ, "bert.embeddings.LayerNorm")?;
let scale = ATTN_SCALE_F32;
let mut scratch = AttnScratch::default();
let mut emb_vals = self
.s
.tensor_values_f32(emb)
.map_err(|e| rerank_err("embed.extract", e))?;
for i in 0..L {
let p = format!("bert.encoder.layer.{i}");
emb_vals =
self.encoder_layer_raw(&emb_vals, total, &offsets, &lens, &p, scale, &mut scratch)?;
}
self.s.truncate_autograd_graph(self.weights_boundary);
let mut out = Vec::with_capacity(n_docs);
for (&off, &len) in offsets.iter().zip(&lens) {
let mut acc = vec![0.0f32; H];
if len > 0 {
let doc = &emb_vals[off * H..(off + len) * H];
for t in 0..len {
let row = &doc[t * H..t * H + H];
for (a, &r) in acc.iter_mut().zip(row) {
*a += r;
}
}
let inv = 1.0 / len as f32;
for a in &mut acc {
*a *= inv;
}
}
let norm = acc.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 0.0 {
let inv = 1.0 / norm;
for a in &mut acc {
*a *= inv;
}
}
out.push(acc);
}
Ok(out)
}
}
pub struct NativeReranker {
inner: Mutex<Model>,
tokenizer: Tokenizer,
max_length: usize,
name: String,
id: String,
}
impl std::fmt::Debug for NativeReranker {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("NativeReranker")
.field("name", &self.name)
.field("max_length", &self.max_length)
.finish_non_exhaustive()
}
}
impl NativeReranker {
pub fn load(model_dir: impl AsRef<Path>) -> SearchResult<Self> {
let dir = model_dir.as_ref();
let tok_path = dir.join(TOKENIZER_JSON);
if !tok_path.is_file() {
return Err(SearchError::ModelNotFound {
name: format!(
"{MODEL_NAME} (missing {TOKENIZER_JSON} in {})",
dir.display()
),
});
}
let mut tokenizer =
Tokenizer::from_file(&tok_path).map_err(|e| SearchError::ModelLoadFailed {
path: tok_path.clone(),
source: format!("tokenizer load failed: {e}").into(),
})?;
tokenizer
.with_truncation(Some(tokenizers::TruncationParams {
max_length: DEFAULT_MAX_LENGTH,
..Default::default()
}))
.map_err(|e| SearchError::ModelLoadFailed {
path: tok_path.clone(),
source: format!("failed to enable truncation: {e}").into(),
})?;
let weights_path = {
let primary = dir.join(SAFETENSORS_PRIMARY);
if primary.is_file() {
primary
} else {
dir.join(SAFETENSORS_FALLBACK)
}
};
if !weights_path.is_file() {
return Err(SearchError::ModelNotFound {
name: format!(
"{MODEL_NAME} (missing {SAFETENSORS_PRIMARY} or {SAFETENSORS_FALLBACK} in {})",
dir.display()
),
});
}
let shared = parse_weights(&weights_path)?;
let model = build_model(&shared)?;
tracing::info!(
model = MODEL_NAME,
linear_int8 = shared.qw.len(),
f32_params = shared.f32_params.len(),
max_length = DEFAULT_MAX_LENGTH,
model_dir = %dir.display(),
"native frankentorch reranker loaded (int8 linear, parallel forward)"
);
Ok(Self {
inner: Mutex::new(model),
tokenizer,
max_length: DEFAULT_MAX_LENGTH,
name: MODEL_NAME.to_owned(),
id: MODEL_NAME.to_owned(),
})
}
}
fn is_linear_weight(name: &str) -> bool {
name.ends_with(".weight") && !name.contains("LayerNorm") && !name.contains("embeddings")
}
pub(crate) struct SharedWeights {
qw: HashMap<String, QLinear>,
f32_params: HashMap<String, (Arc<Vec<f32>>, Vec<usize>)>,
}
pub(crate) fn parse_weights(path: &Path) -> SearchResult<SharedWeights> {
let bytes = fs::read(path).map_err(|e| SearchError::ModelLoadFailed {
path: path.to_path_buf(),
source: format!("read safetensors: {e}").into(),
})?;
if bytes.len() < 8 {
return Err(SearchError::ModelLoadFailed {
path: path.to_path_buf(),
source: "safetensors file too small".into(),
});
}
let header_len = usize::try_from(u64::from_le_bytes(bytes[0..8].try_into().expect("8 bytes")))
.map_err(|_| SearchError::ModelLoadFailed {
path: path.to_path_buf(),
source: "safetensors header length exceeds usize::MAX".into(),
})?;
let header_end = 8usize
.checked_add(header_len)
.filter(|&e| e <= bytes.len())
.ok_or_else(|| SearchError::ModelLoadFailed {
path: path.to_path_buf(),
source: "safetensors header length out of range".into(),
})?;
let header: serde_json::Value = serde_json::from_slice(&bytes[8..header_end]).map_err(|e| {
SearchError::ModelLoadFailed {
path: path.to_path_buf(),
source: format!("safetensors header parse: {e}").into(),
}
})?;
let data = &bytes[header_end..];
let obj = header
.as_object()
.ok_or_else(|| SearchError::ModelLoadFailed {
path: path.to_path_buf(),
source: "safetensors header is not an object".into(),
})?;
let mut raw: HashMap<String, (Vec<f32>, Vec<usize>)> = HashMap::new();
for (name, info) in obj {
if name == "__metadata__" {
continue;
}
let dtype = info
.get("dtype")
.and_then(serde_json::Value::as_str)
.unwrap_or("");
if dtype != "F32" {
continue; }
let shape: Vec<usize> = info
.get("shape")
.and_then(serde_json::Value::as_array)
.map(|a| {
a.iter()
.filter_map(serde_json::Value::as_u64)
.map(|u| {
usize::try_from(u).map_err(|_| SearchError::ModelLoadFailed {
path: path.to_path_buf(),
source: format!(
"safetensors tensor {name} shape dimension exceeds usize::MAX"
)
.into(),
})
})
.collect::<SearchResult<_>>()
})
.transpose()?
.unwrap_or_default();
let offsets = info
.get("data_offsets")
.and_then(serde_json::Value::as_array)
.ok_or_else(|| SearchError::ModelLoadFailed {
path: path.to_path_buf(),
source: format!("safetensors tensor {name} missing data_offsets").into(),
})?;
let start = usize::try_from(
offsets
.first()
.and_then(serde_json::Value::as_u64)
.unwrap_or(0),
)
.map_err(|_| SearchError::ModelLoadFailed {
path: path.to_path_buf(),
source: format!("safetensors tensor {name} start offset exceeds usize::MAX").into(),
})?;
let end = usize::try_from(
offsets
.get(1)
.and_then(serde_json::Value::as_u64)
.unwrap_or(0),
)
.map_err(|_| SearchError::ModelLoadFailed {
path: path.to_path_buf(),
source: format!("safetensors tensor {name} end offset exceeds usize::MAX").into(),
})?;
if start > end || end > data.len() {
return Err(SearchError::ModelLoadFailed {
path: path.to_path_buf(),
source: format!("safetensors tensor {name} has out-of-range offsets").into(),
});
}
let (chunks, _) = data[start..end].as_chunks::<4>();
let vals: Vec<f32> = chunks
.iter()
.map(|bytes| f32::from_le_bytes(*bytes))
.collect();
let key = if name.starts_with("embeddings.") || name.starts_with("encoder.") {
format!("bert.{name}")
} else {
name.clone()
};
raw.insert(key, (vals, shape));
}
if raw.is_empty() {
return Err(SearchError::ModelLoadFailed {
path: path.to_path_buf(),
source: "no F32 tensors found in safetensors".into(),
});
}
let mut qw: HashMap<String, QLinear> = HashMap::new();
let mut f32_params: HashMap<String, (Arc<Vec<f32>>, Vec<usize>)> = HashMap::new();
for (name, (vals, shape)) in &raw {
if is_linear_weight(name) {
let prefix = name.strip_suffix(".weight").expect("ends_with .weight");
let out = *shape.first().unwrap_or(&0);
let in_ = *shape.get(1).unwrap_or(&0);
if out == 0 || in_ == 0 || vals.len() != out * in_ {
return Err(SearchError::ModelLoadFailed {
path: path.to_path_buf(),
source: format!(
"linear weight {name} bad shape {shape:?} for {} values",
vals.len()
)
.into(),
});
}
let (w_i8, w_scales) = quantize_per_output_channel_i8(vals, out, in_);
let packed = cfg!(target_arch = "aarch64") && out % 4 == 0 && in_ % 16 == 0;
let w_i8 = if packed {
ft_api::pack_int8_weights_nr4(&w_i8, out, in_)
} else {
w_i8
};
let bias = raw
.get(&format!("{prefix}.bias"))
.map(|(b, _)| b.clone())
.unwrap_or_else(|| vec![0.0f32; out]);
qw.insert(
prefix.to_string(),
QLinear {
w_i8: Arc::new(w_i8),
w_scales: Arc::new(w_scales),
bias: Arc::new(bias),
out,
in_,
packed,
},
);
} else if name.strip_suffix(".bias").is_some() && !name.contains("LayerNorm") {
} else {
f32_params.insert(name.clone(), (Arc::new(vals.clone()), shape.clone()));
}
}
if qw.is_empty() {
return Err(SearchError::ModelLoadFailed {
path: path.to_path_buf(),
source: "no Linear weights found to quantize".into(),
});
}
for i in 0..L {
let p = format!("bert.encoder.layer.{i}");
let parts = ["query", "key", "value"];
let mut stacked: Vec<f32> = Vec::with_capacity(3 * H * H);
let mut bias: Vec<f32> = Vec::with_capacity(3 * H);
let mut ok = true;
for part in parts {
let wn = format!("{p}.attention.self.{part}.weight");
match raw.get(&wn) {
Some((vals, shape)) if shape.len() == 2 && shape[0] == H && shape[1] == H => {
stacked.extend_from_slice(vals);
let b = raw
.get(&format!("{p}.attention.self.{part}.bias"))
.map(|(b, _)| b.clone())
.unwrap_or_else(|| vec![0.0f32; H]);
bias.extend_from_slice(&b);
}
_ => {
ok = false;
break;
}
}
}
if !ok {
continue;
}
let (out, in_) = (3 * H, H);
let (w_i8, w_scales) = quantize_per_output_channel_i8(&stacked, out, in_);
let packed = cfg!(target_arch = "aarch64") && out % 4 == 0 && in_ % 16 == 0;
let w_i8 = if packed {
ft_api::pack_int8_weights_nr4(&w_i8, out, in_)
} else {
w_i8
};
qw.insert(
format!("{p}.attention.self.qkv"),
QLinear {
w_i8: Arc::new(w_i8),
w_scales: Arc::new(w_scales),
bias: Arc::new(bias),
out,
in_,
packed,
},
);
}
Ok(SharedWeights { qw, f32_params })
}
pub(crate) fn build_model(shared: &SharedWeights) -> SearchResult<Model> {
let mut session = FrankenTorchSession::new(ExecutionMode::Strict);
session.no_grad_enter();
let mut w = HashMap::with_capacity(shared.f32_params.len());
let mut raw_params = HashMap::with_capacity(shared.f32_params.len());
for (name, (vals, shape)) in &shared.f32_params {
let node = session
.tensor_variable_f32(vals.as_ref().clone(), shape.clone(), false)
.map_err(|e| rerank_err("build_model", format!("create f32 tensor {name}: {e}")))?;
w.insert(name.clone(), node);
raw_params.insert(name.clone(), Arc::clone(vals));
}
let weights_boundary = session.autograd_graph_node_count();
Ok(Model {
s: session,
w,
qw: shared.qw.clone(),
raw_params,
weights_boundary,
})
}
#[inline]
fn sigmoid(x: f32) -> f32 {
1.0 / (1.0 + (-x).exp())
}
impl SyncRerank for NativeReranker {
fn rerank_sync(
&self,
query: &str,
documents: &[RerankDocument],
) -> SearchResult<Vec<RerankScore>> {
if documents.is_empty() {
return Ok(Vec::new());
}
let mut encoded: Vec<(Vec<i64>, Vec<i64>)> = Vec::with_capacity(documents.len());
for doc in documents {
let encoding = self
.tokenizer
.encode((query, doc.text.as_str()), true)
.map_err(|e| rerank_err("tokenize", e))?;
let ids = crate::ids_to_truncated_i64(encoding.get_ids(), self.max_length);
let typ = crate::ids_to_truncated_i64(encoding.get_type_ids(), self.max_length);
encoded.push((ids, typ));
}
let mut model = self
.inner
.lock()
.map_err(|e| rerank_err("lock", format!("reranker mutex poisoned: {e}")))?;
let mut logits: Vec<f32> = Vec::with_capacity(documents.len());
let mut chunk_start = 0usize;
while chunk_start < encoded.len() {
let mut chunk_end = chunk_start + 1;
let mut chunk_tokens = encoded[chunk_start].0.len();
while chunk_end < encoded.len()
&& chunk_tokens + encoded[chunk_end].0.len() <= MAX_BATCH_TOKENS
{
chunk_tokens += encoded[chunk_end].0.len();
chunk_end += 1;
}
logits.extend(model.forward_batch(&encoded[chunk_start..chunk_end])?);
chunk_start = chunk_end;
}
drop(model);
let out = documents
.iter()
.zip(logits)
.enumerate()
.map(|(rank, (doc, logit))| {
let (score, raw_logit) = if logit.is_finite() {
(sigmoid(logit), Some(logit))
} else {
(0.0, None)
};
RerankScore {
doc_id: doc.doc_id.clone(),
score,
original_rank: rank,
raw_logit,
}
})
.collect();
Ok(out)
}
fn id(&self) -> &str {
&self.id
}
fn model_name(&self) -> &str {
&self.name
}
fn max_length(&self) -> usize {
self.max_length
}
fn is_available(&self) -> bool {
true
}
}
#[cfg(test)]
mod tests {
use super::*;
const MODEL_DIR: &str = "/private/tmp/ee-reranker-port/model";
const PARITY_TOL: f64 = 0.6;
const CASES: &[(&str, &str, f64)] = &[
(
"how to fix a failing release workflow",
"the release pipeline builds cross platform binaries and uploads them to github",
-9.808_567,
),
(
"how to fix a failing release workflow",
"bananas are a good source of potassium and taste sweet",
-11.332_987,
),
(
"what is the capital of france",
"paris is the capital and most populous city of france",
7.472_003,
),
(
"rust memory safety",
"the borrow checker enforces ownership rules at compile time",
-11.367_251,
),
];
fn model_available() -> bool {
Path::new(MODEL_DIR).join(TOKENIZER_JSON).is_file()
&& (Path::new(MODEL_DIR).join(SAFETENSORS_PRIMARY).is_file()
|| Path::new(MODEL_DIR).join(SAFETENSORS_FALLBACK).is_file())
}
fn doc(id: &str, text: &str) -> RerankDocument {
RerankDocument {
doc_id: id.to_owned(),
text: text.to_owned(),
}
}
#[test]
fn parity_logits_and_ranking_match_reference() {
if !model_available() {
eprintln!("[native_reranker] SKIP parity: model dir {MODEL_DIR} not present");
return;
}
let reranker = NativeReranker::load(MODEL_DIR).expect("load native reranker");
let mut logits = Vec::new();
let mut max_diff = 0.0_f64;
eprintln!("[native_reranker] idx | ft_logit | ref_logit | diff");
for (i, (query, document, ref_logit)) in CASES.iter().enumerate() {
let scored = reranker
.rerank_sync(query, &[doc("d", document)])
.expect("rerank_sync");
assert_eq!(scored.len(), 1, "one doc in, one score out");
let logit = f64::from(scored[0].raw_logit.expect("raw logit present"));
let diff = (logit - ref_logit).abs();
max_diff = max_diff.max(diff);
logits.push(logit);
eprintln!("[native_reranker] {i:3} | {logit:12.6} | {ref_logit:12.6} | {diff:8.5}");
assert!(
diff < PARITY_TOL,
"case {i} logit {logit} differs from reference {ref_logit} by {diff} (>{PARITY_TOL})"
);
}
let mut order: Vec<usize> = (0..logits.len()).collect();
order.sort_by(|&a, &b| logits[b].partial_cmp(&logits[a]).unwrap());
eprintln!(
"[native_reranker] ranking(desc)={order:?} expected=[2, 0, 1, 3] max_diff={max_diff:.6}"
);
assert_eq!(order, vec![2usize, 0, 1, 3], "ranking must match reference");
}
#[test]
fn forward_batch_matches_per_doc() {
if !model_available() {
eprintln!("[native_reranker] SKIP batch-equiv: model dir not present");
return;
}
let reranker = NativeReranker::load(MODEL_DIR).expect("load native reranker");
let query = CASES[0].0;
let mut batch: Vec<(Vec<i64>, Vec<i64>)> = Vec::new();
for (_, document, _) in CASES {
let enc = reranker
.tokenizer
.encode((query, *document), true)
.expect("tokenize");
let ids: Vec<i64> = enc.get_ids().iter().map(|&x| i64::from(x)).collect();
let typ: Vec<i64> = enc.get_type_ids().iter().map(|&x| i64::from(x)).collect();
batch.push((ids, typ));
}
let mut model = reranker.inner.lock().expect("lock");
let per_doc: Vec<f32> = batch
.iter()
.map(|(ids, typ)| model.forward(ids, typ).expect("forward"))
.collect();
let batched = model.forward_batch(&batch).expect("forward_batch");
drop(model);
assert_eq!(batched.len(), per_doc.len());
for (i, (b, p)) in batched.iter().zip(&per_doc).enumerate() {
let diff = (f64::from(*b) - f64::from(*p)).abs();
eprintln!("[native_reranker] doc {i}: batched={b:.6} per_doc={p:.6} diff={diff:.2e}");
assert!(
diff < 1e-3,
"doc {i}: batched {b} vs per-doc {p} diff {diff} too large"
);
}
}
#[test]
fn empty_documents_yield_empty_scores() {
if !model_available() {
eprintln!("[native_reranker] SKIP empty-docs: model dir not present");
return;
}
let reranker = NativeReranker::load(MODEL_DIR).expect("load");
let scored = reranker.rerank_sync("any query", &[]).expect("empty ok");
assert!(scored.is_empty());
eprintln!("[native_reranker] empty-docs -> empty scores OK");
}
#[test]
fn whitespace_and_long_documents_do_not_panic() {
if !model_available() {
eprintln!("[native_reranker] SKIP whitespace/long: model dir not present");
return;
}
let reranker = NativeReranker::load(MODEL_DIR).expect("load");
let ws = reranker
.rerank_sync("q", &[doc("ws", " ")])
.expect("whitespace ok");
assert_eq!(ws.len(), 1);
let long_text = "memory safety ".repeat(400);
let lng = reranker
.rerank_sync("rust", &[doc("long", &long_text)])
.expect("long ok");
assert_eq!(lng.len(), 1);
assert!(lng[0].score.is_finite());
eprintln!(
"[native_reranker] whitespace score={:.6}, truncated-long score={:.6} OK",
ws[0].score, lng[0].score
);
}
#[test]
fn ranking_is_deterministic_across_runs() {
if !model_available() {
eprintln!("[native_reranker] SKIP determinism: model dir not present");
return;
}
let reranker = NativeReranker::load(MODEL_DIR).expect("load");
let docs: Vec<RerankDocument> = CASES
.iter()
.enumerate()
.map(|(i, (_, d, _))| doc(&format!("d{i}"), d))
.collect();
let run1 = reranker
.rerank_sync("what is the capital of france", &docs)
.expect("run1");
let run2 = reranker
.rerank_sync("what is the capital of france", &docs)
.expect("run2");
assert_eq!(run1.len(), run2.len());
for (a, b) in run1.iter().zip(run2.iter()) {
assert_eq!(a.doc_id, b.doc_id);
assert_eq!(a.raw_logit, b.raw_logit, "logits must be deterministic");
}
eprintln!("[native_reranker] determinism across 2 runs OK");
}
#[test]
fn many_documents_rerank_without_deadlock() {
if !model_available() {
eprintln!("[native_reranker] SKIP many-docs: model dir not present");
return;
}
let reranker = NativeReranker::load(MODEL_DIR).expect("load");
let docs: Vec<RerankDocument> = (0..24)
.map(|i| doc(&format!("d{i}"), CASES[i % CASES.len()].1))
.collect();
let scored = reranker
.rerank_sync("what is the capital of france", &docs)
.expect("many-doc rerank completes (no deadlock)");
assert_eq!(scored.len(), docs.len(), "one score per doc");
for (i, s) in scored.iter().enumerate() {
assert_eq!(s.original_rank, i, "original_rank preserves input order");
assert_eq!(s.doc_id, format!("d{i}"));
assert!(s.score.is_finite());
}
let again = reranker
.rerank_sync("what is the capital of france", &docs)
.expect("rerun");
for (a, b) in scored.iter().zip(again.iter()) {
assert_eq!(
a.raw_logit, b.raw_logit,
"parallel rerank must be deterministic"
);
}
eprintln!(
"[native_reranker] {}-doc concurrent rerank OK (no deadlock, deterministic)",
docs.len()
);
}
#[test]
fn load_missing_dir_errors() {
let err = NativeReranker::load("/private/tmp/definitely-not-a-model-dir-xyz");
assert!(err.is_err(), "loading a missing dir must error, not panic");
eprintln!("[native_reranker] missing-dir load error OK: {err:?}");
}
}