Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion .github/workflows/rust-benchmark.yml
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
4 changes: 4 additions & 0 deletions rust/lance-index/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -143,6 +143,10 @@ harness = false
name = "kmeans"
harness = false

[[bench]]
name = "kmeans_recompute"
harness = false

[[bench]]
name = "compute_partition"
harness = false
Expand Down
53 changes: 53 additions & 0 deletions rust/lance-index/benches/kmeans_recompute.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,53 @@
// 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;

fn bench_recompute_centroids(c: &mut Criterion) {
let mut group = c.benchmark_group("kmeans_recompute_centroids");

let cases = [
("default_ivf_high_dim", 16_384, 1024, 64),
("low_sample_high_dim", 512, 1024, 256),
("default_pq_subvector", 65_536, 64, 256),
("max_sample_low_dim", 128 * 1024, 128, 256),
("large_incremental_ivf", 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::<Vec<_>>();
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::<Float32Type>::to_kmeans(
black_box(&data),
dimension,
k,
black_box(&membership),
&mut cluster_sizes,
DistanceType::L2,
0.0,
))
});
},
);
}

group.finish();
}

criterion_group!(benches, bench_recompute_centroids);
criterion_main!(benches);
263 changes: 232 additions & 31 deletions rust/lance-index/src/vector/kmeans.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -331,6 +331,110 @@ where
phantom_data: std::marker::PhantomData<T>,
}

struct MembershipRowIndex {
/// Prefix offsets into `rows`; cluster `i` owns `offsets[i]..offsets[i + 1]`.
offsets: Vec<usize>,
/// Input row numbers grouped by cluster and stable within each cluster.
rows: Vec<usize>,
}

/// Group input rows by their assigned cluster in linear time.
///
/// This counting-sort layout lets centroid owners read each assigned vector
/// without rescanning every membership. Keeping rows stable within a cluster
/// also preserves the prior floating-point accumulation order.
fn build_membership_row_index(
membership: &[Option<u32>],
num_vectors: usize,
k: usize,
) -> MembershipRowIndex {
let membership = &membership[..membership.len().min(num_vectors)];
let mut offsets = vec![0; k + 1];
membership.iter().flatten().for_each(|&cluster_id| {
let cluster_id = cluster_id as usize;
if cluster_id < k {
offsets[cluster_id + 1] += 1;
}
});
for cluster_id in 0..k {
offsets[cluster_id + 1] += offsets[cluster_id];
}

let mut next_offsets = offsets[..k].to_vec();
let mut rows = vec![0; offsets[k]];
membership
.iter()
.enumerate()
.filter_map(|(row, cluster_id)| cluster_id.map(|cluster_id| (row, cluster_id as usize)))
.for_each(|(row, cluster_id)| {
if cluster_id < k {
rows[next_offsets[cluster_id]] = row;
next_offsets[cluster_id] += 1;
}
});

MembershipRowIndex { offsets, rows }
}

fn recompute_float_centroids<T>(
data: &[T],
dimension: usize,
k: usize,
membership: &[Option<u32>],
available_parallelism: usize,
) -> Vec<T>
where
T: Float + AddAssign + Send + Sync,
{
let num_vectors = data.len() / dimension;
let centroid_len = k * dimension;
let mut centroids = vec![T::zero(); centroid_len];
if k == 0 {
return centroids;
}

let available_parallelism = available_parallelism.max(1);
if available_parallelism == 1 || k < available_parallelism || k < 16 {
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 cluster_id < k {
let start = cluster_id * dimension;
let centroid = &mut centroids[start..start + dimension];
centroid.iter_mut().zip(vector).for_each(|(c, v)| *c += *v);
}
});
return centroids;
}

let row_index = build_membership_row_index(membership, num_vectors, k);
let centroids_per_chunk = k / available_parallelism;
centroids
.par_chunks_mut(dimension * centroids_per_chunk)
.enumerate()
.with_max_len(1)
.for_each(|(chunk_idx, centroids)| {
let first_cluster = chunk_idx * centroids_per_chunk;
centroids
.chunks_mut(dimension)
.enumerate()
.for_each(|(local_cluster, centroid)| {
let cluster_id = first_cluster + local_cluster;
row_index.rows
[row_index.offsets[cluster_id]..row_index.offsets[cluster_id + 1]]
.iter()
.for_each(|&row| {
let vector = &data[row * dimension..(row + 1) * dimension];
centroid.iter_mut().zip(vector).for_each(|(c, v)| *c += *v);
});
});
});
centroids
}

impl<T: ArrowNumericType> KMeansAlgo<T::Native> for KMeansAlgoFloat<T>
where
T::Native: Float + Dot + L2 + MulAssign + DivAssign + AddAssign + FromPrimitive + Sync,
Expand Down Expand Up @@ -399,35 +503,13 @@ 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);
data.chunks(dimension)
.zip(membership.iter())
.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 mut centroids = recompute_float_centroids(
data,
dimension,
k,
membership,
get_num_compute_intensive_cpus(),
);

centroids
.par_chunks_mut(dimension)
Expand Down Expand Up @@ -1553,7 +1635,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;
Expand Down Expand Up @@ -1616,6 +1698,125 @@ 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::<Float16Type>::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::<Float16Type>().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::<Float32Type>::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::<Float32Type>().values(),
&[2.0, 4.0, 3.0, 5.0]
);

let mut cluster_sizes = [2, 2];
let kmeans = KMeansAlgoFloat::<Float64Type>::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::<Float64Type>().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::<Float32Type>::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::<Float32Type>()
.values()
.iter()
.all(|value| value.is_finite())
);
}

#[test]
fn test_membership_row_index_is_stable() {
let row_index =
build_membership_row_index(&[Some(2), Some(0), None, Some(2), Some(1), Some(0)], 5, 3);

assert_eq!(row_index.offsets, [0, 1, 2, 4]);
assert_eq!(row_index.rows, [1, 4, 0, 3]);
}

#[test]
fn test_indexed_recompute_matches_serial_accumulation() {
let dimension = 2;
let k = 17;
let num_vectors = 65;
let data = (0..num_vectors * dimension)
.map(|value| value as f32)
.collect::<Vec<_>>();
let membership = (0..num_vectors)
.map(|row| (row != 17).then_some((row % k) as u32))
.collect::<Vec<_>>();

let serial = recompute_float_centroids(&data, dimension, k, &membership, 1);
let indexed = recompute_float_centroids(&data, dimension, k, &membership, 4);

assert_eq!(indexed, serial);
}

#[tokio::test]
async fn test_compute_membership_and_loss() {
const DIM: usize = 256;
Expand Down
Loading