From 7d02be7f14fd372a496e29b4926cf701554017af Mon Sep 17 00:00:00 2001 From: isaac-dasari <74317395+isaac-dasari@users.noreply.github.com> Date: Sun, 16 Aug 2026 14:06:19 -0700 Subject: [PATCH 1/2] perf(index): parallelize kmeans centroid recomputation --- .github/workflows/rust-benchmark.yml | 2 +- rust/lance-index/Cargo.toml | 4 + rust/lance-index/benches/kmeans_recompute.rs | 45 +++++ rust/lance-index/src/vector/kmeans.rs | 187 ++++++++++++++++--- 4 files changed, 214 insertions(+), 24 deletions(-) create mode 100644 rust/lance-index/benches/kmeans_recompute.rs diff --git a/.github/workflows/rust-benchmark.yml b/.github/workflows/rust-benchmark.yml index bb0960148a9..29de01a9dc3 100644 --- a/.github/workflows/rust-benchmark.yml +++ b/.github/workflows/rust-benchmark.yml @@ -50,7 +50,7 @@ jobs: working-directory: ./rust/lance-index run: | # TODO: a few benchmarks are failing. Re-enable everything once they are fixed. - cargo bench --bench sq --bench hnsw --bench inverted --bench pq_dist_table --bench pq_assignment -- --output-format bencher | tee -a ../../output.txt + cargo bench --bench sq --bench hnsw --bench inverted --bench pq_dist_table --bench pq_assignment --bench kmeans_recompute -- --output-format bencher | tee -a ../../output.txt - name: Store benchmark result if: github.event_name != 'pull_request' uses: benchmark-action/github-action-benchmark@a7bc2366eda11037936ea57d811a43b3418d3073 # v1.21.0 diff --git a/rust/lance-index/Cargo.toml b/rust/lance-index/Cargo.toml index a74f06fe903..e453ead47d5 100644 --- a/rust/lance-index/Cargo.toml +++ b/rust/lance-index/Cargo.toml @@ -143,6 +143,10 @@ harness = false name = "kmeans" harness = false +[[bench]] +name = "kmeans_recompute" +harness = false + [[bench]] name = "compute_partition" harness = false diff --git a/rust/lance-index/benches/kmeans_recompute.rs b/rust/lance-index/benches/kmeans_recompute.rs new file mode 100644 index 00000000000..984967930d7 --- /dev/null +++ b/rust/lance-index/benches/kmeans_recompute.rs @@ -0,0 +1,45 @@ +// SPDX-License-Identifier: Apache-2.0 +// SPDX-FileCopyrightText: Copyright The Lance Authors + +use std::hint::black_box; + +use arrow_array::types::Float32Type; +use criterion::{BenchmarkId, Criterion, criterion_group, criterion_main}; +use lance_index::vector::kmeans::{KMeansAlgo, KMeansAlgoFloat}; +use lance_linalg::distance::DistanceType; +use lance_testing::datagen::generate_random_array; + +fn bench_recompute_centroids(c: &mut Criterion) { + let mut group = c.benchmark_group("kmeans_recompute_centroids"); + + let (num_vectors, dimension, k) = (128 * 1024, 128, 256); + let data = generate_random_array(num_vectors * dimension); + let membership = (0..num_vectors) + .map(|row| Some((row % k) as u32)) + .collect::>(); + let cluster_sizes = vec![num_vectors / k; k]; + + group.bench_with_input( + BenchmarkId::new(format!("{dimension}d_{k}k"), num_vectors), + &num_vectors, + |b, _| { + b.iter(|| { + let mut cluster_sizes = cluster_sizes.clone(); + black_box(KMeansAlgoFloat::::to_kmeans( + black_box(data.values()), + dimension, + k, + black_box(&membership), + &mut cluster_sizes, + DistanceType::L2, + 0.0, + )) + }); + }, + ); + + group.finish(); +} + +criterion_group!(benches, bench_recompute_centroids); +criterion_main!(benches); diff --git a/rust/lance-index/src/vector/kmeans.rs b/rust/lance-index/src/vector/kmeans.rs index ceff08c2d45..f1a389a3123 100644 --- a/rust/lance-index/src/vector/kmeans.rs +++ b/rust/lance-index/src/vector/kmeans.rs @@ -331,6 +331,28 @@ where phantom_data: std::marker::PhantomData, } +// Keep the thread-local centroid accumulators from adding hundreds of MiB to +// large IVF training jobs. One accumulator is always allowed because the +// centroid output itself already requires that much memory. +const MAX_CENTROID_ACCUMULATOR_BYTES: usize = 64 * 1024 * 1024; + +fn centroid_accumulator_count( + num_vectors: usize, + centroid_len: usize, + available_parallelism: usize, +) -> usize { + let accumulator_bytes = centroid_len.saturating_mul(std::mem::size_of::()); + let memory_limited_parallelism = MAX_CENTROID_ACCUMULATOR_BYTES + .checked_div(accumulator_bytes) + .unwrap_or(1) + .max(1); + + available_parallelism + .min(memory_limited_parallelism) + .min(num_vectors) + .max(1) +} + impl KMeansAlgo for KMeansAlgoFloat where T::Native: Float + Dot + L2 + MulAssign + DivAssign + AddAssign + FromPrimitive + Sync, @@ -399,34 +421,52 @@ where distance_type: DistanceType, loss: f64, ) -> KMeans { - let mut centroids = vec![T::Native::zero(); k * dimension]; - - let mut num_cpus = get_num_compute_intensive_cpus(); - if k < num_cpus || k < 16 { - num_cpus = 1; - } - let chunk_size = k / num_cpus; - - centroids - .par_chunks_mut(dimension * chunk_size) - .enumerate() - .with_max_len(1) - .for_each(|(i, centroids)| { - let start = i * chunk_size; - let end = ((i + 1) * chunk_size).min(k); + let num_vectors = data.len() / dimension; + let centroid_len = k * dimension; + let available_parallelism = if k < 16 { + 1 + } else { + get_num_compute_intensive_cpus() + }; + let accumulator_count = centroid_accumulator_count::( + num_vectors, + centroid_len, + available_parallelism, + ); + let vectors_per_chunk = num_vectors.div_ceil(accumulator_count); + + // Each worker scans a disjoint portion of the input once and writes to + // a private centroid buffer. This changes the input-read complexity + // from O(num_vectors * workers) to O(num_vectors) without contention. + let mut accumulators = data + .par_chunks(vectors_per_chunk * dimension) + .zip(membership.par_chunks(vectors_per_chunk)) + .map(|(data, membership)| { + let mut local_centroids = vec![T::Native::zero(); centroid_len]; data.chunks(dimension) - .zip(membership.iter()) + .zip(membership) .filter_map(|(vector, cluster_id)| { cluster_id.map(|cluster_id| (vector, cluster_id as usize)) }) .for_each(|(vector, cluster_id)| { - if start <= cluster_id && cluster_id < end { - let local_id = cluster_id - start; - let centroid = - &mut centroids[local_id * dimension..(local_id + 1) * dimension]; - centroid.iter_mut().zip(vector).for_each(|(c, v)| *c += *v); - } + let start = cluster_id * dimension; + let centroid = &mut local_centroids[start..start + dimension]; + centroid.iter_mut().zip(vector).for_each(|(c, v)| *c += *v); }); + local_centroids + }) + .collect::>(); + + let mut centroids = accumulators + .pop() + .unwrap_or_else(|| vec![T::Native::zero(); centroid_len]); + centroids + .par_iter_mut() + .enumerate() + .for_each(|(idx, centroid)| { + accumulators + .iter() + .for_each(|local_centroids| *centroid += local_centroids[idx]); }); centroids @@ -1553,7 +1593,7 @@ mod tests { use std::iter::repeat_n; use arrow_array::Float16Array; - use arrow_array::types::Float16Type; + use arrow_array::types::{Float16Type, Float32Type, Float64Type}; use half::f16; use lance_arrow::*; use lance_testing::datagen::generate_random_array; @@ -1616,6 +1656,107 @@ mod tests { ); } + #[test] + fn test_recompute_float_centroids() { + let membership = [Some(0), Some(1), None, Some(0), Some(1)]; + + let mut cluster_sizes = [2, 2]; + let kmeans = KMeansAlgoFloat::::to_kmeans( + &[ + f16::from_f32(1.0), + f16::from_f32(3.0), + f16::from_f32(2.0), + f16::from_f32(4.0), + f16::from_f32(100.0), + f16::from_f32(100.0), + f16::from_f32(3.0), + f16::from_f32(5.0), + f16::from_f32(4.0), + f16::from_f32(6.0), + ], + 2, + 2, + &membership, + &mut cluster_sizes, + DistanceType::L2, + 0.0, + ); + assert_eq!( + kmeans.centroids.as_primitive::().values(), + &[ + f16::from_f32(2.0), + f16::from_f32(4.0), + f16::from_f32(3.0), + f16::from_f32(5.0), + ] + ); + + let mut cluster_sizes = [2, 2]; + let kmeans = KMeansAlgoFloat::::to_kmeans( + &[1.0, 3.0, 2.0, 4.0, 100.0, 100.0, 3.0, 5.0, 4.0, 6.0], + 2, + 2, + &membership, + &mut cluster_sizes, + DistanceType::L2, + 0.0, + ); + assert_eq!( + kmeans.centroids.as_primitive::().values(), + &[2.0, 4.0, 3.0, 5.0] + ); + + let mut cluster_sizes = [2, 2]; + let kmeans = KMeansAlgoFloat::::to_kmeans( + &[1.0, 3.0, 2.0, 4.0, 100.0, 100.0, 3.0, 5.0, 4.0, 6.0], + 2, + 2, + &membership, + &mut cluster_sizes, + DistanceType::L2, + 0.0, + ); + assert_eq!( + kmeans.centroids.as_primitive::().values(), + &[2.0, 4.0, 3.0, 5.0] + ); + } + + #[test] + fn test_recompute_centroids_splits_empty_cluster() { + let data = [1.0, 3.0, 3.0, 5.0, 5.0, 7.0, 10.0, 12.0]; + let membership = [Some(0), Some(0), Some(0), Some(1)]; + let mut cluster_sizes = [3, 1, 0]; + let kmeans = KMeansAlgoFloat::::to_kmeans( + &data, + 2, + 3, + &membership, + &mut cluster_sizes, + DistanceType::L2, + 0.0, + ); + + assert!(cluster_sizes.iter().all(|size| *size > 0)); + assert!( + kmeans + .centroids + .as_primitive::() + .values() + .iter() + .all(|value| value.is_finite()) + ); + } + + #[test] + fn test_centroid_accumulator_count_respects_memory_limit() { + assert_eq!(centroid_accumulator_count::(1_000_000, 1024, 16), 16); + assert_eq!( + centroid_accumulator_count::(1_000_000, 16 * 1024 * 1024, 16), + 1 + ); + } + #[tokio::test] async fn test_compute_membership_and_loss() { const DIM: usize = 256; From 295c63d4db66f2b6f0dc269d881e4adf5460e758 Mon Sep 17 00:00:00 2001 From: isaac-dasari <74317395+isaac-dasari@users.noreply.github.com> Date: Mon, 17 Aug 2026 01:15:10 -0700 Subject: [PATCH 2/2] perf(index): preserve kmeans parallelism for large grids --- rust/lance-index/benches/kmeans_recompute.rs | 57 ++--- rust/lance-index/src/vector/kmeans.rs | 233 ++++++++++++++----- 2 files changed, 205 insertions(+), 85 deletions(-) diff --git a/rust/lance-index/benches/kmeans_recompute.rs b/rust/lance-index/benches/kmeans_recompute.rs index 984967930d7..48d9d100ec1 100644 --- a/rust/lance-index/benches/kmeans_recompute.rs +++ b/rust/lance-index/benches/kmeans_recompute.rs @@ -7,36 +7,41 @@ use arrow_array::types::Float32Type; use criterion::{BenchmarkId, Criterion, criterion_group, criterion_main}; use lance_index::vector::kmeans::{KMeansAlgo, KMeansAlgoFloat}; use lance_linalg::distance::DistanceType; -use lance_testing::datagen::generate_random_array; fn bench_recompute_centroids(c: &mut Criterion) { let mut group = c.benchmark_group("kmeans_recompute_centroids"); - let (num_vectors, dimension, k) = (128 * 1024, 128, 256); - let data = generate_random_array(num_vectors * dimension); - let membership = (0..num_vectors) - .map(|row| Some((row % k) as u32)) - .collect::>(); - let cluster_sizes = vec![num_vectors / k; k]; - - group.bench_with_input( - BenchmarkId::new(format!("{dimension}d_{k}k"), num_vectors), - &num_vectors, - |b, _| { - b.iter(|| { - let mut cluster_sizes = cluster_sizes.clone(); - black_box(KMeansAlgoFloat::::to_kmeans( - black_box(data.values()), - dimension, - k, - black_box(&membership), - &mut cluster_sizes, - DistanceType::L2, - 0.0, - )) - }); - }, - ); + let cases = [ + ("input_partitioned", 128 * 1024, 128, 256), + ("large_centroid_grid", 65_536, 1024, 4096), + ]; + + for (name, num_vectors, dimension, k) in cases { + let data = vec![1.0_f32; num_vectors * dimension]; + let membership = (0..num_vectors) + .map(|row| Some((row % k) as u32)) + .collect::>(); + let cluster_sizes = vec![num_vectors / k; k]; + + group.bench_with_input( + BenchmarkId::new(name, format!("{num_vectors}x{dimension}d_{k}k")), + &num_vectors, + |b, _| { + b.iter(|| { + let mut cluster_sizes = cluster_sizes.clone(); + black_box(KMeansAlgoFloat::::to_kmeans( + black_box(&data), + dimension, + k, + black_box(&membership), + &mut cluster_sizes, + DistanceType::L2, + 0.0, + )) + }); + }, + ); + } group.finish(); } diff --git a/rust/lance-index/src/vector/kmeans.rs b/rust/lance-index/src/vector/kmeans.rs index f1a389a3123..c0e75da8d0a 100644 --- a/rust/lance-index/src/vector/kmeans.rs +++ b/rust/lance-index/src/vector/kmeans.rs @@ -33,7 +33,7 @@ use lance_linalg::distance::{DistanceType, Normalize, dot_distance_batch}; use lance_linalg::kernels::{argmin_value_float, argmin_value_float_with_bias}; use log::{info, warn}; use num_traits::One; -use num_traits::{AsPrimitive, Float, FromPrimitive, Num, Zero}; +use num_traits::{AsPrimitive, Float, FromPrimitive, Num}; use rand::prelude::*; use rayon::prelude::*; use { @@ -336,21 +336,129 @@ where // centroid output itself already requires that much memory. const MAX_CENTROID_ACCUMULATOR_BYTES: usize = 64 * 1024 * 1024; -fn centroid_accumulator_count( +#[derive(Debug, PartialEq, Eq)] +enum CentroidRecomputeStrategy { + InputPartitioned { accumulator_count: usize }, + CentroidPartitioned { worker_count: usize }, +} + +fn centroid_recompute_strategy( num_vectors: usize, centroid_len: usize, + k: usize, available_parallelism: usize, -) -> usize { - let accumulator_bytes = centroid_len.saturating_mul(std::mem::size_of::()); - let memory_limited_parallelism = MAX_CENTROID_ACCUMULATOR_BYTES - .checked_div(accumulator_bytes) - .unwrap_or(1) - .max(1); - - available_parallelism - .min(memory_limited_parallelism) - .min(num_vectors) - .max(1) +) -> CentroidRecomputeStrategy { + let available_parallelism = available_parallelism.max(1); + let accumulator_count = if k < 16 { + 1 + } else { + available_parallelism.min(num_vectors.max(1)) + }; + let accumulator_bytes = centroid_len + .saturating_mul(std::mem::size_of::()) + .saturating_mul(accumulator_count); + + if accumulator_count == 1 || accumulator_bytes <= MAX_CENTROID_ACCUMULATOR_BYTES { + CentroidRecomputeStrategy::InputPartitioned { accumulator_count } + } else { + // A smaller accumulator count would couple the memory limit to compute + // parallelism. Retain the original disjoint centroid-owner algorithm + // when full input partitioning would exceed the allocation budget. + let worker_count = if k < available_parallelism || k < 16 { + 1 + } else { + available_parallelism + }; + CentroidRecomputeStrategy::CentroidPartitioned { worker_count } + } +} + +fn recompute_float_centroids( + data: &[T], + dimension: usize, + k: usize, + membership: &[Option], + strategy: CentroidRecomputeStrategy, +) -> Vec +where + T: Float + AddAssign + Send + Sync, +{ + let num_vectors = data.len() / dimension; + let centroid_len = k * dimension; + + match strategy { + CentroidRecomputeStrategy::InputPartitioned { accumulator_count } => { + if num_vectors == 0 { + return vec![T::zero(); centroid_len]; + } + + let vectors_per_chunk = num_vectors.div_ceil(accumulator_count); + + // Each worker scans a disjoint portion of the input once and + // writes to a private centroid buffer. + let mut accumulators = data + .par_chunks(vectors_per_chunk * dimension) + .zip(membership.par_chunks(vectors_per_chunk)) + .map(|(data, membership)| { + let mut local_centroids = vec![T::zero(); centroid_len]; + data.chunks(dimension) + .zip(membership) + .filter_map(|(vector, cluster_id)| { + cluster_id.map(|cluster_id| (vector, cluster_id as usize)) + }) + .for_each(|(vector, cluster_id)| { + let start = cluster_id * dimension; + let centroid = &mut local_centroids[start..start + dimension]; + centroid.iter_mut().zip(vector).for_each(|(c, v)| *c += *v); + }); + local_centroids + }) + .collect::>(); + + let mut centroids = accumulators + .pop() + .unwrap_or_else(|| vec![T::zero(); centroid_len]); + centroids + .par_iter_mut() + .enumerate() + .for_each(|(idx, centroid)| { + accumulators + .iter() + .for_each(|local_centroids| *centroid += local_centroids[idx]); + }); + centroids + } + CentroidRecomputeStrategy::CentroidPartitioned { worker_count } => { + let mut centroids = vec![T::zero(); centroid_len]; + if k == 0 { + return centroids; + } + + let centroids_per_chunk = k / worker_count; + centroids + .par_chunks_mut(dimension * centroids_per_chunk) + .enumerate() + .with_max_len(1) + .for_each(|(chunk_idx, centroids)| { + let start = chunk_idx * centroids_per_chunk; + let end = ((chunk_idx + 1) * centroids_per_chunk).min(k); + data.chunks(dimension) + .zip(membership) + .filter_map(|(vector, cluster_id)| { + cluster_id.map(|cluster_id| (vector, cluster_id as usize)) + }) + .for_each(|(vector, cluster_id)| { + if start <= cluster_id && cluster_id < end { + let local_id = cluster_id - start; + let centroid = &mut centroids + [local_id * dimension..(local_id + 1) * dimension]; + centroid.iter_mut().zip(vector).for_each(|(c, v)| *c += *v); + } + }); + }); + centroids + } + } } impl KMeansAlgo for KMeansAlgoFloat @@ -423,51 +531,13 @@ where ) -> KMeans { let num_vectors = data.len() / dimension; let centroid_len = k * dimension; - let available_parallelism = if k < 16 { - 1 - } else { - get_num_compute_intensive_cpus() - }; - let accumulator_count = centroid_accumulator_count::( + let strategy = centroid_recompute_strategy::( num_vectors, centroid_len, - available_parallelism, + k, + get_num_compute_intensive_cpus(), ); - let vectors_per_chunk = num_vectors.div_ceil(accumulator_count); - - // Each worker scans a disjoint portion of the input once and writes to - // a private centroid buffer. This changes the input-read complexity - // from O(num_vectors * workers) to O(num_vectors) without contention. - let mut accumulators = data - .par_chunks(vectors_per_chunk * dimension) - .zip(membership.par_chunks(vectors_per_chunk)) - .map(|(data, membership)| { - let mut local_centroids = vec![T::Native::zero(); centroid_len]; - data.chunks(dimension) - .zip(membership) - .filter_map(|(vector, cluster_id)| { - cluster_id.map(|cluster_id| (vector, cluster_id as usize)) - }) - .for_each(|(vector, cluster_id)| { - let start = cluster_id * dimension; - let centroid = &mut local_centroids[start..start + dimension]; - centroid.iter_mut().zip(vector).for_each(|(c, v)| *c += *v); - }); - local_centroids - }) - .collect::>(); - - let mut centroids = accumulators - .pop() - .unwrap_or_else(|| vec![T::Native::zero(); centroid_len]); - centroids - .par_iter_mut() - .enumerate() - .for_each(|(idx, centroid)| { - accumulators - .iter() - .for_each(|local_centroids| *centroid += local_centroids[idx]); - }); + let mut centroids = recompute_float_centroids(data, dimension, k, membership, strategy); centroids .par_chunks_mut(dimension) @@ -1749,11 +1819,56 @@ mod tests { } #[test] - fn test_centroid_accumulator_count_respects_memory_limit() { - assert_eq!(centroid_accumulator_count::(1_000_000, 1024, 16), 16); + fn test_centroid_recompute_strategies_produce_same_sums() { + let data = [1.0_f32, 3.0, 2.0, 4.0, 100.0, 100.0, 3.0, 5.0, 4.0, 6.0]; + let membership = [Some(0), Some(1), None, Some(0), Some(1)]; + let expected = [4.0, 8.0, 6.0, 10.0]; + assert_eq!( - centroid_accumulator_count::(1_000_000, 16 * 1024 * 1024, 16), - 1 + recompute_float_centroids( + &data, + 2, + 2, + &membership, + CentroidRecomputeStrategy::InputPartitioned { + accumulator_count: 2, + }, + ), + expected + ); + assert_eq!( + recompute_float_centroids( + &data, + 2, + 2, + &membership, + CentroidRecomputeStrategy::CentroidPartitioned { worker_count: 2 }, + ), + expected + ); + } + + #[test] + fn test_centroid_recompute_strategy_respects_memory_without_losing_parallelism() { + assert_eq!( + centroid_recompute_strategy::(1_000_000, 1024, 256, 16), + CentroidRecomputeStrategy::InputPartitioned { + accumulator_count: 16, + } + ); + assert_eq!( + centroid_recompute_strategy::(1_000_000, 1024 * 1024, 4096, 16), + CentroidRecomputeStrategy::InputPartitioned { + accumulator_count: 16, + } + ); + + // This is the large centroid grid from the performance regression: + // private accumulators would require 992 MiB on 62 workers. Keep all + // workers by partitioning centroids instead of reducing parallelism. + assert_eq!( + centroid_recompute_strategy::(65_536, 4096 * 1024, 4096, 62), + CentroidRecomputeStrategy::CentroidPartitioned { worker_count: 62 } ); }