fix: preserve seeded IVF-PQ training contracts

This commit is contained in:
Gatefixer
2026-08-06 07:35:47 +00:00
parent 845868a343
commit e0d6b9a4fa
3 changed files with 522 additions and 92 deletions
+5
View File
@@ -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",
+6 -1
View File
@@ -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):
+511 -91
View File
@@ -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>,
) -> 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::<Vec<_>>();
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::<UInt64Array>();
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::<Vec<_>>();
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<T>(
fn seeded_assignments<T>(
data: &[T::Native],
centroids: &[T::Native],
dimension: usize,
distance_type: LanceDistanceType,
balance_factor: f32,
cluster_sizes: Option<&[usize]>,
) -> Result<Vec<(usize, f32)>>
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::<usize>() 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<N>(
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::<f32>() < 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<T>(
data: &FixedSizeListArray,
num_centroids: usize,
max_iterations: u32,
distance_type: LanceDistanceType,
balance_factor: f32,
seed: u64,
) -> Result<FixedSizeListArray>
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<T>: From<Vec<T::Native>>,
{
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::<T>().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::<T>().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,
&centroids,
dimension,
)),
LanceDistanceType::Dot => argmin_value_float(dot_distance_batch(
vector,
&centroids,
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::<Result<Vec<_>>>()?;
let iteration_balance_factor = adjusted_balance_factor.min(balance_factor);
let assignments = Self::seeded_assignments::<T>(
data_values,
&centroids,
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::<f64>()
+ 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::<Vec<_>>();
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::<f64>();
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<T>(
data: &FixedSizeListArray,
target_centroids: usize,
max_iterations: u32,
distance_type: LanceDistanceType,
balance_factor: f32,
seed: u64,
) -> Result<FixedSizeListArray>
where
T: ArrowPrimitiveType,
T::Native:
Float + FromPrimitive + AddAssign + DivAssign + MulAssign + Dot + L2 + Send + Sync,
PrimitiveArray<T>: From<Vec<T::Native>>,
{
#[derive(Clone)]
struct Cluster<N> {
id: usize,
indices: Vec<usize>,
centroid: Vec<N>,
finalized: bool,
}
impl<N> Eq for Cluster<N> {}
impl<N> PartialEq for Cluster<N> {
fn eq(&self, other: &Self) -> bool {
self.indices.len() == other.indices.len() && self.id == other.id
}
}
impl<N> Ord for Cluster<N> {
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<N> PartialOrd for Cluster<N> {
fn partial_cmp(&self, other: &Self) -> Option<Ordering> {
Some(self.cmp(other))
}
}
const HIERARCHICAL_K: usize = 16;
let dimension = data.value_length() as usize;
let data_values = data.values().as_primitive::<T>().values();
let initial_k = HIERARCHICAL_K.min(target_centroids).min(data.len());
let initial = Self::train_seeded_flat_kmeans_typed::<T>(
data,
initial_k,
max_iterations,
distance_type,
balance_factor,
Self::seeded_domain(seed, 0x4849_4552_5f49_4e49),
)?;
let initial_values = initial.values().as_primitive::<T>().values();
let assignments = Self::seeded_assignments::<T>(
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::<Vec<_>>();
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::<T>::from(sub_values),
dimension as i32,
)?;
let sub_centroids = Self::train_seeded_flat_kmeans_typed::<T>(
&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::<T>().values();
let sub_values = sub_data.values().as_primitive::<T>().values();
let sub_assignments = Self::seeded_assignments::<T>(
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::<Vec<_>>();
Ok(FixedSizeListArray::try_new_from_values(
PrimitiveArray::<T>::from(centroids),
dimension as i32,
)?)
}
fn train_seeded_kmeans_typed<T>(
data: &FixedSizeListArray,
num_centroids: usize,
max_iterations: u32,
distance_type: LanceDistanceType,
balance_factor: f32,
seed: u64,
) -> Result<FixedSizeListArray>
where
T: ArrowPrimitiveType,
T::Native:
Float + FromPrimitive + AddAssign + DivAssign + MulAssign + Dot + L2 + Send + Sync,
PrimitiveArray<T>: From<Vec<T::Native>>,
{
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::<T>(
data,
num_centroids,
max_iterations,
distance_type,
balance_factor,
seed,
)
} else {
Self::train_seeded_flat_kmeans_typed::<T>(
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<FixedSizeListArray> {
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::<Float32Type>(
@@ -473,6 +840,7 @@ impl NativeTable {
num_centroids,
max_iterations,
distance_type,
balance_factor,
seed,
),
DataType::Float64 => Self::train_seeded_kmeans_typed::<Float64Type>(
@@ -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<ArrayRef>
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<T>: From<Vec<T::Native>>,
{
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::<T>().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();