use crate::text::tokenize;
pub const DEMO_DIM: usize = 512;
pub fn embed(text: &str, dim: usize) -> Vec<f32> {
let mut v = vec![0.0f32; dim.max(1)];
let toks = tokenize(text);
for tok in &toks {
add_feature(&mut v, dim, &format!("w:{tok}"), 1.0);
}
for pair in toks.windows(2) {
add_feature(&mut v, dim, &format!("b:{}_{}", pair[0], pair[1]), 0.6);
}
for tok in &toks {
if tok.chars().count() < 4 {
continue;
}
let padded = format!("^{tok}$");
let chars: Vec<char> = padded.chars().collect();
for w in chars.windows(3) {
let tri: String = w.iter().collect();
add_feature(&mut v, dim, &format!("c:{tri}"), 0.2);
}
}
l2_normalize(&mut v);
v
}
fn add_feature(v: &mut [f32], dim: usize, feature: &str, weight: f32) {
let h = fnv1a(feature.as_bytes());
let bucket = (h % dim as u64) as usize;
let sign = if fnv1a_salted(feature.as_bytes(), 0x9E37_79B9) & 1 == 0 {
1.0
} else {
-1.0
};
v[bucket] += sign * weight;
}
fn l2_normalize(v: &mut [f32]) {
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
if norm > 1e-9 {
for x in v.iter_mut() {
*x /= norm;
}
}
}
fn fnv1a(bytes: &[u8]) -> u64 {
let mut h: u64 = 0xcbf2_9ce4_8422_2325;
for &b in bytes {
h ^= b as u64;
h = h.wrapping_mul(0x0000_0100_0000_01b3);
}
h
}
fn fnv1a_salted(bytes: &[u8], salt: u64) -> u64 {
let mut h: u64 = 0xcbf2_9ce4_8422_2325 ^ salt;
for &b in bytes {
h ^= b as u64;
h = h.wrapping_mul(0x0000_0100_0000_01b3);
}
h
}
#[cfg(test)]
mod tests {
use super::*;
fn cos(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b).map(|(x, y)| x * y).sum()
}
#[test]
fn deterministic() {
assert_eq!(
embed("rust memory safety", DEMO_DIM),
embed("rust memory safety", DEMO_DIM)
);
}
#[test]
fn normalized() {
let v = embed("hierarchical navigable small world graphs", DEMO_DIM);
let norm = v.iter().map(|x| x * x).sum::<f32>().sqrt();
assert!((norm - 1.0).abs() < 1e-4, "norm was {norm}");
}
#[test]
fn related_text_is_closer_than_unrelated() {
let q = embed("rust memory safety", DEMO_DIM);
let related = embed("the rust borrow checker guarantees memory safety", DEMO_DIM);
let unrelated = embed("slow simmered tomato pasta sauce", DEMO_DIM);
assert!(
cos(&q, &related) > cos(&q, &unrelated),
"related {} should beat unrelated {}",
cos(&q, &related),
cos(&q, &unrelated)
);
}
#[test]
fn morphology_brings_word_variants_near() {
let q = embed("distant galaxies and telescopes", DEMO_DIM);
let related = embed("a telescope resolving a faraway galaxy", DEMO_DIM);
let unrelated = embed("slow simmered tomato pasta sauce", DEMO_DIM);
assert!(cos(&q, &related) > cos(&q, &unrelated));
}
#[test]
fn empty_is_zero_vector() {
let v = embed("", DEMO_DIM);
assert!(v.iter().all(|&x| x == 0.0));
}
}