From e0d6b9a4fa48cdfd34250b2b81edf6aac6bfe847 Mon Sep 17 00:00:00 2001 From: Gatefixer <313497061+lancedb-gatefixer[bot]@users.noreply.github.com> Date: Thu, 6 Aug 2026 07:35:47 +0000 Subject: [PATCH] fix: preserve seeded IVF-PQ training contracts --- python/python/lancedb/table.py | 5 + python/python/tests/test_index.py | 7 +- rust/lancedb/src/table/create_index.rs | 602 +++++++++++++++++++++---- 3 files changed, 522 insertions(+), 92 deletions(-) diff --git a/python/python/lancedb/table.py b/python/python/lancedb/table.py index ae36bac7a..ba230accc 100644 --- a/python/python/lancedb/table.py +++ b/python/python/lancedb/table.py @@ -2762,6 +2762,11 @@ class LanceTable(Table): if config is not None and hasattr(config, "accelerator"): acc = getattr(config, "accelerator", None) if acc is not None: + if isinstance(config, IvfPq) and config.seed is not None: + raise ValueError( + "IvfPq seed is not supported with accelerator-based " + "index training" + ) # Dispatch to pylance for GPU acceleration index_type_map = { "IvfFlat": "IVF_FLAT", diff --git a/python/python/tests/test_index.py b/python/python/tests/test_index.py index f0b2fee78..5f29ae8ee 100644 --- a/python/python/tests/test_index.py +++ b/python/python/tests/test_index.py @@ -26,7 +26,7 @@ from lancedb.index import ( HnswFlat, FTS, ) -from lancedb.table import IndexStatistics +from lancedb.table import IndexStatistics, LanceTable @pytest_asyncio.fixture @@ -399,6 +399,11 @@ def test_seeded_ivfpq_rejects_accelerator(): with pytest.raises(ValueError, match="seed is not supported with accelerator"): IvfPq(seed=42, accelerator="cuda") + config = IvfPq(seed=42) + config.accelerator = "cuda" + with pytest.raises(ValueError, match="seed is not supported with accelerator"): + object.__new__(LanceTable).create_index("vector", config=config) + @pytest.mark.asyncio async def test_create_ivfrq_index(some_table: AsyncTable): diff --git a/rust/lancedb/src/table/create_index.rs b/rust/lancedb/src/table/create_index.rs index 39c25a8d2..4a0b3d7f6 100644 --- a/rust/lancedb/src/table/create_index.rs +++ b/rust/lancedb/src/table/create_index.rs @@ -8,7 +8,9 @@ //! [`BaseTable::create_index`](super::BaseTable::create_index) implementation //! on `NativeTable`. -use std::ops::{AddAssign, DivAssign}; +use std::cmp::Ordering; +use std::collections::BinaryHeap; +use std::ops::{AddAssign, DivAssign, MulAssign}; use std::sync::Arc; use arrow_array::cast::AsArray; @@ -33,11 +35,11 @@ use lance_index::vector::sq::builder::SQBuildParams; use lance_linalg::distance::{ DistanceType as LanceDistanceType, Dot, L2, dot_distance_batch, l2_distance_batch, }; -use lance_linalg::kernels::{argmin_value_float, normalize_fsl_owned}; +use lance_linalg::kernels::{argmin_value_float_with_bias, normalize_fsl_owned}; use num_traits::{Float, FromPrimitive, Zero}; -use rand::SeedableRng; use rand::rngs::SmallRng; use rand::seq::index::sample; +use rand::{Rng, SeedableRng}; use rayon::prelude::*; use crate::error::{Error, Result}; @@ -239,6 +241,23 @@ impl NativeTable { }) } + fn seeded_sampling_batch_rows( + remaining_vectors: usize, + is_multivector: bool, + vectors_per_row: Option, + ) -> usize { + const MAX_ROWS_PER_TAKE: usize = 8192; + const MIN_MULTIVECTOR_ROWS_PER_TAKE: usize = 128; + + if !is_multivector { + return remaining_vectors.min(MAX_ROWS_PER_TAKE); + } + vectors_per_row + .map(|estimate| remaining_vectors.div_ceil(estimate.max(1))) + .unwrap_or(MIN_MULTIVECTOR_ROWS_PER_TAKE) + .clamp(MIN_MULTIVECTOR_ROWS_PER_TAKE, MAX_ROWS_PER_TAKE) + } + /// Select training vectors in a stable pseudo-random order. /// /// The row permutation is continued when null or non-finite vectors are @@ -258,13 +277,22 @@ impl NativeTable { } let projection = Arc::new(dataset.schema().project(&[column])?); + let is_multivector = matches!( + projection.field(column).map(|field| field.data_type()), + Some(DataType::List(_)) + ); let mut permutation = SeededRowPermutation::new(num_rows, seed)?; let mut arrays = Vec::new(); let mut sampled_vectors = 0; - const TAKE_BATCH_SIZE: usize = 8192; + let mut vectors_per_row = None; while sampled_vectors < sample_size { - let rows_to_read = (sample_size - sampled_vectors).min(TAKE_BATCH_SIZE); + let remaining_vectors = sample_size - sampled_vectors; + let rows_to_read = Self::seeded_sampling_batch_rows( + remaining_vectors, + is_multivector, + vectors_per_row, + ); let indices = permutation.by_ref().take(rows_to_read).collect::>(); if indices.is_empty() { break; @@ -276,10 +304,22 @@ impl NativeTable { message: format!("Vector column `{column}` missing from sampled batch"), })?; let sampled = Self::flatten_vector_array(array)?; + if is_multivector && !sampled.is_empty() { + vectors_per_row = Some(sampled.len().div_ceil(indices.len())); + } let sampled = filter_finite_training_data(sampled)?; - sampled_vectors += sampled.len(); if !sampled.is_empty() { - arrays.push(sampled); + let retained = sampled.len().min(remaining_vectors); + let retained = if retained == sampled.len() { + sampled + } else { + let indices = (0..retained as u64).collect::(); + arrow_select::take::take(&sampled, &indices, None)? + .as_fixed_size_list() + .clone() + }; + sampled_vectors += retained.len(); + arrays.push(retained); } } @@ -293,8 +333,7 @@ impl NativeTable { .map(|array| array as &dyn Array) .collect::>(); let sampled = concat(&array_refs)?; - let sampled = sampled.as_fixed_size_list(); - Ok(sampled.slice(0, sampled.len().min(sample_size))) + Ok(sampled.as_fixed_size_list().clone()) } fn seeded_initial_centroids( @@ -320,67 +359,181 @@ impl NativeTable { .clone()) } - fn train_seeded_kmeans_typed( + fn seeded_assignments( + data: &[T::Native], + centroids: &[T::Native], + dimension: usize, + distance_type: LanceDistanceType, + balance_factor: f32, + cluster_sizes: Option<&[usize]>, + ) -> Result> + where + T: ArrowPrimitiveType, + T::Native: Float + Dot + L2 + Send + Sync, + { + data.par_chunks(dimension) + .map(|vector| { + let bias = || { + cluster_sizes + .map(|sizes| sizes.iter().map(|size| balance_factor * *size as f32)) + }; + let nearest = match distance_type { + LanceDistanceType::L2 => argmin_value_float_with_bias( + l2_distance_batch(vector, centroids, dimension), + bias(), + ), + LanceDistanceType::Dot => argmin_value_float_with_bias( + dot_distance_batch(vector, centroids, dimension), + bias(), + ), + distance_type => { + return Err(Error::InvalidInput { + message: format!( + "Distance type {distance_type} is not supported for seeded kmeans" + ), + }); + } + }; + nearest + .map(|(cluster, distance)| (cluster as usize, distance)) + .ok_or_else(|| Error::InvalidInput { + message: "Could not assign a vector during seeded kmeans".to_string(), + }) + }) + .collect() + } + + fn seeded_cluster_statistics(assignments: &[(usize, f32)], cluster_sizes: &mut [usize]) -> f32 { + cluster_sizes.fill(0); + let mut radii = vec![0.0_f32; cluster_sizes.len()]; + let mut losses = vec![0.0_f64; cluster_sizes.len()]; + let mut max_cluster = 0; + let mut max_cluster_size = 0; + for (cluster, distance) in assignments { + cluster_sizes[*cluster] += 1; + radii[*cluster] = radii[*cluster].max(*distance); + losses[*cluster] += *distance as f64; + if cluster_sizes[*cluster] > max_cluster_size { + max_cluster = *cluster; + max_cluster_size = cluster_sizes[*cluster]; + } + } + if max_cluster_size == 0 { + return 0.0; + } + (radii[max_cluster] - losses[max_cluster] as f32 / max_cluster_size as f32) + / assignments.len() as f32 + } + + fn seeded_balance_loss( + cluster_sizes: &[usize], + num_vectors: usize, + balance_factor: f32, + ) -> f32 { + let size_loss = cluster_sizes.iter().map(|size| size.pow(2)).sum::() as f32; + balance_factor * (size_loss - num_vectors.pow(2) as f32 / cluster_sizes.len() as f32) + } + + /// Lance-compatible empty-cluster splitting with an injected RNG. + fn split_seeded_empty_clusters( + num_vectors: usize, + cluster_sizes: &mut [usize], + centroids: &mut [N], + dimension: usize, + rng: &mut SmallRng, + ) where + N: Float + MulAssign, + { + let epsilon = N::from(1.0 / 1024.0).unwrap(); + for cluster in 0..cluster_sizes.len() { + if cluster_sizes[cluster] != 0 { + continue; + } + let mut donor = 0; + loop { + let probability = (cluster_sizes[donor] as f32 - 1.0) + / (num_vectors - cluster_sizes.len()) as f32; + if rng.random::() < probability { + break; + } + donor = (donor + 1) % cluster_sizes.len(); + } + + cluster_sizes[cluster] = cluster_sizes[donor] / 2; + cluster_sizes[donor] -= cluster_sizes[cluster]; + for value in 0..dimension { + if value % 2 == 0 { + centroids[cluster * dimension + value] = + centroids[donor * dimension + value] * (N::one() + epsilon); + centroids[donor * dimension + value] *= N::one() - epsilon; + } else { + centroids[cluster * dimension + value] = + centroids[donor * dimension + value] * (N::one() - epsilon); + centroids[donor * dimension + value] *= N::one() + epsilon; + } + } + } + } + + fn seeded_domain(seed: u64, domain: u64) -> u64 { + SeededRowPermutation::round(domain, seed, 0) + } + + fn train_seeded_flat_kmeans_typed( data: &FixedSizeListArray, num_centroids: usize, max_iterations: u32, distance_type: LanceDistanceType, + balance_factor: f32, seed: u64, ) -> Result where T: ArrowPrimitiveType, - T::Native: Float + FromPrimitive + AddAssign + DivAssign + Dot + L2 + Send + Sync, + T::Native: + Float + FromPrimitive + AddAssign + DivAssign + MulAssign + Dot + L2 + Send + Sync, PrimitiveArray: From>, { - if max_iterations == 0 { - return Err(Error::InvalidInput { - message: "max_iterations must be greater than zero for seeded IVF PQ training" - .to_string(), - }); - } + // Match Lance's per-kmeans cap. Hierarchical IVF recursively trains + // small clusters, while PQ remains on the flat 256-centroid path. + let data = if data.len() >= num_centroids * 512 { + data.slice(0, num_centroids * 512) + } else { + data.clone() + }; let dimension = data.value_length() as usize; let data_values = data.values().as_primitive::().values(); - let initial = Self::seeded_initial_centroids(data, num_centroids, seed)?; + let initial = Self::seeded_initial_centroids( + &data, + num_centroids, + Self::seeded_domain(seed, 0x494e_4954), + )?; let mut centroids = initial.values().as_primitive::().values().to_vec(); let mut previous_loss = f64::MAX; + let mut cluster_sizes = vec![0usize; num_centroids]; + let mut adjusted_balance_factor = f32::MAX; + let mut split_rng = + SmallRng::seed_from_u64(Self::seeded_domain(seed, 0x454d_5054_595f_5350)); for _ in 0..max_iterations { - // Seeded builds use exact assignment. Lance's large-centroid - // acceleration builds an approximate HNSW graph whose parallel - // insertion order is intentionally not deterministic. - let assignments = data_values - .par_chunks(dimension) - .map(|vector| { - let nearest = match distance_type { - LanceDistanceType::L2 => argmin_value_float(l2_distance_batch( - vector, - ¢roids, - dimension, - )), - LanceDistanceType::Dot => argmin_value_float(dot_distance_batch( - vector, - ¢roids, - dimension, - )), - distance_type => { - return Err(Error::InvalidInput { - message: format!( - "Distance type {distance_type} is not supported for seeded kmeans" - ), - }); - } - }; - nearest - .map(|(cluster, distance)| (cluster as usize, distance)) - .ok_or_else(|| Error::InvalidInput { - message: "Could not assign a vector during seeded kmeans".to_string(), - }) - }) - .collect::>>()?; + let iteration_balance_factor = adjusted_balance_factor.min(balance_factor); + let assignments = Self::seeded_assignments::( + data_values, + ¢roids, + dimension, + distance_type, + iteration_balance_factor, + Some(&cluster_sizes), + )?; + adjusted_balance_factor = + Self::seeded_cluster_statistics(&assignments, &mut cluster_sizes); + let loss = assignments + .iter() + .map(|(_, distance)| *distance as f64) + .sum::() + + Self::seeded_balance_loss(&cluster_sizes, data.len(), iteration_balance_factor) + as f64; let mut next_centroids = vec![T::Native::zero(); num_centroids * dimension]; - let mut cluster_sizes = vec![0usize; num_centroids]; for (row, (cluster, _)) in assignments.iter().enumerate() { - cluster_sizes[*cluster] += 1; let vector = &data_values[row * dimension..(row + 1) * dimension]; let centroid = &mut next_centroids[*cluster * dimension..(*cluster + 1) * dimension]; @@ -400,45 +553,14 @@ impl NativeTable { } } - let empty_clusters = cluster_sizes.iter().filter(|size| **size == 0).count(); - if empty_clusters > 0 { - // Lance's ordinary trainer repairs empty clusters using OS - // randomness. Select only the required farthest rows in linear - // expected time, then sort that bounded top-k for stable ties. - let mut replacement_rows = assignments - .iter() - .enumerate() - .map(|(row, (_, distance))| (row, *distance)) - .collect::>(); - let by_priority = |left: &(usize, f32), right: &(usize, f32)| { - right - .1 - .total_cmp(&left.1) - .then_with(|| left.0.cmp(&right.0)) - }; - if empty_clusters < replacement_rows.len() { - replacement_rows.select_nth_unstable_by(empty_clusters, by_priority); - replacement_rows.truncate(empty_clusters); - } - replacement_rows.sort_by(by_priority); - let mut replacements = replacement_rows.into_iter(); - for (cluster, cluster_size) in cluster_sizes.iter_mut().enumerate() { - if *cluster_size == 0 { - let (row, _) = replacements.next().ok_or_else(|| Error::InvalidInput { - message: "Could not repair an empty seeded kmeans cluster".to_string(), - })?; - let vector = &data_values[row * dimension..(row + 1) * dimension]; - next_centroids[cluster * dimension..(cluster + 1) * dimension] - .copy_from_slice(vector); - *cluster_size = 1; - } - } - } + Self::split_seeded_empty_clusters( + data.len(), + &mut cluster_sizes, + &mut next_centroids, + dimension, + &mut split_rng, + ); - let loss = assignments - .iter() - .map(|(_, distance)| *distance as f64) - .sum::(); let converged = (previous_loss - loss).abs() < 1e-4 * loss; centroids = next_centroids; previous_loss = loss; @@ -453,11 +575,255 @@ impl NativeTable { )?) } + /// Seeded counterpart to Lance's balanced, 16-way hierarchical trainer. + fn train_seeded_hierarchical_kmeans_typed( + data: &FixedSizeListArray, + target_centroids: usize, + max_iterations: u32, + distance_type: LanceDistanceType, + balance_factor: f32, + seed: u64, + ) -> Result + where + T: ArrowPrimitiveType, + T::Native: + Float + FromPrimitive + AddAssign + DivAssign + MulAssign + Dot + L2 + Send + Sync, + PrimitiveArray: From>, + { + #[derive(Clone)] + struct Cluster { + id: usize, + indices: Vec, + centroid: Vec, + finalized: bool, + } + + impl Eq for Cluster {} + + impl PartialEq for Cluster { + fn eq(&self, other: &Self) -> bool { + self.indices.len() == other.indices.len() && self.id == other.id + } + } + + impl Ord for Cluster { + fn cmp(&self, other: &Self) -> Ordering { + match (self.finalized, other.finalized) { + (false, true) => Ordering::Greater, + (true, false) => Ordering::Less, + _ => self + .indices + .len() + .cmp(&other.indices.len()) + .then_with(|| other.id.cmp(&self.id)), + } + } + } + + impl PartialOrd for Cluster { + fn partial_cmp(&self, other: &Self) -> Option { + Some(self.cmp(other)) + } + } + + const HIERARCHICAL_K: usize = 16; + let dimension = data.value_length() as usize; + let data_values = data.values().as_primitive::().values(); + let initial_k = HIERARCHICAL_K.min(target_centroids).min(data.len()); + let initial = Self::train_seeded_flat_kmeans_typed::( + data, + initial_k, + max_iterations, + distance_type, + balance_factor, + Self::seeded_domain(seed, 0x4849_4552_5f49_4e49), + )?; + let initial_values = initial.values().as_primitive::().values(); + let assignments = Self::seeded_assignments::( + data_values, + initial_values, + dimension, + distance_type, + 0.0, + None, + )?; + + let mut heap = BinaryHeap::new(); + let mut next_cluster_id = 0; + for cluster in 0..initial_k { + let indices = assignments + .iter() + .enumerate() + .filter_map(|(row, (assigned, _))| (*assigned == cluster).then_some(row)) + .collect::>(); + if !indices.is_empty() { + heap.push(Cluster { + id: next_cluster_id, + indices, + centroid: initial_values[cluster * dimension..(cluster + 1) * dimension] + .to_vec(), + finalized: false, + }); + next_cluster_id += 1; + } + } + + while heap.len() < target_centroids { + let mut largest = heap.pop().ok_or_else(|| Error::InvalidInput { + message: "No seeded kmeans cluster can be further split".to_string(), + })?; + if largest.finalized || largest.indices.len() <= 1 { + largest.finalized = true; + heap.push(largest); + if heap.iter().all(|cluster| cluster.finalized) { + break; + } + continue; + } + + let remaining_centroids = target_centroids - heap.len(); + let cluster_k = if largest.indices.len() <= HIERARCHICAL_K { + 2.min(remaining_centroids).min(largest.indices.len()) + } else { + (largest.indices.len() / HIERARCHICAL_K) + .min(remaining_centroids) + .clamp(2, HIERARCHICAL_K) + }; + let mut sub_values = Vec::with_capacity(largest.indices.len() * dimension); + for row in &largest.indices { + sub_values + .extend_from_slice(&data_values[*row * dimension..(*row + 1) * dimension]); + } + let sub_data = FixedSizeListArray::try_new_from_values( + PrimitiveArray::::from(sub_values), + dimension as i32, + )?; + let sub_centroids = Self::train_seeded_flat_kmeans_typed::( + &sub_data, + cluster_k, + max_iterations, + distance_type, + balance_factor, + Self::seeded_domain(seed, largest.id as u64 + 1), + )?; + let sub_centroid_values = sub_centroids.values().as_primitive::().values(); + let sub_values = sub_data.values().as_primitive::().values(); + let sub_assignments = Self::seeded_assignments::( + sub_values, + sub_centroid_values, + dimension, + distance_type, + 0.0, + None, + )?; + let mut cluster_assignments = vec![Vec::new(); cluster_k]; + for (local_row, (cluster, _)) in sub_assignments.iter().enumerate() { + cluster_assignments[*cluster].push(largest.indices[local_row]); + } + let non_empty_clusters = cluster_assignments + .iter() + .filter(|indices| !indices.is_empty()) + .count(); + if non_empty_clusters <= 1 { + largest.finalized = true; + heap.push(largest); + continue; + } + + for (cluster, indices) in cluster_assignments.into_iter().enumerate() { + if indices.is_empty() { + continue; + } + heap.push(Cluster { + id: next_cluster_id, + indices, + centroid: sub_centroid_values[cluster * dimension..(cluster + 1) * dimension] + .to_vec(), + finalized: false, + }); + next_cluster_id += 1; + } + } + + if heap.len() < target_centroids { + return Err(Error::InvalidInput { + message: format!( + "Cannot create {target_centroids} IVF partitions: seeded kmeans could only form {} non-empty clusters from {} training vectors", + heap.len(), + data.len() + ), + }); + } + let mut clusters = heap.into_vec(); + clusters.sort_by_key(|cluster| cluster.id); + let centroids = clusters + .into_iter() + .flat_map(|cluster| cluster.centroid) + .collect::>(); + Ok(FixedSizeListArray::try_new_from_values( + PrimitiveArray::::from(centroids), + dimension as i32, + )?) + } + + fn train_seeded_kmeans_typed( + data: &FixedSizeListArray, + num_centroids: usize, + max_iterations: u32, + distance_type: LanceDistanceType, + balance_factor: f32, + seed: u64, + ) -> Result + where + T: ArrowPrimitiveType, + T::Native: + Float + FromPrimitive + AddAssign + DivAssign + MulAssign + Dot + L2 + Send + Sync, + PrimitiveArray: From>, + { + if max_iterations == 0 { + return Err(Error::InvalidInput { + message: "max_iterations must be greater than zero for seeded IVF PQ training" + .to_string(), + }); + } + if data.len() < num_centroids { + return Err(Error::InvalidInput { + message: format!( + "Not enough valid vectors to train {num_centroids} centroids; only {} are available", + data.len() + ), + }); + } + // Lance scales the IVF balance factor by the full training set and + // switches to hierarchical training above 256 centroids. + let balance_factor = balance_factor / data.len() as f32; + if num_centroids > 256 { + Self::train_seeded_hierarchical_kmeans_typed::( + data, + num_centroids, + max_iterations, + distance_type, + balance_factor, + seed, + ) + } else { + Self::train_seeded_flat_kmeans_typed::( + data, + num_centroids, + max_iterations, + distance_type, + balance_factor, + seed, + ) + } + } + fn train_seeded_kmeans( data: &FixedSizeListArray, num_centroids: usize, max_iterations: u32, distance_type: LanceDistanceType, + balance_factor: f32, seed: u64, ) -> Result { match data.value_type() { @@ -466,6 +832,7 @@ impl NativeTable { num_centroids, max_iterations, distance_type, + balance_factor, seed, ), DataType::Float32 => Self::train_seeded_kmeans_typed::( @@ -473,6 +840,7 @@ impl NativeTable { num_centroids, max_iterations, distance_type, + balance_factor, seed, ), DataType::Float64 => Self::train_seeded_kmeans_typed::( @@ -480,6 +848,7 @@ impl NativeTable { num_centroids, max_iterations, distance_type, + balance_factor, seed, ), data_type => Err(Error::InvalidInput { @@ -499,7 +868,8 @@ impl NativeTable { ) -> Result where T: ArrowPrimitiveType, - T::Native: Float + FromPrimitive + AddAssign + DivAssign + Dot + L2 + Send + Sync, + T::Native: + Float + FromPrimitive + AddAssign + DivAssign + MulAssign + Dot + L2 + Send + Sync, PrimitiveArray: From>, { let dimension = data.value_length() as usize; @@ -530,6 +900,7 @@ impl NativeTable { num_centroids, max_iterations, LanceDistanceType::L2, + 0.0, seed.wrapping_add((sub_vector as u64).wrapping_mul(0x9e37_79b9_7f4a_7c15)), )?; codebook.extend_from_slice(centroids.values().as_primitive::().values()); @@ -683,6 +1054,7 @@ impl NativeTable { num_partitions, index.max_iterations, metric_type, + 1.0, seed ^ Self::IVF_INIT_SEED_SALT, )?; ivf_params.centroids = Some(Arc::new(ivf_centroids.clone())); @@ -1549,6 +1921,54 @@ mod tests { assert!(offsets.iter().any(|offset| *offset > 800.0)); } + #[test] + fn test_seeded_multivector_sampling_batch_is_bounded() { + assert_eq!( + super::NativeTable::seeded_sampling_batch_rows(65_536, true, None), + 128 + ); + assert_eq!( + super::NativeTable::seeded_sampling_batch_rows(65_536, true, Some(512)), + 128 + ); + assert_eq!( + super::NativeTable::seeded_sampling_batch_rows(65_536, false, None), + 8192 + ); + } + + #[test] + fn test_seeded_kmeans_hierarchical_path_is_reproducible() { + const NUM_VECTORS: usize = 4096; + const DIMENSION: usize = 2; + const NUM_CENTROIDS: usize = 257; + let values = Float32Array::from_iter_values((0..NUM_VECTORS).flat_map(|row| { + let row = row as f32; + [row, (row * 0.618_034).fract() * 1000.0] + })); + let vectors = create_fixed_size_list(values, DIMENSION as i32).unwrap(); + let first = super::NativeTable::train_seeded_kmeans( + &vectors, + NUM_CENTROIDS, + 3, + super::LanceDistanceType::L2, + 1.0, + 42, + ) + .unwrap(); + let second = super::NativeTable::train_seeded_kmeans( + &vectors, + NUM_CENTROIDS, + 3, + super::LanceDistanceType::L2, + 1.0, + 42, + ) + .unwrap(); + assert_eq!(first.len(), NUM_CENTROIDS); + assert_eq!(first, second); + } + #[tokio::test] async fn test_seeded_ivf_pq_training_supports_multivectors() { let tmp_dir = tempdir().unwrap();