use ternlang_core::trit::Trit;
use ternlang_ml::{TritMatrix, quantize, bitnet_threshold};
use rayon::prelude::*;
use serde::{Serialize, Deserialize};
#[derive(Serialize, Deserialize)]
pub struct RuVectorDB {
pub embeddings: TritMatrix,
pub metadata: Vec<String>,
}
impl RuVectorDB {
pub fn from_f32(embeddings: &[Vec<f32>], metadata: Vec<String>) -> anyhow::Result<Self> {
if embeddings.is_empty() {
return Err(anyhow::anyhow!("Embeddings cannot be empty"));
}
let rows = embeddings.len();
let cols = embeddings[0].len();
let flat_f32: Vec<f32> = embeddings.iter().flatten().cloned().collect();
let threshold = bitnet_threshold(&flat_f32);
let trit_matrix = TritMatrix::from_f32(rows, cols, &flat_f32, threshold);
Ok(Self {
embeddings: trit_matrix,
metadata,
})
}
pub fn search(&self, query_f32: &[f32], top_k: usize) -> Vec<SearchResult> {
let threshold = bitnet_threshold(query_f32);
let query_trits = quantize(query_f32, threshold);
let scores = self.sparse_gemv_similarity(&query_trits);
let mut results: Vec<SearchResult> = scores.into_iter()
.enumerate()
.map(|(i, score)| SearchResult {
index: i,
score,
metadata: self.metadata.get(i).cloned().unwrap_or_default(),
})
.collect();
results.sort_by(|a, b| b.score.partial_cmp(&a.score).unwrap());
results.truncate(top_k);
results
}
fn sparse_gemv_similarity(&self, query: &[Trit]) -> Vec<f32> {
let num_docs = self.embeddings.rows;
let dim = self.embeddings.cols;
let q_flat: Vec<i8> = query.iter().map(|&t| match t {
Trit::Affirm => 1,
Trit::Reject => -1,
Trit::Tend => 0,
}).collect();
let db_flat = self.embeddings.to_i8_vec();
(0..num_docs).into_par_iter().map(|row_idx| {
let row_data = &db_flat[row_idx * dim .. (row_idx + 1) * dim];
let mut acc: i32 = 0;
for i in 0..dim {
let qi = q_flat[i];
if qi == 0 { continue; }
let di = row_data[i];
if di == 0 { continue; }
acc += (qi * di) as i32;
}
acc as f32
}).collect()
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct SearchResult {
pub index: usize,
pub score: f32,
pub metadata: String,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_ruvector_search() {
let embeddings = vec![
vec![1.0, 0.0, -1.0, 0.5],
vec![-1.0, 1.0, 0.0, 0.0],
vec![0.1, 0.1, 0.1, 0.1], ];
let metadata = vec!["Doc A".to_string(), "Doc B".to_string(), "Doc C".to_string()];
let db = RuVectorDB::from_f32(&embeddings, metadata).unwrap();
let query = vec![1.0, 0.0, -1.0, 0.0];
let results = db.search(&query, 2);
assert_eq!(results.len(), 2);
assert_eq!(results[0].metadata, "Doc A");
assert!(results[0].score > results[1].score);
}
}