pub fn blas_cosine_batch(
query: &[f32],
vectors: &[Vec<f32>],
norms: &[f32],
tombstones: &[u8],
k: usize,
) -> Vec<(usize, f32)> {
let query_norm = rustyhdf5_accel::vector_norm(query);
if query_norm == 0.0 || vectors.is_empty() {
return Vec::new();
}
let dim = query.len();
let n = vectors.len();
let mut active_indices: Vec<usize> = Vec::with_capacity(n);
let mut flat: Vec<f32> = Vec::with_capacity(n * dim);
for i in 0..n {
if i < tombstones.len() && tombstones[i] != 0 {
continue;
}
active_indices.push(i);
flat.extend_from_slice(&vectors[i]);
}
let active_n = active_indices.len();
if active_n == 0 {
return Vec::new();
}
let mut scores = vec![0.0f32; active_n];
unsafe {
matrixmultiply::sgemm(
active_n, dim, 1, 1.0, flat.as_ptr(),
dim as isize, 1, query.as_ptr(),
1, 1, 0.0, scores.as_mut_ptr(),
1, 1, );
}
let mut results: Vec<(usize, f32)> = Vec::with_capacity(active_n);
for (j, &orig_idx) in active_indices.iter().enumerate() {
let vec_norm = norms[orig_idx];
let denom = query_norm * vec_norm;
let score = if denom == 0.0 { 0.0 } else { scores[j] / denom };
results.push((orig_idx, score));
}
results.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
results.truncate(k);
results
}
pub fn blas_cosine_batch_flat(
query: &[f32],
vectors_flat: &[f32],
norms: &[f32],
tombstones: &[u8],
dim: usize,
k: usize,
) -> Vec<(usize, f32)> {
let query_norm = rustyhdf5_accel::vector_norm(query);
if query_norm == 0.0 || vectors_flat.is_empty() {
return Vec::new();
}
let n = vectors_flat.len() / dim;
let all_active = tombstones.iter().all(|&t| t == 0);
if all_active {
let mut scores = vec![0.0f32; n];
unsafe {
matrixmultiply::sgemm(
n,
dim,
1,
1.0,
vectors_flat.as_ptr(),
dim as isize,
1,
query.as_ptr(),
1,
1,
0.0,
scores.as_mut_ptr(),
1,
1,
);
}
let mut results: Vec<(usize, f32)> = scores
.iter()
.enumerate()
.map(|(i, &dot)| {
let denom = query_norm * norms[i];
let score = if denom == 0.0 { 0.0 } else { dot / denom };
(i, score)
})
.collect();
results.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
results.truncate(k);
return results;
}
let mut active_indices: Vec<usize> = Vec::with_capacity(n);
let mut flat: Vec<f32> = Vec::with_capacity(n * dim);
for i in 0..n {
if i < tombstones.len() && tombstones[i] != 0 {
continue;
}
active_indices.push(i);
let offset = i * dim;
flat.extend_from_slice(&vectors_flat[offset..offset + dim]);
}
let active_n = active_indices.len();
if active_n == 0 {
return Vec::new();
}
let mut scores = vec![0.0f32; active_n];
unsafe {
matrixmultiply::sgemm(
active_n,
dim,
1,
1.0,
flat.as_ptr(),
dim as isize,
1,
query.as_ptr(),
1,
1,
0.0,
scores.as_mut_ptr(),
1,
1,
);
}
let mut results: Vec<(usize, f32)> = Vec::with_capacity(active_n);
for (j, &orig_idx) in active_indices.iter().enumerate() {
let denom = query_norm * norms[orig_idx];
let score = if denom == 0.0 { 0.0 } else { scores[j] / denom };
results.push((orig_idx, score));
}
results.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap_or(std::cmp::Ordering::Equal));
results.truncate(k);
results
}
pub fn blas_batch_norms(vectors_flat: &[f32], dim: usize) -> Vec<f32> {
if dim == 0 || vectors_flat.is_empty() {
return Vec::new();
}
let n = vectors_flat.len() / dim;
let mut norms = Vec::with_capacity(n);
for i in 0..n {
let offset = i * dim;
let v = &vectors_flat[offset..offset + dim];
norms.push(rustyhdf5_accel::vector_norm(v));
}
norms
}
pub fn blas_distance_matrix(
queries: &[f32],
vectors: &[f32],
dim: usize,
) -> Vec<f32> {
if dim == 0 || queries.is_empty() || vectors.is_empty() {
return Vec::new();
}
let q = queries.len() / dim;
let n = vectors.len() / dim;
let mut result = vec![0.0f32; q * n];
unsafe {
matrixmultiply::sgemm(
q, dim, n, 1.0, queries.as_ptr(),
dim as isize, 1, vectors.as_ptr(),
1, dim as isize, 0.0, result.as_mut_ptr(),
n as isize, 1, );
}
result
}
#[cfg(test)]
mod tests {
use super::*;
fn make_vectors(n: usize, dim: usize, seed: u32) -> Vec<Vec<f32>> {
let mut s = seed;
let mut next = || -> f32 {
s = s.wrapping_mul(1103515245).wrapping_add(12345);
((s >> 16) as f32) / 65536.0 - 0.5
};
(0..n).map(|_| (0..dim).map(|_| next()).collect()).collect()
}
fn compute_norms(vectors: &[Vec<f32>]) -> Vec<f32> {
vectors
.iter()
.map(|v| rustyhdf5_accel::vector_norm(v))
.collect()
}
#[test]
fn blas_matches_simd_scores() {
let dim = 384;
let n = 500;
let vectors = make_vectors(n, dim, 42);
let norms = compute_norms(&vectors);
let tombstones = vec![0u8; n];
let query = vectors[0].clone();
let blas_results = blas_cosine_batch(&query, &vectors, &norms, &tombstones, n);
let simd_results = crate::vector_search::cosine_similarity_batch_prenorm(
&query, &vectors, &norms, &tombstones,
);
assert_eq!(blas_results.len(), simd_results.len());
for (b, s) in blas_results.iter().zip(&simd_results) {
assert_eq!(b.0, s.0, "index mismatch");
assert!(
(b.1 - s.1).abs() < 1e-4,
"score mismatch at idx {}: blas={} vs simd={}",
b.0,
b.1,
s.1,
);
}
}
#[test]
fn blas_ranking_matches_simd() {
let dim = 128;
let n = 200;
let vectors = make_vectors(n, dim, 77);
let norms = compute_norms(&vectors);
let tombstones = vec![0u8; n];
let query = vectors[5].clone();
let blas_top10 = blas_cosine_batch(&query, &vectors, &norms, &tombstones, 10);
let simd_all = crate::vector_search::cosine_similarity_batch_prenorm(
&query, &vectors, &norms, &tombstones,
);
let simd_top10 = crate::vector_search::top_k(simd_all, 10);
let blas_ids: Vec<usize> = blas_top10.iter().map(|r| r.0).collect();
let simd_ids: Vec<usize> = simd_top10.iter().map(|r| r.0).collect();
assert_eq!(blas_ids, simd_ids, "top-10 ranking should match");
}
#[test]
fn blas_batch_norms_match_individual() {
let dim = 384;
let n = 100;
let vectors = make_vectors(n, dim, 42);
let flat: Vec<f32> = vectors.iter().flat_map(|v| v.iter().copied()).collect();
let batch_norms = blas_batch_norms(&flat, dim);
let individual_norms = compute_norms(&vectors);
assert_eq!(batch_norms.len(), individual_norms.len());
for (b, i) in batch_norms.iter().zip(&individual_norms) {
assert!(
(b - i).abs() < 1e-6,
"norm mismatch: batch={b} vs individual={i}"
);
}
}
#[test]
fn blas_excludes_tombstones() {
let query = vec![1.0, 0.0, 0.0];
let vectors = vec![
vec![1.0, 0.0, 0.0], vec![0.0, 1.0, 0.0], vec![0.5, 0.5, 0.0], ];
let norms = compute_norms(&vectors);
let tombstones = vec![0, 1, 0];
let results = blas_cosine_batch(&query, &vectors, &norms, &tombstones, 10);
assert_eq!(results.len(), 2);
assert!(results.iter().all(|(idx, _)| *idx != 1));
assert_eq!(results[0].0, 0); }
#[test]
fn blas_distance_matrix_shape() {
let dim = 4;
let queries: Vec<f32> = vec![1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0]; let vectors: Vec<f32> = vec![
1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, ];
let result = blas_distance_matrix(&queries, &vectors, dim);
assert_eq!(result.len(), 2 * 3);
assert!((result[0] - 1.0).abs() < 1e-6);
assert!(result[1].abs() < 1e-6);
assert!((result[4] - 1.0).abs() < 1e-6);
}
#[test]
fn blas_empty_vectors() {
let query = vec![1.0, 0.0, 0.0];
let vectors: Vec<Vec<f32>> = Vec::new();
let norms: Vec<f32> = Vec::new();
let tombstones: Vec<u8> = Vec::new();
let results = blas_cosine_batch(&query, &vectors, &norms, &tombstones, 10);
assert!(results.is_empty());
}
#[test]
fn blas_zero_query() {
let query = vec![0.0, 0.0, 0.0];
let vectors = vec![vec![1.0, 0.0, 0.0]];
let norms = compute_norms(&vectors);
let tombstones = vec![0u8];
let results = blas_cosine_batch(&query, &vectors, &norms, &tombstones, 10);
assert!(results.is_empty());
}
#[test]
fn blas_all_tombstoned() {
let query = vec![1.0, 0.0, 0.0];
let vectors = vec![vec![1.0, 0.0, 0.0], vec![0.0, 1.0, 0.0]];
let norms = compute_norms(&vectors);
let tombstones = vec![1, 1];
let results = blas_cosine_batch(&query, &vectors, &norms, &tombstones, 10);
assert!(results.is_empty());
}
#[test]
fn blas_identical_vector_score_one() {
let query = vec![1.0, 2.0, 3.0, 4.0];
let vectors = vec![query.clone()];
let norms = compute_norms(&vectors);
let tombstones = vec![0u8];
let results = blas_cosine_batch(&query, &vectors, &norms, &tombstones, 1);
assert_eq!(results.len(), 1);
assert!(
(results[0].1 - 1.0).abs() < 1e-5,
"expected ~1.0, got {}",
results[0].1
);
}
#[test]
fn blas_orthogonal_score_zero() {
let query = vec![1.0, 0.0, 0.0];
let vectors = vec![vec![0.0, 1.0, 0.0]];
let norms = compute_norms(&vectors);
let tombstones = vec![0u8];
let results = blas_cosine_batch(&query, &vectors, &norms, &tombstones, 1);
assert_eq!(results.len(), 1);
assert!(
results[0].1.abs() < 1e-5,
"expected ~0.0, got {}",
results[0].1
);
}
#[test]
fn blas_top_k_truncation() {
let dim = 32;
let n = 100;
let vectors = make_vectors(n, dim, 42);
let norms = compute_norms(&vectors);
let tombstones = vec![0u8; n];
let query = vectors[0].clone();
let results = blas_cosine_batch(&query, &vectors, &norms, &tombstones, 5);
assert_eq!(results.len(), 5);
for w in results.windows(2) {
assert!(w[0].1 >= w[1].1);
}
}
#[test]
fn blas_flat_matches_vec_variant() {
let dim = 64;
let n = 200;
let vectors = make_vectors(n, dim, 42);
let norms = compute_norms(&vectors);
let tombstones = vec![0u8; n];
let query = vectors[3].clone();
let flat: Vec<f32> = vectors.iter().flat_map(|v| v.iter().copied()).collect();
let vec_results = blas_cosine_batch(&query, &vectors, &norms, &tombstones, 10);
let flat_results =
blas_cosine_batch_flat(&query, &flat, &norms, &tombstones, dim, 10);
assert_eq!(vec_results.len(), flat_results.len());
for (v, f) in vec_results.iter().zip(&flat_results) {
assert_eq!(v.0, f.0);
assert!((v.1 - f.1).abs() < 1e-5);
}
}
#[test]
fn blas_flat_with_tombstones() {
let dim = 3;
let vectors = vec![
vec![1.0, 0.0, 0.0],
vec![0.0, 1.0, 0.0],
vec![0.5, 0.5, 0.0],
];
let flat: Vec<f32> = vectors.iter().flat_map(|v| v.iter().copied()).collect();
let norms = compute_norms(&vectors);
let tombstones = vec![0, 1, 0]; let query = vec![1.0, 0.0, 0.0];
let results = blas_cosine_batch_flat(&query, &flat, &norms, &tombstones, dim, 10);
assert_eq!(results.len(), 2);
assert!(results.iter().all(|(idx, _)| *idx != 1));
}
#[test]
fn blas_distance_matrix_empty() {
let result = blas_distance_matrix(&[], &[1.0, 0.0], 2);
assert!(result.is_empty());
let result2 = blas_distance_matrix(&[1.0, 0.0], &[], 2);
assert!(result2.is_empty());
}
#[test]
fn blas_batch_norms_empty() {
let norms = blas_batch_norms(&[], 4);
assert!(norms.is_empty());
}
#[test]
fn blas_large_scale_matches_simd() {
let dim = 384;
let n = 1000;
let vectors = make_vectors(n, dim, 42);
let norms = compute_norms(&vectors);
let mut tombstones = vec![0u8; n];
for i in (0..n).step_by(7) {
tombstones[i] = 1;
}
let query = vectors[1].clone();
let blas_top20 = blas_cosine_batch(&query, &vectors, &norms, &tombstones, 20);
let simd_all = crate::vector_search::cosine_similarity_batch_prenorm(
&query, &vectors, &norms, &tombstones,
);
let simd_top20 = crate::vector_search::top_k(simd_all, 20);
assert_eq!(blas_top20.len(), simd_top20.len());
for (b, s) in blas_top20.iter().zip(&simd_top20) {
assert_eq!(b.0, s.0, "index mismatch in top-20");
assert!(
(b.1 - s.1).abs() < 1e-4,
"score mismatch: blas={} vs simd={}",
b.1,
s.1,
);
}
}
#[test]
fn blas_negative_correlation() {
let query = vec![1.0, 0.0];
let vectors = vec![vec![-1.0, 0.0]];
let norms = compute_norms(&vectors);
let tombstones = vec![0u8];
let results = blas_cosine_batch(&query, &vectors, &norms, &tombstones, 1);
assert_eq!(results.len(), 1);
assert!(
(results[0].1 - (-1.0)).abs() < 1e-5,
"expected ~-1.0, got {}",
results[0].1
);
}
#[test]
fn blas_performance_10k() {
let dim = 384;
let n = 10_000;
let vectors = make_vectors(n, dim, 42);
let norms = compute_norms(&vectors);
let tombstones = vec![0u8; n];
let query = vectors[0].clone();
let start = std::time::Instant::now();
let results = blas_cosine_batch(&query, &vectors, &norms, &tombstones, 10);
let elapsed = start.elapsed();
assert_eq!(results.len(), 10);
assert!(
elapsed.as_millis() < 500,
"BLAS 10K took {}ms, expected < 500ms",
elapsed.as_millis()
);
}
}