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 fd3fd3b fix: stabilize IVF-PQ k-means training (#71)
fd3fd3b is described below
commit fd3fd3b504aec6049b18cf40e8fcf75ab5bc9bc3
Author: jerry <[email protected]>
AuthorDate: Fri Aug 7 15:21:30 2026 +0800
fix: stabilize IVF-PQ k-means training (#71)
---
core/Cargo.toml | 4 +
core/benches/ivfpq_train_bench.rs | 140 ++++++++
core/benches/recall_bench.rs | 47 ++-
core/src/blas.rs | 47 +++
core/src/kmeans.rs | 702 ++++++++++++++++++++++++++++++++++++--
5 files changed, 902 insertions(+), 38 deletions(-)
diff --git a/core/Cargo.toml b/core/Cargo.toml
index 317afa5..37b7fc3 100644
--- a/core/Cargo.toml
+++ b/core/Cargo.toml
@@ -57,3 +57,7 @@ harness = false
[[bench]]
name = "ivfpq_batch_reuse_bench"
harness = false
+
+[[bench]]
+name = "ivfpq_train_bench"
+harness = false
diff --git a/core/benches/ivfpq_train_bench.rs
b/core/benches/ivfpq_train_bench.rs
new file mode 100644
index 0000000..07ce98b
--- /dev/null
+++ b/core/benches/ivfpq_train_bench.rs
@@ -0,0 +1,140 @@
+// Licensed to the Apache Software Foundation (ASF) under one
+// or more contributor license agreements. See the NOTICE file
+// distributed with this work for additional information
+// regarding copyright ownership. The ASF licenses this file
+// to you under the Apache License, Version 2.0 (the
+// "License"); you may not use this file except in compliance
+// with the License. You may obtain a copy of the License at
+//
+// http://www.apache.org/licenses/LICENSE-2.0
+//
+// Unless required by applicable law or agreed to in writing,
+// software distributed under the License is distributed on an
+// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
+// KIND, either express or implied. See the License for the
+// specific language governing permissions and limitations
+// under the License.
+
+//! Single-run IVF-PQ training benchmark.
+//!
+//! `train_total_s` is the authoritative `IVFPQIndex::train` wall time.
+//! The mirrored phase timings (`prep_s`, `coarse_s`, `pq_s`) replay the
+//! InnerProduct/no-OPQ training path with public APIs for attribution only;
+//! never sum them as a total.
+//!
+//! Set `TRAIN_TARGET=1` to run the `target-768` scenario. Its dimensions
+//! can be overridden with `TRAIN_N`, `TRAIN_D`, `TRAIN_NLIST`, and
`TRAIN_PQ_M`.
+//! Thread count follows the global Rayon pool (`RAYON_NUM_THREADS`).
+
+use paimon_vindex_core::distance::MetricType;
+use paimon_vindex_core::ivfpq::IVFPQIndex;
+use paimon_vindex_core::kmeans::{self, KMeansConfig};
+use paimon_vindex_core::pq::ProductQuantizer;
+use rand::rngs::StdRng;
+use rand::{Rng, SeedableRng};
+use std::time::Instant;
+
+struct Scenario {
+ name: &'static str,
+ n: usize,
+ d: usize,
+ nlist: usize,
+ pq_m: usize,
+}
+
+fn main() {
+ println!("=== Paimon IVF-PQ Training Benchmark ===");
+ println!(
+ "metric=InnerProduct, use_opq=false, threads={}",
+ rayon::current_num_threads()
+ );
+ println!();
+ println!(
+ "{:<11} {:>8} {:>5} {:>6} {:>5} {:>8} {:>9} {:>8} {:>14}",
+ "scenario", "n", "d", "nlist", "pq_m", "prep_s", "coarse_s", "pq_s",
"train_total_s"
+ );
+ println!(
+ "{:<11} {:>8} {:>5} {:>6} {:>5} {:>8} {:>9} {:>8} {:>14}",
+ "---------", "------", "---", "-----", "----", "------", "-------",
"----", "-------------"
+ );
+
+ run_scenario(&Scenario {
+ name: "small",
+ n: 50_000,
+ d: 128,
+ nlist: 1024,
+ pq_m: 16,
+ });
+
+ if std::env::var("TRAIN_TARGET").as_deref() == Ok("1") {
+ run_scenario(&Scenario {
+ name: "target-768",
+ n: env_usize("TRAIN_N", 244_606),
+ d: env_usize("TRAIN_D", 768),
+ nlist: env_usize("TRAIN_NLIST", 1024),
+ pq_m: env_usize("TRAIN_PQ_M", 96),
+ });
+ }
+}
+
+fn run_scenario(s: &Scenario) {
+ let data = generate_normalized_vectors(s.n, s.d, 20260806);
+
+ // Authoritative total: real IVFPQIndex::train.
+ let mut index = IVFPQIndex::new(s.d, s.nlist, s.pq_m,
MetricType::InnerProduct, false);
+ let t_train = Instant::now();
+ index.train(&data, s.n);
+ let train_secs = t_train.elapsed().as_secs_f64();
+
+ // Mirrored phases for attribution only.
+ let t_prep = Instant::now();
+ let effective_data = data[..s.n * s.d].to_vec();
+ let prep_secs = t_prep.elapsed().as_secs_f64();
+
+ let km_config = KMeansConfig::default();
+ let t_coarse = Instant::now();
+ let centroids = kmeans::kmeans_train(&km_config, &effective_data, s.n,
s.d, s.nlist);
+ let coarse_secs = t_coarse.elapsed().as_secs_f64();
+
+ let mut pq = ProductQuantizer::new(s.d, s.pq_m);
+ let t_pq = Instant::now();
+ pq.train(&effective_data, s.n);
+ let pq_secs = t_pq.elapsed().as_secs_f64();
+
+ // Keep results observable so nothing is optimized away.
+ let checksum: f32 =
+ centroids.iter().take(8).sum::<f32>() +
pq.centroids.iter().take(8).sum::<f32>();
+
+ println!(
+ "{:<11} {:>8} {:>5} {:>6} {:>5} {:>8.3} {:>9.3} {:>8.3} {:>14.3}",
+ s.name, s.n, s.d, s.nlist, s.pq_m, prep_secs, coarse_secs, pq_secs,
train_secs
+ );
+ debug_assert!(checksum.is_finite());
+}
+
+fn generate_normalized_vectors(n: usize, d: usize, seed: u64) -> Vec<f32> {
+ let mut rng = StdRng::seed_from_u64(seed);
+ let mut data = vec![0.0f32; n * d];
+ for row in data.chunks_mut(d) {
+ let mut norm_sq = 0.0f32;
+ for v in row.iter_mut() {
+ *v = rng.gen::<f32>() * 2.0 - 1.0;
+ norm_sq += *v * *v;
+ }
+ let inv = 1.0 / norm_sq.sqrt().max(1e-12);
+ for v in row.iter_mut() {
+ *v *= inv;
+ }
+ }
+ data
+}
+
+fn env_usize(name: &str, default: usize) -> usize {
+ std::env::var(name)
+ .ok()
+ .map(|v| {
+ v.parse()
+ .unwrap_or_else(|_| panic!("invalid {}: {}", name, v))
+ })
+ .unwrap_or(default)
+}
diff --git a/core/benches/recall_bench.rs b/core/benches/recall_bench.rs
index d2ec25a..bc900f8 100644
--- a/core/benches/recall_bench.rs
+++ b/core/benches/recall_bench.rs
@@ -35,6 +35,7 @@ fn main() {
nlist: 64,
pq_m: 8,
nprobes: &[1, 4, 8, 16, 32, 64],
+ metric: MetricType::L2,
});
println!();
@@ -48,6 +49,23 @@ fn main() {
nlist: 8,
pq_m: 8,
nprobes: &[1, 2, 4, 8],
+ metric: MetricType::L2,
+ });
+
+ println!();
+
+ // Exercises the hierarchical coarse k-means path (nlist > 256) with the
+ // target workload's InnerProduct metric.
+ run_scenario(Scenario {
+ name: "inner-product-hierarchical",
+ d: 64,
+ n: 100_000,
+ nq: 50,
+ k: 10,
+ nlist: 1024,
+ pq_m: 8,
+ nprobes: &[8, 16, 32, 64],
+ metric: MetricType::InnerProduct,
});
}
@@ -60,44 +78,54 @@ struct Scenario<'a> {
nlist: usize,
pq_m: usize,
nprobes: &'a [usize],
+ metric: MetricType,
}
fn run_scenario(s: Scenario<'_>) {
println!("=== IVF Recall Attribution Benchmark ===");
println!(
- "scenario: {}, n={}, nq={}, d={}, nlist={}, avg_list={}, k={},
metric=L2",
+ "scenario: {}, n={}, nq={}, d={}, nlist={}, avg_list={}, k={},
metric={:?}",
s.name,
s.n,
s.nq,
s.d,
s.nlist,
s.n / s.nlist,
- s.k
+ s.k,
+ s.metric
);
- let data = generate_clustered_data(s.n, s.d, 32, 42);
+ let mut data = generate_clustered_data(s.n, s.d, 32, 42);
+ if s.metric == MetricType::InnerProduct {
+ for row in data.chunks_mut(s.d) {
+ let norm = row.iter().map(|v| v *
v).sum::<f32>().sqrt().max(1e-12);
+ for v in row.iter_mut() {
+ *v /= norm;
+ }
+ }
+ }
let ids: Vec<i64> = (0..s.n as i64).collect();
- let queries = &data[..s.nq * s.d];
+ let queries = &data[..s.nq * s.d].to_vec();
let start = Instant::now();
- let ground_truth = brute_force_ground_truth(&data, queries, s.n, s.nq,
s.d, s.k);
+ let ground_truth = brute_force_ground_truth(&data, queries, s.n, s.nq,
s.d, s.k, s.metric);
println!("ground truth: {:.2}s", start.elapsed().as_secs_f64());
let start = Instant::now();
- let mut ivfpq = IVFPQIndex::new(s.d, s.nlist, s.pq_m, MetricType::L2,
false);
+ let mut ivfpq = IVFPQIndex::new(s.d, s.nlist, s.pq_m, s.metric, false);
ivfpq.train(&data, s.n);
ivfpq.add(&data, &ids, s.n);
ivfpq.build_precomputed_table();
println!("build IVF-PQ: {:.2}s", start.elapsed().as_secs_f64());
let start = Instant::now();
- let mut ivfflat = IVFFlatIndex::new(s.d, s.nlist, MetricType::L2);
+ let mut ivfflat = IVFFlatIndex::new(s.d, s.nlist, s.metric);
ivfflat.train(&data, s.n);
ivfflat.add(&data, &ids, s.n);
println!("build IVF-FLAT: {:.2}s", start.elapsed().as_secs_f64());
let start = Instant::now();
- let mut ivfsq = IVFSQIndex::new(s.d, s.nlist, MetricType::L2);
+ let mut ivfsq = IVFSQIndex::new(s.d, s.nlist, s.metric);
ivfsq.train(&data, s.n);
ivfsq.add(&data, &ids, s.n);
println!("build IVF-SQ scan: {:.2}s", start.elapsed().as_secs_f64());
@@ -201,6 +229,7 @@ fn brute_force_ground_truth(
nq: usize,
d: usize,
k: usize,
+ metric: MetricType,
) -> Vec<Vec<i64>> {
(0..nq)
.map(|qi| {
@@ -208,7 +237,7 @@ fn brute_force_ground_truth(
let mut distances: Vec<(f32, i64)> = (0..n)
.map(|i| {
let vector = &data[i * d..(i + 1) * d];
- (fvec_distance(query, vector, MetricType::L2), i as i64)
+ (fvec_distance(query, vector, metric), i as i64)
})
.collect();
distances.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap());
diff --git a/core/src/blas.rs b/core/src/blas.rs
index ddbe9a2..315a7dc 100644
--- a/core/src/blas.rs
+++ b/core/src/blas.rs
@@ -32,6 +32,17 @@ pub fn sgemm_a_bt(
beta: f32,
c: &mut [f32],
) {
+ assert!(
+ k <= isize::MAX as usize && n <= isize::MAX as usize,
+ "SGEMM dimensions exceed isize"
+ );
+ let a_len = m.checked_mul(k).expect("a shape overflows usize");
+ let b_len = n.checked_mul(k).expect("b shape overflows usize");
+ let c_len = m.checked_mul(n).expect("c shape overflows usize");
+ assert!(a.len() >= a_len, "a is shorter than m * k");
+ assert!(b.len() >= b_len, "b is shorter than n * k");
+ assert!(c.len() >= c_len, "c is shorter than m * n");
+
unsafe {
matrixmultiply::sgemm(
m,
@@ -88,4 +99,40 @@ mod tests {
assert!((c[0] - 32.0).abs() < 1e-5);
assert!((c[1] - 50.0).abs() < 1e-5);
}
+
+ #[test]
+ fn test_sgemm_accepts_larger_backing_slices() {
+ let a = [1.0f32, 2.0, 99.0];
+ let b = [3.0f32, 4.0, 99.0];
+ let mut c = [0.0f32, 99.0];
+
+ sgemm_a_bt(1, 1, 2, 1.0, &a, &b, 0.0, &mut c);
+
+ assert_eq!(c, [11.0, 99.0]);
+ }
+
+ #[test]
+ fn test_sgemm_rejects_short_slices() {
+ let short_a = std::panic::catch_unwind(|| {
+ let mut c = [0.0f32];
+ sgemm_a_bt(1, 1, 2, 1.0, &[1.0], &[2.0, 3.0], 0.0, &mut c);
+ });
+ let short_b = std::panic::catch_unwind(|| {
+ let mut c = [0.0f32];
+ sgemm_a_bt(1, 1, 2, 1.0, &[1.0, 2.0], &[3.0], 0.0, &mut c);
+ });
+ let short_c = std::panic::catch_unwind(|| {
+ sgemm_a_bt(1, 1, 2, 1.0, &[1.0, 2.0], &[3.0, 4.0], 0.0, &mut []);
+ });
+
+ assert!(short_a.is_err());
+ assert!(short_b.is_err());
+ assert!(short_c.is_err());
+ }
+
+ #[test]
+ #[should_panic(expected = "SGEMM dimensions exceed isize")]
+ fn test_sgemm_rejects_stride_overflow() {
+ sgemm_a_bt(0, 0, usize::MAX, 1.0, &[], &[], 0.0, &mut []);
+ }
}
diff --git a/core/src/kmeans.rs b/core/src/kmeans.rs
index ede8cb7..485a431 100644
--- a/core/src/kmeans.rs
+++ b/core/src/kmeans.rs
@@ -19,6 +19,47 @@ use crate::blas::sgemm_a_bt;
use crate::distance::{fvec_l2sqr, fvec_norm_l2sqr};
use rand::rngs::StdRng;
use rand::{Rng, SeedableRng};
+use rayon::prelude::*;
+
+/// Cap aggregate concurrent ip_matrix scratch to ~16MB (4M f32 elements).
+const MAX_MATRIX_ELEMS: usize = 4 * 1024 * 1024;
+const MIN_BLOCK_ROWS: usize = 32;
+const MIN_BLOCK_FLOPS: usize = 4_000_000;
+const TARGET_BLOCKS: usize = 64;
+
+fn checked_matrix_len(name: &str, rows: usize, cols: usize) -> usize {
+ rows.checked_mul(cols)
+ .unwrap_or_else(|| panic!("{name} shape overflows usize"))
+}
+
+fn assert_data_shape(data: &[f32], n: usize, d: usize) {
+ let expected = checked_matrix_len("data", n, d);
+ assert_eq!(data.len(), expected, "data length does not match n * d");
+}
+
+fn assignment_block_plan(n: usize, d: usize, k: usize, threads: usize) ->
(usize, bool) {
+ let max_rows = (MAX_MATRIX_ELEMS / k.max(1)).max(1);
+ if threads <= 1 {
+ return (max_rows, false);
+ }
+
+ let row_flops = k.saturating_mul(d).saturating_mul(2).max(1);
+ let min_rows = MIN_BLOCK_ROWS
+ .max(MIN_BLOCK_FLOPS.div_ceil(row_flops))
+ .min(max_rows);
+ let budget_rows = (MAX_MATRIX_ELEMS / threads /
k.max(1)).max(1).min(max_rows);
+ if budget_rows < min_rows {
+ return (budget_rows, false);
+ }
+
+ (
+ n.div_ceil(TARGET_BLOCKS)
+ .max(min_rows)
+ .min(max_rows)
+ .min(budget_rows),
+ true,
+ )
+}
pub struct KMeansConfig {
pub niter: usize,
@@ -49,6 +90,8 @@ const EPS: f32 = 1.0 / 1024.0;
const HIERARCHICAL_THRESHOLD: usize = 256;
pub fn kmeans_train(config: &KMeansConfig, data: &[f32], n: usize, d: usize,
k: usize) -> Vec<f32> {
+ assert_data_shape(data, n, d);
+ checked_matrix_len("centroid", k, d);
if k > HIERARCHICAL_THRESHOLD && n > k {
kmeans_train_hierarchical(config, data, n, d, k)
} else {
@@ -132,13 +175,13 @@ fn kmeans_train_hierarchical(
heap.push(Cluster { centroid, indices });
}
- // Step 2: Iteratively split the largest cluster
+ // Step 2: Iteratively split the largest cluster.
let mut finalized: Vec<Vec<f32>> = Vec::new();
let split_k = 2; // Split into 2 each time
while finalized.len() + heap.len() < target_k {
let largest = match heap.pop() {
- Some(c) => c,
+ Some(cluster) => cluster,
None => break,
};
@@ -147,7 +190,6 @@ fn kmeans_train_hierarchical(
continue;
}
- // Extract sub-data for this cluster
let sub_n = largest.indices.len();
let mut sub_data = vec![0.0f32; sub_n * d];
for (new_idx, &orig_idx) in largest.indices.iter().enumerate() {
@@ -155,7 +197,6 @@ fn kmeans_train_hierarchical(
.copy_from_slice(&train_data[orig_idx * d..(orig_idx + 1) *
d]);
}
- // Run k-means to split
let sub_config = KMeansConfig {
niter: 10,
seed: config.seed + finalized.len() as u64,
@@ -163,7 +204,6 @@ fn kmeans_train_hierarchical(
};
let sub_centroids = kmeans_train_with_init(&sub_config, &sub_data,
sub_n, d, split_k, None);
- // Reassign points in this cluster
let mut sub_assignments = vec![0usize; sub_n];
assign_clusters_fast(
&sub_data,
@@ -175,17 +215,29 @@ fn kmeans_train_hierarchical(
0.0,
);
- for sc in 0..split_k {
- let sub_indices: Vec<usize> = (0..sub_n)
- .filter(|&i| sub_assignments[i] == sc)
- .map(|i| largest.indices[i])
- .collect();
- let centroid = sub_centroids[sc * d..(sc + 1) * d].to_vec();
- if !sub_indices.is_empty() {
- heap.push(Cluster {
- centroid,
- indices: sub_indices,
- });
+ let children: Vec<Cluster> = (0..split_k)
+ .filter_map(|sc| {
+ let sub_indices: Vec<usize> = (0..sub_n)
+ .filter(|&i| sub_assignments[i] == sc)
+ .map(|i| largest.indices[i])
+ .collect();
+ if sub_indices.is_empty() {
+ None
+ } else {
+ Some(Cluster {
+ centroid: sub_centroids[sc * d..(sc + 1) * d].to_vec(),
+ indices: sub_indices,
+ })
+ }
+ })
+ .collect();
+
+ // Finalize a degenerate split instead of re-queuing it forever.
+ if children.len() < 2 {
+ finalized.push(largest.centroid);
+ } else {
+ for child in children {
+ heap.push(child);
}
}
}
@@ -202,7 +254,16 @@ fn kmeans_train_hierarchical(
}
}
- // Pad if needed
+ // If the hierarchy exhausted before reaching target_k (e.g. highly
+ // duplicated data), pad by repeating valid centroids. Zero padding would
+ // fabricate origin centroids that exist nowhere in the data.
+ if result.len() < target_k * d && !result.is_empty() {
+ let produced = result.len() / d;
+ for i in produced..target_k {
+ let src = (i % produced) * d;
+ result.extend_from_within(src..src + d);
+ }
+ }
result.resize(target_k * d, 0.0);
result
}
@@ -215,8 +276,17 @@ pub fn kmeans_train_with_init(
k: usize,
initial_centroids: Option<&[f32]>,
) -> Vec<f32> {
+ assert_data_shape(data, n, d);
+ let centroid_len = checked_matrix_len("centroid", k, d);
+ if let Some(init) = initial_centroids {
+ assert_eq!(
+ init.len(),
+ centroid_len,
+ "initial_centroids length does not match k * d"
+ );
+ }
if n == 0 || k == 0 {
- return vec![0.0; k * d];
+ return vec![0.0; centroid_len];
}
let mut rng = StdRng::seed_from_u64(config.seed);
@@ -230,7 +300,7 @@ pub fn kmeans_train_with_init(
};
if train_n <= k {
- let mut centroids = vec![0.0f32; k * d];
+ let mut centroids = vec![0.0f32; centroid_len];
for i in 0..k {
let src = i % train_n;
centroids[i * d..(i + 1) * d].copy_from_slice(&train_data[src *
d..(src + 1) * d]);
@@ -238,7 +308,7 @@ pub fn kmeans_train_with_init(
return centroids;
}
- let mut best_centroids = vec![0.0f32; k * d];
+ let mut best_centroids = vec![0.0f32; centroid_len];
let mut best_obj = f32::MAX;
let nredo = if initial_centroids.is_some() {
@@ -336,6 +406,12 @@ fn kmeans_plusplus_init(data: &[f32], n: usize, d: usize,
k: usize, rng: &mut St
/// Fast assignment using sgemm: ||x-c||² = ||x||² + ||c||² - 2·x·cᵀ.
/// Supports balance_factor to penalize large clusters.
+///
+/// balance_factor == 0: rows are processed as independent Rayon blocks. Each
+/// row's result does not depend on block boundaries; the objective keeps the
+/// historical serial chunk order so Rayon pool sizes reproduce bitwise.
+/// balance_factor > 0: keeps the historical serial chunking because cluster
+/// size penalties are computed from each chunk's incoming assignments.
fn assign_clusters_fast(
data: &[f32],
n: usize,
@@ -345,15 +421,151 @@ fn assign_clusters_fast(
assignments: &mut [usize],
balance_factor: f32,
) -> f32 {
- // Cap ip_matrix size to ~16MB. Chunk if n*k would be too large.
- const MAX_MATRIX_ELEMS: usize = 4 * 1024 * 1024; // 16MB / 4 bytes
+ if balance_factor > 0.0 {
+ return assign_clusters_balanced_serial(
+ data,
+ n,
+ d,
+ centroids,
+ k,
+ assignments,
+ balance_factor,
+ );
+ }
+ if n == 0 {
+ return 0.0;
+ }
+ if d == 0 {
+ // Degenerate dimension: every distance is zero. Matches the serial
+ // path instead of panicking in par_chunks(0).
+ assignments.fill(0);
+ return 0.0;
+ }
+
+ let c_norms: Vec<f32> = (0..k)
+ .map(|c| fvec_norm_l2sqr(¢roids[c * d..(c + 1) * d]))
+ .collect();
+
+ let max_rows = (MAX_MATRIX_ELEMS / k.max(1)).max(1);
+ let (block_rows, parallel) = assignment_block_plan(n, d, k,
rayon::current_num_threads());
+ if n <= block_rows {
+ return assign_block(data, n, d, centroids, k, &c_norms, assignments,
&mut []);
+ } else if !parallel && block_rows == max_rows {
+ return data
+ .chunks(block_rows * d)
+ .zip(assignments.chunks_mut(block_rows))
+ .map(|(block_data, block_assign)| {
+ assign_block(
+ block_data,
+ block_assign.len(),
+ d,
+ centroids,
+ k,
+ &c_norms,
+ block_assign,
+ &mut [],
+ )
+ })
+ .sum();
+ }
+
+ let mut row_objs = vec![0.0f32; n];
+ if parallel {
+ data.par_chunks(block_rows * d)
+ .zip(assignments.par_chunks_mut(block_rows))
+ .zip(row_objs.par_chunks_mut(block_rows))
+ .for_each(|((block_data, block_assign), block_objs)| {
+ assign_block(
+ block_data,
+ block_assign.len(),
+ d,
+ centroids,
+ k,
+ &c_norms,
+ block_assign,
+ block_objs,
+ );
+ });
+ } else {
+ data.chunks(block_rows * d)
+ .zip(assignments.chunks_mut(block_rows))
+ .zip(row_objs.chunks_mut(block_rows))
+ .for_each(|((block_data, block_assign), block_objs)| {
+ assign_block(
+ block_data,
+ block_assign.len(),
+ d,
+ centroids,
+ k,
+ &c_norms,
+ block_assign,
+ block_objs,
+ );
+ });
+ }
+
+ row_objs
+ .chunks(max_rows)
+ .map(|chunk| chunk.iter().sum::<f32>())
+ .sum()
+}
+
+/// Assign one row block: sgemm inner products + per-row argmin.
+fn assign_block(
+ data: &[f32],
+ n: usize,
+ d: usize,
+ centroids: &[f32],
+ k: usize,
+ c_norms: &[f32],
+ assignments: &mut [usize],
+ row_objs: &mut [f32],
+) -> f32 {
+ let mut ip_matrix = vec![0.0f32; n * k];
+ sgemm_a_bt(n, k, d, 1.0, data, centroids, 0.0, &mut ip_matrix);
+
+ let mut objective = 0.0f32;
+ for i in 0..n {
+ let x_norm = fvec_norm_l2sqr(&data[i * d..(i + 1) * d]);
+ let mut best = 0;
+ let mut best_dist = f32::MAX;
+ let row = i * k;
+ for c in 0..k {
+ let dist = x_norm + c_norms[c] - 2.0 * ip_matrix[row + c];
+ if dist < best_dist {
+ best_dist = dist;
+ best = c;
+ }
+ }
+ assignments[i] = best;
+ if row_objs.is_empty() {
+ objective += best_dist;
+ } else {
+ row_objs[i] = best_dist;
+ }
+ }
+ objective
+}
+
+/// Historical serial path for balance_factor > 0. Chunk boundaries are part of
+/// the observable behavior: cluster sizes are recomputed from each chunk's
+/// incoming assignments.
+fn assign_clusters_balanced_serial(
+ data: &[f32],
+ n: usize,
+ d: usize,
+ centroids: &[f32],
+ k: usize,
+ assignments: &mut [usize],
+ balance_factor: f32,
+) -> f32 {
if n * k > MAX_MATRIX_ELEMS {
let chunk_n = MAX_MATRIX_ELEMS / k;
let mut total_obj = 0.0f32;
let mut offset = 0;
while offset < n {
let cn = (n - offset).min(chunk_n);
- total_obj += assign_clusters_fast(
+ total_obj += assign_clusters_balanced_serial(
&data[offset * d..(offset + cn) * d],
cn,
d,
@@ -379,11 +591,9 @@ fn assign_clusters_fast(
// Compute cluster sizes for balance penalty
let mut cluster_sizes = vec![0u32; k];
- if balance_factor > 0.0 {
- for &a in assignments.iter() {
- if a < k {
- cluster_sizes[a] += 1;
- }
+ for &a in assignments.iter() {
+ if a < k {
+ cluster_sizes[a] += 1;
}
}
@@ -395,7 +605,7 @@ fn assign_clusters_fast(
for c in 0..k {
let mut dist = x_norms[i] + c_norms[c] - 2.0 * ip_matrix[row + c];
// Balance penalty: prefer smaller clusters
- if balance_factor > 0.0 && cluster_sizes[c] > 0 {
+ if cluster_sizes[c] > 0 {
dist += balance_factor * (cluster_sizes[c] as f32).ln();
}
if dist < best_dist {
@@ -773,6 +983,287 @@ fn subsample(data: &[f32], n: usize, d: usize, target_n:
usize, rng: &mut StdRng
mod tests {
use super::*;
+ /// Sequential scalar reference for cluster assignment. Mirrors the
+ /// balance-penalty semantics of a single (unchunked) call.
+ fn assign_clusters_reference(
+ data: &[f32],
+ n: usize,
+ d: usize,
+ centroids: &[f32],
+ k: usize,
+ assignments: &mut [usize],
+ balance_factor: f32,
+ ) -> f32 {
+ let mut cluster_sizes = vec![0u32; k];
+ if balance_factor > 0.0 {
+ for &a in assignments.iter() {
+ if a < k {
+ cluster_sizes[a] += 1;
+ }
+ }
+ }
+ let mut total_obj = 0.0f32;
+ for i in 0..n {
+ let mut best = 0;
+ let mut best_dist = f32::MAX;
+ for c in 0..k {
+ let mut dist =
+ fvec_l2sqr(&data[i * d..(i + 1) * d], ¢roids[c * d..(c
+ 1) * d]);
+ if balance_factor > 0.0 && cluster_sizes[c] > 0 {
+ dist += balance_factor * (cluster_sizes[c] as f32).ln();
+ }
+ if dist < best_dist {
+ best_dist = dist;
+ best = c;
+ }
+ }
+ assignments[i] = best;
+ total_obj += best_dist;
+ }
+ total_obj
+ }
+
+ fn deterministic_data(n: usize, d: usize, seed: u64) -> Vec<f32> {
+ let mut rng = StdRng::seed_from_u64(seed);
+ (0..n * d).map(|_| rng.gen::<f32>() * 2.0 - 1.0).collect()
+ }
+
+ fn pool(threads: usize) -> rayon::ThreadPool {
+ rayon::ThreadPoolBuilder::new()
+ .num_threads(threads)
+ .build()
+ .unwrap()
+ }
+
+ #[test]
+ fn test_assign_clusters_matches_reference_shapes() {
+ // Shapes chosen to cover: tiny, uneven final row block, and the
+ // chunked path (n * k > MAX_MATRIX_ELEMS).
+ let shapes: &[(usize, usize, usize)] = &[(17, 4, 5), (1003, 7, 9),
(70_000, 64, 8)];
+ for &(n, k, d) in shapes {
+ let data = deterministic_data(n, d, 7);
+ let centroids = deterministic_data(k, d, 11);
+
+ let mut fast = vec![0usize; n];
+ let obj_fast = assign_clusters_fast(&data, n, d, ¢roids, k,
&mut fast, 0.0);
+
+ let mut reference = vec![0usize; n];
+ let obj_ref =
+ assign_clusters_reference(&data, n, d, ¢roids, k, &mut
reference, 0.0);
+
+ assert_eq!(
+ fast, reference,
+ "assignments diverge for shape ({n},{k},{d})"
+ );
+ let rel = (obj_fast - obj_ref).abs() / obj_ref.max(1e-10);
+ assert!(
+ rel < 1e-5,
+ "objective rel err {rel} for shape ({n},{k},{d})"
+ );
+ }
+ }
+
+ #[test]
+ fn test_assign_clusters_target_boundary_shape() {
+ // The target coarse assignment shape n=244606, k=16 sits just below
+ // MAX_MATRIX_ELEMS. Use a small d instead of 768 to keep memory low.
+ let (n, k, d) = (244_606, 16, 4);
+ let data = deterministic_data(n, d, 13);
+ let centroids = deterministic_data(k, d, 17);
+
+ let mut fast = vec![0usize; n];
+ let obj_fast = assign_clusters_fast(&data, n, d, ¢roids, k, &mut
fast, 0.0);
+
+ let mut reference = vec![0usize; n];
+ let obj_ref = assign_clusters_reference(&data, n, d, ¢roids, k,
&mut reference, 0.0);
+
+ assert_eq!(fast, reference);
+ let rel = (obj_fast - obj_ref).abs() / obj_ref.max(1e-10);
+ assert!(rel < 1e-5, "objective rel err {rel}");
+ }
+
+ #[test]
+ fn test_assign_clusters_cross_thread_objective() {
+ // Fixed block boundaries make assignments and the objective
+ // bitwise reproducible across Rayon pool sizes.
+ for &(n, k, d) in &[
+ (1003usize, 7usize, 9usize),
+ (70_000, 64, 8),
+ (244_606, 16, 4),
+ (524_289, 2, 1),
+ ] {
+ let data = deterministic_data(n, d, 19);
+ let centroids = deterministic_data(k, d, 23);
+
+ let mut a1 = vec![0usize; n];
+ let obj1 =
+ pool(1).install(|| assign_clusters_fast(&data, n, d,
¢roids, k, &mut a1, 0.0));
+
+ let mut a8 = vec![0usize; n];
+ let obj8 =
+ pool(8).install(|| assign_clusters_fast(&data, n, d,
¢roids, k, &mut a8, 0.0));
+
+ assert_eq!(a1, a8, "assignments diverge across pools for
({n},{k},{d})");
+ assert_eq!(
+ obj1.to_bits(),
+ obj8.to_bits(),
+ "objective diverges across pools for ({n},{k},{d})"
+ );
+ }
+ }
+
+ #[test]
+ fn test_parallel_assignment_respects_aggregate_scratch_budget() {
+ let (rows, parallel) = assignment_block_plan(262_144, 768, 1024, 16);
+ assert!(parallel);
+ assert_eq!(rows, 256);
+ assert!(rows * 1024 * 16 <= MAX_MATRIX_ELEMS);
+
+ let (rows, parallel) = assignment_block_plan(2_000_000, 1, 2, 16);
+ assert!(!parallel);
+ assert_eq!(rows, 131_072);
+ assert!(rows * 2 * 16 <= MAX_MATRIX_ELEMS);
+ }
+
+ #[test]
+ fn test_assign_clusters_preserves_serial_chunk_objective() {
+ let (n, k, d) = (70_000usize, 64usize, 8usize);
+ let data = deterministic_data(n, d, 27);
+ let centroids = deterministic_data(k, d, 29);
+
+ let mut fast = vec![0usize; n];
+ let fast_obj =
+ pool(8).install(|| assign_clusters_fast(&data, n, d, ¢roids,
k, &mut fast, 0.0));
+
+ let c_norms: Vec<f32> =
centroids.chunks(d).map(fvec_norm_l2sqr).collect();
+ let max_rows = MAX_MATRIX_ELEMS / k;
+ let mut serial = vec![0usize; n];
+ let mut serial_rows = vec![0.0f32; n];
+ let serial_obj: f32 = data
+ .chunks(max_rows * d)
+ .zip(serial.chunks_mut(max_rows))
+ .zip(serial_rows.chunks_mut(max_rows))
+ .map(|((block_data, block_assign), block_objs)| {
+ assign_block(
+ block_data,
+ block_assign.len(),
+ d,
+ ¢roids,
+ k,
+ &c_norms,
+ block_assign,
+ block_objs,
+ );
+ block_objs.iter().sum::<f32>()
+ })
+ .sum();
+
+ assert_eq!(fast, serial);
+ assert_eq!(fast_obj.to_bits(), serial_obj.to_bits());
+ }
+
+ #[test]
+ fn test_split_training_shape_uses_single_sgemm() {
+ let (n, k, d) = (512usize, 2usize, 768usize);
+ let data = deterministic_data(n, d, 29);
+ let centroids = deterministic_data(k, d, 31);
+
+ let mut fast = vec![0usize; n];
+ let fast_obj =
+ pool(8).install(|| assign_clusters_fast(&data, n, d, ¢roids,
k, &mut fast, 0.0));
+
+ let c_norms: Vec<f32> =
centroids.chunks(d).map(fvec_norm_l2sqr).collect();
+ let mut single = vec![0usize; n];
+ let mut single_rows = vec![0.0f32; n];
+ assign_block(
+ &data,
+ n,
+ d,
+ ¢roids,
+ k,
+ &c_norms,
+ &mut single,
+ &mut single_rows,
+ );
+ let single_obj: f32 = single_rows.iter().sum();
+
+ assert_eq!(fast, single);
+ assert_eq!(fast_obj.to_bits(), single_obj.to_bits());
+ }
+
+ #[test]
+ fn test_small_assignments_bitwise_reproducible_across_pools() {
+ let (n, k, d) = (256usize, 2usize, 64usize);
+ let data = deterministic_data(n, d, 47);
+ let centroids = deterministic_data(k, d, 53);
+
+ let mut a1 = vec![0usize; n];
+ let obj1 =
+ pool(1).install(|| assign_clusters_fast(&data, n, d, ¢roids,
k, &mut a1, 0.0));
+
+ let mut a8 = vec![0usize; n];
+ let obj8 =
+ pool(8).install(|| assign_clusters_fast(&data, n, d, ¢roids,
k, &mut a8, 0.0));
+
+ assert_eq!(a1, a8);
+ assert_eq!(obj1.to_bits(), obj8.to_bits());
+ }
+
+ #[test]
+ fn test_assign_clusters_balanced_stays_serial() {
+ // balance_factor > 0 keeps the serial chunked path, so results must be
+ // bitwise identical regardless of the Rayon pool size. n*k exceeds
+ // MAX_MATRIX_ELEMS to exercise the chunk boundaries.
+ let (n, k, d) = (70_000usize, 64usize, 8usize);
+ let data = deterministic_data(n, d, 29);
+ let centroids = deterministic_data(k, d, 31);
+ let seed_assign: Vec<usize> = (0..n).map(|i| i % k).collect();
+
+ let mut a1 = seed_assign.clone();
+ let obj1 =
+ pool(1).install(|| assign_clusters_fast(&data, n, d, ¢roids,
k, &mut a1, 0.1));
+
+ let mut a8 = seed_assign.clone();
+ let obj8 =
+ pool(8).install(|| assign_clusters_fast(&data, n, d, ¢roids,
k, &mut a8, 0.1));
+
+ assert_eq!(a1, a8);
+ assert_eq!(obj1.to_bits(), obj8.to_bits());
+ }
+
+ #[test]
+ fn test_assign_clusters_bitwise_reproducible_fixed_pool() {
+ // Same data, seed, and pool size must reproduce k-means bitwise.
+ let n = 3000;
+ let d = 6;
+ let data = deterministic_data(n, d, 37);
+ let config = KMeansConfig::default();
+
+ let run = || {
+ pool(4).install(|| {
+ let flat = kmeans_train_with_init(&config, &data, n, d, 24,
None);
+ let hier = kmeans_train(&config, &data, n, d, 300);
+ let mut assignments = vec![0usize; n];
+ let obj = assign_clusters_fast(&data, n, d, &flat, 24, &mut
assignments, 0.0);
+ (flat, hier, assignments, obj)
+ })
+ };
+
+ let (flat_a, hier_a, assign_a, obj_a) = run();
+ let (flat_b, hier_b, assign_b, obj_b) = run();
+
+ assert!(flat_a
+ .iter()
+ .zip(&flat_b)
+ .all(|(x, y)| x.to_bits() == y.to_bits()));
+ assert!(hier_a
+ .iter()
+ .zip(&hier_b)
+ .all(|(x, y)| x.to_bits() == y.to_bits()));
+ assert_eq!(assign_a, assign_b);
+ assert_eq!(obj_a.to_bits(), obj_b.to_bits());
+ }
+
#[test]
fn test_two_clusters() {
let mut data = Vec::new();
@@ -811,6 +1302,47 @@ mod tests {
assert_eq!(indices[0], 0);
}
+ #[test]
+ fn test_kmeans_rejects_invalid_data_shapes_before_early_return() {
+ let config = KMeansConfig::default();
+
+ let short =
+ std::panic::catch_unwind(|| kmeans_train_with_init(&config, &[0.0;
3], 2, 2, 0, None));
+ let long = std::panic::catch_unwind(|| kmeans_train(&config, &[0.0;
5], 2, 2, 0));
+ let overflow = std::panic::catch_unwind(|| {
+ kmeans_train_with_init(&config, &[], usize::MAX, 2, 0, None)
+ });
+
+ assert!(short.is_err());
+ assert!(long.is_err());
+ assert!(overflow.is_err());
+ }
+
+ #[test]
+ fn test_kmeans_rejects_invalid_centroid_shapes_before_early_return() {
+ let config = KMeansConfig::default();
+
+ let short = std::panic::catch_unwind(|| {
+ kmeans_train_with_init(&config, &[], 0, 2, 1, Some(&[0.0]))
+ });
+ let long = std::panic::catch_unwind(|| {
+ kmeans_train_with_init(&config, &[], 0, 2, 1, Some(&[0.0; 3]))
+ });
+ let overflow = std::panic::catch_unwind(|| {
+ kmeans_train_with_init(&config, &[], 0, 2, usize::MAX, None)
+ })
+ .unwrap_err();
+ let overflow_message = overflow
+ .downcast_ref::<String>()
+ .map(String::as_str)
+ .or_else(|| overflow.downcast_ref::<&str>().copied())
+ .unwrap_or_default();
+
+ assert!(short.is_err());
+ assert!(long.is_err());
+ assert!(overflow_message.contains("centroid shape overflows usize"));
+ }
+
#[test]
fn test_find_topk_batch_matches_full_sort_with_ties() {
let d = 2;
@@ -946,6 +1478,118 @@ mod tests {
assert!(diverse, "Streaming centroids are not diverse");
}
+ #[test]
+ fn test_hierarchical_exact_k() {
+ // Requested k must be returned exactly, including non-power-of-two k.
+ let d = 4;
+ let n = 4000;
+ let data = deterministic_data(n, d, 41);
+ let config = KMeansConfig::default();
+ for &k in &[257usize, 1000, 1024] {
+ let centroids = kmeans_train(&config, &data, n, d, k);
+ assert_eq!(centroids.len(), k * d, "wrong centroid count for
k={k}");
+ for &v in ¢roids {
+ assert!(v.is_finite());
+ }
+ }
+ }
+
+ #[test]
+ fn test_hierarchical_deterministic_same_seed() {
+ let d = 4;
+ let n = 3000;
+ let data = deterministic_data(n, d, 43);
+ let config = KMeansConfig::default();
+ let a = kmeans_train(&config, &data, n, d, 300);
+ let b = kmeans_train(&config, &data, n, d, 300);
+ assert!(a.iter().zip(&b).all(|(x, y)| x.to_bits() == y.to_bits()));
+ }
+
+ #[test]
+ fn test_hierarchical_strict_largest_first_fixture() {
+ let d = 2;
+ let k = 257;
+ let mut data = Vec::new();
+
+ for i in 0..1024 {
+ data.push((i % 32) as f32 * 0.01);
+ data.push((i / 32) as f32 * 0.01);
+ }
+ for cluster in 1..16 {
+ for i in 0..64 {
+ data.push(cluster as f32 * 1000.0 + (i % 8) as f32 * 0.01);
+ data.push((i / 8) as f32 * 0.01);
+ }
+ }
+
+ let centroids = kmeans_train(&KMeansConfig::default(), &data,
data.len() / d, d, k);
+ let largest_cluster_centroids = centroids
+ .chunks_exact(d)
+ .filter(|centroid| centroid[0] < 500.0)
+ .count();
+
+ // Strict pop/split/reinsert assigns 232 centroids here; batched parent
+ // pops assign 234 and therefore change the trained index.
+ assert_eq!(
+ largest_cluster_centroids, 232,
+ "hierarchical split order changed"
+ );
+ }
+
+ #[test]
+ fn test_hierarchical_tiny_split_candidates() {
+ // Highly duplicated data creates tiny/empty split candidates; the
+ // hierarchy must still return exactly k centroids.
+ let d = 4;
+ let k = 300;
+ let n = 600;
+ let mut data = vec![0.0f32; n * d];
+ for i in 0..n {
+ let v = (i % 5) as f32;
+ for j in 0..d {
+ data[i * d + j] = v;
+ }
+ }
+ let config = KMeansConfig::default();
+ let centroids = kmeans_train(&config, &data, n, d, k);
+ assert_eq!(centroids.len(), k * d);
+ for &v in ¢roids {
+ assert!(v.is_finite());
+ }
+ }
+
+ #[test]
+ fn test_hierarchical_duplicate_data_pads_with_valid_centroids() {
+ // All-duplicate non-zero data exhausts the split hierarchy early.
+ // Padding must repeat valid centroids, never fabricate zeros.
+ let d = 4;
+ let k = 300;
+ let n = 2000;
+ let data = vec![10.0f32; n * d];
+ let config = KMeansConfig::default();
+ let centroids = kmeans_train(&config, &data, n, d, k);
+ assert_eq!(centroids.len(), k * d);
+ for c in 0..k {
+ let row = ¢roids[c * d..(c + 1) * d];
+ // Empty-cluster handling perturbs donors by ±EPS, so allow a small
+ // relative band around 10.0; zero padding would land far outside.
+ assert!(
+ row.iter().all(|&v| (9.0..=11.0).contains(&v)),
+ "centroid {c} is not derived from the data: {row:?}"
+ );
+ }
+ }
+
+ #[test]
+ fn test_assign_clusters_zero_dimension_does_not_panic() {
+ let n = 5;
+ let k = 3;
+ let mut assignments = vec![7usize; n];
+ let obj = assign_clusters_fast(&[], n, 0, &[], k, &mut assignments,
0.0);
+ assert_eq!(assignments, vec![0usize; n]);
+ assert_eq!(obj, 0.0);
+ }
+
#[test]
fn test_hierarchical_kmeans() {
let n = 2000;