#![doc(html_root_url = "https://docs.rs/slate-simd")]
#![deny(unsafe_op_in_unsafe_fn)]
#![cfg_attr(feature = "portable_simd", feature(portable_simd))]
#[cfg(target_arch = "x86_64")]
mod avx2;
#[cfg(target_arch = "x86_64")]
mod avx512;
mod dispatch;
#[cfg(target_arch = "aarch64")]
mod neon;
pub mod scalar;
pub use dispatch::{active_tier, detect_tier, Tier};
use slate_core::{Error, Result};
#[inline]
fn check(a: &[f32], b: &[f32]) -> Result<()> {
if a.len() != b.len() {
return Err(Error::DimensionMismatch {
expected: a.len(),
got: b.len(),
});
}
Ok(())
}
#[inline]
pub fn l2_sq(a: &[f32], b: &[f32]) -> Result<f32> {
check(a, b)?;
Ok(dispatch::l2_sq_kernel()(a, b))
}
#[inline]
pub fn inner_product(a: &[f32], b: &[f32]) -> Result<f32> {
check(a, b)?;
Ok(-dispatch::dot_kernel()(a, b))
}
#[inline]
pub fn dot(a: &[f32], b: &[f32]) -> Result<f32> {
check(a, b)?;
Ok(dispatch::dot_kernel()(a, b))
}
#[inline]
pub fn cosine(a: &[f32], b: &[f32]) -> Result<f32> {
check(a, b)?;
let (d, na, nb) = dispatch::cosine_parts(a, b);
let denom = (na * nb).sqrt();
if denom == 0.0 {
Ok(1.0)
} else {
Ok(1.0 - d / denom)
}
}
#[inline]
pub fn cosine_normalized(a: &[f32], b: &[f32]) -> Result<f32> {
check(a, b)?;
Ok(1.0 - dispatch::dot_kernel()(a, b))
}
#[inline]
pub fn distance(metric: slate_core::Metric, a: &[f32], b: &[f32]) -> Result<f32> {
use slate_core::Metric;
match metric {
Metric::L2 => l2_sq(a, b),
Metric::InnerProduct => inner_product(a, b),
Metric::Cosine => cosine(a, b),
}
}
#[inline]
pub fn distance_f16(metric: slate_core::Metric, query: &[f32], stored: &[u8]) -> Result<f32> {
use slate_core::Metric;
if stored.len() != 2 * query.len() {
return Err(Error::DimensionMismatch {
expected: 2 * query.len(),
got: stored.len(),
});
}
match metric {
Metric::L2 => Ok(dispatch::l2_sq_f16(query, stored)),
Metric::InnerProduct => Ok(-dispatch::dot_f16(query, stored)),
Metric::Cosine => {
let (d, nq, ns) = dispatch::cosine_parts_f16(query, stored);
let denom = (nq * ns).sqrt();
if denom == 0.0 {
Ok(1.0)
} else {
Ok(1.0 - d / denom)
}
}
}
}
#[inline]
pub fn distance_i8(
metric: slate_core::Metric,
query: &[f32],
scale: f32,
codes: &[i8],
) -> Result<f32> {
use slate_core::Metric;
if codes.len() != query.len() {
return Err(Error::DimensionMismatch {
expected: query.len(),
got: codes.len(),
});
}
match metric {
Metric::L2 => Ok(dispatch::l2_sq_i8(query, scale, codes)),
Metric::InnerProduct => Ok(-dispatch::dot_i8(query, scale, codes)),
Metric::Cosine => {
let (d, nq, ns) = dispatch::cosine_parts_i8(query, scale, codes);
let denom = (nq * ns).sqrt();
if denom == 0.0 {
Ok(1.0)
} else {
Ok(1.0 - d / denom)
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn dimension_mismatch_is_reported() {
let a = [1.0f32, 2.0, 3.0];
let b = [1.0f32, 2.0];
assert!(matches!(
l2_sq(&a, &b),
Err(Error::DimensionMismatch { expected: 3, got: 2 })
));
}
#[test]
fn l2_sq_known_value() {
let a = [0.0f32, 0.0, 0.0];
let b = [1.0f32, 2.0, 2.0];
assert!((l2_sq(&a, &b).unwrap() - 9.0).abs() < 1e-6);
}
#[test]
fn inner_product_is_negated() {
let a = [1.0f32, 2.0, 3.0];
let b = [1.0f32, 1.0, 1.0];
assert!((inner_product(&a, &b).unwrap() + 6.0).abs() < 1e-6);
assert!((dot(&a, &b).unwrap() - 6.0).abs() < 1e-6);
}
#[test]
fn cosine_identical_is_zero() {
let a = [1.0f32, 2.0, 3.0, 4.0];
assert!(cosine(&a, &a).unwrap().abs() < 1e-6);
}
#[test]
fn cosine_orthogonal_is_one() {
let a = [1.0f32, 0.0];
let b = [0.0f32, 1.0];
assert!((cosine(&a, &b).unwrap() - 1.0).abs() < 1e-6);
}
#[test]
fn cosine_zero_norm_is_one() {
let a = [0.0f32, 0.0, 0.0];
let b = [1.0f32, 2.0, 3.0];
assert!((cosine(&a, &b).unwrap() - 1.0).abs() < 1e-6);
}
#[test]
fn cosine_normalized_matches_cosine_on_unit_vectors() {
let a = [0.6f32, 0.8];
let b = [1.0f32, 0.0];
let raw = cosine(&a, &b).unwrap();
let norm = cosine_normalized(&a, &b).unwrap();
assert!((raw - norm).abs() < 1e-6);
}
#[test]
fn active_tier_is_reported() {
let t = active_tier();
assert_eq!(t, active_tier());
println!("active tier: {}", t.as_str());
}
#[test]
fn distance_f16_equals_decode_then_distance() {
use half::f16;
use slate_core::Metric;
let query = [0.5f32, -1.25, 3.0, 0.0, -2.5, 7.5, -0.125, 4.0, 1.0];
let raw = [0.4f32, -1.0, 3.25, 0.5, -2.0, 7.0, -0.25, 4.5, 0.75];
let stored: Vec<u8> = raw
.iter()
.flat_map(|&x| f16::from_f32(x).to_le_bytes())
.collect();
let decoded: Vec<f32> = raw.iter().map(|&x| f16::from_f32(x).to_f32()).collect();
for metric in [Metric::L2, Metric::InnerProduct, Metric::Cosine] {
let native = distance_f16(metric, &query, &stored).unwrap();
let reference = distance(metric, &query, &decoded).unwrap();
assert!(
(native - reference).abs() <= 1e-6,
"metric={metric:?} native={native} reference={reference}"
);
}
}
#[test]
fn distance_i8_equals_decode_then_distance() {
use slate_core::Metric;
let query = [0.5f32, -1.25, 3.0, 0.0, -2.5, 7.5, -0.125, 4.0, 1.0];
let scale = 0.05f32;
let codes = [10i8, -20, 60, 0, -50, 127, -3, 80, 15];
let decoded: Vec<f32> = codes.iter().map(|&c| f32::from(c) * scale).collect();
for metric in [Metric::L2, Metric::InnerProduct, Metric::Cosine] {
let native = distance_i8(metric, &query, scale, &codes).unwrap();
let reference = distance(metric, &query, &decoded).unwrap();
assert!(
(native - reference).abs() <= 1e-6,
"metric={metric:?} native={native} reference={reference}"
);
}
}
#[test]
fn narrow_distance_rejects_wrong_length() {
use slate_core::Metric;
let query = [1.0f32, 2.0, 3.0];
assert!(matches!(
distance_f16(Metric::L2, &query, &[0u8; 4]),
Err(Error::DimensionMismatch { expected: 6, got: 4 })
));
assert!(matches!(
distance_i8(Metric::L2, &query, 1.0, &[0i8; 2]),
Err(Error::DimensionMismatch { expected: 3, got: 2 })
));
}
}
#[cfg(test)]
mod proptests {
use super::*;
use proptest::prelude::*;
fn approx_eq_scaled(got: f32, want: f32, scale: f32) -> bool {
let tol = 1e-4 * scale.max(1.0);
(got - want).abs() <= tol
}
fn dot_scale(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b).map(|(x, y)| (x * y).abs()).sum()
}
fn l2_scale(a: &[f32], b: &[f32]) -> f32 {
a.iter().zip(b).map(|(x, y)| (x - y) * (x - y)).sum()
}
prop_compose! {
fn vec_pair()(len in 0usize..=257)
(a in prop::collection::vec(-10.0f32..10.0, len),
b in prop::collection::vec(-10.0f32..10.0, len))
-> (Vec<f32>, Vec<f32>) {
(a, b)
}
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(2000))]
#[test]
fn l2_sq_matches_oracle((a, b) in vec_pair()) {
let got = l2_sq(&a, &b).unwrap();
let want = scalar::l2_sq(&a, &b);
let scale = l2_scale(&a, &b);
prop_assert!(approx_eq_scaled(got, want, scale),
"got={got} want={want} scale={scale} len={}", a.len());
}
#[test]
fn dot_matches_oracle((a, b) in vec_pair()) {
let got = dot(&a, &b).unwrap();
let want = scalar::dot(&a, &b);
let scale = dot_scale(&a, &b);
prop_assert!(approx_eq_scaled(got, want, scale),
"got={got} want={want} scale={scale} len={}", a.len());
}
#[test]
fn inner_product_is_negated_dot((a, b) in vec_pair()) {
let got = inner_product(&a, &b).unwrap();
let want = -scalar::dot(&a, &b);
let scale = dot_scale(&a, &b);
prop_assert!(approx_eq_scaled(got, want, scale));
}
#[test]
fn cosine_matches_oracle((a, b) in vec_pair()) {
let got = cosine(&a, &b).unwrap();
let want = scalar::cosine_distance(&a, &b);
prop_assert!((got - want).abs() <= 1e-4,
"got={got} want={want} len={}", a.len());
}
}
use half::f16;
use slate_core::Metric;
fn encode_f16(v: &[f32]) -> Vec<u8> {
let mut out = Vec::with_capacity(2 * v.len());
for &x in v {
out.extend_from_slice(&f16::from_f32(x).to_le_bytes());
}
out
}
fn encode_i8(v: &[f32]) -> (f32, Vec<i8>) {
let max_abs = v.iter().fold(0.0f32, |m, &x| m.max(x.abs()));
let scale = if max_abs == 0.0 { 0.0 } else { max_abs / 127.0 };
let codes = v
.iter()
.map(|&x| {
if scale == 0.0 {
0i8
} else {
(x / scale).round().clamp(-127.0, 127.0) as i8
}
})
.collect();
(scale, codes)
}
fn scalar_distance_f16(metric: Metric, query: &[f32], stored: &[u8]) -> f32 {
match metric {
Metric::L2 => scalar::l2_sq_f16(query, stored),
Metric::InnerProduct => -scalar::dot_f16(query, stored),
Metric::Cosine => {
let (d, nq, ns) = scalar::cosine_parts_f16(query, stored);
let denom = (nq * ns).sqrt();
if denom == 0.0 { 1.0 } else { 1.0 - d / denom }
}
}
}
fn scalar_distance_i8(metric: Metric, query: &[f32], scale: f32, codes: &[i8]) -> f32 {
match metric {
Metric::L2 => scalar::l2_sq_i8(query, scale, codes),
Metric::InnerProduct => -scalar::dot_i8(query, scale, codes),
Metric::Cosine => {
let (d, nq, ns) = scalar::cosine_parts_i8(query, scale, codes);
let denom = (nq * ns).sqrt();
if denom == 0.0 { 1.0 } else { 1.0 - d / denom }
}
}
}
fn dot_scale_decoded(query: &[f32], stored: &[f32]) -> f32 {
query.iter().zip(stored).map(|(x, y)| (x * y).abs()).sum()
}
fn l2_scale_decoded(query: &[f32], stored: &[f32]) -> f32 {
query.iter().zip(stored).map(|(x, y)| (x - y) * (x - y)).sum()
}
proptest! {
#![proptest_config(ProptestConfig::with_cases(2000))]
#[test]
fn l2_sq_f16_matches_oracle((q, v) in vec_pair()) {
let stored = encode_f16(&v);
let decoded: Vec<f32> = v.iter().map(|&x| f16::from_f32(x).to_f32()).collect();
let got = distance_f16(Metric::L2, &q, &stored).unwrap();
let want = scalar_distance_f16(Metric::L2, &q, &stored);
let scale = l2_scale_decoded(&q, &decoded);
prop_assert!(approx_eq_scaled(got, want, scale),
"got={got} want={want} scale={scale} len={}", q.len());
}
#[test]
fn dot_f16_matches_oracle((q, v) in vec_pair()) {
let stored = encode_f16(&v);
let decoded: Vec<f32> = v.iter().map(|&x| f16::from_f32(x).to_f32()).collect();
let got = distance_f16(Metric::InnerProduct, &q, &stored).unwrap();
let want = scalar_distance_f16(Metric::InnerProduct, &q, &stored);
let scale = dot_scale_decoded(&q, &decoded);
prop_assert!(approx_eq_scaled(got, want, scale),
"got={got} want={want} scale={scale} len={}", q.len());
}
#[test]
fn cosine_f16_matches_oracle((q, v) in vec_pair()) {
let stored = encode_f16(&v);
let got = distance_f16(Metric::Cosine, &q, &stored).unwrap();
let want = scalar_distance_f16(Metric::Cosine, &q, &stored);
prop_assert!((got - want).abs() <= 1e-4,
"got={got} want={want} len={}", q.len());
}
#[test]
fn l2_sq_i8_matches_oracle((q, v) in vec_pair()) {
let (scale_q, codes) = encode_i8(&v);
let decoded: Vec<f32> = codes.iter().map(|&c| f32::from(c) * scale_q).collect();
let got = distance_i8(Metric::L2, &q, scale_q, &codes).unwrap();
let want = scalar_distance_i8(Metric::L2, &q, scale_q, &codes);
let scale = l2_scale_decoded(&q, &decoded);
prop_assert!(approx_eq_scaled(got, want, scale),
"got={got} want={want} scale={scale} len={}", q.len());
}
#[test]
fn dot_i8_matches_oracle((q, v) in vec_pair()) {
let (scale_q, codes) = encode_i8(&v);
let decoded: Vec<f32> = codes.iter().map(|&c| f32::from(c) * scale_q).collect();
let got = distance_i8(Metric::InnerProduct, &q, scale_q, &codes).unwrap();
let want = scalar_distance_i8(Metric::InnerProduct, &q, scale_q, &codes);
let scale = dot_scale_decoded(&q, &decoded);
prop_assert!(approx_eq_scaled(got, want, scale),
"got={got} want={want} scale={scale} len={}", q.len());
}
#[test]
fn cosine_i8_matches_oracle((q, v) in vec_pair()) {
let (scale_q, codes) = encode_i8(&v);
let got = distance_i8(Metric::Cosine, &q, scale_q, &codes).unwrap();
let want = scalar_distance_i8(Metric::Cosine, &q, scale_q, &codes);
prop_assert!((got - want).abs() <= 1e-4,
"got={got} want={want} len={}", q.len());
}
}
}