use crate::Result;
pub const DEFAULT_EMBEDDING_DIM: usize = 384;
#[derive(Debug, Clone, PartialEq)]
#[cfg_attr(feature = "serde", derive(serde::Serialize, serde::Deserialize))]
#[cfg_attr(feature = "serde", serde(transparent))]
pub struct Embedding(pub Vec<f32>);
impl Embedding {
pub fn dim(&self) -> usize {
self.0.len()
}
pub fn as_slice(&self) -> &[f32] {
&self.0
}
}
pub trait Embedder: Send + Sync {
fn dim(&self) -> usize;
fn embed(&self, text: &str) -> Result<Embedding>;
fn embed_batch(&self, texts: &[&str]) -> Result<Vec<Embedding>> {
texts.iter().map(|t| self.embed(t)).collect()
}
fn model_id(&self) -> &str {
"unknown"
}
}
#[derive(Debug, Clone)]
pub struct HashEmbedder {
pub dims: usize,
}
impl Default for HashEmbedder {
fn default() -> Self {
Self {
dims: DEFAULT_EMBEDDING_DIM,
}
}
}
impl Embedder for HashEmbedder {
fn dim(&self) -> usize {
self.dims
}
fn embed(&self, text: &str) -> Result<Embedding> {
use std::hash::{Hash, Hasher};
let mut vec = vec![0.0f32; self.dims];
for (lane, slot) in vec.iter_mut().enumerate() {
let mut h = std::collections::hash_map::DefaultHasher::new();
lane.hash(&mut h);
text.hash(&mut h);
let raw = h.finish();
*slot = ((raw >> 11) as f64 / (1u64 << 52) as f64 - 0.5) as f32 * 2.0;
}
let norm = vec.iter().map(|v| v * v).sum::<f32>().sqrt();
if norm > f32::EPSILON {
for v in &mut vec {
*v /= norm;
}
}
Ok(Embedding(vec))
}
fn model_id(&self) -> &str {
"hash-embedder"
}
}
#[cfg(test)]
mod tests {
use super::*;
struct ConstEmbedder;
impl Embedder for ConstEmbedder {
fn dim(&self) -> usize {
2
}
fn embed(&self, text: &str) -> Result<Embedding> {
Ok(Embedding(vec![text.len() as f32, 0.0]))
}
}
#[test]
fn default_dim_matches_mempalace_for_migration() {
assert_eq!(DEFAULT_EMBEDDING_DIM, 384);
}
#[test]
fn embed_batch_defaults_to_per_item_loop() {
let e = ConstEmbedder;
let got = e.embed_batch(&["a", "bb", "ccc"]).expect("must embed");
assert_eq!(got.len(), 3);
assert_eq!(got[0].dim(), 2);
assert_eq!(got[2].as_slice(), &[3.0, 0.0]);
}
}