coremlit 0.1.2

Safe, synchronous CoreML runtime for macOS (CPU/GPU/Neural Engine) with opt-in on-device multimodal pipelines: speech (Whisper STT, forced alignment, speaker diarization, Silero VAD), AudioSet sound-event tagging, and audio/text/image embeddings (CLAP, granite, SigLIP)
//! The shared 512-dim L2-normalized CLAP [`Embedding`].
//!
//! The type and its numeric contract (f64-accumulated normalization,
//! `is_close` / `is_close_cosine`, no `PartialEq`) mirror textclap's `Embedding`
//! deliberately, so cross-crate oracle comparisons are direct.

use core::{fmt, ops::Deref};

use crate::embeddings::clap::error::{EmbeddingDimMismatch, Error, Result};

/// Dimensionality of a CLAP joint embedding (both towers project to 512).
/// Pinned from the converted graphs' output contract (`tests/clap/model_io.rs` /
/// `tests/clap/text_model_io.rs`).
pub const EMBEDDING_DIM: usize = 512;

/// Norm-tolerance budget for the trusted-path unit-norm check
/// ([`Embedding::try_from_unit_slice`]) — matches textclap's `NORM_BUDGET`
/// (worst case `512 · ulp(1) ≈ 6.1e-5`, rounded to `1e-4`) so the two crates'
/// trust-path guards agree.
pub(crate) const NORM_BUDGET: f32 = 1e-4;

/// A 512-dim L2-normalized CLAP embedding.
///
/// Returned by every `embed*` call. The unit-norm invariant holds within fp32
/// ULP.
///
/// # Compile-fail contracts
///
/// `Embedding` exposes no `DIM` associated const (use the module-level
/// [`EMBEDDING_DIM`]):
///
/// ```compile_fail
/// let _ = coremlit::embeddings::clap::Embedding::DIM;
/// ```
///
/// `Embedding` does not implement `PartialEq` — f32 outputs of an ML model are
/// not bit-stable across runs / threads / OSes; use [`Embedding::is_close`] or
/// [`Embedding::is_close_cosine`]:
///
/// ```compile_fail
/// # let mut s = [0.0_f32; 512]; s[0] = 1.0;
/// # let a = coremlit::embeddings::clap::Embedding::from_slice_normalizing(&s).unwrap();
/// # let b = a.clone();
/// let _ = a == b;
/// ```
#[derive(Clone)]
#[repr(transparent)]
pub struct Embedding {
  inner: [f32; EMBEDDING_DIM],
}

impl Embedding {
  /// Length of the embedding (512).
  #[inline]
  pub const fn dim(&self) -> usize {
    self.inner.len()
  }

  /// Borrow the embedding as a slice.
  #[inline]
  pub const fn as_slice(&self) -> &[f32] {
    self.inner.as_slice()
  }

  /// Owned conversion to a `Vec<f32>`. Allocates.
  #[inline]
  pub fn to_vec(&self) -> Vec<f32> {
    self.inner.to_vec()
  }

  /// Reconstruct from a stored unit vector. Validates length, finiteness, AND
  /// unit-norm (`(norm² − 1).abs() ≤ NORM_BUDGET`), then re-normalizes the
  /// accepted vector through the f64 path, so the stored embedding is unit-norm
  /// to fp32 ULP and [`Self::cosine`] stays in `[−1, 1]` (within fp rounding)
  /// for every successfully constructed pair — a budget-edge vector stored raw
  /// would let `cosine` escape `[−1, 1]`.
  ///
  /// # Errors
  /// [`Error::EmbeddingDimMismatch`] if `s.len() != `[`EMBEDDING_DIM`];
  /// [`Error::NonFiniteEmbedding`] on any non-finite component;
  /// [`Error::EmbeddingNotUnitNorm`] if the norm is outside the budget.
  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));
    }
    // Within budget: rebuild through the f64 normalization path so the STORED
    // vector is unit-norm to fp32 ULP regardless of where in the budget the
    // input fell — a budget-edge vector copied raw makes `cosine` escape
    // [−1, 1]. `EmbeddingZero` is unreachable here (norm² ≥ 1 − NORM_BUDGET > 0)
    // and the dim/finite re-checks inside are already-proven cheap passes.
    Self::from_slice_normalizing(s)
  }

  /// Construct from any non-zero finite slice, re-normalizing to unit length.
  ///
  /// The norm is accumulated in f64 so any finite f32 input normalizes without
  /// intermediate overflow (e.g. `f32::MAX`) or underflow-to-`+Inf` (e.g.
  /// subnormal magnitudes) — matching textclap's `from_slice_normalizing`.
  ///
  /// # Errors
  /// [`Error::EmbeddingDimMismatch`] if `s.len() != `[`EMBEDDING_DIM`];
  /// [`Error::NonFiniteEmbedding`] on any non-finite component;
  /// [`Error::EmbeddingZero`] if the input has zero magnitude.
  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));
      }
    }
    // f64 accumulation: for any finite f32 (|x| ≤ ~3.4e38), x² ≤ ~1.16e77 and
    // 512 terms sum to at most ~5.9e79, well inside f64's ~1.8e308 range.
    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();
    // Multiply per-component in f64 then cast: casting inv_norm to f32 first
    // would overflow to +Inf for subnormal-magnitude inputs.
    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 })
  }

  /// Inner product. For two unit vectors this equals [`Self::cosine`] to fp32
  /// ULP.
  pub fn dot(&self, other: &Embedding) -> f32 {
    self
      .inner
      .iter()
      .zip(other.inner.iter())
      .map(|(a, b)| a * b)
      .sum()
  }

  /// Cosine similarity. For unit vectors, equivalent to [`Self::dot`].
  pub fn cosine(&self, other: &Embedding) -> f32 {
    self.dot(other)
  }

  /// Approximate equality — max-abs metric. `true` iff
  /// `(self − other).max_abs() ≤ tol` (inclusive, so `is_close(self, 0.0)` is
  /// always true).
  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
  }

  /// Approximate equality — semantic (cosine) metric. `true` iff
  /// `1 − cosine(other) ≤ tol`, computed as `0.5·‖a − b‖² ≤ tol` to avoid
  /// catastrophic cancellation near identity (valid because both operands are
  /// unit-norm to fp32 ULP).
  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],
    )
  }
}

/// The windit aggregation seam: exposes the embedding's stored `f32` scalars and
/// rebuilds a unit-norm [`Embedding`] from windit's `f64` compute-domain output,
/// so the long-audio pipeline can aggregate per-window embeddings through
/// windit's engine.
///
/// The inward error map is lossy by design and documented only here: a dimension
/// mismatch keeps its `got`/`expected` payload, while both
/// [`Error::NonFiniteEmbedding`] (dropping its component index) and
/// [`Error::EmbeddingZero`] collapse into [`WinditError::NonFinite`](windit::WinditError::NonFinite)
/// — exactly windit's documented meaning, "no finite unit direction".
impl windit::windowed::Vector for Embedding {
  type Scalar = f32;

  fn as_slice(&self) -> &[f32] {
    // UFCS: the inherent `as_slice` shadows this trait method, so this is a call
    // to the inherent accessor, not unbounded self-recursion.
    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,
      });
    }
    // windit hands back an already-unit-normalized vector, so every component is
    // in [-1, 1] and the f64→f32 narrowing cannot overflow; a pathological
    // non-finite lands in `from_slice_normalizing`'s finite/zero guards.
    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,
    })
  }
}

/// Scans a raw model-output projection — the copied CoreML tensor, before it is
/// normalized into an [`Embedding`] — for the first non-finite (NaN/±∞)
/// component, classifying it as MODEL corruption ([`Error::NonFiniteOutput`]).
/// This is the counterpart to the caller-data corruption
/// ([`Error::NonFiniteEmbedding`]) that [`Embedding::from_slice_normalizing`]
/// raises for a caller's own slice: the audio and text towers run this on the
/// model output *before* normalizing, so a NaN the runtime produced is reported
/// as model-output corruption rather than mislabeled as caller-supplied
/// embedding data. That is the workspace convention (it mirrors speakerkit's
/// identically shaped `check_finite_output`). Extracted so the classification is
/// hermetically testable without a loaded model.
///
/// # Errors
/// [`Error::NonFiniteOutput`] carrying the flat index of the first non-finite
/// component.
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;