use std::collections::{HashMap, HashSet};
use crate::{all_finite, lcg_f32, Error};
#[derive(Debug)]
pub(crate) struct DenseSimHash {
embedding_dim: usize,
num_bits: usize,
hyperplanes: Vec<Vec<f32>>,
}
impl DenseSimHash {
pub(crate) fn new(embedding_dim: usize, num_bits: usize) -> Result<Self, Error> {
if embedding_dim == 0 {
return Err(Error::InvalidParam("embedding_dim must be >= 1"));
}
if num_bits == 0 || num_bits > 64 {
return Err(Error::InvalidParam("num_bits must be in [1, 64]"));
}
let mut hyperplanes = Vec::with_capacity(num_bits);
let mut rng_state = 0x12345678u64;
for _ in 0..num_bits {
let mut plane = Vec::with_capacity(embedding_dim);
for _ in 0..embedding_dim {
plane.push(lcg_f32(&mut rng_state));
}
hyperplanes.push(plane);
}
Ok(Self {
embedding_dim,
num_bits,
hyperplanes,
})
}
pub(crate) fn fingerprint(&self, embedding: &[f32]) -> Result<u64, Error> {
if embedding.len() != self.embedding_dim {
return Err(Error::DimensionMismatch {
expected: self.embedding_dim,
got: embedding.len(),
});
}
if !all_finite(embedding) {
return Err(Error::NonFiniteInput);
}
let mut hash = 0u64;
for (i, plane) in self.hyperplanes.iter().enumerate() {
let dot: f32 = plane.iter().zip(embedding.iter()).map(|(a, b)| a * b).sum();
if dot > 0.0 {
hash |= 1u64 << i;
}
}
Ok(hash)
}
pub(crate) fn num_bits(&self) -> usize {
self.num_bits
}
}
#[derive(Debug)]
pub struct DenseSimHashLSH {
simhash: DenseSimHash,
buckets: HashMap<u64, Vec<usize>>,
ids: Vec<String>,
}
impl DenseSimHashLSH {
pub fn new(embedding_dim: usize, num_bits: usize) -> Result<Self, Error> {
Ok(Self {
simhash: DenseSimHash::new(embedding_dim, num_bits)?,
buckets: HashMap::new(),
ids: Vec::new(),
})
}
pub fn insert(&mut self, id: impl Into<String>, embedding: &[f32]) -> Result<usize, Error> {
let idx = self.ids.len();
let fp = self.simhash.fingerprint(embedding)?;
self.buckets.entry(fp).or_default().push(idx);
self.ids.push(id.into());
Ok(idx)
}
pub fn query(&self, embedding: &[f32]) -> Result<Vec<usize>, Error> {
let fp = self.simhash.fingerprint(embedding)?;
let mut candidates: HashSet<usize> = HashSet::new();
if let Some(indices) = self.buckets.get(&fp) {
candidates.extend(indices.iter().copied());
}
for bit in 0..self.simhash.num_bits() {
let neighbor = fp ^ (1u64 << bit);
if let Some(indices) = self.buckets.get(&neighbor) {
candidates.extend(indices.iter().copied());
}
}
let mut v: Vec<usize> = candidates.into_iter().collect();
v.sort_unstable();
Ok(v)
}
pub fn get_id(&self, idx: usize) -> Option<&str> {
self.ids.get(idx).map(|s| s.as_str())
}
pub fn len(&self) -> usize {
self.ids.len()
}
pub fn is_empty(&self) -> bool {
self.ids.is_empty()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn smoke_query_returns_candidates() {
let mut lsh = DenseSimHashLSH::new(8, 64).unwrap();
let v1: Vec<f32> = (0..8).map(|i| (i as f32).sin()).collect();
let v2: Vec<f32> = (0..8).map(|i| (i as f32).sin() + 0.01).collect();
let v3: Vec<f32> = (0..8).map(|i| (i as f32).cos()).collect();
lsh.insert("1", &v1).unwrap();
lsh.insert("2", &v2).unwrap();
lsh.insert("3", &v3).unwrap();
let candidates = lsh.query(&v1).unwrap();
assert!(!candidates.is_empty());
}
#[test]
fn rejects_num_bits_over_64() {
assert!(DenseSimHashLSH::new(8, 65).is_err());
assert!(DenseSimHashLSH::new(8, 128).is_err());
}
#[test]
fn rejects_zero_num_bits() {
assert!(DenseSimHashLSH::new(8, 0).is_err());
}
#[test]
fn dense_simhash_fingerprint_determinism() {
let dsh = DenseSimHash::new(4, 16).unwrap();
let fp = dsh.fingerprint(&[1.0, -0.5, 0.3, 0.8]).unwrap();
assert_eq!(fp, 7679);
}
}