use core::{fmt, ops::Deref};
use crate::embeddings::granite::error::{EmbeddingDimMismatch, Error, Result};
pub const EMBEDDING_DIM: usize = 384;
pub(crate) const NORM_BUDGET: f32 = 1e-4;
#[derive(Clone)]
#[repr(transparent)]
pub struct Embedding {
inner: [f32; EMBEDDING_DIM],
}
impl Embedding {
#[inline]
pub const fn dim(&self) -> usize {
self.inner.len()
}
#[inline]
pub const fn as_slice(&self) -> &[f32] {
self.inner.as_slice()
}
#[inline]
pub fn to_vec(&self) -> Vec<f32> {
self.inner.to_vec()
}
pub fn try_from_unit_slice(s: &[f32]) -> Result<Self> {
if s.len() != EMBEDDING_DIM {
return Err(Error::EmbeddingDimMismatch(EmbeddingDimMismatch::new(
EMBEDDING_DIM,
s.len(),
)));
}
for (i, &v) in s.iter().enumerate() {
if !v.is_finite() {
return Err(Error::NonFiniteEmbedding(i));
}
}
let norm_sq: f32 = s.iter().map(|x| x * x).sum();
let dev = (norm_sq - 1.0).abs();
if dev > NORM_BUDGET {
return Err(Error::EmbeddingNotUnitNorm(dev));
}
Self::from_slice_normalizing(s)
}
pub fn from_slice_normalizing(s: &[f32]) -> Result<Self> {
if s.len() != EMBEDDING_DIM {
return Err(Error::EmbeddingDimMismatch(EmbeddingDimMismatch::new(
EMBEDDING_DIM,
s.len(),
)));
}
for (i, &v) in s.iter().enumerate() {
if !v.is_finite() {
return Err(Error::NonFiniteEmbedding(i));
}
}
let norm_sq_f64: f64 = s.iter().map(|&x| (x as f64) * (x as f64)).sum();
if norm_sq_f64 == 0.0 {
return Err(Error::EmbeddingZero);
}
let inv_norm_f64 = 1.0_f64 / norm_sq_f64.sqrt();
let mut inner = [0.0f32; EMBEDDING_DIM];
for (out, &v) in inner.iter_mut().zip(s.iter()) {
*out = ((v as f64) * inv_norm_f64) as f32;
}
Ok(Self { inner })
}
pub fn dot(&self, other: &Embedding) -> f32 {
self
.inner
.iter()
.zip(other.inner.iter())
.map(|(a, b)| a * b)
.sum()
}
pub fn cosine(&self, other: &Embedding) -> f32 {
self.dot(other)
}
pub fn is_close(&self, other: &Embedding, tol: f32) -> bool {
self
.inner
.iter()
.zip(other.inner.iter())
.map(|(a, b)| (a - b).abs())
.fold(0.0f32, f32::max)
<= tol
}
pub fn is_close_cosine(&self, other: &Embedding, tol: f32) -> bool {
let sq: f32 = self
.inner
.iter()
.zip(other.inner.iter())
.map(|(a, b)| {
let d = a - b;
d * d
})
.sum();
(sq * 0.5) <= tol
}
}
impl AsRef<[f32]> for Embedding {
#[inline]
fn as_ref(&self) -> &[f32] {
self.as_slice()
}
}
impl Deref for Embedding {
type Target = [f32];
#[inline]
fn deref(&self) -> &[f32] {
&self.inner
}
}
impl fmt::Debug for Embedding {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(
f,
"Embedding {{ dim: {}, head: [{:.4}, {:.4}, {:.4}, ..] }}",
self.dim(),
self.inner[0],
self.inner[1],
self.inner[2],
)
}
}
impl windit::windowed::Vector for Embedding {
type Scalar = f32;
fn as_slice(&self) -> &[f32] {
Embedding::as_slice(self)
}
fn from_unnormalized(v: &[f64]) -> core::result::Result<Self, windit::WinditError> {
if v.len() != EMBEDDING_DIM {
return Err(windit::WinditError::DimMismatch {
got: v.len(),
expected: EMBEDDING_DIM,
});
}
let narrowed: [f32; EMBEDDING_DIM] = core::array::from_fn(|i| v[i] as f32);
Embedding::from_slice_normalizing(&narrowed).map_err(|e| match e {
Error::EmbeddingDimMismatch(dim) => windit::WinditError::DimMismatch {
got: dim.got(),
expected: dim.expected(),
},
_ => windit::WinditError::NonFinite,
})
}
}
pub(crate) fn check_finite_output(values: &[f32]) -> Result<()> {
if let Some(index) = values.iter().position(|v| !v.is_finite()) {
return Err(Error::NonFiniteOutput(index));
}
Ok(())
}
#[cfg(test)]
mod tests;