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(&centroids[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], &centroids[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, &centroids, k, 
&mut fast, 0.0);
+
+            let mut reference = vec![0usize; n];
+            let obj_ref =
+                assign_clusters_reference(&data, n, d, &centroids, 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, &centroids, k, &mut 
fast, 0.0);
+
+        let mut reference = vec![0usize; n];
+        let obj_ref = assign_clusters_reference(&data, n, d, &centroids, 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, 
&centroids, k, &mut a1, 0.0));
+
+            let mut a8 = vec![0usize; n];
+            let obj8 =
+                pool(8).install(|| assign_clusters_fast(&data, n, d, 
&centroids, 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, &centroids, 
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,
+                    &centroids,
+                    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, &centroids, 
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,
+            &centroids,
+            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, &centroids, 
k, &mut a1, 0.0));
+
+        let mut a8 = vec![0usize; n];
+        let obj8 =
+            pool(8).install(|| assign_clusters_fast(&data, n, d, &centroids, 
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, &centroids, 
k, &mut a1, 0.1));
+
+        let mut a8 = seed_assign.clone();
+        let obj8 =
+            pool(8).install(|| assign_clusters_fast(&data, n, d, &centroids, 
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 &centroids {
+                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 &centroids {
+            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 = &centroids[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;

Reply via email to