Skip to content
Open
1 change: 1 addition & 0 deletions java/lance-jni/src/utils.rs
Original file line number Diff line number Diff line change
Expand Up @@ -468,6 +468,7 @@ pub fn get_vector_index_params(
kmeans_redos,
codebook,
sample_rate,
kmeans_seed: None,
})
},
)?;
Expand Down
1 change: 1 addition & 0 deletions java/lance-jni/src/vector_trainer.rs
Original file line number Diff line number Diff line change
Expand Up @@ -78,6 +78,7 @@ fn build_pq_params_from_java(
kmeans_redos,
codebook: None,
sample_rate,
kmeans_seed: None,
})
}

Expand Down
7 changes: 7 additions & 0 deletions rust/lance-index/src/vector/ivf/builder.rs
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,12 @@ pub struct IvfBuildParams {

pub sample_rate: usize,

/// Optional seed for k-means centroid initialization.
///
/// `None` initializes from OS entropy, while `Some(seed)` makes centroid
/// selection reproducible for the same training data and parameters.
pub kmeans_seed: Option<u64>,

/// Optional per-step sample rate for streaming IVF kmeans training.
///
/// When set, IVF training loads at most `num_partitions * streaming_sample_rate`
Expand Down Expand Up @@ -94,6 +100,7 @@ impl Default for IvfBuildParams {
centroids: None,
retrain: false,
sample_rate: 256, // See faiss
kmeans_seed: None,
streaming_sample_rate: None,
streaming_coreset_rate: None,
streaming_refine_passes: 0,
Expand Down
46 changes: 44 additions & 2 deletions rust/lance-index/src/vector/kmeans.rs
Original file line number Diff line number Diff line change
Expand Up @@ -92,6 +92,12 @@ pub struct KMeansParams {

/// Optional sync callback for iteration progress: (current_iteration, max_iterations).
pub on_progress: Option<Arc<dyn Fn(u32, u32) + Send + Sync>>,

/// Optional seed for random centroid initialization.
///
/// `None` initializes from OS entropy, while `Some(seed)` makes centroid
/// selection reproducible for the same training data and parameters.
pub seed: Option<u64>,
}

impl std::fmt::Debug for KMeansParams {
Expand All @@ -105,6 +111,7 @@ impl std::fmt::Debug for KMeansParams {
.field("balance_factor", &self.balance_factor)
.field("hierarchical_k", &self.hierarchical_k)
.field("on_progress", &self.on_progress.as_ref().map(|_| "..."))
.field("seed", &self.seed)
.finish()
}
}
Expand All @@ -120,6 +127,7 @@ impl Default for KMeansParams {
balance_factor: 0.0,
hierarchical_k: 16,
on_progress: None,
seed: None,
}
}
}
Expand Down Expand Up @@ -159,6 +167,12 @@ impl KMeansParams {
self
}

/// Set the seed used for random centroid initialization.
pub fn with_seed(mut self, seed: u64) -> Self {
self.seed = Some(seed);
self
}

/// Set the number of clusters to train in each hierarchical level.
///
/// Higher would split the clusters more aggressively, which would be more accurate but slower.
Expand Down Expand Up @@ -927,8 +941,10 @@ impl KMeans {
let mut cluster_sizes = vec![0; k];
let mut adjusted_balance_factor = f32::MAX;

// TODO: use seed for Rng.
let mut rng = SmallRng::from_os_rng();
let mut rng = match params.seed {
Some(seed) => SmallRng::seed_from_u64(seed),
None => SmallRng::from_os_rng(),
};
for redo in 1..=params.redos {
let mut kmeans: Self = match &params.init {
KMeanInit::Random => Self::init_random::<T>(
Expand Down Expand Up @@ -1850,6 +1866,32 @@ mod tests {
);
}

#[test]
fn test_seeded_initialization_is_reproducible() {
const DIM: usize = 4;
const K: usize = 8;
const NUM_ROWS: usize = 64;

let values = Float32Array::from_iter_values(
(0..NUM_ROWS * DIM).map(|value| ((value * 37) % 101) as f32),
);
let data = FixedSizeListArray::try_new_from_values(values, DIM as i32).unwrap();
let train = || {
// Keep this to one iteration to isolate the seeded initialization
// from floating-point ordering in later parallel reductions.
let params = KMeansParams::new(None, 1, 1, DistanceType::L2).with_seed(42);
KMeans::new_with_params(&data, K, &params).unwrap()
};

let first = train();
let second = train();
assert_eq!(
first.centroids.as_primitive::<Float32Type>().values(),
second.centroids.as_primitive::<Float32Type>().values()
);
assert_eq!(first.loss, second.loss);
}

#[tokio::test]
async fn test_compute_membership_and_loss() {
const DIM: usize = 256;
Expand Down
12 changes: 11 additions & 1 deletion rust/lance-index/src/vector/pq/builder.rs
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,12 @@ pub struct PQBuildParams {

/// Sample rate to train PQ codebook.
pub sample_rate: usize,

/// Optional seed for k-means centroid initialization.
///
/// `None` initializes from OS entropy, while `Some(seed)` makes centroid
/// selection reproducible for the same training data and parameters.
pub kmeans_seed: Option<u64>,
}

impl From<&PQBuildParams> for crate::pb::vector_index_details::ProductQuantization {
Expand All @@ -62,6 +68,7 @@ impl Default for PQBuildParams {
kmeans_redos: 1,
codebook: None,
sample_rate: 256,
kmeans_seed: None,
}
}
}
Expand Down Expand Up @@ -148,7 +155,7 @@ impl PQBuildParams {
.into_iter()
.enumerate()
.map(|(sub_vec_idx, sub_vec)| {
let params = KMeansParams::new(
let mut params = KMeansParams::new(
self.codebook.as_ref().map(|cb| {
let sub_vec_centroids = FixedSizeListArray::try_new_from_values(
cb.as_fixed_size_list().values().as_primitive::<T>().slice(
Expand All @@ -164,6 +171,9 @@ impl PQBuildParams {
self.kmeans_redos,
distance_type,
);
if let Some(seed) = self.kmeans_seed {
params = params.with_seed(seed);
}
train_kmeans::<T>(
&sub_vec,
params,
Expand Down
2 changes: 2 additions & 0 deletions rust/lance/src/index/vector.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1984,6 +1984,7 @@ fn derive_ivf_params(ivf_model: &IvfModel) -> IvfBuildParams {
#[allow(deprecated)]
retrain: false, // Don't retrain since we have centroids
sample_rate: 256, // Default
kmeans_seed: None,
streaming_sample_rate: None,
streaming_coreset_rate: None,
streaming_refine_passes: 0,
Expand All @@ -2005,6 +2006,7 @@ fn derive_pq_params(pq_quantizer: &ProductQuantizer) -> PQBuildParams {
kmeans_redos: 1, // Default
codebook: Some(Arc::new(pq_quantizer.codebook.clone())),
sample_rate: 256, // Default
kmeans_seed: None,
}
}

Expand Down
74 changes: 37 additions & 37 deletions rust/lance/src/index/vector/ivf.rs
Original file line number Diff line number Diff line change
Expand Up @@ -3044,9 +3044,13 @@ where
let _ = progress_tx.send(total);
})
};
let kmeans_params = KMeansParams::new(centroids, params.max_iters as u32, REDOS, metric_type)
.with_balance_factor(1.0)
.with_on_progress(on_progress);
let mut kmeans_params =
KMeansParams::new(centroids, params.max_iters as u32, REDOS, metric_type)
.with_balance_factor(1.0)
.with_on_progress(on_progress);
if let Some(seed) = params.kmeans_seed {
kmeans_params = kmeans_params.with_seed(seed);
}
let kmeans = lance_index::vector::kmeans::train_kmeans::<T>(
data,
kmeans_params,
Expand Down Expand Up @@ -3422,6 +3426,7 @@ struct KMeansStepOptions {
sample_rate: usize,
max_iters: usize,
on_progress: KMeansProgressCallback,
kmeans_seed: Option<u64>,
}

fn train_ivf_kmeans_step<T: ArrowPrimitiveType>(
Expand All @@ -3438,6 +3443,9 @@ where
KMeansParams::new(centroids, options.max_iters as u32, 1, options.metric_type)
.with_balance_factor(1.0)
.with_on_progress(options.on_progress.clone());
if let Some(seed) = options.kmeans_seed {
kmeans_params = kmeans_params.with_seed(seed);
}
if has_centroids {
// Incremental refinement already has the full centroid set. The
// hierarchical trainer bootstraps a smaller tree and is only suitable
Expand All @@ -3456,37 +3464,24 @@ where
fn train_ivf_kmeans_step_arrow_array_no_loss(
centroids: Option<Arc<FixedSizeListArray>>,
data: &FixedSizeListArray,
metric_type: MetricType,
num_partitions: usize,
sample_rate: usize,
max_iters: usize,
on_progress: Arc<dyn Fn(u32, u32) + Send + Sync>,
options: KMeansStepOptions,
) -> Result<KMeans> {
let dimension = data.value_length() as usize;
let values = data.values();
let step_options = KMeansStepOptions {
dimension,
metric_type,
num_partitions,
sample_rate,
max_iters,
on_progress,
};
let kmeans = match (values.data_type(), metric_type) {
let kmeans = match (values.data_type(), options.metric_type) {
(DataType::Float16, _) => train_ivf_kmeans_step::<Float16Type>(
centroids,
values.as_primitive::<Float16Type>(),
&step_options,
&options,
)?,
(DataType::Float32, _) => train_ivf_kmeans_step::<Float32Type>(
centroids,
values.as_primitive::<Float32Type>(),
&step_options,
&options,
)?,
(DataType::Float64, _) => train_ivf_kmeans_step::<Float64Type>(
centroids,
values.as_primitive::<Float64Type>(),
&step_options,
&options,
)?,
(DataType::Int8, DistanceType::L2)
| (DataType::Int8, DistanceType::Dot)
Expand All @@ -3495,18 +3490,18 @@ fn train_ivf_kmeans_step_arrow_array_no_loss(
train_ivf_kmeans_step::<Float32Type>(
centroids,
data.values().as_primitive::<Float32Type>(),
&step_options,
&options,
)?
}
(DataType::UInt8, DistanceType::Hamming) => train_ivf_kmeans_step::<UInt8Type>(
centroids,
values.as_primitive::<UInt8Type>(),
&step_options,
&options,
)?,
_ => Err(Error::index(format!(
"KMeans: can not train data type {} with distance type: {}",
values.data_type(),
metric_type
options.metric_type
)))?,
};
Ok(kmeans)
Expand Down Expand Up @@ -4064,18 +4059,20 @@ fn append_local_coreset(
local_k: usize,
max_iters: usize,
on_progress: Arc<dyn Fn(u32, u32) + Send + Sync>,
kmeans_seed: Option<u64>,
) -> Result<()> {
let dimension = data.value_length() as usize;
let sample_rate = data.len().div_ceil(local_k).max(1);
let kmeans = train_ivf_kmeans_step_arrow_array_no_loss(
None,
data,
let options = KMeansStepOptions {
dimension,
metric_type,
local_k,
num_partitions: local_k,
sample_rate,
max_iters,
on_progress,
)?;
kmeans_seed,
};
let kmeans = train_ivf_kmeans_step_arrow_array_no_loss(None, data, options)?;
let centroids = FixedSizeListArray::try_new_from_values(kmeans.centroids, dimension as i32)?;
let kmeans =
KMeans::with_centroids(centroids.values().clone(), dimension, metric_type, f64::MAX);
Expand Down Expand Up @@ -4447,6 +4444,7 @@ async fn train_streaming_coreset_ivf_model(
local_k,
params.max_iters,
on_progress.clone(),
params.kmeans_seed,
)?;
coreset.append(chunk_coreset);
coreset.reduce_to_budget(dimension, coreset_budget);
Expand Down Expand Up @@ -4622,15 +4620,17 @@ async fn train_streaming_ivf_model(
);
}

let kmeans = train_ivf_kmeans_step_arrow_array_no_loss(
centroids.clone(),
&training_data,
mt,
let options = KMeansStepOptions {
dimension,
metric_type: mt,
num_partitions,
step_sample_rate,
params.max_iters,
on_progress.clone(),
)?;
sample_rate: step_sample_rate,
max_iters: params.max_iters,
on_progress: on_progress.clone(),
kmeans_seed: params.kmeans_seed,
};
let kmeans =
train_ivf_kmeans_step_arrow_array_no_loss(centroids.clone(), &training_data, options)?;
let trained_centroids = Arc::new(FixedSizeListArray::try_new_from_values(
kmeans.centroids,
dimension as i32,
Expand Down
Loading
Loading