use anyhow::Result;
use std::path::{Path, PathBuf};
use std::sync::OnceLock;
use crate::search::simd::dot_f32;
use ort::session::{Session, builder::GraphOptimizationLevel};
use ort::value::Tensor;
use parking_lot::Mutex;
use tokenizers::Tokenizer;
pub(crate) const EMBEDDING_DIM: usize = 768;
const MAX_CTX: usize = 8192;
pub(crate) const QUERY_PREFIX: &str = "Represent this query for searching relevant code: ";
const INPUT_IDS: &str = "input_ids";
const ATTENTION_MASK: &str = "attention_mask";
const POOLED_OUTPUT_NAMES: &[&str] = &["sentence_embedding"];
const TOKEN_OUTPUT_NAMES: &[&str] = &["token_embeddings", "last_hidden_state"];
pub(crate) fn l2_normalize(v: &mut [f32]) {
let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 1e-12 {
let inv = 1.0 / norm;
for x in v.iter_mut() {
*x *= inv;
}
}
}
pub(crate) fn quantize_int8(v: &[f32]) -> (Vec<i8>, f32) {
let max_abs = v.iter().fold(0.0f32, |m, &x| m.max(x.abs()));
let scale = (max_abs / 127.0).max(1e-8);
let q = v
.iter()
.map(|&x| {
let scaled = (x / scale).round();
scaled.clamp(-127.0, 127.0) as i8
})
.collect();
(q, scale)
}
pub(crate) fn dequantize_int8(q: &[i8], scale: f32) -> Vec<f32> {
q.iter().map(|&x| f32::from(x) * scale).collect()
}
pub(crate) fn prefix_query(q: &str) -> String {
format!("{QUERY_PREFIX}{q}")
}
pub(crate) fn prefix_document(d: &str) -> &str {
d
}
fn query_threads() -> usize {
crate::config::env_usize("SEMANTEX_ORT_THREADS", 4)
}
fn default_index_threads(cores: usize, max_concurrent_builds: usize) -> usize {
if max_concurrent_builds <= 1 {
cores.clamp(2, 8)
} else {
(cores / 2).clamp(2, 8)
}
}
fn index_threads() -> usize {
let cores = std::thread::available_parallelism().map_or(4, std::num::NonZeroUsize::get);
let default = default_index_threads(cores, crate::index::gate::max_concurrent_builds());
crate::config::env_usize("SEMANTEX_INDEX_ORT_THREADS", default)
}
pub struct SingleVectorEmbedder {
model_dir: PathBuf,
threads: usize,
use_coreml: bool,
session: OnceLock<Mutex<Session>>,
tokenizer: OnceLock<Tokenizer>,
build_lock: std::sync::Mutex<()>,
}
static GLOBAL: OnceLock<SingleVectorEmbedder> = OnceLock::new();
static GLOBAL_INIT_LOCK: parking_lot::Mutex<()> = parking_lot::Mutex::new(());
impl SingleVectorEmbedder {
pub fn embedding_dim() -> usize {
EMBEDDING_DIM
}
pub fn new(model_dir: &Path) -> Result<Self> {
Self::with_threads(model_dir, query_threads())
}
pub fn for_indexing(model_dir: &Path) -> Result<Self> {
Self::with_threads(model_dir, index_threads())
}
fn with_threads(model_dir: &Path, threads: usize) -> Result<Self> {
if !model_dir.exists() {
anyhow::bail!(
"CodeRankEmbed model dir does not exist: {}",
model_dir.display()
);
}
Ok(Self {
model_dir: model_dir.to_path_buf(),
threads: threads.max(1),
use_coreml: std::env::var("SEMANTEX_COREML").is_ok_and(|v| v == "1"),
session: OnceLock::new(),
tokenizer: OnceLock::new(),
build_lock: std::sync::Mutex::new(()),
})
}
pub fn global(model_dir: &Path) -> Result<&'static SingleVectorEmbedder> {
if let Some(e) = GLOBAL.get() {
return Ok(e);
}
let _g = GLOBAL_INIT_LOCK.lock();
if let Some(e) = GLOBAL.get() {
return Ok(e);
}
let e = Self::new(model_dir)?;
let _ = GLOBAL.set(e);
Ok(GLOBAL.get().expect("just set"))
}
pub fn is_initialized(&self) -> bool {
self.session.get().is_some()
}
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;
providers.push(ort::ep::CPU::default().build());
providers
}
fn build_session(&self) -> Result<Session> {
let runtime_root = crate::config::SemantexConfig::semantex_home().join("runtime");
let _ = crate::embedding::runtime_manager::ensure_onnxruntime(&runtime_root);
let model_path = self
.model_dir
.join(crate::embedding::single_vector_model::CODERANK_ONNX);
let session = Session::builder()?
.with_execution_providers(self.execution_providers())?
.with_intra_threads(self.threads)?
.with_optimization_level(GraphOptimizationLevel::Level3)?
.commit_from_file(&model_path)?;
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 tokenizer(&self) -> Result<&Tokenizer> {
if let Some(t) = self.tokenizer.get() {
return Ok(t);
}
let path = self.model_dir.join("tokenizer.json");
let tok = Tokenizer::from_file(&path)
.map_err(|e| anyhow::anyhow!("failed to load tokenizer {}: {e}", path.display()))?;
let _ = self.tokenizer.set(tok);
Ok(self.tokenizer.get().expect("tokenizer set above"))
}
pub fn encode_document(&self, text: &str) -> Result<Vec<f32>> {
self.encode_text(prefix_document(text))
}
pub fn encode_document_with_context(&self, context: &str, code: &str) -> Result<Vec<f32>> {
if context.is_empty() {
return self.encode_document(code);
}
self.encode_text(&format!("{context}\n{code}"))
}
pub fn encode_query(&self, text: &str) -> Result<Vec<f32>> {
self.encode_text(&prefix_query(text))
}
fn encode_text(&self, text: &str) -> Result<Vec<f32>> {
let encoding = self
.tokenizer()?
.encode(text, true)
.map_err(|e| anyhow::anyhow!("tokenize failed: {e}"))?;
let ids_u32 = encoding.get_ids();
let mask_u32 = encoding.get_attention_mask();
let len = ids_u32.len().min(MAX_CTX);
let ids: Vec<i64> = ids_u32[..len].iter().map(|&x| i64::from(x)).collect();
let mask: Vec<i64> = mask_u32[..len].iter().map(|&x| i64::from(x)).collect();
let seq = ids.len() as i64;
let id_tensor = Tensor::from_array((vec![1_i64, seq], ids))?;
let mask_tensor = Tensor::from_array((vec![1_i64, seq], mask.clone()))?;
let session = self.session()?;
let mut guard = session.lock();
let outputs = guard.run(ort::inputs![
INPUT_IDS => id_tensor,
ATTENTION_MASK => mask_tensor,
])?;
let dim = EMBEDDING_DIM;
let mut pooled = if let Some(name) = POOLED_OUTPUT_NAMES
.iter()
.find(|n| outputs.contains_key(**n))
{
let (shape, data) = outputs[*name].try_extract_tensor::<f32>()?;
let n: usize = shape.iter().map(|&d| d as usize).product();
anyhow::ensure!(
n == dim,
"unexpected pooled output `{name}` shape {shape:?} (expected {dim} values)"
);
data[..dim].to_vec()
} else {
let name = TOKEN_OUTPUT_NAMES
.iter()
.find(|n| outputs.contains_key(**n))
.ok_or_else(|| {
anyhow::anyhow!(
"encoder produced no known output (looked for {POOLED_OUTPUT_NAMES:?} / {TOKEN_OUTPUT_NAMES:?})"
)
})?;
let (shape, data) = outputs[*name].try_extract_tensor::<f32>()?;
anyhow::ensure!(
shape.len() == 3 && shape[2] as usize == dim,
"unexpected token output `{name}` shape {shape:?} (expected [1, seq, {dim}])"
);
let seq_len = shape[1] as usize;
let mut pooled = vec![0.0f32; dim];
let mut count = 0.0f32;
for t in 0..seq_len {
if mask.get(t).copied().unwrap_or(0) == 0 {
continue;
}
count += 1.0;
let row = &data[t * dim..(t + 1) * dim];
for (p, &x) in pooled.iter_mut().zip(row) {
*p += x;
}
}
if count > 0.0 {
for p in &mut pooled {
*p /= count;
}
}
pooled
};
l2_normalize(&mut pooled);
debug_assert!(dot_f32(&pooled, &pooled) <= 1.0 + 1e-3);
Ok(pooled)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn default_index_threads_uses_all_cores_when_only_one_build_slot_exists() {
assert_eq!(default_index_threads(4, 1), 4);
assert_eq!(default_index_threads(2, 1), 2, "clamped to floor of 2");
assert_eq!(default_index_threads(16, 1), 8, "clamped to ceiling of 8");
assert_eq!(default_index_threads(1, 1), 2, "clamped to floor of 2");
}
#[test]
fn default_index_threads_halves_cores_when_multiple_build_slots_exist() {
assert_eq!(default_index_threads(8, 2), 4);
assert_eq!(default_index_threads(32, 4), 8, "clamped to ceiling of 8");
assert_eq!(default_index_threads(4, 2), 2, "clamped to floor of 2");
}
#[test]
fn l2_normalize_unit_length() {
let mut v = vec![3.0f32, 4.0];
l2_normalize(&mut v);
let norm = (v[0] * v[0] + v[1] * v[1]).sqrt();
assert!((norm - 1.0).abs() < 1e-6, "expected unit norm, got {norm}");
assert!((v[0] - 0.6).abs() < 1e-6 && (v[1] - 0.8).abs() < 1e-6);
}
#[test]
fn l2_normalize_zero_vector_is_safe() {
let mut v = vec![0.0f32, 0.0, 0.0];
l2_normalize(&mut v); assert!(v.iter().all(|x| x.is_finite()));
}
#[test]
fn int8_round_trip_preserves_direction() {
let mut v = vec![0.1f32, -0.5, 0.8, 0.2, -0.3];
l2_normalize(&mut v);
let (q, scale) = quantize_int8(&v);
assert_eq!(q.len(), v.len());
assert!(scale > 0.0);
let back = dequantize_int8(&q, scale);
let dot: f32 = v.iter().zip(&back).map(|(a, b)| a * b).sum();
let nb: f32 = back.iter().map(|x| x * x).sum::<f32>().sqrt();
let cos = dot / nb.max(1e-12);
assert!(cos > 0.99, "int8 round-trip cosine too low: {cos}");
}
#[test]
fn int8_all_zero_vector_has_nonzero_scale() {
let (q, scale) = quantize_int8(&[0.0, 0.0, 0.0]);
assert!(q.iter().all(|&x| x == 0));
assert!(
scale > 0.0,
"scale must never be zero (avoids div-by-zero on dequant)"
);
}
#[test]
fn embedder_new_is_lazy_no_session_at_construction() {
let tmp = tempfile::TempDir::new().unwrap();
let emb = SingleVectorEmbedder::new(tmp.path())
.expect("constructor must succeed for any existing dir");
assert!(
!emb.is_initialized(),
"session must not build at construction"
);
}
#[test]
fn embedder_rejects_missing_dir() {
let res = SingleVectorEmbedder::new(Path::new("/nonexistent/s2/model/dir"));
assert!(res.is_err(), "missing model dir must fail at construction");
}
#[test]
fn query_prefix_is_applied_document_is_raw() {
let doc = "fn add(a:i32,b:i32)->i32{a+b}";
assert_eq!(prefix_document(doc), doc, "documents get NO prefix");
let q = "add two integers";
assert_eq!(prefix_query(q), format!("{QUERY_PREFIX}{q}"));
assert!(
QUERY_PREFIX.ends_with(' '),
"RECORDED prefix keeps its trailing space"
);
}
#[test]
fn embedding_dim_is_recorded_constant() {
assert_eq!(SingleVectorEmbedder::embedding_dim(), EMBEDDING_DIM);
const { assert!(EMBEDDING_DIM > 0) };
}
}