This is an automated email from the ASF dual-hosted git repository.
JingsongLi pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/paimon-vector-index.git
The following commit(s) were added to refs/heads/main by this push:
new b974335 Optimize SIMD hot paths in vector indexes (#43)
b974335 is described below
commit b974335df10db1b49d2175e727b8762c6a1e51b1
Author: Jingsong Lee <[email protected]>
AuthorDate: Fri Jun 12 19:15:34 2026 +0800
Optimize SIMD hot paths in vector indexes (#43)
---
core/src/distance.rs | 324 +++++++++++++++++++++++--
core/src/hnsw.rs | 221 +++++++++++++-----
core/src/hnsw_search.rs | 12 +-
core/src/ivfflat.rs | 15 +-
core/src/ivfhnswflat.rs | 5 +-
core/src/ivfhnswflat_io.rs | 17 +-
core/src/ivfhnswsq.rs | 27 +--
core/src/ivfhnswsq_io.rs | 34 +--
core/src/ivfpq.rs | 71 +++---
core/src/kmeans.rs | 37 +++
core/src/opq.rs | 42 +++-
core/src/pq.rs | 151 +++++++-----
core/src/sq.rs | 572 +++++++++++++++++++++++++++++++++++++++++++--
13 files changed, 1272 insertions(+), 256 deletions(-)
diff --git a/core/src/distance.rs b/core/src/distance.rs
index e14381c..0fe983b 100644
--- a/core/src/distance.rs
+++ b/core/src/distance.rs
@@ -15,6 +15,8 @@
// specific language governing permissions and limitations
// under the License.
+use crate::blas::sgemm_a_bt;
+
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
#[repr(u32)]
pub enum MetricType {
@@ -154,17 +156,41 @@ unsafe fn fvec_l2sqr_neon(a: &[f32], b: &[f32]) -> f32 {
/// Squared L2 distance on sub-vectors.
pub fn fvec_l2sqr_sub(a: &[f32], a_off: usize, b: &[f32], b_off: usize, len:
usize) -> f32 {
- let mut sum = 0.0f32;
- for i in 0..len {
- let d = a[a_off + i] - b[b_off + i];
- sum += d * d;
- }
- sum
+ fvec_l2sqr(&a[a_off..a_off + len], &b[b_off..b_off + len])
}
/// Inner product of two vectors.
+#[inline]
pub fn fvec_inner_product(a: &[f32], b: &[f32]) -> f32 {
debug_assert_eq!(a.len(), b.len());
+ fvec_inner_product_simd(a, b)
+}
+
+#[cfg(target_arch = "x86_64")]
+#[inline]
+fn fvec_inner_product_simd(a: &[f32], b: &[f32]) -> f32 {
+ if is_x86_feature_detected!("avx2") && a.len() >= 8 {
+ unsafe { fvec_inner_product_avx2(a, b) }
+ } else {
+ fvec_inner_product_scalar(a, b)
+ }
+}
+
+#[cfg(target_arch = "aarch64")]
+#[inline]
+fn fvec_inner_product_simd(a: &[f32], b: &[f32]) -> f32 {
+ unsafe { fvec_inner_product_neon(a, b) }
+}
+
+#[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
+#[inline]
+fn fvec_inner_product_simd(a: &[f32], b: &[f32]) -> f32 {
+ fvec_inner_product_scalar(a, b)
+}
+
+#[inline]
+#[cfg(not(target_arch = "aarch64"))]
+fn fvec_inner_product_scalar(a: &[f32], b: &[f32]) -> f32 {
let mut dot = 0.0f32;
for i in 0..a.len() {
dot += a[i] * b[i];
@@ -172,8 +198,95 @@ pub fn fvec_inner_product(a: &[f32], b: &[f32]) -> f32 {
dot
}
+#[cfg(target_arch = "x86_64")]
+#[target_feature(enable = "avx2")]
+unsafe fn fvec_inner_product_avx2(a: &[f32], b: &[f32]) -> f32 {
+ use std::arch::x86_64::*;
+
+ let n = a.len();
+ let mut sum = _mm256_setzero_ps();
+ let mut i = 0;
+ while i + 8 <= n {
+ let va = unsafe { _mm256_loadu_ps(a.as_ptr().add(i)) };
+ let vb = unsafe { _mm256_loadu_ps(b.as_ptr().add(i)) };
+ sum = _mm256_add_ps(sum, _mm256_mul_ps(va, vb));
+ i += 8;
+ }
+
+ let hi = _mm256_extractf128_ps::<1>(sum);
+ let lo = _mm256_castps256_ps128(sum);
+ let sum128 = _mm_add_ps(lo, hi);
+ let sum64 = _mm_add_ps(sum128, _mm_movehl_ps(sum128, sum128));
+ let sum32 = _mm_add_ss(sum64, _mm_shuffle_ps::<1>(sum64, sum64));
+ let mut result = _mm_cvtss_f32(sum32);
+
+ while i < n {
+ result += unsafe { *a.get_unchecked(i) * *b.get_unchecked(i) };
+ i += 1;
+ }
+ result
+}
+
+#[cfg(target_arch = "aarch64")]
+#[target_feature(enable = "neon")]
+unsafe fn fvec_inner_product_neon(a: &[f32], b: &[f32]) -> f32 {
+ use std::arch::aarch64::*;
+
+ let n = a.len();
+ let mut sum0 = vdupq_n_f32(0.0);
+ let mut sum1 = vdupq_n_f32(0.0);
+ let mut i = 0;
+ while i + 8 <= n {
+ let va0 = unsafe { vld1q_f32(a.as_ptr().add(i)) };
+ let vb0 = unsafe { vld1q_f32(b.as_ptr().add(i)) };
+ sum0 = vmlaq_f32(sum0, va0, vb0);
+
+ let va1 = unsafe { vld1q_f32(a.as_ptr().add(i + 4)) };
+ let vb1 = unsafe { vld1q_f32(b.as_ptr().add(i + 4)) };
+ sum1 = vmlaq_f32(sum1, va1, vb1);
+
+ i += 8;
+ }
+
+ let mut result = vaddvq_f32(vaddq_f32(sum0, sum1));
+ while i < n {
+ result += unsafe { *a.get_unchecked(i) * *b.get_unchecked(i) };
+ i += 1;
+ }
+ result
+}
+
/// Squared L2 norm of a vector.
+#[inline]
pub fn fvec_norm_l2sqr(a: &[f32]) -> f32 {
+ fvec_norm_l2sqr_simd(a)
+}
+
+#[cfg(target_arch = "x86_64")]
+#[inline]
+fn fvec_norm_l2sqr_simd(a: &[f32]) -> f32 {
+ if is_x86_feature_detected!("avx2") && a.len() >= 8 {
+ unsafe { fvec_norm_l2sqr_avx2(a) }
+ } else {
+ fvec_norm_l2sqr_scalar(a)
+ }
+}
+
+#[cfg(target_arch = "aarch64")]
+#[inline]
+fn fvec_norm_l2sqr_simd(a: &[f32]) -> f32 {
+ unsafe { fvec_norm_l2sqr_neon(a) }
+}
+
+#[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
+#[inline]
+fn fvec_norm_l2sqr_simd(a: &[f32]) -> f32 {
+ fvec_norm_l2sqr_scalar(a)
+}
+
+#[inline]
+#[cfg(not(target_arch = "aarch64"))]
+fn fvec_norm_l2sqr_scalar(a: &[f32]) -> f32 {
let mut sum = 0.0f32;
for &v in a {
sum += v * v;
@@ -181,6 +294,63 @@ pub fn fvec_norm_l2sqr(a: &[f32]) -> f32 {
sum
}
+#[cfg(target_arch = "x86_64")]
+#[target_feature(enable = "avx2")]
+unsafe fn fvec_norm_l2sqr_avx2(a: &[f32]) -> f32 {
+ use std::arch::x86_64::*;
+
+ let n = a.len();
+ let mut sum = _mm256_setzero_ps();
+ let mut i = 0;
+ while i + 8 <= n {
+ let va = unsafe { _mm256_loadu_ps(a.as_ptr().add(i)) };
+ sum = _mm256_add_ps(sum, _mm256_mul_ps(va, va));
+ i += 8;
+ }
+
+ let hi = _mm256_extractf128_ps::<1>(sum);
+ let lo = _mm256_castps256_ps128(sum);
+ let sum128 = _mm_add_ps(lo, hi);
+ let sum64 = _mm_add_ps(sum128, _mm_movehl_ps(sum128, sum128));
+ let sum32 = _mm_add_ss(sum64, _mm_shuffle_ps::<1>(sum64, sum64));
+ let mut result = _mm_cvtss_f32(sum32);
+
+ while i < n {
+ let v = unsafe { *a.get_unchecked(i) };
+ result += v * v;
+ i += 1;
+ }
+ result
+}
+
+#[cfg(target_arch = "aarch64")]
+#[target_feature(enable = "neon")]
+unsafe fn fvec_norm_l2sqr_neon(a: &[f32]) -> f32 {
+ use std::arch::aarch64::*;
+
+ let n = a.len();
+ let mut sum0 = vdupq_n_f32(0.0);
+ let mut sum1 = vdupq_n_f32(0.0);
+ let mut i = 0;
+ while i + 8 <= n {
+ let va0 = unsafe { vld1q_f32(a.as_ptr().add(i)) };
+ sum0 = vmlaq_f32(sum0, va0, va0);
+
+ let va1 = unsafe { vld1q_f32(a.as_ptr().add(i + 4)) };
+ sum1 = vmlaq_f32(sum1, va1, va1);
+
+ i += 8;
+ }
+
+ let mut result = vaddvq_f32(vaddq_f32(sum0, sum1));
+ while i < n {
+ let v = unsafe { *a.get_unchecked(i) };
+ result += v * v;
+ i += 1;
+ }
+ result
+}
+
/// Normalize a vector in-place to unit length. Returns the original norm.
pub fn fvec_normalize(v: &mut [f32]) -> f32 {
let norm = fvec_norm_l2sqr(v).sqrt();
@@ -199,19 +369,72 @@ pub fn fvec_distance(query: &[f32], vector: &[f32],
metric: MetricType) -> f32 {
MetricType::L2 => fvec_l2sqr(query, vector),
MetricType::InnerProduct => -fvec_inner_product(query, vector),
MetricType::Cosine => {
- let dot = fvec_inner_product(query, vector);
let nq = fvec_norm_l2sqr(query).sqrt();
let nv = fvec_norm_l2sqr(vector).sqrt();
- let denom = nq * nv;
- if denom > 0.0 {
- 1.0 - dot / denom
- } else {
- 1.0
+ fvec_cosine_distance_with_norms(query, vector, nq, nv)
+ }
+ }
+}
+
+pub(crate) fn fvec_distance_with_norms(
+ a: &[f32],
+ b: &[f32],
+ metric: MetricType,
+ a_norm: f32,
+ b_norm: f32,
+) -> f32 {
+ match metric {
+ MetricType::L2 => fvec_l2sqr(a, b),
+ MetricType::InnerProduct => -fvec_inner_product(a, b),
+ MetricType::Cosine => fvec_cosine_distance_with_norms(a, b, a_norm,
b_norm),
+ }
+}
+
+#[derive(Debug, Clone, Copy)]
+pub(crate) struct QueryDistance<'a> {
+ query: &'a [f32],
+ metric: MetricType,
+ query_norm: f32,
+}
+
+impl<'a> QueryDistance<'a> {
+ #[inline]
+ pub(crate) fn new(query: &'a [f32], metric: MetricType) -> Self {
+ let query_norm = if metric == MetricType::Cosine {
+ fvec_norm_l2sqr(query).sqrt()
+ } else {
+ 0.0
+ };
+ Self {
+ query,
+ metric,
+ query_norm,
+ }
+ }
+
+ #[inline]
+ pub(crate) fn distance_to(&self, vector: &[f32], vector_norm: Option<f32>)
-> f32 {
+ match self.metric {
+ MetricType::L2 => fvec_l2sqr(self.query, vector),
+ MetricType::InnerProduct => -fvec_inner_product(self.query,
vector),
+ MetricType::Cosine => {
+ let vector_norm = vector_norm.unwrap_or_else(||
fvec_norm_l2sqr(vector).sqrt());
+ fvec_cosine_distance_with_norms(self.query, vector,
self.query_norm, vector_norm)
}
}
}
}
+#[inline]
+fn fvec_cosine_distance_with_norms(a: &[f32], b: &[f32], a_norm: f32, b_norm:
f32) -> f32 {
+ let denom = a_norm * b_norm;
+ if denom > 0.0 {
+ 1.0 - fvec_inner_product(a, b) / denom
+ } else {
+ 1.0
+ }
+}
+
pub fn preprocess_vectors(data: &[f32], n: usize, d: usize, metric:
MetricType) -> Vec<f32> {
let mut processed = data[..n * d].to_vec();
if metric == MetricType::Cosine {
@@ -324,8 +547,23 @@ pub fn fvec_l2sqr_batch(
ksub: usize,
result: &mut [f32],
) {
- for j in 0..ksub {
- result[j] = fvec_l2sqr_sub(query_sub, 0, centroids, j * dsub, dsub);
+ debug_assert!(query_sub.len() >= dsub);
+ debug_assert!(centroids.len() >= ksub * dsub);
+ debug_assert!(result.len() >= ksub);
+
+ if dsub >= 4 && ksub >= 8 {
+ fvec_ip_batch(query_sub, centroids, dsub, ksub, result);
+ let q_norm = fvec_norm_l2sqr(&query_sub[..dsub]);
+ for j in 0..ksub {
+ let c_off = j * dsub;
+ let c_norm = fvec_norm_l2sqr(¢roids[c_off..c_off + dsub]);
+ result[j] = (q_norm + c_norm - 2.0 * result[j]).max(0.0);
+ }
+ } else {
+ for j in 0..ksub {
+ let c_off = j * dsub;
+ result[j] = fvec_l2sqr(&query_sub[..dsub], ¢roids[c_off..c_off
+ dsub]);
+ }
}
}
@@ -337,12 +575,26 @@ pub fn fvec_ip_batch(
ksub: usize,
result: &mut [f32],
) {
- for j in 0..ksub {
- let mut dot = 0.0f32;
- for d in 0..dsub {
- dot += query_sub[d] * centroids[j * dsub + d];
+ debug_assert!(query_sub.len() >= dsub);
+ debug_assert!(centroids.len() >= ksub * dsub);
+ debug_assert!(result.len() >= ksub);
+
+ if dsub >= 4 && ksub >= 8 {
+ sgemm_a_bt(
+ 1,
+ ksub,
+ dsub,
+ 1.0,
+ &query_sub[..dsub],
+ ¢roids[..ksub * dsub],
+ 0.0,
+ &mut result[..ksub],
+ );
+ } else {
+ for j in 0..ksub {
+ let c_off = j * dsub;
+ result[j] = fvec_inner_product(&query_sub[..dsub],
¢roids[c_off..c_off + dsub]);
}
- result[j] = dot;
}
}
@@ -568,9 +820,10 @@ pub fn fvec_distances_batch(
metric: MetricType,
distances: &mut [f32],
) {
+ let distance = QueryDistance::new(query, metric);
for i in 0..n {
let vec = &vectors[i * d..(i + 1) * d];
- distances[i] = fvec_distance(query, vec, metric);
+ distances[i] = distance.distance_to(vec, None);
}
}
@@ -600,6 +853,37 @@ mod tests {
assert!((fvec_inner_product(&a, &b) - 32.0).abs() < 1e-6);
}
+ #[test]
+ fn test_inner_product_and_norm_large_vector() {
+ let a: Vec<f32> = (0..37).map(|i| i as f32 * 0.25 - 3.0).collect();
+ let b: Vec<f32> = (0..37).map(|i| 2.0 - i as f32 * 0.125).collect();
+
+ let expected_dot: f32 = a.iter().zip(&b).map(|(&x, &y)| x * y).sum();
+ let expected_norm: f32 = a.iter().map(|&x| x * x).sum();
+
+ assert!((fvec_inner_product(&a, &b) - expected_dot).abs() < 1e-4);
+ assert!((fvec_norm_l2sqr(&a) - expected_norm).abs() < 1e-4);
+ }
+
+ #[test]
+ fn test_batch_distance_helpers_match_scalar() {
+ let dsub = 5;
+ let ksub = 9;
+ let query: Vec<f32> = (0..dsub).map(|i| i as f32 * 0.3 -
0.7).collect();
+ let centroids: Vec<f32> = (0..ksub * dsub).map(|i| i as f32 * 0.07 -
1.2).collect();
+
+ let mut l2 = vec![0.0f32; ksub];
+ let mut ip = vec![0.0f32; ksub];
+ fvec_l2sqr_batch(&query, ¢roids, dsub, ksub, &mut l2);
+ fvec_ip_batch(&query, ¢roids, dsub, ksub, &mut ip);
+
+ for j in 0..ksub {
+ let c = ¢roids[j * dsub..(j + 1) * dsub];
+ assert!((l2[j] - fvec_l2sqr(&query, c)).abs() < 1e-5);
+ assert!((ip[j] - fvec_inner_product(&query, c)).abs() < 1e-5);
+ }
+ }
+
#[test]
fn test_fvec_distance_by_metric() {
let a = [1.0, 0.0];
diff --git a/core/src/hnsw.rs b/core/src/hnsw.rs
index 37374db..0d81fc8 100644
--- a/core/src/hnsw.rs
+++ b/core/src/hnsw.rs
@@ -15,7 +15,9 @@
// specific language governing permissions and limitations
// under the License.
-use crate::distance::{fvec_distance, MetricType};
+use crate::distance::{
+ fvec_distance, fvec_distance_with_norms, fvec_norm_l2sqr, MetricType,
QueryDistance,
+};
use rayon::prelude::*;
use std::cmp::Reverse;
use std::collections::BinaryHeap;
@@ -59,6 +61,7 @@ pub struct HnswGraph {
d: usize,
metric: MetricType,
vectors: Vec<f32>,
+ vector_norms: Option<Vec<f32>>,
levels: Vec<usize>,
neighbors: Vec<Vec<Vec<usize>>>,
entry_point: usize,
@@ -117,10 +120,12 @@ impl HnswGraph {
return Ok(Self::build_parallel(vectors, n, d, metric, params));
}
+ let vector_norms = vector_norms_for(metric, &vectors, n, d);
let mut graph = HnswGraph {
d,
metric,
vectors,
+ vector_norms,
levels: Vec::with_capacity(n),
neighbors: Vec::with_capacity(n),
entry_point: 0,
@@ -142,6 +147,7 @@ impl HnswGraph {
metric: MetricType,
params: HnswBuildParams,
) -> Self {
+ let vector_norms = vector_norms_for(metric, &vectors, n, d);
let levels = parallel_build_levels(n, params);
let max_observed_level = levels.iter().copied().max().unwrap_or(0);
let nodes = levels
@@ -154,6 +160,7 @@ impl HnswGraph {
d,
metric,
vectors: &vectors,
+ vector_norms: vector_norms.as_deref(),
levels: &levels,
nodes: &nodes,
params,
@@ -182,6 +189,7 @@ impl HnswGraph {
d,
metric,
vectors,
+ vector_norms,
levels,
neighbors,
entry_point: 0,
@@ -221,11 +229,13 @@ impl HnswGraph {
"graph level metadata does not match vector count",
));
}
+ let vector_norms = vector_norms_for(metric, &vectors, n, d);
if n == 0 {
return Ok(Self {
d,
metric,
vectors,
+ vector_norms,
levels,
neighbors,
entry_point: 0,
@@ -280,6 +290,7 @@ impl HnswGraph {
d,
metric,
vectors,
+ vector_norms,
levels,
neighbors,
entry_point,
@@ -289,42 +300,45 @@ impl HnswGraph {
}
pub fn search(&self, query: &[f32], k: usize, ef: usize) -> Vec<(usize,
f32)> {
- let mut visited = Vec::new();
- let mut visit_mark = 1usize;
- self.search_with_workspace(query, k, ef, &mut visited, &mut visit_mark)
+ let mut workspace = HnswSearchWorkspace::new(ef.max(k));
+ self.search_with_reusable_workspace(query, k, ef, &mut workspace)
+ .to_vec()
}
- pub(crate) fn search_with_workspace(
+ pub(crate) fn search_with_reusable_workspace<'a>(
&self,
query: &[f32],
k: usize,
ef: usize,
- visited: &mut Vec<usize>,
- visit_mark: &mut usize,
- ) -> Vec<(usize, f32)> {
+ workspace: &'a mut HnswSearchWorkspace,
+ ) -> &'a [(usize, f32)] {
+ workspace.output_pairs.clear();
if self.levels.is_empty() || k == 0 {
- return Vec::new();
- }
- if visited.len() < self.levels.len() {
- visited.resize(self.levels.len(), 0);
+ return &workspace.output_pairs;
}
+ let ef = ef.max(k);
+ workspace.prepare(self.levels.len(), ef);
+ let query_distance = QueryDistance::new(query, self.metric);
let mut ep = self.entry_point;
- let mut ep_dist = self.distance_to_query(query, ep);
+ let mut ep_dist = self.distance_to_query(&query_distance, ep);
for level in (1..=self.max_observed_level).rev() {
- let (next, dist) = self.greedy_search_query(query, ep, ep_dist,
level);
+ let (next, dist) = self.greedy_search_query(&query_distance, ep,
ep_dist, level);
ep = next;
ep_dist = dist;
}
- let current_mark = *visit_mark;
- let candidates = self.search_layer_query(query, ep, ef.max(k), 0,
visited, current_mark);
- *visit_mark = advance_visit_mark(visited, current_mark);
- candidates
- .into_iter()
- .take(k)
- .map(|n| (n.id, n.dist))
- .collect()
+ let current_mark = workspace.visit_mark;
+ self.search_layer_query_into(&query_distance, ep, ef, 0, current_mark,
workspace);
+ workspace.visit_mark = advance_visit_mark(&mut workspace.visited,
current_mark);
+ workspace.output_pairs.extend(
+ workspace
+ .output
+ .iter()
+ .take(k)
+ .map(|node| (node.id, node.dist)),
+ );
+ &workspace.output_pairs
}
pub fn len(&self) -> usize {
@@ -479,7 +493,7 @@ impl HnswGraph {
fn greedy_search_query(
&self,
- query: &[f32],
+ distance: &QueryDistance<'_>,
mut current: usize,
mut current_dist: f32,
level: usize,
@@ -488,7 +502,7 @@ impl HnswGraph {
let mut best = current;
let mut best_dist = current_dist;
for &neighbor in self.neighbors_at(current, level) {
- let dist = self.distance_to_query(query, neighbor);
+ let dist = self.distance_to_query(distance, neighbor);
if dist < best_dist {
best = neighbor;
best_dist = dist;
@@ -527,18 +541,26 @@ impl HnswGraph {
}
}
- fn search_layer_query(
+ fn search_layer_query_into(
&self,
- query: &[f32],
+ distance: &QueryDistance<'_>,
entry: usize,
ef: usize,
level: usize,
- visited: &mut [usize],
visit_mark: usize,
- ) -> Vec<ScoredNode> {
- self.search_layer(entry, ef, level, visited, visit_mark, |id| {
- self.distance_to_query(query, id)
- })
+ workspace: &mut HnswSearchWorkspace,
+ ) {
+ self.search_layer_into(
+ entry,
+ ef,
+ level,
+ &mut workspace.visited,
+ visit_mark,
+ &mut workspace.candidates,
+ &mut workspace.results,
+ &mut workspace.output,
+ |id| self.distance_to_query(distance, id),
+ );
}
fn search_layer_node_with_workspace(
@@ -563,32 +585,6 @@ impl HnswGraph {
workspace.visit_mark = advance_visit_mark(&mut workspace.visited,
visit_mark);
}
- fn search_layer(
- &self,
- entry: usize,
- ef: usize,
- level: usize,
- visited: &mut [usize],
- visit_mark: usize,
- distance: impl FnMut(usize) -> f32,
- ) -> Vec<ScoredNode> {
- let mut candidates = BinaryHeap::with_capacity(ef);
- let mut results = BinaryHeap::with_capacity(ef);
- let mut output = Vec::with_capacity(ef);
- self.search_layer_into(
- entry,
- ef,
- level,
- visited,
- visit_mark,
- &mut candidates,
- &mut results,
- &mut output,
- distance,
- );
- output
- }
-
#[allow(clippy::too_many_arguments)]
fn search_layer_into(
&self,
@@ -674,12 +670,30 @@ impl HnswGraph {
fn distance_between(&self, a: usize, b: usize) -> f32 {
let va = &self.vectors[a * self.d..(a + 1) * self.d];
let vb = &self.vectors[b * self.d..(b + 1) * self.d];
- fvec_distance(va, vb, self.metric)
+ match self.metric {
+ MetricType::Cosine => fvec_distance_with_norms(
+ va,
+ vb,
+ self.metric,
+ self.vector_norm(a),
+ self.vector_norm(b),
+ ),
+ _ => fvec_distance(va, vb, self.metric),
+ }
}
- fn distance_to_query(&self, query: &[f32], id: usize) -> f32 {
+ fn distance_to_query(&self, query_distance: &QueryDistance<'_>, id: usize)
-> f32 {
let vector = &self.vectors[id * self.d..(id + 1) * self.d];
- fvec_distance(query, vector, self.metric)
+ query_distance.distance_to(vector,
self.vector_norms.as_ref().map(|norms| norms[id]))
+ }
+
+ fn vector_norm(&self, id: usize) -> f32 {
+ self.vector_norms
+ .as_ref()
+ .map(|norms| norms[id])
+ .unwrap_or_else(|| {
+ fvec_norm_l2sqr(&self.vectors[id * self.d..(id + 1) *
self.d]).sqrt()
+ })
}
}
@@ -709,6 +723,42 @@ impl Ord for HeapNode {
}
}
+pub(crate) struct HnswSearchWorkspace {
+ visited: Vec<usize>,
+ visit_mark: usize,
+ candidates: BinaryHeap<Reverse<HeapNode>>,
+ results: BinaryHeap<HeapNode>,
+ output: Vec<ScoredNode>,
+ output_pairs: Vec<(usize, f32)>,
+}
+
+impl HnswSearchWorkspace {
+ pub(crate) fn new(ef: usize) -> Self {
+ Self {
+ visited: Vec::new(),
+ visit_mark: 1,
+ candidates: BinaryHeap::with_capacity(ef),
+ results: BinaryHeap::with_capacity(ef),
+ output: Vec::with_capacity(ef),
+ output_pairs: Vec::with_capacity(ef),
+ }
+ }
+
+ fn prepare(&mut self, graph_len: usize, ef: usize) {
+ if self.visited.len() < graph_len {
+ self.visited.resize(graph_len, 0);
+ }
+ self.candidates
+ .reserve(ef.saturating_sub(self.candidates.capacity()));
+ self.results
+ .reserve(ef.saturating_sub(self.results.capacity()));
+ self.output
+ .reserve(ef.saturating_sub(self.output.capacity()));
+ self.output_pairs
+ .reserve(ef.saturating_sub(self.output_pairs.capacity()));
+ }
+}
+
struct HnswBuildWorkspace {
visited: Vec<usize>,
visit_mark: usize,
@@ -749,6 +799,7 @@ struct ParallelHnswBuilder<'a> {
d: usize,
metric: MetricType,
vectors: &'a [f32],
+ vector_norms: Option<&'a [f32]>,
levels: &'a [usize],
nodes: &'a [RwLock<ParallelBuildNode>],
params: HnswBuildParams,
@@ -974,7 +1025,22 @@ impl ParallelHnswBuilder<'_> {
fn distance_between(&self, a: usize, b: usize) -> f32 {
let va = &self.vectors[a * self.d..(a + 1) * self.d];
let vb = &self.vectors[b * self.d..(b + 1) * self.d];
- fvec_distance(va, vb, self.metric)
+ match self.metric {
+ MetricType::Cosine => fvec_distance_with_norms(
+ va,
+ vb,
+ self.metric,
+ self.vector_norm(a),
+ self.vector_norm(b),
+ ),
+ _ => fvec_distance(va, vb, self.metric),
+ }
+ }
+
+ fn vector_norm(&self, id: usize) -> f32 {
+ self.vector_norms.map(|norms| norms[id]).unwrap_or_else(|| {
+ fvec_norm_l2sqr(&self.vectors[id * self.d..(id + 1) *
self.d]).sqrt()
+ })
}
}
@@ -1042,6 +1108,17 @@ fn select_neighbors_sorted_into(
}
}
+fn vector_norms_for(metric: MetricType, vectors: &[f32], n: usize, d: usize)
-> Option<Vec<f32>> {
+ if metric != MetricType::Cosine {
+ return None;
+ }
+ Some(
+ (0..n)
+ .map(|id| fvec_norm_l2sqr(&vectors[id * d..(id + 1) * d]).sqrt())
+ .collect(),
+ )
+}
+
fn random_level(node: usize, m: usize, max_level: usize) -> usize {
if node == 0 || max_level <= 1 {
// Keep the first insertion deterministic. Later higher-level nodes
replace
@@ -1270,12 +1347,32 @@ mod tests {
)
.unwrap();
- let (next, dist) = graph.greedy_search_query(&[2.0], 0, 4.0, 0);
+ let distance = QueryDistance::new(&[2.0], MetricType::L2);
+ let (next, dist) = graph.greedy_search_query(&distance, 0, 4.0, 0);
assert_eq!(next, 2);
assert_eq!(dist, 0.0);
}
+ #[test]
+ fn test_hnsw_cosine_distance_uses_vector_norms() {
+ let graph = HnswGraph::from_parts(
+ vec![2.0, 0.0, 4.0, 0.0, 0.0, 3.0],
+ 3,
+ 2,
+ MetricType::Cosine,
+ vec![0, 0, 0],
+ vec![vec![vec![]], vec![vec![]], vec![vec![]]],
+ 0,
+ 0,
+ HnswBuildParams::default(),
+ )
+ .unwrap();
+
+ assert!((graph.distance_between(0, 1) - 0.0).abs() < 1e-6);
+ assert!((graph.distance_between(0, 2) - 1.0).abs() < 1e-6);
+ }
+
fn exact_topk(data: &[f32], n: usize, d: usize, query: &[f32], k: usize)
-> Vec<usize> {
let mut distances: Vec<(f32, usize)> = (0..n)
.map(|i| {
diff --git a/core/src/hnsw_search.rs b/core/src/hnsw_search.rs
index a975c1c..0bff2ed 100644
--- a/core/src/hnsw_search.rs
+++ b/core/src/hnsw_search.rs
@@ -15,7 +15,7 @@
// specific language governing permissions and limitations
// under the License.
-use crate::hnsw::HnswGraph;
+use crate::hnsw::{HnswGraph, HnswSearchWorkspace};
use crate::ivfpq::RowIdFilter;
use crate::topk::TopKHeap;
@@ -37,8 +37,7 @@ where
F: FnMut(&HnswSearchList<'a, P>, &mut TopKHeap),
{
let mut heap = TopKHeap::new(k);
- let mut visited = Vec::new();
- let mut visit_mark = 1usize;
+ let mut workspace = HnswSearchWorkspace::new(ef_search.max(k));
let force_scan = filter
.map(|f| count_filtered(lists, f) <= ef_search.max(k))
.unwrap_or(false);
@@ -49,14 +48,13 @@ where
continue;
}
if let Some(graph) = list.graph {
- let local_results = graph.search_with_workspace(
+ let local_results = graph.search_with_reusable_workspace(
query,
ef_search.max(k),
ef_search.max(k),
- &mut visited,
- &mut visit_mark,
+ &mut workspace,
);
- for (local_id, dist) in local_results {
+ for &(local_id, dist) in local_results {
let row_id = list.ids[local_id];
if filter.map(|f| f.contains(row_id)).unwrap_or(true) {
heap.push(dist, row_id);
diff --git a/core/src/ivfflat.rs b/core/src/ivfflat.rs
index 601bb3d..796dbc2 100644
--- a/core/src/ivfflat.rs
+++ b/core/src/ivfflat.rs
@@ -15,7 +15,7 @@
// specific language governing permissions and limitations
// under the License.
-use crate::distance::{fvec_distance, preprocess_vectors, MetricType};
+use crate::distance::{preprocess_vectors, MetricType, QueryDistance};
use crate::ivfpq::RowIdFilter;
use crate::kmeans::{self, KMeansConfig};
@@ -49,10 +49,16 @@ impl IVFFlatIndex {
pub fn add(&mut self, data: &[f32], ids: &[i64], n: usize) {
let processed = self.preprocess_vectors(data, n);
+ let list_ids = kmeans::find_nearest_batch(
+ &processed,
+ n,
+ &self.quantizer_centroids,
+ self.nlist,
+ self.d,
+ );
for i in 0..n {
let vector = &processed[i * self.d..(i + 1) * self.d];
- let list_id =
- kmeans::find_nearest(vector, &self.quantizer_centroids,
self.nlist, self.d);
+ let list_id = list_ids[i];
self.ids[list_id].push(ids[i]);
self.vectors[list_id].extend_from_slice(vector);
}
@@ -104,6 +110,7 @@ impl IVFFlatIndex {
for qi in 0..nq {
let query = &processed_queries[qi * self.d..(qi + 1) * self.d];
+ let distance = QueryDistance::new(query, self.metric);
let mut heap = FlatTopKHeap::new(k);
for &list_id in &all_probe_indices[qi] {
@@ -116,7 +123,7 @@ impl IVFFlatIndex {
}
}
let vector = &vectors[local_idx * self.d..(local_idx + 1)
* self.d];
- heap.push(fvec_distance(query, vector, self.metric), id);
+ heap.push(distance.distance_to(vector, None), id);
}
}
diff --git a/core/src/ivfhnswflat.rs b/core/src/ivfhnswflat.rs
index 5ab508a..e534dd6 100644
--- a/core/src/ivfhnswflat.rs
+++ b/core/src/ivfhnswflat.rs
@@ -15,7 +15,7 @@
// specific language governing permissions and limitations
// under the License.
-use crate::distance::{fvec_distance, MetricType};
+use crate::distance::{MetricType, QueryDistance};
use crate::hnsw::{HnswBuildParams, HnswGraph};
use crate::hnsw_search::{search_hnsw_lists, HnswSearchList};
use crate::ivfflat::IVFFlatIndex;
@@ -150,6 +150,7 @@ impl IVFHNSWFlatIndex {
filter: Option<&dyn RowIdFilter>,
heap: &mut TopKHeap,
) {
+ let distance = QueryDistance::new(query, self.flat.metric);
for (local_id, &row_id) in self.flat.ids[list_id].iter().enumerate() {
if let Some(f) = filter {
if !f.contains(row_id) {
@@ -158,7 +159,7 @@ impl IVFHNSWFlatIndex {
}
let vector =
&self.flat.vectors[list_id][local_id * self.flat.d..(local_id
+ 1) * self.flat.d];
- heap.push(fvec_distance(query, vector, self.flat.metric), row_id);
+ heap.push(distance.distance_to(vector, None), row_id);
}
}
}
diff --git a/core/src/ivfhnswflat_io.rs b/core/src/ivfhnswflat_io.rs
index f87e413..679fe1c 100644
--- a/core/src/ivfhnswflat_io.rs
+++ b/core/src/ivfhnswflat_io.rs
@@ -15,8 +15,8 @@
// specific language governing permissions and limitations
// under the License.
-use crate::distance::{fvec_distance, preprocess_vectors, MetricType};
-use crate::hnsw::{HnswBuildParams, HnswGraph};
+use crate::distance::{preprocess_vectors, MetricType, QueryDistance};
+use crate::hnsw::{HnswBuildParams, HnswGraph, HnswSearchWorkspace};
use crate::hnsw_search::{search_hnsw_lists, HnswSearchList};
use crate::index_io_util::{
bytes_to_f32_vec, checked_list_bytes, checked_list_offset,
checked_section_size,
@@ -684,6 +684,7 @@ pub fn search_batch_ivfhnswflat_reader_filter<R: SeekRead>(
}
let mut heaps: Vec<TopKHeap> = (0..nq).map(|_| TopKHeap::new(k)).collect();
+ let mut search_workspace = HnswSearchWorkspace::new(ef_search.max(k));
let mut query_filtered_counts = vec![0usize; nq];
let mut loaded_lists = Vec::with_capacity(unique_lists.len());
for (list_id, list) in reader.read_graph_lists_coalesced(&unique_lists)? {
@@ -724,8 +725,13 @@ pub fn search_batch_ivfhnswflat_reader_filter<R: SeekRead>(
&mut heaps[qi],
);
} else {
- let local_results = list.graph.search(query, ef_search.max(k),
ef_search.max(k));
- for (local_id, dist) in local_results {
+ let local_results = list.graph.search_with_reusable_workspace(
+ query,
+ ef_search.max(k),
+ ef_search.max(k),
+ &mut search_workspace,
+ );
+ for &(local_id, dist) in local_results {
let row_id = list.ids[local_id];
if filter.map(|f| f.contains(row_id)).unwrap_or(true) {
heaps[qi].push(dist, row_id);
@@ -841,12 +847,13 @@ fn scan_flat_list(
filter: Option<&dyn RowIdFilter>,
heap: &mut TopKHeap,
) {
+ let distance = QueryDistance::new(query, metric);
for (local_id, &row_id) in ids.iter().enumerate() {
if filter.map(|f| !f.contains(row_id)).unwrap_or(false) {
continue;
}
let vector = &vectors[local_id * d..(local_id + 1) * d];
- heap.push(fvec_distance(query, vector, metric), row_id);
+ heap.push(distance.distance_to(vector, None), row_id);
}
}
diff --git a/core/src/ivfhnswsq.rs b/core/src/ivfhnswsq.rs
index a2cc8f2..a9f5292 100644
--- a/core/src/ivfhnswsq.rs
+++ b/core/src/ivfhnswsq.rs
@@ -15,7 +15,7 @@
// specific language governing permissions and limitations
// under the License.
-use crate::distance::{preprocess_vectors, MetricType};
+use crate::distance::{fvec_madd, preprocess_vectors, MetricType};
use crate::hnsw::{HnswBuildParams, HnswGraph};
use crate::hnsw_search::{search_hnsw_lists, HnswSearchList};
use crate::ivfpq::RowIdFilter;
@@ -198,28 +198,25 @@ impl IVFHNSWSQIndex {
pub(crate) fn decode_list_vectors(&self, list_id: usize, count: usize) ->
Vec<f32> {
let mut vectors = vec![0.0f32; count * self.d];
- self.list_sq(list_id)
- .decode_batch(&self.codes[list_id], count, &mut vectors);
let centroid = self.list_centroid(list_id);
- for vector in vectors.chunks_exact_mut(self.d) {
- for i in 0..self.d {
- vector[i] += centroid[i];
- }
- }
+ self.list_sq(list_id).decode_batch_with_offset(
+ &self.codes[list_id],
+ count,
+ centroid,
+ &mut vectors,
+ );
vectors
}
fn assign_residuals(&self, processed: &[f32], n: usize) -> (Vec<usize>,
Vec<f32>) {
- let mut list_ids = Vec::with_capacity(n);
+ let list_ids =
+ kmeans::find_nearest_batch(processed, n,
&self.quantizer_centroids, self.nlist, self.d);
let mut residuals = vec![0.0f32; n * self.d];
for i in 0..n {
let vector = &processed[i * self.d..(i + 1) * self.d];
- let list_id =
- kmeans::find_nearest(vector, &self.quantizer_centroids,
self.nlist, self.d);
- list_ids.push(list_id);
self.write_residual(
vector,
- list_id,
+ list_ids[i],
&mut residuals[i * self.d..(i + 1) * self.d],
);
}
@@ -244,9 +241,7 @@ impl IVFHNSWSQIndex {
fn write_residual(&self, vector: &[f32], list_id: usize, out: &mut [f32]) {
let centroid = self.list_centroid(list_id);
- for i in 0..self.d {
- out[i] = vector[i] - centroid[i];
- }
+ fvec_madd(vector, centroid, -1.0, out);
}
fn list_centroid(&self, list_id: usize) -> &[f32] {
diff --git a/core/src/ivfhnswsq_io.rs b/core/src/ivfhnswsq_io.rs
index af159bc..c7416d6 100644
--- a/core/src/ivfhnswsq_io.rs
+++ b/core/src/ivfhnswsq_io.rs
@@ -16,7 +16,7 @@
// under the License.
use crate::distance::{preprocess_vectors, MetricType};
-use crate::hnsw::{HnswBuildParams, HnswGraph};
+use crate::hnsw::{HnswBuildParams, HnswGraph, HnswSearchWorkspace};
use crate::hnsw_search::{search_hnsw_lists, HnswSearchList};
use crate::index_io_util::{
checked_list_bytes, checked_list_offset, checked_section_size,
decode_delta_varint_ids,
@@ -578,14 +578,13 @@ impl<R: SeekRead> IVFHNSWSQIndexReader<R> {
let ids = decode_delta_varint_ids(base_id,
&payload[base_header_len..ids_end], meta.count)?;
let codes = payload[ids_end..codes_end].to_vec();
let mut vectors = vec![0.0f32; meta.count * self.d];
- self.list_sq(meta.list_id)
- .decode_batch(&codes, meta.count, &mut vectors);
let centroid = self.list_centroid(meta.list_id).to_vec();
- for vector in vectors.chunks_exact_mut(self.d) {
- for i in 0..self.d {
- vector[i] += centroid[i];
- }
- }
+ self.list_sq(meta.list_id).decode_batch_with_offset(
+ &codes,
+ meta.count,
+ ¢roid,
+ &mut vectors,
+ );
let graph = decode_graph(
&payload[codes_end..],
vectors,
@@ -737,6 +736,7 @@ pub fn search_batch_ivfhnswsq_reader_filter<R: SeekRead>(
}
let mut heaps: Vec<TopKHeap> = (0..nq).map(|_| TopKHeap::new(k)).collect();
+ let mut search_workspace = HnswSearchWorkspace::new(ef_search.max(k));
let mut query_filtered_counts = vec![0usize; nq];
let mut loaded_lists = Vec::with_capacity(unique_lists.len());
for (list_id, list) in reader.read_graph_lists_coalesced(&unique_lists)? {
@@ -783,8 +783,13 @@ pub fn search_batch_ivfhnswsq_reader_filter<R: SeekRead>(
&mut heaps[qi],
);
} else {
- let local_results = list.graph.search(query, ef_search.max(k),
ef_search.max(k));
- for (local_id, dist) in local_results {
+ let local_results = list.graph.search_with_reusable_workspace(
+ query,
+ ef_search.max(k),
+ ef_search.max(k),
+ &mut search_workspace,
+ );
+ for &(local_id, dist) in local_results {
let row_id = list.ids[local_id];
if filter.map(|f| f.contains(row_id)).unwrap_or(true) {
heaps[qi].push(dist, row_id);
@@ -1114,15 +1119,10 @@ fn build_sorted_sq_graph_list(
}
let mut vectors = vec![0.0f32; count * index.d];
+ let centroid = &index.quantizer_centroids[list_id * index.d..(list_id + 1)
* index.d];
index
.list_sq(list_id)
- .decode_batch(&codes, count, &mut vectors);
- let centroid = &index.quantizer_centroids[list_id * index.d..(list_id + 1)
* index.d];
- for vector in vectors.chunks_exact_mut(index.d) {
- for i in 0..index.d {
- vector[i] += centroid[i];
- }
- }
+ .decode_batch_with_offset(&codes, count, centroid, &mut vectors);
let old_to_new = old_to_new_order(&order);
let source_graph = index.graphs[list_id].as_ref().ok_or_else(|| {
io::Error::new(
diff --git a/core/src/ivfpq.rs b/core/src/ivfpq.rs
index 9971998..4e663f2 100644
--- a/core/src/ivfpq.rs
+++ b/core/src/ivfpq.rs
@@ -16,7 +16,8 @@
// under the License.
use crate::distance::{
- fvec_madd, fvec_normalize, pq_distance_four_codes, pq_distance_from_table,
MetricType,
+ fvec_inner_product, fvec_madd, fvec_normalize, pq_distance_four_codes,
pq_distance_from_table,
+ MetricType,
};
use crate::io::{IVFPQIndexReader, SeekRead};
use crate::kmeans::{self, KMeansConfig};
@@ -218,17 +219,8 @@ impl IVFPQIndex {
let processed = self.preprocess_queries(data, n);
// Step 2: Batch assign to coarse centroids (uses sgemm)
- let assignments: Vec<usize> = (0..n)
- .into_par_iter()
- .map(|i| {
- kmeans::find_nearest(
- &processed[i * d..(i + 1) * d],
- &self.quantizer_centroids,
- self.nlist,
- d,
- )
- })
- .collect();
+ let assignments =
+ kmeans::find_nearest_batch(&processed, n,
&self.quantizer_centroids, self.nlist, d);
// Step 3: Batch compute residuals (parallel)
let to_encode = if self.by_residual {
@@ -238,9 +230,12 @@ impl IVFPQIndex {
.enumerate()
.for_each(|(i, res)| {
let list_id = assignments[i];
- for j in 0..d {
- res[j] = processed[i * d + j] -
self.quantizer_centroids[list_id * d + j];
- }
+ fvec_madd(
+ &processed[i * d..(i + 1) * d],
+ &self.quantizer_centroids[list_id * d..(list_id + 1) *
d],
+ -1.0,
+ res,
+ );
});
residuals
} else {
@@ -315,10 +310,10 @@ impl IVFPQIndex {
for j in 0..ksub {
let pq_off = pq_base + j * self.pq.dsub;
- let mut ip = 0.0f32;
- for dd in 0..self.pq.dsub {
- ip += sub_centroid[dd] * self.pq.centroids[pq_off
+ dd];
- }
+ let ip = fvec_inner_product(
+ sub_centroid,
+ &self.pq.centroids[pq_off..pq_off + self.pq.dsub],
+ );
table[tab_base + sub * ksub + j] = pq_norms[sub * ksub
+ j] + 2.0 * ip;
}
}
@@ -506,9 +501,12 @@ impl IVFPQIndex {
let d = self.d;
if self.by_residual {
let mut residual_query = vec![0.0f32; d];
- for j in 0..d {
- residual_query[j] = query[j] -
self.quantizer_centroids[list_id * d + j];
- }
+ fvec_madd(
+ query,
+ &self.quantizer_centroids[list_id * d..(list_id + 1) * d],
+ -1.0,
+ &mut residual_query,
+ );
self.pq
.compute_distance_table(&residual_query, self.metric,
sim_table);
} else {
@@ -1125,9 +1123,12 @@ fn scan_reader_list(entry: &PreReadList, ctx:
&ReaderSearchContext<'_>, heap: &m
);
} else if ctx.by_residual {
let mut residual_query = vec![0.0f32; d];
- for j in 0..d {
- residual_query[j] = ctx.q[j] -
ctx.quantizer_centroids[entry.list_id * d + j];
- }
+ fvec_madd(
+ ctx.q,
+ &ctx.quantizer_centroids[entry.list_id * d..(entry.list_id + 1) *
d],
+ -1.0,
+ &mut residual_query,
+ );
ctx.pq
.compute_distance_table(&residual_query, metric, &mut sim_table);
} else {
@@ -1426,13 +1427,19 @@ fn compute_residuals(
nlist: usize,
) -> Vec<f32> {
let mut residuals = vec![0.0f32; n * d];
- for i in 0..n {
- let point = &data[i * d..(i + 1) * d];
- let list_id = kmeans::find_nearest(point, centroids, nlist, d);
- for j in 0..d {
- residuals[i * d + j] = point[j] - centroids[list_id * d + j];
- }
- }
+ let assignments = kmeans::find_nearest_batch(data, n, centroids, nlist, d);
+ residuals
+ .par_chunks_mut(d)
+ .enumerate()
+ .for_each(|(i, residual)| {
+ let list_id = assignments[i];
+ fvec_madd(
+ &data[i * d..(i + 1) * d],
+ ¢roids[list_id * d..(list_id + 1) * d],
+ -1.0,
+ residual,
+ );
+ });
residuals
}
diff --git a/core/src/kmeans.rs b/core/src/kmeans.rs
index 72d8069..2e1ee6a 100644
--- a/core/src/kmeans.rs
+++ b/core/src/kmeans.rs
@@ -488,6 +488,25 @@ pub fn find_nearest(point: &[f32], centroids: &[f32], k:
usize, d: usize) -> usi
best
}
+pub(crate) fn find_nearest_batch(
+ data: &[f32],
+ n: usize,
+ centroids: &[f32],
+ k: usize,
+ d: usize,
+) -> Vec<usize> {
+ if n == 0 {
+ return Vec::new();
+ }
+ if n == 1 {
+ return vec![find_nearest(&data[..d], centroids, k, d)];
+ }
+
+ let mut assignments = vec![0usize; n];
+ assign_clusters_fast(data, n, d, centroids, k, &mut assignments, 0.0);
+ assignments
+}
+
pub fn find_topk(
point: &[f32],
centroids: &[f32],
@@ -760,6 +779,24 @@ mod tests {
assert_eq!(indices[0], 0);
}
+ #[test]
+ fn test_find_nearest_batch_matches_scalar() {
+ let d = 5;
+ let k = 4;
+ let n = 17;
+ let centroids: Vec<f32> = (0..k * d).map(|i| i as f32 * 0.25 -
2.0).collect();
+ let data: Vec<f32> = (0..n * d)
+ .map(|i| ((i * 13 % 29) as f32) * 0.1 - 1.0)
+ .collect();
+
+ let batch = find_nearest_batch(&data, n, ¢roids, k, d);
+ let scalar: Vec<usize> = (0..n)
+ .map(|i| find_nearest(&data[i * d..(i + 1) * d], ¢roids, k, d))
+ .collect();
+
+ assert_eq!(batch, scalar);
+ }
+
#[test]
fn test_hot_start_converges_faster() {
let mut rng = StdRng::seed_from_u64(42);
diff --git a/core/src/opq.rs b/core/src/opq.rs
index fc00b7a..0ebf46b 100644
--- a/core/src/opq.rs
+++ b/core/src/opq.rs
@@ -15,6 +15,8 @@
// specific language governing permissions and limitations
// under the License.
+use crate::blas::sgemm_a_bt;
+use crate::distance::fvec_inner_product;
use crate::kmeans::KMeansConfig;
use crate::pq::ProductQuantizer;
use nalgebra::{DMatrix, SVD};
@@ -163,21 +165,17 @@ impl OPQMatrix {
pub fn apply(&self, x: &[f32], y: &mut [f32]) {
let d = self.d;
for i in 0..d {
- let mut sum = 0.0f32;
- for j in 0..d {
- sum += self.rotation[i * d + j] * x[j];
- }
- y[i] = sum;
+ y[i] = fvec_inner_product(&self.rotation[i * d..(i + 1) * d], x);
}
}
/// Apply rotation to a batch of vectors.
pub fn apply_batch(&self, data: &[f32], out: &mut [f32], n: usize) {
- for i in 0..n {
- self.apply(
- &data[i * self.d..(i + 1) * self.d],
- &mut out[i * self.d..(i + 1) * self.d],
- );
+ let d = self.d;
+ if n == 1 {
+ self.apply(&data[..d], &mut out[..d]);
+ } else if n > 1 {
+ sgemm_a_bt(n, d, d, 1.0, data, &self.rotation, 0.0, out);
}
}
@@ -254,4 +252,28 @@ mod tests {
assert!((x[i] - x_back[i]).abs() < 1e-6);
}
}
+
+ #[test]
+ fn test_apply_batch_matches_apply() {
+ let d = 4;
+ let mut opq = OPQMatrix::new(d, 2);
+ opq.rotation = vec![
+ 0.5, 0.0, -0.5, 1.0, 1.0, 0.25, 0.0, -0.25, 0.0, 1.5, 0.5, 0.0,
-1.0, 0.0, 0.75, 0.25,
+ ];
+
+ let n = 3;
+ let data = vec![
+ 1.0, 2.0, 3.0, 4.0, -2.0, 0.5, 1.25, 3.5, 0.0, -1.0, 2.0, 0.75,
+ ];
+ let mut batch = vec![0.0f32; n * d];
+ opq.apply_batch(&data, &mut batch, n);
+
+ for i in 0..n {
+ let mut single = vec![0.0f32; d];
+ opq.apply(&data[i * d..(i + 1) * d], &mut single);
+ for j in 0..d {
+ assert!((batch[i * d + j] - single[j]).abs() < 1e-5);
+ }
+ }
+ }
}
diff --git a/core/src/pq.rs b/core/src/pq.rs
index bb2e292..5c7e0d5 100644
--- a/core/src/pq.rs
+++ b/core/src/pq.rs
@@ -17,8 +17,7 @@
use crate::blas::sgemm_a_bt;
use crate::distance::{
- fvec_ip_batch, fvec_l2sqr_batch, fvec_l2sqr_sub, fvec_norm_l2sqr,
pq_distance_from_table,
- MetricType,
+ fvec_ip_batch, fvec_l2sqr_batch, fvec_norm_l2sqr, pq_distance_from_table,
MetricType,
};
use crate::kmeans::{self, KMeansConfig};
use rayon::prelude::*;
@@ -164,91 +163,79 @@ impl ProductQuantizer {
/// For nbits=8: codes has length M (one byte per sub-quantizer).
/// For nbits=4: codes has length M/2 (two nibbles per byte).
pub fn encode(&self, x: &[f32], codes: &mut [u8]) {
+ let mut distances = vec![0.0f32; self.ksub];
+ self.encode_with_distances(x, codes, &mut distances);
+ }
+
+ fn encode_with_distances(&self, x: &[f32], codes: &mut [u8], distances:
&mut [f32]) {
+ debug_assert!(distances.len() >= self.ksub);
if self.nbits == 4 {
- self.encode_4bit(x, codes);
+ self.encode_4bit(x, codes, distances);
} else {
- self.encode_8bit(x, codes);
+ self.encode_8bit(x, codes, distances);
}
}
- fn encode_8bit(&self, x: &[f32], codes: &mut [u8]) {
+ fn encode_8bit(&self, x: &[f32], codes: &mut [u8], distances: &mut [f32]) {
for sub in 0..self.m {
- let x_off = sub * self.dsub;
- let c_base = sub * self.ksub * self.dsub;
-
- let mut best = 0u8;
- let mut best_dist = f32::MAX;
- for j in 0..self.ksub {
- let c_off = c_base + j * self.dsub;
- let dist = fvec_l2sqr_sub(x, x_off, &self.centroids, c_off,
self.dsub);
- if dist < best_dist {
- best_dist = dist;
- best = j as u8;
- }
- }
- codes[sub] = best;
+ self.compute_sub_l2_distances(x, sub, distances);
+ codes[sub] = argmin_code(&distances[..self.ksub]);
}
}
- fn encode_4bit(&self, x: &[f32], codes: &mut [u8]) {
+ fn encode_4bit(&self, x: &[f32], codes: &mut [u8], distances: &mut [f32]) {
for pair in 0..self.m / 2 {
let sub_lo = pair * 2;
let sub_hi = pair * 2 + 1;
- let mut best_lo = 0u8;
- let mut best_dist_lo = f32::MAX;
- let x_off_lo = sub_lo * self.dsub;
- let c_base_lo = sub_lo * self.ksub * self.dsub;
- for j in 0..self.ksub {
- let dist = fvec_l2sqr_sub(
- x,
- x_off_lo,
- &self.centroids,
- c_base_lo + j * self.dsub,
- self.dsub,
- );
- if dist < best_dist_lo {
- best_dist_lo = dist;
- best_lo = j as u8;
- }
- }
+ self.compute_sub_l2_distances(x, sub_lo, distances);
+ let best_lo = argmin_code(&distances[..self.ksub]);
- let mut best_hi = 0u8;
- let mut best_dist_hi = f32::MAX;
- let x_off_hi = sub_hi * self.dsub;
- let c_base_hi = sub_hi * self.ksub * self.dsub;
- for j in 0..self.ksub {
- let dist = fvec_l2sqr_sub(
- x,
- x_off_hi,
- &self.centroids,
- c_base_hi + j * self.dsub,
- self.dsub,
- );
- if dist < best_dist_hi {
- best_dist_hi = dist;
- best_hi = j as u8;
- }
- }
+ self.compute_sub_l2_distances(x, sub_hi, distances);
+ let best_hi = argmin_code(&distances[..self.ksub]);
// Pack: low nibble + high nibble
codes[pair] = best_lo | (best_hi << 4);
}
}
+ fn compute_sub_l2_distances(&self, x: &[f32], sub: usize, distances: &mut
[f32]) {
+ let x_off = sub * self.dsub;
+ let c_base = sub * self.ksub * self.dsub;
+ let query_sub = &x[x_off..x_off + self.dsub];
+ let centroids = &self.centroids[c_base..c_base + self.ksub *
self.dsub];
+
+ if self.dsub >= 4 && self.ksub >= 8 {
+ fvec_ip_batch(query_sub, centroids, self.dsub, self.ksub,
distances);
+ let q_norm = fvec_norm_l2sqr(query_sub);
+ let norms_base = sub * self.ksub;
+ for j in 0..self.ksub {
+ let c_norm = if !self.centroid_norms_cache.is_empty() {
+ self.centroid_norms_cache[norms_base + j]
+ } else {
+ let c_off = j * self.dsub;
+ fvec_norm_l2sqr(¢roids[c_off..c_off + self.dsub])
+ };
+ distances[j] = (q_norm + c_norm - 2.0 * distances[j]).max(0.0);
+ }
+ } else {
+ fvec_l2sqr_batch(query_sub, centroids, self.dsub, self.ksub,
distances);
+ }
+ }
+
/// Encode multiple vectors in parallel.
pub fn encode_batch(&self, data: &[f32], n: usize, codes: &mut [u8]) {
let d = self.d;
let cs = self.code_size();
- codes
- .par_chunks_mut(cs)
- .enumerate()
- .for_each(|(i, code_chunk)| {
+ codes.par_chunks_mut(cs).enumerate().for_each_init(
+ || vec![0.0f32; self.ksub],
+ |distances, (i, code_chunk)| {
if i < n {
- self.encode(&data[i * d..(i + 1) * d], code_chunk);
+ self.encode_with_distances(&data[i * d..(i + 1) * d],
code_chunk, distances);
}
- });
+ },
+ );
}
/// Decode PQ codes back to an approximate vector.
@@ -323,7 +310,7 @@ impl ProductQuantizer {
let c_off = c_base + j * self.dsub;
fvec_norm_l2sqr(&self.centroids[c_off..c_off +
self.dsub])
};
- table[t_base + j] = q_norm + c_norm - 2.0 *
table[t_base + j];
+ table[t_base + j] = (q_norm + c_norm - 2.0 *
table[t_base + j]).max(0.0);
}
}
MetricType::InnerProduct => {
@@ -431,9 +418,25 @@ impl ProductQuantizer {
}
}
+#[inline]
+fn argmin_code(distances: &[f32]) -> u8 {
+ debug_assert!(distances.len() <= 256);
+
+ let mut best = 0usize;
+ let mut best_dist = f32::MAX;
+ for (j, &dist) in distances.iter().enumerate() {
+ if dist < best_dist {
+ best_dist = dist;
+ best = j;
+ }
+ }
+ best as u8
+}
+
#[cfg(test)]
mod tests {
use super::*;
+ use crate::distance::fvec_l2sqr_sub;
use rand::rngs::StdRng;
use rand::{Rng, SeedableRng};
@@ -539,4 +542,30 @@ mod tests {
// Verify codes are non-trivial (not all zeros)
assert!(codes.iter().any(|&b| b != 0));
}
+
+ #[test]
+ fn test_sgemm_distance_table_clamps_negative_self_distance() {
+ let d = 128;
+ let m = 1;
+ let ksub = 256;
+ let mut pq = ProductQuantizer::new(d, m);
+ pq.centroids = vec![0.0; ksub * d];
+
+ let mut query = vec![0.0; d];
+ for i in 0..d {
+ let value = if i.is_multiple_of(2) {
+ 1.0e10_f32 + i as f32
+ } else {
+ -1.0e10_f32 + i as f32
+ };
+ query[i] = value;
+ pq.centroids[i] = value;
+ }
+ pq.rebuild_norms_cache();
+
+ let mut table = vec![0.0; m * ksub];
+ pq.compute_distance_table(&query, MetricType::L2, &mut table);
+
+ assert_eq!(table[0], 0.0);
+ }
}
diff --git a/core/src/sq.rs b/core/src/sq.rs
index fe29509..112f887 100644
--- a/core/src/sq.rs
+++ b/core/src/sq.rs
@@ -81,12 +81,7 @@ impl ScalarQuantizer {
self.ensure_bounds_len();
self.mins.fill(f32::INFINITY);
self.maxs.fill(f32::NEG_INFINITY);
- for vector in values.chunks_exact(self.d) {
- for i in 0..self.d {
- self.mins[i] = self.mins[i].min(vector[i]);
- self.maxs[i] = self.maxs[i].max(vector[i]);
- }
- }
+ update_bounds_batch(values, n, self.d, &mut self.mins, &mut self.maxs);
self.refresh_global_bounds();
}
@@ -99,20 +94,7 @@ impl ScalarQuantizer {
assert!(data.len() >= len);
assert!(codes.len() >= len);
- for row in 0..n {
- let base = row * self.d;
- for dim in 0..self.d {
- let min = self.mins[dim];
- let max = self.maxs[dim];
- let out = base + dim;
- codes[out] = if min >= max {
- 0
- } else {
- let scaled = ((data[out] - min) * 255.0 / (max -
min)).clamp(0.0, 255.0);
- scaled.round() as u8
- };
- }
- }
+ encode_batch_simd(data, n, self.d, &self.mins, &self.maxs, codes);
}
pub fn encode(&self, vector: &[f32], code: &mut [u8]) {
@@ -132,6 +114,21 @@ impl ScalarQuantizer {
}
}
+ pub fn decode_batch_with_offset(
+ &self,
+ codes: &[u8],
+ n: usize,
+ offset: &[f32],
+ vectors: &mut [f32],
+ ) {
+ let len = n * self.d;
+ assert!(codes.len() >= len);
+ assert!(offset.len() >= self.d);
+ assert!(vectors.len() >= len);
+
+ decode_batch_with_offset_simd(codes, n, self.d, &self.mins,
&self.maxs, offset, vectors);
+ }
+
pub fn decode(&self, code: &[u8], vector: &mut [f32]) {
self.decode_batch(code, 1, vector);
}
@@ -344,6 +341,491 @@ impl ScalarQuantizer {
}
}
+fn update_bounds_batch(data: &[f32], n: usize, d: usize, mins: &mut [f32],
maxs: &mut [f32]) {
+ update_bounds_batch_simd(data, n, d, mins, maxs);
+}
+
+#[cfg(target_arch = "x86_64")]
+#[inline]
+fn update_bounds_batch_simd(data: &[f32], n: usize, d: usize, mins: &mut
[f32], maxs: &mut [f32]) {
+ if is_x86_feature_detected!("avx2") && d >= 8 {
+ unsafe { update_bounds_batch_avx2(data, n, d, mins, maxs) };
+ } else {
+ update_bounds_batch_scalar(data, n, d, mins, maxs);
+ }
+}
+
+#[cfg(target_arch = "aarch64")]
+#[inline]
+fn update_bounds_batch_simd(data: &[f32], n: usize, d: usize, mins: &mut
[f32], maxs: &mut [f32]) {
+ unsafe { update_bounds_batch_neon(data, n, d, mins, maxs) };
+}
+
+#[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
+#[inline]
+fn update_bounds_batch_simd(data: &[f32], n: usize, d: usize, mins: &mut
[f32], maxs: &mut [f32]) {
+ update_bounds_batch_scalar(data, n, d, mins, maxs);
+}
+
+#[cfg(not(target_arch = "aarch64"))]
+fn update_bounds_batch_scalar(
+ data: &[f32],
+ n: usize,
+ d: usize,
+ mins: &mut [f32],
+ maxs: &mut [f32],
+) {
+ for vector in data[..n * d].chunks_exact(d) {
+ for i in 0..d {
+ mins[i] = mins[i].min(vector[i]);
+ maxs[i] = maxs[i].max(vector[i]);
+ }
+ }
+}
+
+#[cfg(target_arch = "x86_64")]
+#[target_feature(enable = "avx2")]
+unsafe fn update_bounds_batch_avx2(
+ data: &[f32],
+ n: usize,
+ d: usize,
+ mins: &mut [f32],
+ maxs: &mut [f32],
+) {
+ use std::arch::x86_64::*;
+
+ for row in 0..n {
+ let base = row * d;
+ let mut dim = 0;
+ while dim + 8 <= d {
+ let values = unsafe { _mm256_loadu_ps(data.as_ptr().add(base +
dim)) };
+ let current_min = unsafe { _mm256_loadu_ps(mins.as_ptr().add(dim))
};
+ let current_max = unsafe { _mm256_loadu_ps(maxs.as_ptr().add(dim))
};
+ unsafe {
+ _mm256_storeu_ps(
+ mins.as_mut_ptr().add(dim),
+ _mm256_min_ps(current_min, values),
+ );
+ _mm256_storeu_ps(
+ maxs.as_mut_ptr().add(dim),
+ _mm256_max_ps(current_max, values),
+ );
+ }
+ dim += 8;
+ }
+ while dim < d {
+ let value = unsafe { *data.get_unchecked(base + dim) };
+ let min_ref = unsafe { mins.get_unchecked_mut(dim) };
+ *min_ref = min_ref.min(value);
+ let max_ref = unsafe { maxs.get_unchecked_mut(dim) };
+ *max_ref = max_ref.max(value);
+ dim += 1;
+ }
+ }
+}
+
+#[cfg(target_arch = "aarch64")]
+#[target_feature(enable = "neon")]
+unsafe fn update_bounds_batch_neon(
+ data: &[f32],
+ n: usize,
+ d: usize,
+ mins: &mut [f32],
+ maxs: &mut [f32],
+) {
+ use std::arch::aarch64::*;
+
+ for row in 0..n {
+ let base = row * d;
+ let mut dim = 0;
+ while dim + 4 <= d {
+ let values = unsafe { vld1q_f32(data.as_ptr().add(base + dim)) };
+ let current_min = unsafe { vld1q_f32(mins.as_ptr().add(dim)) };
+ let current_max = unsafe { vld1q_f32(maxs.as_ptr().add(dim)) };
+ unsafe {
+ vst1q_f32(mins.as_mut_ptr().add(dim), vminq_f32(current_min,
values));
+ vst1q_f32(maxs.as_mut_ptr().add(dim), vmaxq_f32(current_max,
values));
+ }
+ dim += 4;
+ }
+ while dim < d {
+ let value = unsafe { *data.get_unchecked(base + dim) };
+ let min_ref = unsafe { mins.get_unchecked_mut(dim) };
+ *min_ref = min_ref.min(value);
+ let max_ref = unsafe { maxs.get_unchecked_mut(dim) };
+ *max_ref = max_ref.max(value);
+ dim += 1;
+ }
+ }
+}
+
+fn encode_batch_simd(
+ data: &[f32],
+ n: usize,
+ d: usize,
+ mins: &[f32],
+ maxs: &[f32],
+ codes: &mut [u8],
+) {
+ encode_batch_simd_impl(data, n, d, mins, maxs, codes);
+}
+
+#[cfg(target_arch = "x86_64")]
+#[inline]
+fn encode_batch_simd_impl(
+ data: &[f32],
+ n: usize,
+ d: usize,
+ mins: &[f32],
+ maxs: &[f32],
+ codes: &mut [u8],
+) {
+ if is_x86_feature_detected!("avx2") && d >= 8 {
+ unsafe { encode_batch_avx2(data, n, d, mins, maxs, codes) };
+ } else {
+ encode_batch_scalar(data, n, d, mins, maxs, codes);
+ }
+}
+
+#[cfg(target_arch = "aarch64")]
+#[inline]
+fn encode_batch_simd_impl(
+ data: &[f32],
+ n: usize,
+ d: usize,
+ mins: &[f32],
+ maxs: &[f32],
+ codes: &mut [u8],
+) {
+ unsafe { encode_batch_neon(data, n, d, mins, maxs, codes) };
+}
+
+#[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
+#[inline]
+fn encode_batch_simd_impl(
+ data: &[f32],
+ n: usize,
+ d: usize,
+ mins: &[f32],
+ maxs: &[f32],
+ codes: &mut [u8],
+) {
+ encode_batch_scalar(data, n, d, mins, maxs, codes);
+}
+
+#[cfg(not(target_arch = "aarch64"))]
+fn encode_batch_scalar(
+ data: &[f32],
+ n: usize,
+ d: usize,
+ mins: &[f32],
+ maxs: &[f32],
+ codes: &mut [u8],
+) {
+ for row in 0..n {
+ let base = row * d;
+ for dim in 0..d {
+ codes[base + dim] = encode_value(data[base + dim], mins[dim],
maxs[dim]);
+ }
+ }
+}
+
+#[cfg(target_arch = "x86_64")]
+#[target_feature(enable = "avx2")]
+unsafe fn encode_batch_avx2(
+ data: &[f32],
+ n: usize,
+ d: usize,
+ mins: &[f32],
+ maxs: &[f32],
+ codes: &mut [u8],
+) {
+ use std::arch::x86_64::*;
+
+ let zero = _mm256_setzero_ps();
+ let one = _mm256_set1_ps(1.0);
+ let max_code = _mm256_set1_ps(255.0);
+ let mut scaled = [0.0f32; 8];
+ for row in 0..n {
+ let base = row * d;
+ let mut dim = 0;
+ while dim + 8 <= d {
+ let values = unsafe { _mm256_loadu_ps(data.as_ptr().add(base +
dim)) };
+ let minv = unsafe { _mm256_loadu_ps(mins.as_ptr().add(dim)) };
+ let maxv = unsafe { _mm256_loadu_ps(maxs.as_ptr().add(dim)) };
+ let range = _mm256_sub_ps(maxv, minv);
+ let valid = _mm256_cmp_ps::<_CMP_GT_OQ>(maxv, minv);
+ let safe_range = _mm256_blendv_ps(one, range, valid);
+ let scale = _mm256_blendv_ps(zero, _mm256_div_ps(max_code,
safe_range), valid);
+ let encoded = _mm256_min_ps(
+ max_code,
+ _mm256_max_ps(zero, _mm256_mul_ps(_mm256_sub_ps(values, minv),
scale)),
+ );
+ unsafe { _mm256_storeu_ps(scaled.as_mut_ptr(), encoded) };
+ for lane in 0..8 {
+ codes[base + dim + lane] = scaled[lane].round() as u8;
+ }
+ dim += 8;
+ }
+ while dim < d {
+ codes[base + dim] = encode_value(data[base + dim], mins[dim],
maxs[dim]);
+ dim += 1;
+ }
+ }
+}
+
+#[cfg(target_arch = "aarch64")]
+#[target_feature(enable = "neon")]
+unsafe fn encode_batch_neon(
+ data: &[f32],
+ n: usize,
+ d: usize,
+ mins: &[f32],
+ maxs: &[f32],
+ codes: &mut [u8],
+) {
+ use std::arch::aarch64::*;
+
+ let zero = vdupq_n_f32(0.0);
+ let one = vdupq_n_f32(1.0);
+ let max_code = vdupq_n_f32(255.0);
+ let mut scaled = [0.0f32; 4];
+ for row in 0..n {
+ let base = row * d;
+ let mut dim = 0;
+ while dim + 4 <= d {
+ let values = unsafe { vld1q_f32(data.as_ptr().add(base + dim)) };
+ let minv = unsafe { vld1q_f32(mins.as_ptr().add(dim)) };
+ let maxv = unsafe { vld1q_f32(maxs.as_ptr().add(dim)) };
+ let range = vsubq_f32(maxv, minv);
+ let valid = vcgtq_f32(maxv, minv);
+ let safe_range = vbslq_f32(valid, range, one);
+ let scale = vbslq_f32(valid, vdivq_f32(max_code, safe_range),
zero);
+ let encoded = vminq_f32(
+ max_code,
+ vmaxq_f32(zero, vmulq_f32(vsubq_f32(values, minv), scale)),
+ );
+ unsafe { vst1q_f32(scaled.as_mut_ptr(), encoded) };
+ for lane in 0..4 {
+ codes[base + dim + lane] = scaled[lane].round() as u8;
+ }
+ dim += 4;
+ }
+ while dim < d {
+ codes[base + dim] = encode_value(data[base + dim], mins[dim],
maxs[dim]);
+ dim += 1;
+ }
+ }
+}
+
+#[inline]
+fn encode_value(value: f32, min: f32, max: f32) -> u8 {
+ if min >= max {
+ 0
+ } else {
+ ((value - min) * 255.0 / (max - min))
+ .clamp(0.0, 255.0)
+ .round() as u8
+ }
+}
+
+fn decode_batch_with_offset_simd(
+ codes: &[u8],
+ n: usize,
+ d: usize,
+ mins: &[f32],
+ maxs: &[f32],
+ offset: &[f32],
+ vectors: &mut [f32],
+) {
+ decode_batch_with_offset_simd_impl(codes, n, d, mins, maxs, offset,
vectors);
+}
+
+#[cfg(target_arch = "x86_64")]
+#[inline]
+fn decode_batch_with_offset_simd_impl(
+ codes: &[u8],
+ n: usize,
+ d: usize,
+ mins: &[f32],
+ maxs: &[f32],
+ offset: &[f32],
+ vectors: &mut [f32],
+) {
+ if is_x86_feature_detected!("avx2") && d >= 8 {
+ unsafe { decode_batch_with_offset_avx2(codes, n, d, mins, maxs,
offset, vectors) };
+ } else {
+ decode_batch_with_offset_scalar(codes, n, d, mins, maxs, offset,
vectors);
+ }
+}
+
+#[cfg(target_arch = "aarch64")]
+#[inline]
+fn decode_batch_with_offset_simd_impl(
+ codes: &[u8],
+ n: usize,
+ d: usize,
+ mins: &[f32],
+ maxs: &[f32],
+ offset: &[f32],
+ vectors: &mut [f32],
+) {
+ unsafe { decode_batch_with_offset_neon(codes, n, d, mins, maxs, offset,
vectors) };
+}
+
+#[cfg(not(any(target_arch = "x86_64", target_arch = "aarch64")))]
+#[inline]
+fn decode_batch_with_offset_simd_impl(
+ codes: &[u8],
+ n: usize,
+ d: usize,
+ mins: &[f32],
+ maxs: &[f32],
+ offset: &[f32],
+ vectors: &mut [f32],
+) {
+ decode_batch_with_offset_scalar(codes, n, d, mins, maxs, offset, vectors);
+}
+
+#[cfg(not(target_arch = "aarch64"))]
+fn decode_batch_with_offset_scalar(
+ codes: &[u8],
+ n: usize,
+ d: usize,
+ mins: &[f32],
+ maxs: &[f32],
+ offset: &[f32],
+ vectors: &mut [f32],
+) {
+ for row in 0..n {
+ let base = row * d;
+ for dim in 0..d {
+ vectors[base + dim] =
+ decode_value(codes[base + dim], mins[dim], maxs[dim]) +
offset[dim];
+ }
+ }
+}
+
+#[cfg(target_arch = "x86_64")]
+#[target_feature(enable = "avx2")]
+unsafe fn decode_batch_with_offset_avx2(
+ codes: &[u8],
+ n: usize,
+ d: usize,
+ mins: &[f32],
+ maxs: &[f32],
+ offset: &[f32],
+ vectors: &mut [f32],
+) {
+ use std::arch::x86_64::*;
+
+ let inv_255 = _mm256_set1_ps(1.0 / 255.0);
+ for row in 0..n {
+ let base = row * d;
+ let mut dim = 0;
+ while dim + 8 <= d {
+ let code_bytes = unsafe { _mm_loadl_epi64(codes.as_ptr().add(base
+ dim).cast()) };
+ let code_i32 = _mm256_cvtepu8_epi32(code_bytes);
+ let code_f32 = _mm256_cvtepi32_ps(code_i32);
+ let minv = unsafe { _mm256_loadu_ps(mins.as_ptr().add(dim)) };
+ let maxv = unsafe { _mm256_loadu_ps(maxs.as_ptr().add(dim)) };
+ let offsetv = unsafe { _mm256_loadu_ps(offset.as_ptr().add(dim)) };
+ let decoded = _mm256_add_ps(
+ offsetv,
+ _mm256_add_ps(
+ minv,
+ _mm256_mul_ps(code_f32, _mm256_mul_ps(_mm256_sub_ps(maxv,
minv), inv_255)),
+ ),
+ );
+ let constant = _mm256_cmp_ps::<_CMP_GE_OQ>(minv, maxv);
+ let constant_decoded = _mm256_add_ps(minv, offsetv);
+ let result = _mm256_blendv_ps(decoded, constant_decoded, constant);
+ unsafe { _mm256_storeu_ps(vectors.as_mut_ptr().add(base + dim),
result) };
+ dim += 8;
+ }
+ while dim < d {
+ vectors[base + dim] =
+ decode_value(codes[base + dim], mins[dim], maxs[dim]) +
offset[dim];
+ dim += 1;
+ }
+ }
+}
+
+#[cfg(target_arch = "aarch64")]
+#[target_feature(enable = "neon")]
+unsafe fn decode_batch_with_offset_neon(
+ codes: &[u8],
+ n: usize,
+ d: usize,
+ mins: &[f32],
+ maxs: &[f32],
+ offset: &[f32],
+ vectors: &mut [f32],
+) {
+ use std::arch::aarch64::*;
+
+ let inv_255 = vdupq_n_f32(1.0 / 255.0);
+ for row in 0..n {
+ let base = row * d;
+ let mut dim = 0;
+ while dim + 8 <= d {
+ let code_u8 = unsafe { vld1_u8(codes.as_ptr().add(base + dim)) };
+ let code_u16 = vmovl_u8(code_u8);
+ let low_u32 = vmovl_u16(vget_low_u16(code_u16));
+ let high_u32 = vmovl_u16(vget_high_u16(code_u16));
+
+ let min0 = unsafe { vld1q_f32(mins.as_ptr().add(dim)) };
+ let max0 = unsafe { vld1q_f32(maxs.as_ptr().add(dim)) };
+ let offset0 = unsafe { vld1q_f32(offset.as_ptr().add(dim)) };
+ let decoded0 = vaddq_f32(
+ offset0,
+ vaddq_f32(
+ min0,
+ vmulq_f32(
+ vcvtq_f32_u32(low_u32),
+ vmulq_f32(vsubq_f32(max0, min0), inv_255),
+ ),
+ ),
+ );
+ let constant0 = vcgeq_f32(min0, max0);
+ let result0 = vbslq_f32(constant0, vaddq_f32(min0, offset0),
decoded0);
+ unsafe { vst1q_f32(vectors.as_mut_ptr().add(base + dim), result0)
};
+
+ let min1 = unsafe { vld1q_f32(mins.as_ptr().add(dim + 4)) };
+ let max1 = unsafe { vld1q_f32(maxs.as_ptr().add(dim + 4)) };
+ let offset1 = unsafe { vld1q_f32(offset.as_ptr().add(dim + 4)) };
+ let decoded1 = vaddq_f32(
+ offset1,
+ vaddq_f32(
+ min1,
+ vmulq_f32(
+ vcvtq_f32_u32(high_u32),
+ vmulq_f32(vsubq_f32(max1, min1), inv_255),
+ ),
+ ),
+ );
+ let constant1 = vcgeq_f32(min1, max1);
+ let result1 = vbslq_f32(constant1, vaddq_f32(min1, offset1),
decoded1);
+ unsafe { vst1q_f32(vectors.as_mut_ptr().add(base + dim + 4),
result1) };
+ dim += 8;
+ }
+ while dim < d {
+ vectors[base + dim] =
+ decode_value(codes[base + dim], mins[dim], maxs[dim]) +
offset[dim];
+ dim += 1;
+ }
+ }
+}
+
+#[inline]
+fn decode_value(code: u8, min: f32, max: f32) -> f32 {
+ if min >= max {
+ min
+ } else {
+ min + code as f32 * (max - min) / 255.0
+ }
+}
+
impl ScalarQuantizerDecodeLut {
#[inline]
pub fn decode_value(&self, code: u8, dim: usize) -> f32 {
@@ -447,6 +929,42 @@ mod tests {
assert!((decoded[3] - 100.0).abs() < 1e-6);
}
+ #[test]
+ fn test_scalar_quantizer_wide_batch_paths() {
+ let d = 9;
+ let n = 5;
+ let data: Vec<f32> = (0..n * d)
+ .map(|i| ((i * 7 % 23) as f32) * 0.5 - 3.0)
+ .collect();
+ let mut sq = ScalarQuantizer::new(d);
+
+ sq.train(&data, n);
+
+ for dim in 0..d {
+ let expected_min = (0..n)
+ .map(|row| data[row * d + dim])
+ .fold(f32::INFINITY, f32::min);
+ let expected_max = (0..n)
+ .map(|row| data[row * d + dim])
+ .fold(f32::NEG_INFINITY, f32::max);
+ assert_eq!(sq.mins[dim], expected_min);
+ assert_eq!(sq.maxs[dim], expected_max);
+ }
+
+ let mut codes = vec![0u8; n * d];
+ sq.encode_batch(&data, n, &mut codes);
+ let mut decoded = vec![0.0f32; n * d];
+ let offset: Vec<f32> = (0..d).map(|dim| dim as f32 * 0.25).collect();
+ sq.decode_batch_with_offset(&codes, n, &offset, &mut decoded);
+
+ for row in 0..n {
+ for dim in 0..d {
+ let expected = sq.decode_value(codes[row * d + dim], dim) +
offset[dim];
+ assert!((decoded[row * d + dim] - expected).abs() < 1e-6);
+ }
+ }
+ }
+
#[test]
fn test_scalar_quantizer_distance_to_code() {
let sq = ScalarQuantizer::with_bounds(2, 0.0, 1.0);
@@ -457,4 +975,18 @@ mod tests {
assert!(dist < 1e-6);
}
+
+ #[test]
+ fn test_scalar_quantizer_decode_batch_with_offset() {
+ let sq = ScalarQuantizer::with_dimension_bounds(2, vec![0.0, -1.0],
vec![1.0, 1.0]);
+ let codes = vec![255, 0, 0, 255];
+ let mut decoded = vec![0.0f32; 4];
+
+ sq.decode_batch_with_offset(&codes, 2, &[10.0, 20.0], &mut decoded);
+
+ assert!((decoded[0] - 11.0).abs() < 1e-6);
+ assert!((decoded[1] - 19.0).abs() < 1e-6);
+ assert!((decoded[2] - 10.0).abs() < 1e-6);
+ assert!((decoded[3] - 21.0).abs() < 1e-6);
+ }
}