use half::f16;
#[inline]
#[must_use]
pub fn l2_sq(a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(a.len(), b.len());
let mut acc = 0.0f32;
for i in 0..a.len() {
let d = a[i] - b[i];
acc += d * d;
}
acc
}
#[inline]
#[must_use]
pub fn dot(a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(a.len(), b.len());
let mut acc = 0.0f32;
for i in 0..a.len() {
acc += a[i] * b[i];
}
acc
}
#[inline]
#[must_use]
pub fn inner_product_distance(a: &[f32], b: &[f32]) -> f32 {
-dot(a, b)
}
#[inline]
#[must_use]
pub fn cosine_distance(a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(a.len(), b.len());
let mut dot_acc = 0.0f32;
let mut na = 0.0f32;
let mut nb = 0.0f32;
for i in 0..a.len() {
dot_acc += a[i] * b[i];
na += a[i] * a[i];
nb += b[i] * b[i];
}
let denom = (na * nb).sqrt();
if denom == 0.0 {
1.0
} else {
1.0 - dot_acc / denom
}
}
#[inline]
#[must_use]
pub fn cosine_distance_normalized(a: &[f32], b: &[f32]) -> f32 {
1.0 - dot(a, b)
}
#[inline]
#[must_use]
fn decode_f16_chunk(chunk: &[u8]) -> f32 {
let bits = u16::from_le_bytes([chunk[0], chunk[1]]);
f16::from_bits(bits).to_f32()
}
#[inline]
#[must_use]
pub fn l2_sq_f16(query: &[f32], stored: &[u8]) -> f32 {
debug_assert_eq!(stored.len(), 2 * query.len());
let mut acc = 0.0f32;
for (&q, chunk) in query.iter().zip(stored.chunks_exact(2)) {
let d = q - decode_f16_chunk(chunk);
acc += d * d;
}
acc
}
#[inline]
#[must_use]
pub fn dot_f16(query: &[f32], stored: &[u8]) -> f32 {
debug_assert_eq!(stored.len(), 2 * query.len());
let mut acc = 0.0f32;
for (&q, chunk) in query.iter().zip(stored.chunks_exact(2)) {
acc += q * decode_f16_chunk(chunk);
}
acc
}
#[inline]
#[must_use]
pub fn cosine_parts_f16(query: &[f32], stored: &[u8]) -> (f32, f32, f32) {
debug_assert_eq!(stored.len(), 2 * query.len());
let mut dot_acc = 0.0f32;
let mut nq = 0.0f32;
let mut ns = 0.0f32;
for (&q, chunk) in query.iter().zip(stored.chunks_exact(2)) {
let s = decode_f16_chunk(chunk);
dot_acc += q * s;
nq += q * q;
ns += s * s;
}
(dot_acc, nq, ns)
}
#[inline]
#[must_use]
pub fn l2_sq_i8(query: &[f32], scale: f32, codes: &[i8]) -> f32 {
debug_assert_eq!(codes.len(), query.len());
let mut acc = 0.0f32;
for (&q, &c) in query.iter().zip(codes) {
let d = q - f32::from(c) * scale;
acc += d * d;
}
acc
}
#[inline]
#[must_use]
pub fn dot_i8(query: &[f32], scale: f32, codes: &[i8]) -> f32 {
debug_assert_eq!(codes.len(), query.len());
let mut acc = 0.0f32;
for (&q, &c) in query.iter().zip(codes) {
acc += q * (f32::from(c) * scale);
}
acc
}
#[inline]
#[must_use]
pub fn cosine_parts_i8(query: &[f32], scale: f32, codes: &[i8]) -> (f32, f32, f32) {
debug_assert_eq!(codes.len(), query.len());
let mut dot_acc = 0.0f32;
let mut nq = 0.0f32;
let mut ns = 0.0f32;
for (&q, &c) in query.iter().zip(codes) {
let s = f32::from(c) * scale;
dot_acc += q * s;
nq += q * q;
ns += s * s;
}
(dot_acc, nq, ns)
}