use lance_core::utils::row_addr_remap::RowAddrRemap;
use std::collections::HashSet;
use std::sync::{Arc, Mutex};
use std::{
collections::{BTreeMap, HashMap},
pin::Pin,
};
use arrow::array::{AsArray as _, PrimitiveBuilder, UInt64Builder};
use arrow::compute::sort_to_indices;
use arrow::datatypes::{self};
use arrow::datatypes::{Float16Type, Float64Type, UInt8Type, UInt64Type};
use arrow_array::types::Float32Type;
use arrow_array::{
Array, ArrayRef, ArrowPrimitiveType, BooleanArray, FixedSizeListArray, PrimitiveArray,
RecordBatch, UInt32Array, UInt64Array,
};
use arrow_schema::{DataType, Field, Fields};
use futures::{FutureExt, SinkExt, stream};
use futures::{
Stream,
prelude::stream::{StreamExt, TryStreamExt},
};
use lance_arrow::{FixedSizeListArrayExt, RecordBatchExt};
use lance_core::ROW_ID;
use lance_core::datatypes::Schema;
use lance_core::utils::tempfile::TempStdDir;
use lance_core::utils::tokio::{get_num_compute_intensive_cpus, spawn_cpu};
use lance_core::{Error, ROW_ID_FIELD, Result};
use lance_file::version::ConcreteFileVersion;
use lance_file::versions as file_versions;
use lance_file::writer::FileWriterOptions;
use lance_index::frag_reuse::{CompactFragReuseIndex, CompactFragReuseIndexHandle};
use lance_index::metrics::NoOpMetricsCollector;
use lance_index::optimize::OptimizeOptions;
use lance_index::progress::{IndexBuildProgress, NoopIndexBuildProgress};
use lance_index::scalar::RowIdRemapper;
use lance_index::vector::bq::storage::{RABIT_CODE_COLUMN, unpack_codes};
use lance_index::vector::kmeans::KMeansParams;
use lance_index::vector::pq::storage::transpose;
use lance_index::vector::quantizer::{
QuantizationMetadata, QuantizationType, QuantizerBuildParams,
};
use lance_index::vector::quantizer::{QuantizerMetadata, QuantizerStorage};
use lance_index::vector::shared::{SupportedIvfIndexType, write_unified_ivf_and_index_metadata};
use lance_index::vector::storage::STORAGE_METADATA_KEY;
use lance_index::vector::transform::Flatten;
use lance_index::vector::v3::shuffler::{
DEFAULT_PARTITION_WINDOW_BYTES, EmptyReader, IvfShufflerReader, create_ivf_shuffler,
};
use lance_index::vector::v3::subindex::SubIndexType;
use lance_index::vector::{LOSS_METADATA_KEY, PART_ID_COLUMN, PQ_CODE_COLUMN, VectorIndex};
use lance_index::vector::{PART_ID_FIELD, ivf::IvfTransformer, ivf::storage::IvfModel};
use lance_index::{
INDEX_AUXILIARY_FILE_NAME, INDEX_FILE_NAME, pb,
vector::{
DISTANCE_TYPE_KEY,
ivf::{IvfBuildParams, storage::IVF_METADATA_KEY},
quantizer::Quantization,
storage::{StorageBuilder, VectorStore},
transform::Transformer,
v3::{
shuffler::{ShuffleReader, Shuffler},
subindex::IvfSubIndex,
},
},
};
use lance_index::{
INDEX_METADATA_SCHEMA_KEY, IndexMetadata, IndexType, MAX_PARTITION_SIZE_FACTOR,
MIN_PARTITION_SIZE_PERCENT, scalar::OldIndexDataFilter,
};
use lance_io::local::to_local_path;
use lance_io::stream::RecordBatchStream;
use lance_io::{object_store::ObjectStore, stream::RecordBatchStreamAdapter};
use lance_linalg::distance::{DistanceType, Dot, L2, Normalize};
use lance_linalg::kernels::normalize_fsl;
use lance_table::format::IndexFile;
use log::info;
use object_store::path::Path;
use prost::Message;
use roaring::RoaringBitmap;
use tokio::sync::{OnceCell, OwnedSemaphorePermit, Semaphore};
use tracing::{Level, instrument, span};
use crate::Dataset;
use crate::dataset::ProjectionRequest;
use crate::dataset::index::dataset_format_version;
use crate::index::append::build_old_data_filter;
use crate::index::vector::bounded_partition_stream::{
BoundedPartitionStream, Budgeted, OrderedPartitionResults, WeightedJob,
};
use crate::index::vector::utils::infer_vector_dim;
use super::v2::IVFIndex;
use super::{
ivf::load_precomputed_partitions_if_available,
utils::{self, get_vector_type},
};
const REASSIGN_RANGE: usize = 64;
const SPLIT_SAMPLE_RATE: usize = 256;
const MAX_SPLIT_WAYS: usize = 1024;
const JOIN_FETCH_BYTES: usize = 32 * 1024 * 1024;
const JOIN_MULTIVECTOR_FETCHES_IN_FLIGHT: usize = 4;
const JOIN_VECTORS_PER_BATCH: usize = 1024;
fn new_partition_ids_len(split_partitions: &[(usize, Vec<usize>)]) -> usize {
split_partitions.iter().map(|(_, ids)| ids.len()).sum()
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
struct PartitionSplit {
partition: usize,
ways: usize,
}
const PARTITION_BUILD_BUDGET_BYTES: usize = 512 * 1024 * 1024;
const PARTITION_BUILD_ENTRIES_PER_WORKER: usize = 2;
#[derive(Debug, Clone, Copy)]
struct FreshPartitionBuildLimits {
window_bytes: usize,
decoded_budget_bytes: usize,
}
impl Default for FreshPartitionBuildLimits {
fn default() -> Self {
Self {
window_bytes: DEFAULT_PARTITION_WINDOW_BYTES,
decoded_budget_bytes: PARTITION_BUILD_BUDGET_BYTES,
}
}
}
fn apply_centroid_splits(
original: &FixedSizeListArray,
splits: &[(usize, Vec<ArrayRef>)],
) -> Result<FixedSizeListArray> {
let mut new_centroids: Vec<ArrayRef> = original.iter().map(|v| v.unwrap()).collect();
for (part_idx, centroids) in splits {
let (first, rest) = centroids.split_first().ok_or_else(|| {
Error::invalid_input(format!(
"split of partition {part_idx} produced no centroids"
))
})?;
new_centroids[*part_idx] = first.clone();
new_centroids.extend(rest.iter().cloned());
}
let refs: Vec<&dyn Array> = new_centroids.iter().map(|a| a.as_ref()).collect();
let concatenated = arrow::compute::concat(&refs)?;
Ok(FixedSizeListArray::try_new_from_values(
concatenated,
original.value_length(),
)?)
}
#[derive(Clone)]
pub struct ExistingIndex {
pub index: Arc<dyn VectorIndex>,
coverage: Option<Arc<SegmentCoverage>>,
}
struct SegmentCoverage {
dataset: Dataset,
effective_frags: RoaringBitmap,
deleted_frags: RoaringBitmap,
filter: OnceCell<Option<OldIndexDataFilter>>,
}
impl ExistingIndex {
pub fn unfiltered(index: Arc<dyn VectorIndex>) -> Self {
Self {
index,
coverage: None,
}
}
pub fn with_coverage(
index: Arc<dyn VectorIndex>,
dataset: Dataset,
effective_frags: RoaringBitmap,
deleted_frags: RoaringBitmap,
) -> Self {
Self {
index,
coverage: Some(Arc::new(SegmentCoverage {
dataset,
effective_frags,
deleted_frags,
filter: OnceCell::new(),
})),
}
}
#[cfg(test)]
pub(crate) fn filter_is_built(&self) -> bool {
self.coverage
.as_deref()
.is_some_and(|coverage| coverage.filter.initialized())
}
pub(crate) async fn old_data_filter(&self) -> Result<Option<&OldIndexDataFilter>> {
let Some(coverage) = self.coverage.as_deref() else {
return Ok(None);
};
let filter = coverage
.filter
.get_or_try_init(|| {
build_old_data_filter(
&coverage.dataset,
&coverage.effective_frags,
&coverage.deleted_frags,
)
})
.await?;
Ok(filter.as_ref())
}
}
pub struct IvfIndexBuilder<S: IvfSubIndex, Q: Quantization> {
store: ObjectStore,
column: String,
index_dir: Path,
distance_type: DistanceType,
dataset: Option<Dataset>,
shuffler: Option<Arc<dyn Shuffler>>,
ivf_params: Option<IvfBuildParams>,
quantizer_params: Option<Q::BuildParams>,
sub_index_params: Option<S::BuildParams>,
_temp_dir: TempStdDir, temp_dir: Path,
ivf: Option<IvfModel>,
quantizer: Option<Q>,
shuffle_reader: Option<Arc<dyn ShuffleReader>>,
shuffle_data_input: Mutex<Option<UnindexedStream>>,
existing_indices: Vec<ExistingIndex>,
frag_reuse_index: Option<Arc<CompactFragReuseIndex>>,
fragment_filter: Option<Vec<u32>>,
optimize_options: Option<OptimizeOptions>,
merged_num: usize,
target_partition_size: Option<usize>,
transpose_codes: bool,
format_version: ConcreteFileVersion,
progress: Arc<dyn IndexBuildProgress>,
}
type BuildStream<S, Q> =
Pin<Box<dyn Stream<Item = Result<Budgeted<PartitionBuildResult<S, Q>>>> + Send>>;
type FreshWindowBuildStream<S, Q> =
Pin<Box<dyn Stream<Item = Result<(PartitionBuildResult<S, Q>, OwnedSemaphorePermit)>> + Send>>;
type PartitionInputAdmissionStream<T> =
Pin<Box<dyn Stream<Item = Result<(T, OwnedSemaphorePermit)>> + Send>>;
fn admit_partition_inputs<T: Send + 'static>(
inputs: Vec<T>,
entry_permits: Arc<Semaphore>,
) -> PartitionInputAdmissionStream<T> {
stream::iter(inputs)
.then(move |input| {
let entry_permits = entry_permits.clone();
async move {
let entry_permit = entry_permits
.acquire_owned()
.await
.map_err(|_| Error::internal("partition build entry semaphore was closed"))?;
Ok((input, entry_permit))
}
})
.boxed()
}
fn partition_window_entry_limit(
partition_range: &std::ops::Range<usize>,
num_partitions: usize,
max_entries: usize,
concurrency: usize,
) -> usize {
if partition_range.start == 0 && partition_range.end == num_partitions {
max_entries
} else {
max_entries.div_ceil(concurrency)
}
}
struct PartitionBuildResult<S: IvfSubIndex, Q: Quantization> {
partition_id: usize,
built: Option<(Q::Storage, S, f64)>,
}
struct FreshPartitionInput {
partition_id: usize,
batches: Vec<RecordBatch>,
loss: f64,
}
type UnindexedStream = Box<dyn Stream<Item = Result<RecordBatch>> + Send + Unpin + 'static>;
pub struct VectorIndexBuildSummary {
pub indices_merged: usize,
pub files: Vec<IndexFile>,
}
impl<S: IvfSubIndex + 'static, Q: Quantization + 'static> IvfIndexBuilder<S, Q> {
#[allow(clippy::too_many_arguments)]
pub fn new(
dataset: Dataset,
column: String,
index_dir: Path,
distance_type: DistanceType,
shuffler: Box<dyn Shuffler>,
ivf_params: Option<IvfBuildParams>,
quantizer_params: Option<Q::BuildParams>,
sub_index_params: S::BuildParams,
frag_reuse_index: Option<Arc<CompactFragReuseIndex>>,
) -> Result<Self> {
let temp_dir = TempStdDir::default();
let temp_dir_path = Path::from_filesystem_path(&temp_dir)?;
let format_version = dataset_format_version(&dataset);
Ok(Self {
store: dataset.object_store.as_ref().clone(),
column,
index_dir,
distance_type,
dataset: Some(dataset),
shuffler: Some(shuffler.into()),
ivf_params,
quantizer_params,
sub_index_params: Some(sub_index_params),
_temp_dir: temp_dir,
temp_dir: temp_dir_path,
ivf: None,
quantizer: None,
shuffle_reader: None,
shuffle_data_input: Mutex::new(None),
existing_indices: Vec::new(),
frag_reuse_index,
fragment_filter: None,
optimize_options: None,
merged_num: 0,
target_partition_size: None,
transpose_codes: true,
format_version,
progress: Arc::new(NoopIndexBuildProgress),
})
}
#[allow(clippy::too_many_arguments)]
pub fn new_incremental(
dataset: Dataset,
column: String,
index_dir: Path,
distance_type: DistanceType,
shuffler: Box<dyn Shuffler>,
sub_index_params: S::BuildParams,
frag_reuse_index: Option<Arc<CompactFragReuseIndex>>,
optimize_options: OptimizeOptions,
) -> Result<Self> {
let mut builder = Self::new(
dataset,
column,
index_dir,
distance_type,
shuffler,
None,
None,
sub_index_params,
frag_reuse_index,
)?;
builder.optimize_options = Some(optimize_options);
Ok(builder)
}
pub fn new_remapper(
dataset: Dataset,
column: String,
index_dir: Path,
index: Arc<dyn VectorIndex>,
) -> Result<Self> {
let ivf_index = index
.as_any()
.downcast_ref::<IVFIndex<S, Q>>()
.ok_or(Error::invalid_input("existing index is not IVF index"))?;
let temp_dir = TempStdDir::default();
let temp_dir_path = Path::from_filesystem_path(&temp_dir)?;
let format_version = dataset_format_version(&dataset);
Ok(Self {
store: dataset.object_store.as_ref().clone(),
column,
index_dir,
distance_type: ivf_index.metric_type(),
dataset: Some(dataset),
shuffler: None,
ivf_params: None,
quantizer_params: None,
sub_index_params: None,
_temp_dir: temp_dir,
temp_dir: temp_dir_path,
ivf: Some(ivf_index.ivf_model().clone()),
quantizer: Some(ivf_index.quantizer().try_into()?),
shuffle_reader: None,
shuffle_data_input: Mutex::new(None),
existing_indices: vec![ExistingIndex::unfiltered(index)],
frag_reuse_index: None,
fragment_filter: None,
optimize_options: None,
merged_num: 0,
target_partition_size: None,
transpose_codes: true,
format_version,
progress: Arc::new(NoopIndexBuildProgress),
})
}
pub async fn build(&mut self) -> Result<VectorIndexBuildSummary> {
let progress = self.progress.clone();
let max_iters = self.ivf_params.as_ref().map(|p| p.max_iters as u64);
progress
.stage_start("train_ivf", max_iters, "iterations")
.await?;
self.with_ivf(self.load_or_build_ivf().boxed().await?);
progress.stage_complete("train_ivf").await?;
progress.stage_start("train_quantizer", None, "").await?;
self.with_quantizer(self.load_or_build_quantizer().await?);
progress.stage_complete("train_quantizer").await?;
if self.shuffle_reader.is_none() {
let num_rows = self.num_rows_to_shuffle().await?;
progress.stage_start("shuffle", num_rows, "rows").await?;
let input = self.shuffle_data_input.lock().unwrap().take();
if let Some(input) = input {
self.shuffle_data(Some(input)).boxed().await?;
} else {
self.shuffle_dataset().boxed().await?;
}
progress.stage_complete("shuffle").await?;
}
let num_partitions = self.ivf.as_ref().map(|ivf| ivf.num_partitions() as u64);
progress
.stage_start("merge_partitions", num_partitions, "partitions")
.await?;
let build_idx_stream = self.build_partitions().boxed().await?;
let files = self.merge_partitions(build_idx_stream).await?;
progress.stage_complete("merge_partitions").await?;
Ok(VectorIndexBuildSummary {
indices_merged: self.merged_num,
files,
})
}
pub async fn remap(&mut self, mapping: &RowAddrRemap) -> Result<Vec<IndexFile>> {
if self.existing_indices.is_empty() {
return Err(Error::invalid_input(
"No existing indices available for remapping",
));
}
let Some(ivf) = self.ivf.as_ref() else {
return Err(Error::invalid_input("IVF model not set before remapping"));
};
log::info!("remap {} partitions", ivf.num_partitions());
let existing_index = self.existing_indices[0].index.clone();
let mapping = Arc::new(mapping.clone());
let build_iter = (0..ivf.num_partitions()).map(move |part_id| {
let existing_index = existing_index.clone();
let mapping = mapping.clone();
async move {
let ivf_index = existing_index
.as_any()
.downcast_ref::<IVFIndex<S, Q>>()
.ok_or(Error::invalid_input("existing index is not IVF index"))?;
let part = ivf_index
.load_partition(part_id, false, &NoOpMetricsCollector)
.await?;
let storage = part.storage.remap(&mapping)?;
let index = part.index.remap(&mapping, &storage)?;
Result::Ok(Budgeted::untracked(PartitionBuildResult {
partition_id: part_id,
built: Some((storage, index, 0.0)),
}))
}
});
let files = self
.merge_partitions(
stream::iter(build_iter)
.buffered(get_num_compute_intensive_cpus())
.boxed(),
)
.await?;
Ok(files)
}
pub fn with_ivf(&mut self, ivf: IvfModel) -> &mut Self {
self.ivf = Some(ivf);
self
}
pub fn with_quantizer(&mut self, quantizer: Q) -> &mut Self {
self.quantizer = Some(quantizer);
self
}
pub fn with_existing_indices(&mut self, indices: Vec<Arc<dyn VectorIndex>>) -> &mut Self {
self.existing_indices = indices.into_iter().map(ExistingIndex::unfiltered).collect();
self
}
pub fn with_existing_index_sources(&mut self, sources: Vec<ExistingIndex>) -> &mut Self {
self.existing_indices = sources;
self
}
pub fn with_fragment_filter(&mut self, fragment_ids: Vec<u32>) -> &mut Self {
self.fragment_filter = Some(Dataset::normalize_fragment_ids(&fragment_ids));
self
}
pub fn with_optional_fragment_filter(&mut self, fragment_ids: Option<&[u32]>) -> &mut Self {
if let Some(fragment_ids) = fragment_ids {
self.fragment_filter = Some(Dataset::normalize_fragment_ids(fragment_ids));
}
self
}
pub fn with_transpose(&mut self, transpose: bool) -> &mut Self {
self.transpose_codes = transpose;
self
}
pub fn with_target_partition_size(
&mut self,
target_partition_size: Option<usize>,
) -> &mut Self {
self.target_partition_size = target_partition_size;
self
}
pub fn with_progress(&mut self, progress: Arc<dyn IndexBuildProgress>) -> &mut Self {
self.progress = progress;
self
}
#[instrument(name = "load_or_build_ivf", level = "debug", skip_all)]
async fn load_or_build_ivf(&self) -> Result<IvfModel> {
match &self.ivf {
Some(ivf) => Ok(ivf.clone()),
None => {
let Some(dataset) = self.dataset.as_ref() else {
return Err(Error::invalid_input(
"dataset not set before loading or building IVF",
));
};
let dim = utils::get_vector_dim(dataset.schema(), &self.column)?;
let ivf_params = self
.ivf_params
.as_ref()
.ok_or(Error::invalid_input("IVF build params not set"))?;
super::build_ivf_model(
dataset,
&self.column,
dim,
self.distance_type,
ivf_params,
self.fragment_filter.as_deref(),
self.progress.clone(),
)
.await
}
}
}
#[instrument(name = "load_or_build_quantizer", level = "debug", skip_all)]
async fn load_or_build_quantizer(&self) -> Result<Q> {
if self.quantizer.is_some() {
return Ok(self.quantizer.clone().unwrap());
}
let Some(dataset) = self.dataset.as_ref() else {
return Err(Error::invalid_input(
"dataset not set before loading or building quantizer",
));
};
let sample_size_hint = match &self.quantizer_params {
Some(params) => params.try_sample_size()?,
None => 256 * 256, };
let start = std::time::Instant::now();
info!(
"loading training data for quantizer. sample size: {}",
sample_size_hint
);
let training_data = utils::maybe_sample_training_data(
dataset,
&self.column,
sample_size_hint,
self.fragment_filter.as_deref(),
)
.await?;
info!(
"Finished loading training data in {:02} seconds",
start.elapsed().as_secs_f32()
);
let training_data = if self.distance_type == DistanceType::Cosine {
lance_linalg::kernels::normalize_fsl_owned(training_data)?
} else {
training_data
};
let training_data = utils::filter_finite_training_data(training_data)?;
let training_data = match (self.ivf.as_ref(), Q::use_residual(self.distance_type)) {
(Some(ivf), true) => {
let ivf_transformer = lance_index::vector::ivf::new_ivf_transformer(
ivf.centroids.clone().unwrap(),
DistanceType::L2,
vec![],
);
span!(Level::INFO, "compute residual for PQ training")
.in_scope(|| ivf_transformer.compute_residual(&training_data))?
}
_ => training_data,
};
info!("Start to train quantizer");
let start = std::time::Instant::now();
let quantizer = match &self.quantizer {
Some(q) => q.clone(),
None => {
let quantizer_params = self
.quantizer_params
.as_ref()
.ok_or(Error::invalid_input("quantizer build params not set"))?;
Q::build(&training_data, DistanceType::L2, quantizer_params)?
}
};
info!(
"Trained quantizer in {:02} seconds",
start.elapsed().as_secs_f32()
);
Ok(quantizer)
}
fn rename_row_id(
stream: impl RecordBatchStream + Unpin + 'static,
row_id_idx: usize,
) -> impl RecordBatchStream + Unpin + 'static {
let new_schema = Arc::new(arrow_schema::Schema::new(
stream
.schema()
.fields
.iter()
.enumerate()
.map(|(field_idx, field)| {
if field_idx == row_id_idx {
arrow_schema::Field::new(
ROW_ID,
field.data_type().clone(),
field.is_nullable(),
)
} else {
field.as_ref().clone()
}
})
.collect::<Fields>(),
));
RecordBatchStreamAdapter::new(
new_schema.clone(),
stream.map_ok(move |batch| {
RecordBatch::try_new(new_schema.clone(), batch.columns().to_vec()).unwrap()
}),
)
}
async fn num_rows_to_shuffle(&self) -> Result<Option<u64>> {
let Some(dataset) = self.dataset.as_ref() else {
return Ok(None);
};
match &self.fragment_filter {
Some(fragment_ids) => Ok(Some(
dataset
.count_rows_in_existing_fragments(fragment_ids)
.await? as u64,
)),
None => Ok(Some(dataset.count_rows(None).await? as u64)),
}
}
async fn shuffle_dataset(&mut self) -> Result<()> {
let Some(dataset) = self.dataset.as_ref() else {
return Err(Error::invalid_input("dataset not set before shuffling"));
};
let stream = match self
.ivf_params
.as_ref()
.and_then(|p| p.precomputed_shuffle_buffers.as_ref())
{
Some((uri, _)) => {
let uri = to_local_path(uri);
let uri = uri.trim_end_matches("data");
log::info!("shuffle with precomputed shuffle buffers from {}", uri);
let ds = Dataset::open(uri).await?;
ds.scan().try_into_stream().await?
}
_ => {
log::info!("shuffle column {} over dataset", self.column);
let mut builder = dataset.scan();
builder
.batch_readahead(get_num_compute_intensive_cpus())
.project(&[self.column.as_str()])?
.with_row_id();
if let Some(fragment_ids) = &self.fragment_filter {
log::info!(
"applying fragment filter for distributed indexing: {:?}",
fragment_ids
);
builder.with_fragments(
dataset.get_existing_fragment_metadata_from_ids(fragment_ids),
);
}
let (vector_type, _) = get_vector_type(dataset.schema(), &self.column)?;
let is_multivector = matches!(vector_type, datatypes::DataType::List(_));
if is_multivector {
builder.batch_size(64);
}
builder.try_into_stream().await?
}
};
if let Some((row_id_idx, _)) = stream.schema().column_with_name("row_id") {
self.shuffle_data(Some(Self::rename_row_id(stream, row_id_idx)))
.await?;
} else {
self.shuffle_data(Some(stream)).await?;
}
Ok(())
}
pub fn shuffle_data_input(
&mut self,
data: Option<impl RecordBatchStream + Unpin + 'static>,
) -> &mut Self {
match data {
Some(d) => {
*self.shuffle_data_input.lock().unwrap() = Some(Box::new(d) as UnindexedStream);
}
None => {
self.shuffle_reader = Some(Arc::new(EmptyReader));
}
}
self
}
pub async fn shuffle_data(
&mut self,
data: Option<impl Stream<Item = Result<RecordBatch>> + Unpin + Send + 'static>,
) -> Result<&mut Self> {
let Some(ivf) = self.ivf.as_ref() else {
return Err(Error::invalid_input("IVF not set before shuffle data"));
};
let Some(data) = data else {
self.shuffle_reader = Some(Arc::new(EmptyReader));
return Ok(self);
};
let Some(quantizer) = self.quantizer.clone() else {
return Err(Error::invalid_input(
"quantizer not set before shuffle data",
));
};
let Some(shuffler) = self.shuffler.as_ref() else {
return Err(Error::invalid_input("shuffler not set before shuffle data"));
};
let code_column = quantizer.column();
let transformer = Arc::new(
lance_index::vector::ivf::new_ivf_transformer_with_quantizer(
ivf.centroids.clone().unwrap(),
self.distance_type,
&self.column,
quantizer.into(),
None,
)?,
);
let precomputed_partitions = if let Some(params) = self.ivf_params.as_ref() {
load_precomputed_partitions_if_available(params)
.await?
.unwrap_or_default()
} else {
HashMap::new()
};
let partition_map = Arc::new(precomputed_partitions);
let mut transformed_stream = Box::pin(
data.map(move |batch| {
let partition_map = partition_map.clone();
let ivf_transformer = transformer.clone();
tokio::spawn(async move {
let mut batch = batch?;
if !partition_map.is_empty() {
let row_ids = &batch[ROW_ID];
let part_ids = UInt32Array::from_iter(
row_ids
.as_primitive::<UInt64Type>()
.values()
.iter()
.map(|row_id| partition_map.get(row_id).copied()),
);
let part_ids = UInt32Array::from(part_ids);
batch = batch
.try_with_column(PART_ID_FIELD.clone(), Arc::new(part_ids.clone()))
.expect("failed to add part id column");
if part_ids.null_count() > 0 {
log::info!(
"Filter out rows without valid partition IDs: null_count={}",
part_ids.null_count()
);
let indices = UInt32Array::from_iter(
part_ids
.iter()
.enumerate()
.filter_map(|(idx, v)| v.map(|_| idx as u32)),
);
assert_eq!(indices.len(), batch.num_rows() - part_ids.null_count());
batch = batch.take(&indices)?;
}
}
match batch.schema().column_with_name(code_column) {
Some(_) => {
Ok(batch)
}
None => ivf_transformer.transform(&batch),
}
})
})
.buffered(get_num_compute_intensive_cpus())
.map(|x| x.unwrap())
.peekable(),
);
let batch = transformed_stream.as_mut().peek_mut().await;
let schema = match batch {
Some(Ok(b)) => b.schema(),
Some(Err(e)) => return Err(std::mem::replace(e, Error::Stop)),
None => {
log::info!("no data to shuffle");
self.shuffle_reader = Some(Arc::new(IvfShufflerReader::new(
Arc::new(self.store.clone()),
self.temp_dir.clone(),
vec![0; ivf.num_partitions()],
0.0,
)));
return Ok(self);
}
};
self.shuffle_reader = Some(
shuffler
.shuffle(Box::new(RecordBatchStreamAdapter::new(
schema,
transformed_stream,
)))
.await?
.into(),
);
Ok(self)
}
#[instrument(name = "build_partitions", level = "debug", skip_all)]
async fn build_partitions(&mut self) -> Result<BuildStream<S, Q>> {
let Some(ivf) = self.ivf.as_ref() else {
return Err(Error::invalid_input(
"IVF not set before building partitions",
));
};
let Some(quantizer) = self.quantizer.clone() else {
return Err(Error::invalid_input(
"quantizer not set before building partition",
));
};
let Some(sub_index_params) = self.sub_index_params.clone() else {
return Err(Error::invalid_input(
"sub index params not set before building partition",
));
};
let Some(reader) = self.shuffle_reader.as_ref() else {
return Err(Error::invalid_input(
"shuffle reader not set before building partitions",
));
};
let reader = reader.clone();
let num_indices_to_merge = self
.optimize_options
.as_ref()
.and_then(|opt| opt.num_indices_to_merge);
let no_partition_adjustment = || {
let is_retrain = self
.optimize_options
.as_ref()
.map(|opt| opt.retrain)
.unwrap_or(false);
let num_to_merge = match is_retrain {
true => self.existing_indices.len(), false => num_indices_to_merge.unwrap_or(0),
};
let indices_to_merge = self.existing_indices
[self.existing_indices.len().saturating_sub(num_to_merge)..]
.to_vec();
(ivf.num_partitions(), Arc::new(indices_to_merge), None)
};
let (num_partitions, merge_indices, partition_adjustment) = if num_indices_to_merge
.is_some()
|| self.optimize_options.is_none()
{
no_partition_adjustment()
} else {
let target_partition_size = self.effective_target_partition_size()?;
let PartitionAdjustmentPlan {
splits,
joins,
partition_sizes,
} = Self::check_partition_adjustment(
ivf,
reader.as_ref(),
&self.existing_indices,
target_partition_size,
)?;
let split_result = if splits.is_empty() {
None
} else {
log::info!(
"split partitions {:?} (target partition size {}), will merge all {} delta indices",
splits,
target_partition_size,
self.existing_indices.len()
);
self.split_partitions_streaming(&splits, ivf)
.boxed()
.await?
};
if let Some(split_result) = split_result {
let Some(ivf) = self.ivf.as_mut() else {
return Err(Error::invalid_input(
"IVF not set before building partitions",
));
};
ivf.centroids = Some(split_result.new_centroids);
(
ivf.num_partitions(),
Arc::new(self.existing_indices.clone()),
Some(PartitionAdjustment::Split {
affected_partitions: split_result.affected_partitions,
split_shuffle_reader: split_result.shuffle_reader,
}),
)
} else if !joins.is_empty() {
log::info!(
"join partitions {:?} (target partition size {}), will merge all {} delta indices",
joins,
target_partition_size,
self.existing_indices.len()
);
let results = self
.join_partitions(
&joins,
ivf,
&partition_sizes,
MAX_PARTITION_SIZE_FACTOR * target_partition_size,
)
.boxed()
.await?;
let Some(ivf) = self.ivf.as_mut() else {
return Err(Error::invalid_input(
"IVF model not set before joining partitions",
));
};
ivf.centroids = Some(results.new_centroids);
(
results.kept_partitions.len(),
Arc::new(self.existing_indices.clone()),
Some(PartitionAdjustment::Join {
kept_partitions: results.kept_partitions,
reindexed_row_ids: results.reindexed_row_ids,
join_shuffle_reader: results.join_shuffle_reader,
}),
)
} else {
no_partition_adjustment()
}
};
self.merged_num = merge_indices.len();
log::info!(
"merge {}/{} delta indices",
self.merged_num,
self.existing_indices.len()
);
let distance_type = self.distance_type;
let column = self.column.clone();
let frag_reuse_index = self.frag_reuse_index.clone();
if self.optimize_options.is_none()
&& self.existing_indices.is_empty()
&& partition_adjustment.is_none()
{
return Self::build_fresh_partitions_windowed(
reader,
num_partitions,
distance_type,
quantizer,
sub_index_params,
column,
frag_reuse_index,
FreshPartitionBuildLimits::default(),
);
}
let partition_adjustment = Arc::new(partition_adjustment);
let build_iter = (0..num_partitions).map(move |partition| {
let output_partition_id = partition;
let reader = reader.clone();
let indices = merge_indices.clone();
let distance_type = distance_type;
let quantizer = quantizer.clone();
let sub_index_params = sub_index_params.clone();
let column = column.clone();
let frag_reuse_index = frag_reuse_index.clone();
let partition_adjustment = partition_adjustment.clone();
async move {
let (is_affected, split_reader) = match partition_adjustment.as_ref() {
Some(PartitionAdjustment::Split {
affected_partitions,
split_shuffle_reader,
}) => (
affected_partitions.contains(&partition),
Some(split_shuffle_reader.clone()),
),
_ => (false, None),
};
let partition = match partition_adjustment.as_ref() {
Some(PartitionAdjustment::Join {
kept_partitions, ..
}) => kept_partitions[partition],
_ => partition,
};
let (mut batches, mut loss) = if is_affected {
Self::take_partition_batches(
partition,
&[],
Some(split_reader.as_ref().unwrap().as_ref()),
)
.await?
} else {
Self::take_partition_batches(partition, indices.as_ref(), Some(reader.as_ref()))
.await?
};
if !is_affected && let Some(sr) = split_reader.as_ref() {
let (extra, extra_loss) =
Self::take_partition_batches(partition, &[], Some(sr.as_ref())).await?;
batches.extend(extra);
loss += extra_loss;
}
if let Some(PartitionAdjustment::Join {
reindexed_row_ids,
join_shuffle_reader,
..
}) = partition_adjustment.as_ref()
{
if !reindexed_row_ids.is_empty() {
for batch in batches.iter_mut() {
let row_ids = batch[ROW_ID].as_primitive::<UInt64Type>();
let mask = BooleanArray::from_iter(row_ids.iter().map(|row_id| {
row_id.map(|row_id| !reindexed_row_ids.contains(&row_id))
}));
*batch = arrow::compute::filter_record_batch(batch, &mask)?;
}
}
let (extra, extra_loss) = Self::take_partition_batches(
partition,
&[],
Some(join_shuffle_reader.as_ref()),
)
.await?;
batches.extend(extra);
loss += extra_loss;
}
spawn_cpu(move || {
let num_rows = batches.iter().map(|b| b.num_rows()).sum::<usize>();
if num_rows == 0 {
return Ok(Budgeted::untracked(PartitionBuildResult {
partition_id: output_partition_id,
built: None,
}));
}
let (storage, sub_index) = Self::build_index(
distance_type,
quantizer,
sub_index_params,
batches,
column,
frag_reuse_index,
)?;
Ok(Budgeted::untracked(PartitionBuildResult {
partition_id: output_partition_id,
built: Some((storage, sub_index, loss)),
}))
})
.await
}
});
Ok(stream::iter(build_iter)
.buffered(get_num_compute_intensive_cpus())
.boxed())
}
#[allow(clippy::too_many_arguments)]
fn build_fresh_partitions_windowed(
reader: Arc<dyn ShuffleReader>,
num_partitions: usize,
distance_type: DistanceType,
quantizer: Q,
sub_index_params: S::BuildParams,
column: String,
frag_reuse_index: Option<Arc<CompactFragReuseIndex>>,
limits: FreshPartitionBuildLimits,
) -> Result<BuildStream<S, Q>> {
let concurrency = get_num_compute_intensive_cpus().max(1);
let max_entries = concurrency.saturating_mul(PARTITION_BUILD_ENTRIES_PER_WORKER);
let cpu_permits = Arc::new(Semaphore::new(concurrency));
let jobs = stream::try_unfold(0usize, move |next_partition_id| {
let reader = reader.clone();
let quantizer = quantizer.clone();
let sub_index_params = sub_index_params.clone();
let column = column.clone();
let frag_reuse_index = frag_reuse_index.clone();
let cpu_permits = cpu_permits.clone();
async move {
if next_partition_id == num_partitions {
return Ok(None);
}
let plan = reader.plan_partition_window(
next_partition_id,
limits.window_bytes,
)?;
if plan.partition_range.start != next_partition_id
|| plan.partition_range.end <= plan.partition_range.start
|| plan.partition_range.end > num_partitions
{
return Err(Error::internal(format!(
"shuffle reader planned invalid partition window {:?}; expected a non-empty window starting at {} within {} partitions",
plan.partition_range, next_partition_id, num_partitions
)));
}
let next_partition_id = plan.partition_range.end;
let planned_range = plan.partition_range;
let window_entry_limit = partition_window_entry_limit(
&planned_range,
num_partitions,
max_entries,
concurrency,
);
let job = WeightedJob::with_permit(
plan.estimated_decoded_bytes,
move |mut admission| async move {
let mut window = reader
.read_partition_window(
planned_range.start,
limits.window_bytes,
)
.await?;
if window.partition_range != planned_range
|| window.partitions.len() != planned_range.len()
{
return Err(Error::internal(format!(
"shuffle reader returned partition window {:?} with {} entries after planning {:?}",
window.partition_range,
window.partitions.len(),
planned_range
)));
}
for (expected_partition_id, partition) in
planned_range.clone().zip(&window.partitions)
{
if partition.partition_id != expected_partition_id {
return Err(Error::internal(format!(
"shuffle reader window {:?} returned partition id {} at position {}",
planned_range,
partition.partition_id,
expected_partition_id - planned_range.start
)));
}
}
let count_stream_bytes = window.materialized_decoded_bytes.is_none();
let mut decoded_bytes =
window.materialized_decoded_bytes.unwrap_or_default();
let mut inputs = Vec::with_capacity(window.partitions.len());
for mut partition in window.partitions.drain(..) {
let mut batches = Vec::new();
let mut loss = 0.0;
if let Some(mut data) = partition.data.take() {
while let Some(batch) = data.try_next().await? {
loss += batch
.metadata()
.get(LOSS_METADATA_KEY)
.map(|value| value.parse::<f64>().unwrap_or(0.0))
.unwrap_or(0.0);
if count_stream_bytes {
decoded_bytes = batch.columns().iter().try_fold(
decoded_bytes,
|total, array| {
total
.checked_add(array.get_array_memory_size())
.ok_or_else(|| {
Error::internal(format!(
"decoded byte count overflow for partition {}",
partition.partition_id
))
})
},
)?;
}
batches.push(batch.drop_column(PART_ID_COLUMN)?);
}
}
inputs.push(FreshPartitionInput {
partition_id: partition.partition_id,
batches,
loss,
});
}
admission.reconcile(decoded_bytes);
let entry_permits = Arc::new(Semaphore::new(window_entry_limit));
let builds = admit_partition_inputs(inputs, entry_permits)
.map_ok(move |(input, entry_permit)| {
let quantizer = quantizer.clone();
let sub_index_params = sub_index_params.clone();
let column = column.clone();
let frag_reuse_index = frag_reuse_index.clone();
let cpu_permits = cpu_permits.clone();
async move {
let partition_id = input.partition_id;
let loss = input.loss;
let _cpu_permit =
cpu_permits.acquire_owned().await.map_err(|_| {
Error::internal(
"partition build CPU semaphore was closed",
)
})?;
let built = spawn_cpu(move || -> Result<_> {
let num_rows = input
.batches
.iter()
.map(|batch| batch.num_rows())
.sum::<usize>();
if num_rows == 0 {
return Ok(None);
}
let (storage, sub_index) = Self::build_index(
distance_type,
quantizer,
sub_index_params,
input.batches,
column,
frag_reuse_index,
)?;
Ok(Some((storage, sub_index, loss)))
})
.await?;
Ok::<_, Error>((
PartitionBuildResult {
partition_id,
built,
},
entry_permit,
))
}
})
.try_buffer_unordered(concurrency)
.boxed();
Ok::<(FreshWindowBuildStream<S, Q>, _), Error>((builds, admission))
},
);
Ok(Some((job, next_partition_id)))
}
})
.boxed();
let windows = BoundedPartitionStream::try_new(
jobs,
concurrency,
limits.decoded_budget_bytes,
concurrency,
)?;
Ok(windows
.map_ok(|window| {
let Budgeted {
value: builds,
permit,
entry_permit,
} = window;
debug_assert!(entry_permit.is_none());
builds.map_ok(move |(value, entry_permit)| Budgeted {
value,
permit: permit.clone(),
entry_permit: Some(entry_permit),
})
})
.try_flatten_unordered(Some(concurrency))
.boxed())
}
#[instrument(name = "build_index", level = "debug", skip_all)]
#[allow(clippy::too_many_arguments)]
fn build_index(
distance_type: DistanceType,
quantizer: Q,
sub_index_params: S::BuildParams,
batches: Vec<RecordBatch>,
column: String,
frag_reuse_index: Option<Arc<CompactFragReuseIndex>>,
) -> Result<(Q::Storage, S)> {
let frag_reuse_index = frag_reuse_index
.map(|index| Arc::new(CompactFragReuseIndexHandle(index)) as Arc<dyn RowIdRemapper>);
let storage =
StorageBuilder::new_with_remapper(column, distance_type, quantizer, frag_reuse_index)?
.build(batches)?;
let sub_index = S::index_vectors(&storage, sub_index_params)?;
Ok((storage, sub_index))
}
#[instrument(name = "take_partition_batches", level = "debug", skip_all)]
async fn take_partition_batches(
part_id: usize,
existing_indices: &[ExistingIndex],
reader: Option<&dyn ShuffleReader>,
) -> Result<(Vec<RecordBatch>, f64)> {
let mut batches = Vec::new();
for source in existing_indices.iter() {
let existing_index = source
.index
.as_any()
.downcast_ref::<IVFIndex<S, Q>>()
.ok_or(Error::invalid_input("existing index is not IVF index"))?;
if part_id >= existing_index.ivf_model().num_partitions() {
continue;
}
let old_data_filter = source.old_data_filter().await?;
let part_storage = existing_index.load_partition_storage(part_id, None).await?;
let mut part_batches = part_storage.to_batches()?.collect::<Vec<_>>();
match Q::quantization_type() {
QuantizationType::Product => {
for batch in part_batches.iter_mut() {
if batch.num_rows() == 0 {
continue;
}
let codes = batch[PQ_CODE_COLUMN]
.as_fixed_size_list()
.values()
.as_primitive::<datatypes::UInt8Type>();
let codes_num_bytes = codes.len() / batch.num_rows();
let original_codes = transpose(codes, codes_num_bytes, batch.num_rows());
let original_codes = FixedSizeListArray::try_new_from_values(
original_codes,
codes_num_bytes as i32,
)?;
*batch = batch
.replace_column_by_name(PQ_CODE_COLUMN, Arc::new(original_codes))?
.drop_column(PART_ID_COLUMN)?;
}
}
QuantizationType::Rabit => {
for batch in part_batches.iter_mut() {
if batch.num_rows() == 0 {
continue;
}
let codes = batch[RABIT_CODE_COLUMN].as_fixed_size_list();
let original_codes = unpack_codes(codes);
*batch = batch
.replace_column_by_name(RABIT_CODE_COLUMN, Arc::new(original_codes))?
.drop_column(PART_ID_COLUMN)?;
}
}
_ => {}
}
if let Some(filter) = old_data_filter {
for batch in part_batches.iter_mut() {
let keep = filter.filter_row_ids(batch[ROW_ID].as_primitive::<UInt64Type>());
if keep.true_count() < batch.num_rows() {
*batch = arrow::compute::filter_record_batch(batch, &keep)?;
}
}
}
batches.extend(part_batches);
}
let mut loss = 0.0;
if let Some(reader) = reader
&& reader.partition_size(part_id)? > 0
{
let mut partition_data =
reader
.read_partition(part_id)
.await?
.ok_or(Error::invalid_input(format!(
"partition {} is empty",
part_id
)))?;
while let Some(batch) = partition_data.try_next().await? {
loss += batch
.metadata()
.get(LOSS_METADATA_KEY)
.map(|s| s.parse::<f64>().unwrap_or(0.0))
.unwrap_or(0.0);
batches.push(batch.drop_column(PART_ID_COLUMN)?);
}
}
Ok((batches, loss))
}
#[instrument(name = "merge_partitions", level = "debug", skip_all)]
async fn merge_partitions(
&mut self,
mut build_stream: BuildStream<S, Q>,
) -> Result<Vec<IndexFile>> {
let Some(ivf) = self.ivf.as_ref() else {
return Err(Error::invalid_input("IVF not set before merge partitions"));
};
let Some(quantizer) = self.quantizer.clone() else {
return Err(Error::invalid_input(
"quantizer not set before merge partitions",
));
};
let quantization_type = Q::quantization_type();
let is_pq = quantization_type == QuantizationType::Product;
let is_rq = quantization_type == QuantizationType::Rabit;
let is_flat = quantization_type == QuantizationType::Flat;
let storage_path = self.index_dir.clone().join(INDEX_AUXILIARY_FILE_NAME);
let index_path = self.index_dir.clone().join(INDEX_FILE_NAME);
let writer_options = FileWriterOptions::default();
let mut storage_writer = if is_flat {
None
} else {
let mut fields = vec![ROW_ID_FIELD.clone(), quantizer.field()];
fields.extend(quantizer.extra_fields());
let storage_schema: Schema = (&arrow_schema::Schema::new(fields)).try_into()?;
Some(file_versions::create_writer(
self.format_version,
self.store.create(&storage_path).await?,
storage_schema,
writer_options.clone(),
)?)
};
let mut index_writer = file_versions::create_writer(
self.format_version,
self.store.create(&index_path).await?,
S::schema().as_ref().try_into()?,
writer_options.clone(),
)?;
let mut storage_ivf = IvfModel::empty();
let mut index_ivf = IvfModel::new(ivf.centroids.clone().unwrap(), ivf.loss);
let mut partition_index_metadata = Vec::with_capacity(ivf.num_partitions());
let num_partitions = ivf.num_partitions();
let mut ordered_results = OrderedPartitionResults::new(num_partitions);
let mut total_loss = 0.0;
let progress = self.progress.clone();
log::info!("merging {} partitions", num_partitions);
while let Some(result) = build_stream.try_next().await? {
let partition_id = result.value.partition_id;
ordered_results.push(partition_id, result)?;
while let Some((partition_id, result)) = ordered_results.pop_next() {
let Budgeted {
value: PartitionBuildResult { built: part, .. },
permit: _permit,
entry_permit: _entry_permit,
} = result;
let completed_partitions = partition_id + 1;
progress
.stage_progress("merge_partitions", completed_partitions as u64)
.await?;
let Some((storage, index, loss)) = part else {
log::warn!("partition {} is empty, skipping", partition_id);
storage_ivf.add_partition(0);
index_ivf.add_partition(0);
partition_index_metadata.push(String::new());
continue;
};
total_loss += loss;
if storage.len() == 0 {
storage_ivf.add_partition(0);
} else {
for mut batch in storage.to_batches()? {
if is_pq
&& !self.transpose_codes
&& batch.num_rows() > 0
&& batch.column_by_name(PQ_CODE_COLUMN).is_some()
{
let codes_fsl = batch
.column_by_name(PQ_CODE_COLUMN)
.unwrap()
.as_fixed_size_list();
let num_rows = batch.num_rows();
let bytes_per_code = codes_fsl.value_length() as usize;
let codes = codes_fsl.values().as_primitive::<datatypes::UInt8Type>();
let original_codes = transpose(codes, bytes_per_code, num_rows);
let original_fsl = Arc::new(FixedSizeListArray::try_new_from_values(
original_codes,
bytes_per_code as i32,
)?);
batch = batch.replace_column_by_name(PQ_CODE_COLUMN, original_fsl)?;
}
if is_rq
&& !self.transpose_codes
&& batch.num_rows() > 0
&& batch.column_by_name(RABIT_CODE_COLUMN).is_some()
{
let codes_fsl = batch
.column_by_name(RABIT_CODE_COLUMN)
.unwrap()
.as_fixed_size_list();
let unpacked = Arc::new(unpack_codes(codes_fsl));
batch = batch.replace_column_by_name(RABIT_CODE_COLUMN, unpacked)?;
}
if storage_writer.is_none() {
let storage_schema: Schema = batch.schema_ref().as_ref().try_into()?;
storage_writer = Some(file_versions::create_writer(
self.format_version,
self.store.create(&storage_path).await?,
storage_schema,
writer_options.clone(),
)?);
}
storage_writer
.as_mut()
.expect("storage writer must be initialized before write")
.write_batch(&batch)
.await?;
storage_ivf.add_partition(batch.num_rows() as u32);
}
}
let index_batch = index.to_batch()?;
if index_batch.num_rows() == 0 {
index_ivf.add_partition(0);
partition_index_metadata.push(String::new());
} else {
index_writer.write_batch(&index_batch).await?;
index_ivf.add_partition(index_batch.num_rows() as u32);
partition_index_metadata.push(
index_batch
.schema()
.metadata
.get(S::metadata_key())
.cloned()
.unwrap_or_default(),
);
}
}
}
ordered_results.finish()?;
match self.shuffle_reader.as_ref() {
Some(reader) => {
if let Some(loss) = reader.total_loss() {
total_loss += loss;
}
index_ivf.loss = Some(total_loss);
}
None => {
}
}
if storage_writer.is_none() {
let Some(centroids) = ivf.centroids.as_ref() else {
return Err(Error::invalid_input(
"flat storage writer could not infer schema from empty partitions without IVF centroids",
));
};
let flat_schema = arrow_schema::Schema::new(vec![
ROW_ID_FIELD.as_ref().clone(),
arrow_schema::Field::new(
lance_index::vector::flat::storage::FLAT_COLUMN,
DataType::FixedSizeList(
Arc::new(arrow_schema::Field::new(
"item",
centroids.value_type(),
true,
)),
centroids.value_length(),
),
true,
),
]);
let storage_schema: Schema = (&flat_schema).try_into()?;
storage_writer = Some(file_versions::create_writer(
self.format_version,
self.store.create(&storage_path).await?,
storage_schema,
writer_options.clone(),
)?);
}
let storage_writer = storage_writer
.as_mut()
.expect("storage writer must be initialized before final metadata write");
let storage_ivf_pb = pb::Ivf::try_from(&storage_ivf)?;
storage_writer.add_schema_metadata(DISTANCE_TYPE_KEY, self.distance_type.to_string());
let ivf_buffer_pos = storage_writer
.add_global_buffer(storage_ivf_pb.encode_to_vec().into())
.await?;
storage_writer.add_schema_metadata(IVF_METADATA_KEY, ivf_buffer_pos.to_string());
let transposed = match quantization_type {
QuantizationType::Product | QuantizationType::Rabit => self.transpose_codes,
_ => false,
};
let mut metadata = quantizer.metadata(Some(QuantizationMetadata {
codebook_position: Some(0),
codebook: None,
transposed,
}));
if let Some(extra_metadata) = metadata.extra_metadata()? {
let idx = storage_writer.add_global_buffer(extra_metadata).await?;
metadata.set_buffer_index(idx);
}
let metadata = serde_json::to_string(&metadata)?;
let storage_partition_metadata = vec![metadata];
storage_writer.add_schema_metadata(
STORAGE_METADATA_KEY,
serde_json::to_string(&storage_partition_metadata)?,
);
let index_type_str = index_type_string(S::name().try_into()?, Q::quantization_type());
if let Some(idx_type) = SupportedIvfIndexType::from_index_type_str(&index_type_str) {
write_unified_ivf_and_index_metadata(
&mut index_writer,
&index_ivf,
self.distance_type,
idx_type,
)
.await?;
} else {
let index_ivf_pb = pb::Ivf::try_from(&index_ivf)?;
let index_metadata = IndexMetadata {
index_type: index_type_str,
distance_type: self.distance_type.to_string(),
};
index_writer.add_schema_metadata(
INDEX_METADATA_SCHEMA_KEY,
serde_json::to_string(&index_metadata)?,
);
let ivf_buffer_pos = index_writer
.add_global_buffer(index_ivf_pb.encode_to_vec().into())
.await?;
index_writer.add_schema_metadata(IVF_METADATA_KEY, ivf_buffer_pos.to_string());
}
index_writer.add_schema_metadata(
S::metadata_key(),
serde_json::to_string(&partition_index_metadata)?,
);
let storage_summary = storage_writer.finish().await?;
let index_summary = index_writer.finish().await?;
log::info!("merging {} partitions done", ivf.num_partitions());
Ok(vec![
IndexFile {
path: INDEX_AUXILIARY_FILE_NAME.to_string(),
size_bytes: storage_summary.size_bytes,
},
IndexFile {
path: INDEX_FILE_NAME.to_string(),
size_bytes: index_summary.size_bytes,
},
])
}
async fn take_vectors(
dataset: &Dataset,
column: &str,
store: &ObjectStore,
row_ids: &[u64],
) -> Result<Vec<RecordBatch>> {
Self::take_vectors_stream(
dataset,
column,
row_ids,
store.block_size(),
store.io_parallelism(),
)
.await?
.try_collect()
.await
}
async fn take_vectors_stream(
dataset: &Dataset,
column: &str,
row_ids: &[u64],
rows_per_chunk: usize,
prefetch: usize,
) -> Result<impl Stream<Item = Result<RecordBatch>> + Send + 'static> {
let projection = Arc::new(dataset.schema().project(&[column])?);
let row_ids = dataset.filter_deleted_ids(row_ids).await?;
let chunks: Vec<Vec<u64>> = row_ids
.chunks(rows_per_chunk.max(1))
.map(|chunk| chunk.to_vec())
.collect();
let dataset = dataset.clone();
Ok(stream::iter(chunks)
.map(move |chunk| {
let dataset = dataset.clone();
let projection = projection.clone();
async move {
let batch = dataset
.take_rows(&chunk, ProjectionRequest::Schema(projection))
.await?;
if batch.num_rows() != chunk.len() {
return Err(Error::invalid_input(format!(
"batch.num_rows() != chunk.len() ({} != {})",
batch.num_rows(),
chunk.len()
)));
}
Ok(batch.try_with_column(
ROW_ID_FIELD.clone(),
Arc::new(UInt64Array::from(chunk)),
)?)
}
})
.buffered(prefetch.max(1)))
}
fn flatten_raw_vectors(
&self,
batch: &RecordBatch,
) -> Result<(UInt64Array, FixedSizeListArray)> {
let batch = Flatten::new(&self.column).transform(batch)?;
let row_ids = batch[ROW_ID].as_primitive::<UInt64Type>().clone();
let vectors = batch
.column_by_qualified_name(&self.column)
.ok_or(Error::invalid_input(format!(
"vector column {} not found in batch {}",
self.column,
batch.schema()
)))?
.as_fixed_size_list()
.clone();
Ok((row_ids, vectors))
}
fn effective_target_partition_size(&self) -> Result<usize> {
if let Some(size) = self.target_partition_size {
return Ok(size);
}
let index_type = IndexType::try_from(
index_type_string(S::name().try_into()?, Q::quantization_type()).as_str(),
)?;
Ok(index_type.target_partition_size())
}
fn check_partition_adjustment(
ivf: &IvfModel,
reader: &dyn ShuffleReader,
existing_indices: &[ExistingIndex],
target_partition_size: usize,
) -> Result<PartitionAdjustmentPlan> {
let mut partition_sizes = Vec::with_capacity(ivf.num_partitions());
for partition in 0..ivf.num_partitions() {
let mut num_rows = reader.partition_size(partition)?;
for source in existing_indices.iter() {
num_rows += source.index.partition_size(partition);
}
partition_sizes.push(num_rows);
}
let (splits, joins) = plan_partition_adjustment(&partition_sizes, target_partition_size);
Ok(PartitionAdjustmentPlan {
splits,
joins,
partition_sizes,
})
}
async fn split_partitions_streaming(
&self,
splits: &[PartitionSplit],
ivf: &IvfModel,
) -> Result<Option<SplitResult>> {
let Some(dataset) = self.dataset.as_ref() else {
return Err(Error::invalid_input(
"dataset not set before split partition",
));
};
let (_, element_type) = get_vector_type(dataset.schema(), &self.column)?;
let (new_centroids, split_partitions) = match element_type {
DataType::Float16 => {
self.compute_split_centroids::<Float16Type>(splits, ivf)
.await?
}
DataType::Float32 => {
self.compute_split_centroids::<Float32Type>(splits, ivf)
.await?
}
DataType::Float64 => {
self.compute_split_centroids::<Float64Type>(splits, ivf)
.await?
}
DataType::UInt8 => {
self.compute_split_centroids::<UInt8Type>(splits, ivf)
.await?
}
dt => {
return Err(Error::invalid_input(format!(
"vectors must be float16, float32, float64 or uint8, but got {:?}",
dt
)));
}
};
if split_partitions.is_empty() {
return Ok(None);
}
let mut affected_partitions = HashSet::new();
let mut neighbors = 0usize;
for (part_idx, new_partitions) in &split_partitions {
affected_partitions.extend(new_partitions.iter().copied());
let c0 = ivf.centroid(*part_idx).ok_or(Error::invalid_input(format!(
"centroid not found for partition {part_idx}",
)))?;
let (neighbor_ids, _) = select_reassign_candidates_impl(
self.distance_type,
ivf,
*part_idx,
&c0,
&HashSet::new(),
)?;
for id in neighbor_ids.values() {
if affected_partitions.insert(*id as usize) {
neighbors += 1;
}
}
}
log::info!(
"split {} partitions into {} partitions; {} neighbor partitions reassigned, {} total affected partitions",
split_partitions.len(),
new_partition_ids_len(&split_partitions),
neighbors,
affected_partitions.len(),
);
let split_shuffle_reader = self
.reshuffle_partitions(&affected_partitions, &new_centroids)
.await?;
Ok(Some(SplitResult {
new_centroids,
affected_partitions,
shuffle_reader: split_shuffle_reader.into(),
}))
}
async fn compute_split_centroids<T: ArrowPrimitiveType>(
&self,
splits: &[PartitionSplit],
ivf: &IvfModel,
) -> Result<(FixedSizeListArray, Vec<(usize, Vec<usize>)>)>
where
T::Native: Dot + L2 + Normalize,
PrimitiveArray<T>: From<Vec<T::Native>>,
{
let centroids = ivf.centroids_array().unwrap();
let trained_centroids = stream::iter(splits.iter().copied())
.map(|split| async move { self.train_split_centroids::<T>(split, ivf).await })
.buffered(get_num_compute_intensive_cpus())
.try_collect::<Vec<_>>()
.await?;
let mut applied = Vec::new();
let mut split_partitions = Vec::new();
let mut next_partition = centroids.len();
for (split, trained) in splits.iter().zip(trained_centroids) {
let Some(trained) = trained else {
continue;
};
let appended = next_partition..next_partition + trained.len() - 1;
next_partition = appended.end;
let mut partitions = vec![split.partition];
partitions.extend(appended);
split_partitions.push((split.partition, partitions));
applied.push((split.partition, trained));
}
if applied.is_empty() {
return Ok((centroids.clone(), vec![]));
}
let new_centroids = apply_centroid_splits(centroids, &applied)?;
Ok((new_centroids, split_partitions))
}
async fn reshuffle_partitions(
&self,
affected_partitions: &HashSet<usize>,
new_centroids: &FixedSizeListArray,
) -> Result<Box<dyn ShuffleReader>> {
let Some(dataset) = self.dataset.as_ref() else {
return Err(Error::invalid_input("dataset not set before reshuffle"));
};
let Some(quantizer) = self.quantizer.clone() else {
return Err(Error::invalid_input("quantizer not set before reshuffle"));
};
let mut all_row_ids = Vec::new();
for &part_idx in affected_partitions {
let mut row_ids = self.partition_row_ids(part_idx).await?;
all_row_ids.append(&mut row_ids);
}
all_row_ids.sort();
all_row_ids.dedup();
let projection = Arc::new(dataset.schema().project(&[self.column.as_str()])?);
let row_ids = dataset.filter_deleted_ids(&all_row_ids).await?;
let block_size = self.store.block_size();
let io_parallelism = self.store.io_parallelism();
let column = self.column.clone();
let dataset_clone = dataset.clone();
let projection_clone = projection.clone();
let raw_stream = stream::iter(
row_ids
.chunks(block_size)
.map(|c| c.to_vec())
.collect::<Vec<_>>(),
)
.map(move |chunk| {
let dataset = dataset_clone.clone();
let projection = projection_clone.clone();
let column = column.clone();
async move {
let batch = dataset
.take_rows(&chunk, ProjectionRequest::Schema(projection))
.await?;
let batch = batch
.try_with_column(ROW_ID_FIELD.clone(), Arc::new(UInt64Array::from(chunk)))?;
Flatten::new(&column).transform(&batch)
}
})
.buffered(io_parallelism)
.boxed();
let transformer = Arc::new(
lance_index::vector::ivf::new_ivf_transformer_with_quantizer(
new_centroids.clone(),
self.distance_type,
&self.column,
quantizer.into(),
None,
)?,
);
let mut transformed_stream = Box::pin(
raw_stream
.map(move |batch| {
let ivf_transformer = transformer.clone();
tokio::spawn(async move { ivf_transformer.transform(&batch?) })
})
.buffered(get_num_compute_intensive_cpus())
.map(|x| x.unwrap())
.peekable(),
);
let schema = match transformed_stream.as_mut().peek_mut().await {
Some(Ok(b)) => b.schema(),
Some(Err(e)) => return Err(std::mem::replace(e, Error::Stop)),
None => {
log::info!("no vectors to reshuffle");
let empty_reader: Box<dyn ShuffleReader> = Box::new(IvfShufflerReader::new(
Arc::new(self.store.clone()),
self.temp_dir.clone().join("split_shuffle"),
vec![0; new_centroids.len()],
0.0,
));
return Ok(empty_reader);
}
};
let transformed_stream =
Box::new(RecordBatchStreamAdapter::new(schema, transformed_stream));
let split_shuffle_dir = self.temp_dir.clone().join("split_shuffle");
let shuffler = create_ivf_shuffler(
split_shuffle_dir,
new_centroids.len(),
self.format_version,
None,
);
shuffler.shuffle(transformed_stream).await
}
async fn sample_partition_raw_vectors(
&self,
part_idx: usize,
sample_size: usize,
) -> Result<Option<FixedSizeListArray>> {
let Some(dataset) = self.dataset.as_ref() else {
return Err(Error::invalid_input(
"dataset not set before sample partition",
));
};
let mut row_ids = self.partition_row_ids(part_idx).await?;
if !row_ids.is_sorted() {
row_ids.sort();
}
row_ids.dedup();
if row_ids.is_empty() {
return Ok(None);
}
if row_ids.len() > sample_size {
let step = row_ids.len() / sample_size;
row_ids = row_ids
.iter()
.copied()
.step_by(step.max(1))
.take(sample_size)
.collect();
}
let batches = Self::take_vectors(dataset, &self.column, &self.store, &row_ids).await?;
if batches.is_empty() {
return Ok(None);
}
let batch = arrow::compute::concat_batches(&batches[0].schema(), batches.iter())?;
let batch = Flatten::new(&self.column).transform(&batch)?;
let vectors = batch
.column_by_qualified_name(&self.column)
.ok_or(Error::invalid_input(format!(
"vector column {} not found in batch",
self.column,
)))?
.as_fixed_size_list()
.clone();
Ok(Some(vectors))
}
async fn train_split_centroids<T: ArrowPrimitiveType>(
&self,
split: PartitionSplit,
ivf: &IvfModel,
) -> Result<Option<Vec<ArrayRef>>>
where
T::Native: Dot + L2 + Normalize,
PrimitiveArray<T>: From<Vec<T::Native>>,
{
let Some(vectors) = self
.sample_partition_raw_vectors(split.partition, SPLIT_SAMPLE_RATE * split.ways)
.await?
else {
return Ok(None);
};
let ways = split.ways.min(vectors.len());
if ways < 2 {
return Ok(None);
}
let dimension = infer_vector_dim(vectors.data_type())?;
let (normalized_dist_type, normalized_vectors) = match self.distance_type {
DistanceType::Cosine => {
let vectors = normalize_fsl(&vectors)?;
(DistanceType::L2, vectors)
}
_ => (self.distance_type, vectors),
};
let params =
KMeansParams::new(None, 50, 1, normalized_dist_type).with_seed(split.partition as u64);
let values = normalized_vectors.values().as_primitive::<T>();
let kmeans = match lance_index::vector::kmeans::train_kmeans::<T>(
values,
params.clone(),
dimension,
ways,
SPLIT_SAMPLE_RATE,
) {
Ok(kmeans) => kmeans,
Err(err) => {
log::info!(
"retrying the {ways}-way split of partition {} with flat k-means: {err}",
split.partition
);
lance_index::vector::kmeans::train_kmeans::<T>(
values,
params.with_hierarchical_k(1),
dimension,
ways,
SPLIT_SAMPLE_RATE,
)?
}
};
let trained_centroids =
FixedSizeListArray::try_new_from_values(kmeans.centroids.clone(), dimension as i32)?;
let (membership, new_dists) = lance_index::vector::kmeans::compute_partitions_arrow_array(
&trained_centroids,
&normalized_vectors,
normalized_dist_type,
)?;
let c0 = ivf
.centroid(split.partition)
.ok_or(Error::invalid_input("original centroid not found"))?;
let (_, neighbor_centroids) =
self.select_reassign_candidates(ivf, split.partition, &c0, &HashSet::new())?;
let (_, old_dists) = lance_index::vector::kmeans::compute_partitions_arrow_array(
&neighbor_centroids,
&normalized_vectors,
normalized_dist_type,
)?;
let mut attracts_rows = vec![false; ways];
for ((cluster, new_dist), old_dist) in membership.into_iter().zip(new_dists).zip(old_dists)
{
let (Some(cluster), Some(new_dist)) = (cluster, new_dist) else {
continue;
};
if old_dist.is_none_or(|old_dist| new_dist < old_dist) {
attracts_rows[cluster as usize] = true;
}
}
let centroids: Vec<ArrayRef> = (0..ways)
.filter(|&i| attracts_rows[i])
.map(|i| kmeans.centroids.slice(i * dimension, dimension))
.collect();
if centroids.len() < 2 {
log::warn!(
"partition {} is oversized but its sampled rows are too alike to split; leaving it as is",
split.partition
);
return Ok(None);
}
Ok(Some(centroids))
}
async fn join_partitions(
&self,
partitions: &[usize],
ivf: &IvfModel,
partition_sizes: &[usize],
split_threshold: usize,
) -> Result<JoinResult> {
let Some(dataset) = self.dataset.as_ref() else {
return Err(Error::invalid_input(
"dataset not set before joining partitions",
));
};
let mut reindexed_row_ids = HashSet::new();
let mut rows_per_partition = Vec::with_capacity(partitions.len());
for &partition in partitions {
let mut row_ids = self.partition_row_ids(partition).await?;
row_ids.sort_unstable();
row_ids.dedup();
row_ids.retain(|row_id| reindexed_row_ids.insert(*row_id));
rows_per_partition.push(row_ids);
}
let reindexed_row_ids = Arc::new(reindexed_row_ids);
let removed: HashSet<usize> = partitions.iter().copied().collect();
let kept_partitions: Vec<usize> = (0..ivf.num_partitions())
.filter(|partition| !removed.contains(partition))
.collect();
let centroids = ivf
.centroids_array()
.ok_or_else(|| Error::invalid_input("IVF model has no centroids"))?;
let kept_ids = UInt32Array::from_iter_values(
kept_partitions.iter().map(|&partition| partition as u32),
);
let new_centroids = arrow::compute::take(centroids, &kept_ids, None)?
.as_fixed_size_list()
.clone();
let column_type = dataset
.schema()
.field(&self.column)
.map(|field| field.data_type())
.ok_or_else(|| Error::invalid_input(format!("column {} not found", self.column)))?;
let multivector = matches!(column_type, DataType::List(_) | DataType::LargeList(_));
let mut room = vec![0usize; ivf.num_partitions()];
for &partition in &kept_partitions {
let mut load = partition_sizes[partition];
if multivector {
load -= self
.partition_row_ids(partition)
.await?
.iter()
.filter(|row_id| reindexed_row_ids.contains(row_id))
.count();
}
room[partition] = split_threshold.saturating_sub(load);
}
let (_, element_type) = get_vector_type(dataset.schema(), &self.column)?;
let join_shuffle_reader = match element_type {
DataType::Float16 => {
self.join_partitions_impl::<Float16Type>(
partitions,
&rows_per_partition,
ivf,
&removed,
&mut room,
multivector,
)
.await?
}
DataType::Float32 => {
self.join_partitions_impl::<Float32Type>(
partitions,
&rows_per_partition,
ivf,
&removed,
&mut room,
multivector,
)
.await?
}
DataType::Float64 => {
self.join_partitions_impl::<Float64Type>(
partitions,
&rows_per_partition,
ivf,
&removed,
&mut room,
multivector,
)
.await?
}
DataType::UInt8 => {
self.join_partitions_impl::<UInt8Type>(
partitions,
&rows_per_partition,
ivf,
&removed,
&mut room,
multivector,
)
.await?
}
dt => {
return Err(Error::invalid_input(format!(
"vectors must be float16, float32, float64 or uint8, but got {:?}",
dt
)));
}
};
Ok(JoinResult {
new_centroids,
kept_partitions,
reindexed_row_ids,
join_shuffle_reader,
})
}
async fn join_partitions_impl<T: ArrowPrimitiveType>(
&self,
partitions: &[usize],
rows_per_partition: &[Vec<u64>],
ivf: &IvfModel,
removed: &HashSet<usize>,
room: &mut [usize],
multivector: bool,
) -> Result<Arc<dyn ShuffleReader>>
where
T::Native: Dot + L2 + Normalize,
PrimitiveArray<T>: From<Vec<T::Native>>,
{
let Some(dataset) = self.dataset.as_ref() else {
return Err(Error::invalid_input(
"dataset not set before joining partitions",
));
};
let old_centroids = ivf
.centroids_array()
.ok_or_else(|| Error::invalid_input("IVF model has no centroids"))?;
let (transformer, vector_field) = self.assign_transformer(old_centroids)?;
let num_partitions = ivf.num_partitions();
let shuffle_dir = self.temp_dir.clone().join("join_shuffle");
let shuffler = create_ivf_shuffler(
shuffle_dir.clone(),
num_partitions,
self.format_version,
None,
);
let store = Arc::new(self.store.clone());
let (mut sender, mut receiver) = futures::channel::mpsc::channel::<RecordBatch>(2);
let shuffle = async move {
let Some(first) = receiver.next().await else {
let empty: Box<dyn ShuffleReader> = Box::new(IvfShufflerReader::new(
store,
shuffle_dir,
vec![0; num_partitions],
0.0,
));
return Ok(empty);
};
let schema = first.schema();
let batches = stream::iter(std::iter::once(Ok(first))).chain(receiver.map(Ok));
shuffler
.shuffle(Box::new(RecordBatchStreamAdapter::new(schema, batches)))
.await
};
let (rows_per_fetch, prefetch) = if multivector {
(1, JOIN_MULTIVECTOR_FETCHES_IN_FLIGHT)
} else {
let row_bytes = ivf.dimension() * vector_field_element_width(&vector_field);
(
(JOIN_FETCH_BYTES / row_bytes.max(1)).clamp(1, self.store.block_size()),
2,
)
};
let route = async move {
for (&part_idx, row_ids) in partitions.iter().zip(rows_per_partition) {
if row_ids.is_empty() {
continue;
}
let c0 = ivf
.centroid(part_idx)
.ok_or(Error::invalid_input("original centroid not found"))?;
let partition_window =
self.select_reassign_candidates(ivf, part_idx, &c0, removed)?;
let mut overflowed = 0usize;
let fetched = Self::take_vectors_stream(
dataset,
&self.column,
row_ids,
rows_per_fetch,
prefetch,
)
.await?;
let mut chunks = Box::pin(regroup_by_bytes(fetched, JOIN_FETCH_BYTES));
while let Some(chunk) = chunks.try_next().await? {
let (row_ids, vectors) = self.flatten_raw_vectors(&chunk)?;
for start in (0..row_ids.len()).step_by(JOIN_VECTORS_PER_BATCH) {
let end = (start + JOIN_VECTORS_PER_BATCH).min(row_ids.len());
let mut assign_ops: BTreeMap<u32, Vec<AssignOp>> = BTreeMap::new();
for i in start..end {
let vector = vectors.value(i);
let vector_window;
let (window_ids, window_centroids) = if multivector {
vector_window =
window_for_vector(self.distance_type, ivf, &vector, removed)?;
&vector_window
} else {
&partition_window
};
let (target, had_room) = choose_join_destination(
self.distance_type,
vector.as_primitive::<T>(),
window_ids,
window_centroids,
ivf,
removed,
room,
)?;
if had_room {
room[target as usize] -= 1;
} else {
overflowed += 1;
}
assign_ops
.entry(target)
.or_default()
.push(AssignOp::Add((row_ids.value(i), vector)));
}
for (target, ops) in assign_ops {
let batch = Self::build_assign_batch::<T>(
&transformer,
&vector_field,
target,
&ops,
)?;
sender.send(batch).await.map_err(|_| {
Error::internal("the join shuffler stopped taking batches")
})?;
}
}
}
if overflowed > 0 {
log::warn!(
"{overflowed} rows of joined partition {part_idx} found no partition below the split threshold and went to the nearest full one"
);
}
}
drop(sender);
Ok(())
};
let (reader, ()) = futures::try_join!(shuffle, route)?;
Ok(reader.into())
}
fn assign_transformer(
&self,
centroids: &FixedSizeListArray,
) -> Result<(IvfTransformer, Field)> {
let Some(dataset) = self.dataset.as_ref() else {
return Err(Error::invalid_input(
"dataset not set before building assign batch",
));
};
let Some(quantizer) = self.quantizer.clone() else {
return Err(Error::invalid_input(
"quantizer not set before building assign batch",
));
};
let Some(vector_field) =
dataset
.schema()
.field(&self.column)
.map(|f| match f.data_type() {
DataType::List(inner) | DataType::LargeList(inner) => {
Field::new(self.column.as_str(), inner.data_type().clone(), true)
}
_ => f.into(),
})
else {
return Err(Error::invalid_input(
"vector field not found in dataset schema",
));
};
let transformer = lance_index::vector::ivf::new_ivf_transformer_with_quantizer(
centroids.clone(),
self.distance_type,
vector_field.name().as_str(),
quantizer.into(),
None,
)?;
Ok((transformer, vector_field))
}
fn build_assign_batch<T: ArrowPrimitiveType>(
transformer: &IvfTransformer,
vector_field: &Field,
part_id: u32,
ops: &[AssignOp],
) -> Result<RecordBatch> {
let dimension = infer_vector_dim(vector_field.data_type())?;
let num_rows = ops.len();
let mut row_ids_builder = UInt64Builder::with_capacity(num_rows);
let mut vector_builder = PrimitiveBuilder::<T>::with_capacity(num_rows * dimension);
for AssignOp::Add((row_id, vector)) in ops {
row_ids_builder.append_value(*row_id);
vector_builder.append_array(vector.as_primitive::<T>());
}
let row_ids = row_ids_builder.finish();
let vector =
FixedSizeListArray::try_new_from_values(vector_builder.finish(), dimension as i32)?;
let part_ids = UInt32Array::from(vec![part_id; num_rows]);
let schema = arrow_schema::Schema::new(vec![
ROW_ID_FIELD.clone(),
vector_field.clone(),
PART_ID_FIELD.clone(),
]);
let batch = RecordBatch::try_new(
Arc::new(schema),
vec![Arc::new(row_ids), Arc::new(vector), Arc::new(part_ids)],
)?;
transformer.transform(&batch)
}
async fn partition_row_ids(&self, part_idx: usize) -> Result<Vec<u64>> {
let mut row_ids = Vec::new();
for source in self.existing_indices.iter() {
let index = &source.index;
if part_idx >= index.ivf_model().num_partitions() {
log::warn!(
"partition index is {} but the number of partitions is {}, skip loading it",
part_idx,
index.ivf_model().num_partitions()
);
continue;
}
let mut reader = index
.partition_reader(part_idx, false, &NoOpMetricsCollector)
.await?;
let old_data_filter = source.old_data_filter().await?;
while let Some(batch) = reader.try_next().await? {
let batch_row_ids = batch[ROW_ID].as_primitive::<UInt64Type>();
match old_data_filter {
Some(filter) => row_ids.extend(
batch_row_ids
.values()
.iter()
.zip(filter.filter_row_ids(batch_row_ids).values().iter())
.filter_map(|(row_id, keep)| keep.then_some(row_id)),
),
None => row_ids.extend(batch_row_ids.values()),
}
}
}
if let Some(reader) = self.shuffle_reader.as_ref() {
if let Some(mut reader) = reader.read_partition(part_idx).await? {
while let Some(batch) = reader.try_next().await? {
row_ids.extend(batch[ROW_ID].as_primitive::<UInt64Type>().values());
}
}
}
Ok(row_ids)
}
fn select_reassign_candidates(
&self,
ivf: &IvfModel,
part_idx: usize,
c0: &ArrayRef,
excluded: &HashSet<usize>,
) -> Result<(UInt32Array, FixedSizeListArray)> {
select_reassign_candidates_impl(self.distance_type, ivf, part_idx, c0, excluded)
}
}
fn plan_partition_adjustment(
partition_sizes: &[usize],
target_partition_size: usize,
) -> (Vec<PartitionSplit>, Vec<usize>) {
let split_threshold = MAX_PARTITION_SIZE_FACTOR * target_partition_size;
let join_threshold = MIN_PARTITION_SIZE_PERCENT * target_partition_size / 100;
let num_partitions = partition_sizes.len();
let splits = partition_sizes
.iter()
.enumerate()
.filter(|&(_, &num_rows)| num_rows > split_threshold)
.map(|(partition, &num_rows)| PartitionSplit {
partition,
ways: num_rows
.div_ceil(target_partition_size)
.clamp(2, MAX_SPLIT_WAYS),
})
.collect();
if num_partitions < 2 {
return (splits, Vec::new());
}
let total_rows: usize = partition_sizes.iter().sum();
let max_joins = num_partitions.saturating_sub(total_rows.div_ceil(split_threshold).max(1));
let mut candidates: Vec<(usize, usize)> = partition_sizes
.iter()
.enumerate()
.filter(|&(_, &num_rows)| num_rows < join_threshold)
.map(|(partition, &num_rows)| (num_rows, partition))
.collect();
candidates.sort_unstable();
let mut joins: Vec<usize> = candidates
.into_iter()
.take(max_joins)
.map(|(_, partition)| partition)
.collect();
joins.sort_unstable();
(splits, joins)
}
fn nearest_candidate<T: ArrowPrimitiveType>(
distance_type: DistanceType,
vector: &PrimitiveArray<T>,
candidate_ids: &UInt32Array,
candidate_centroids: &FixedSizeListArray,
accept: impl Fn(u32) -> bool,
) -> Result<Option<u32>>
where
T::Native: Dot + L2 + Normalize,
{
let dists = distance_type.arrow_batch_func()(vector, candidate_centroids)?;
Ok(dists
.values()
.iter()
.enumerate()
.filter(|&(i, _)| accept(candidate_ids.value(i)))
.min_by(|(_, a), (_, b)| a.total_cmp(b))
.map(|(i, _)| candidate_ids.value(i)))
}
fn choose_join_destination<T: ArrowPrimitiveType>(
distance_type: DistanceType,
vector: &PrimitiveArray<T>,
window_ids: &UInt32Array,
window_centroids: &FixedSizeListArray,
ivf: &IvfModel,
removed: &HashSet<usize>,
room: &[usize],
) -> Result<(u32, bool)>
where
T::Native: Dot + L2 + Normalize,
{
let has_room = |id: u32| room[id as usize] > 0;
if let Some(target) = nearest_candidate(
distance_type,
vector,
window_ids,
window_centroids,
has_room,
)? {
return Ok((target, true));
}
let anywhere_ids = UInt32Array::from(
(0..ivf.num_partitions())
.filter(|&id| !removed.contains(&id) && room[id] > 0)
.map(|id| id as u32)
.collect::<Vec<_>>(),
);
if !anywhere_ids.is_empty() {
let centroids = ivf
.centroids_array()
.ok_or_else(|| Error::invalid_input("IVF model has no centroids"))?;
let anywhere_centroids = arrow::compute::take(centroids, &anywhere_ids, None)?
.as_fixed_size_list()
.clone();
if let Some(target) = nearest_candidate(
distance_type,
vector,
&anywhere_ids,
&anywhere_centroids,
|_| true,
)? {
return Ok((target, true));
}
}
let target = nearest_candidate(distance_type, vector, window_ids, window_centroids, |_| {
true
})?
.ok_or_else(|| Error::internal("a joined partition has no neighbor to receive its rows"))?;
Ok((target, false))
}
fn regroup_by_bytes<S>(batches: S, budget: usize) -> impl Stream<Item = Result<RecordBatch>>
where
S: Stream<Item = Result<RecordBatch>> + Unpin,
{
stream::try_unfold(
(batches, Vec::<RecordBatch>::new(), 0usize),
move |(mut batches, mut pending, mut pending_bytes)| async move {
loop {
let Some(batch) = batches.try_next().await? else {
if pending.is_empty() {
return Ok(None);
}
let grouped = arrow::compute::concat_batches(&pending[0].schema(), &pending)?;
return Ok(Some((grouped, (batches, Vec::new(), 0))));
};
let bytes = batch.get_array_memory_size();
if !pending.is_empty() && pending_bytes + bytes > budget {
let grouped = arrow::compute::concat_batches(&pending[0].schema(), &pending)?;
return Ok(Some((grouped, (batches, vec![batch], bytes))));
}
pending.push(batch);
pending_bytes += bytes;
}
},
)
}
fn vector_field_element_width(field: &Field) -> usize {
match field.data_type() {
DataType::FixedSizeList(item, _) => item.data_type().primitive_width().unwrap_or(4),
_ => 4,
}
}
fn window_for_vector(
distance_type: DistanceType,
ivf: &IvfModel,
vector: &ArrayRef,
removed: &HashSet<usize>,
) -> Result<(UInt32Array, FixedSizeListArray)> {
select_reassign_candidates_impl(distance_type, ivf, usize::MAX, vector, removed)
}
fn select_reassign_candidates_impl(
distance_type: DistanceType,
ivf: &IvfModel,
part_idx: usize,
c0: &ArrayRef,
excluded: &HashSet<usize>,
) -> Result<(UInt32Array, FixedSizeListArray)> {
let centroids = ivf.centroids_array().unwrap();
let centroid_dists = distance_type.arrow_batch_func()(c0, centroids)?;
let fetch = (REASSIGN_RANGE + excluded.len() + 1).min(ivf.num_partitions());
let nearest = sort_to_indices(centroid_dists.as_ref(), None, Some(fetch))?;
let filtered_ids = nearest
.values()
.iter()
.copied()
.filter(|&idx| idx as usize != part_idx && !excluded.contains(&(idx as usize)))
.take(REASSIGN_RANGE)
.collect::<Vec<_>>();
let reassign_candidate_ids = UInt32Array::from(filtered_ids);
let reassign_candidate_centroids =
arrow::compute::take(centroids, &reassign_candidate_ids, None)?;
Ok((
reassign_candidate_ids,
reassign_candidate_centroids.as_fixed_size_list().clone(),
))
}
struct PartitionAdjustmentPlan {
splits: Vec<PartitionSplit>,
joins: Vec<usize>,
partition_sizes: Vec<usize>,
}
struct JoinResult {
new_centroids: FixedSizeListArray,
kept_partitions: Vec<usize>,
reindexed_row_ids: Arc<HashSet<u64>>,
join_shuffle_reader: Arc<dyn ShuffleReader>,
}
struct SplitResult {
new_centroids: FixedSizeListArray,
affected_partitions: HashSet<usize>,
shuffle_reader: Arc<dyn ShuffleReader>,
}
#[derive(Debug, Clone)]
enum AssignOp {
Add((u64, ArrayRef)),
}
enum PartitionAdjustment {
Split {
affected_partitions: HashSet<usize>,
split_shuffle_reader: Arc<dyn ShuffleReader>,
},
Join {
kept_partitions: Vec<usize>,
reindexed_row_ids: Arc<HashSet<u64>>,
join_shuffle_reader: Arc<dyn ShuffleReader>,
},
}
impl std::fmt::Debug for PartitionAdjustment {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::Split {
affected_partitions,
..
} => f
.debug_struct("Split")
.field("affected_partitions", affected_partitions)
.finish(),
Self::Join {
kept_partitions,
reindexed_row_ids,
..
} => f
.debug_struct("Join")
.field("kept_partitions", &kept_partitions.len())
.field("reindexed_rows", &reindexed_row_ids.len())
.finish(),
}
}
}
pub(crate) fn index_type_string(sub_index: SubIndexType, quantizer: QuantizationType) -> String {
let quantizer = match quantizer {
QuantizationType::FlatBin => QuantizationType::Flat,
other => other,
};
match (sub_index, quantizer) {
(SubIndexType::Flat, quantization_type) => format!("IVF_{}", quantization_type),
(sub_index_type, quantization_type) => {
if sub_index_type.to_string() == quantization_type.to_string() {
format!("IVF_{}", sub_index_type)
} else {
format!("IVF_{}_{}", sub_index_type, quantization_type)
}
}
}
}
#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicUsize, Ordering};
use super::*;
use arrow_array::{Array, Float32Array, NullArray};
use lance_index::vector::flat::index::{FlatIndex, FlatQuantizer};
use lance_index::vector::v3::shuffler::{
ShufflePartition, ShufflePartitionWindow, ShufflePartitionWindowPlan,
};
struct SingleBatchReader {
batch: RecordBatch,
partition_id: usize,
}
#[async_trait::async_trait]
impl ShuffleReader for SingleBatchReader {
async fn read_partition(
&self,
partition_id: usize,
) -> Result<Option<Box<dyn RecordBatchStream + Unpin + 'static>>> {
if partition_id != self.partition_id || self.batch.num_rows() == 0 {
return Ok(None);
}
let schema = self.batch.schema();
let stream = stream::iter(vec![Ok(self.batch.clone())]);
Ok(Some(Box::new(RecordBatchStreamAdapter::new(
schema, stream,
))))
}
fn partition_size(&self, partition_id: usize) -> Result<usize> {
Ok(if partition_id == self.partition_id {
self.batch.num_rows()
} else {
0
})
}
fn total_loss(&self) -> Option<f64> {
None
}
}
struct WindowedBatchReader {
batches: Vec<RecordBatch>,
windows_read: Arc<AtomicUsize>,
}
#[async_trait::async_trait]
impl ShuffleReader for WindowedBatchReader {
async fn read_partition(
&self,
partition_id: usize,
) -> Result<Option<Box<dyn RecordBatchStream + Unpin + 'static>>> {
let Some(batch) = self.batches.get(partition_id) else {
return Ok(None);
};
Ok(Some(Box::new(RecordBatchStreamAdapter::new(
batch.schema(),
stream::iter(vec![Ok(batch.clone())]),
))))
}
fn plan_partition_window(
&self,
start_partition_id: usize,
max_decoded_bytes: usize,
) -> Result<ShufflePartitionWindowPlan> {
if max_decoded_bytes == 0 {
return Err(Error::invalid_input(
"max_decoded_bytes must be greater than 0",
));
}
if start_partition_id >= self.batches.len() {
return Err(Error::invalid_input(format!(
"start_partition_id={} is out of range [0, {})",
start_partition_id,
self.batches.len()
)));
}
let end_partition_id = start_partition_id
.saturating_add(max_decoded_bytes)
.min(self.batches.len());
Ok(ShufflePartitionWindowPlan {
partition_range: start_partition_id..end_partition_id,
estimated_decoded_bytes: end_partition_id - start_partition_id,
})
}
async fn read_partition_window(
&self,
start_partition_id: usize,
max_decoded_bytes: usize,
) -> Result<ShufflePartitionWindow> {
let plan = self.plan_partition_window(start_partition_id, max_decoded_bytes)?;
self.windows_read.fetch_add(1, Ordering::Relaxed);
if start_partition_id == 0 {
tokio::time::timeout(std::time::Duration::from_secs(1), async {
while self.windows_read.load(Ordering::Relaxed) < 2 {
tokio::task::yield_now().await;
}
})
.await
.map_err(|_| Error::internal("second partition window was not admitted"))?;
}
let partitions = plan
.partition_range
.clone()
.map(|partition_id| {
let batch = self.batches[partition_id].clone();
ShufflePartition {
partition_id,
data: Some(Box::new(RecordBatchStreamAdapter::new(
batch.schema(),
stream::iter(vec![Ok(batch)]),
))),
}
})
.collect();
Ok(ShufflePartitionWindow {
materialized_decoded_bytes: Some(plan.partition_range.len()),
partition_range: plan.partition_range,
partitions,
})
}
fn partition_size(&self, partition_id: usize) -> Result<usize> {
Ok(self
.batches
.get(partition_id)
.map(RecordBatch::num_rows)
.unwrap_or(0))
}
fn total_loss(&self) -> Option<f64> {
None
}
}
fn flat_partition_batch(partition_id: usize) -> RecordBatch {
let vectors = FixedSizeListArray::try_new_from_values(
Float32Array::from(vec![partition_id as f32, partition_id as f32 + 0.5]),
2,
)
.unwrap();
RecordBatch::try_new(
Arc::new(arrow_schema::Schema::new(vec![
ROW_ID_FIELD.clone(),
Field::new("vector", vectors.data_type().clone(), false),
])),
vec![
Arc::new(UInt64Array::from(vec![partition_id as u64])),
Arc::new(vectors),
],
)
.unwrap()
}
fn centroid_values(arr: &FixedSizeListArray, i: usize) -> Vec<f32> {
arr.value(i).as_primitive::<Float32Type>().values().to_vec()
}
#[tokio::test]
async fn partition_entry_admission_preserves_input_order() {
let entry_permits = Arc::new(Semaphore::new(1));
let held_permit = entry_permits.clone().acquire_owned().await.unwrap();
let mut admitted = admit_partition_inputs(vec![0, 1, 2], entry_permits);
let first = admitted.next();
tokio::pin!(first);
assert!(
tokio::time::timeout(std::time::Duration::from_millis(10), &mut first)
.await
.is_err()
);
drop(held_permit);
let (partition_id, first_permit) =
tokio::time::timeout(std::time::Duration::from_millis(100), &mut first)
.await
.unwrap()
.unwrap()
.unwrap();
assert_eq!(partition_id, 0);
let second = admitted.next();
tokio::pin!(second);
assert!(
tokio::time::timeout(std::time::Duration::from_millis(10), &mut second)
.await
.is_err()
);
drop(first_permit);
let (partition_id, _second_permit) =
tokio::time::timeout(std::time::Duration::from_millis(100), &mut second)
.await
.unwrap()
.unwrap()
.unwrap();
assert_eq!(partition_id, 1);
}
#[test]
fn single_partition_window_uses_full_entry_budget() {
assert_eq!(partition_window_entry_limit(&(0..64), 64, 32, 16), 32);
assert_eq!(partition_window_entry_limit(&(0..32), 64, 32, 16), 2);
assert_eq!(partition_window_entry_limit(&(32..64), 64, 32, 16), 2);
}
#[tokio::test]
async fn fresh_partition_build_runs_multiple_windows_end_to_end() {
let num_partitions = 6;
let windows_read = Arc::new(AtomicUsize::new(0));
let reader = Arc::new(WindowedBatchReader {
batches: (0..num_partitions).map(flat_partition_batch).collect(),
windows_read: windows_read.clone(),
});
let mut build_stream =
IvfIndexBuilder::<FlatIndex, FlatQuantizer>::build_fresh_partitions_windowed(
reader,
num_partitions,
DistanceType::L2,
FlatQuantizer::new(2, DistanceType::L2),
(),
"vector".to_string(),
None,
FreshPartitionBuildLimits {
window_bytes: 2,
decoded_budget_bytes: 4,
},
)
.unwrap();
let mut ordered_results = OrderedPartitionResults::new(num_partitions);
let mut merged_partition_ids = Vec::with_capacity(num_partitions);
while let Some(result) = build_stream.try_next().await.unwrap() {
ordered_results
.push(result.value.partition_id, result)
.unwrap();
while let Some((partition_id, result)) = ordered_results.pop_next() {
assert!(result.value.built.is_some());
merged_partition_ids.push(partition_id);
}
}
ordered_results.finish().unwrap();
assert_eq!(windows_read.load(Ordering::Relaxed), 3);
assert_eq!(
merged_partition_ids,
(0..num_partitions).collect::<Vec<_>>()
);
}
#[test]
fn apply_centroid_splits_correct_count_and_ordering() {
let original = FixedSizeListArray::try_new_from_values(
Float32Array::from(vec![0.0_f32, 0.0, 1.0, 1.0, 2.0, 2.0, 3.0, 3.0]),
2,
)
.unwrap();
let c1_for_1: ArrayRef = Arc::new(Float32Array::from(vec![1.1_f32, 1.1]));
let c2_for_1: ArrayRef = Arc::new(Float32Array::from(vec![0.9_f32, 0.9]));
let c3_for_1: ArrayRef = Arc::new(Float32Array::from(vec![1.2_f32, 1.2]));
let c1_for_3: ArrayRef = Arc::new(Float32Array::from(vec![3.1_f32, 3.1]));
let c2_for_3: ArrayRef = Arc::new(Float32Array::from(vec![2.9_f32, 2.9]));
let splits = vec![
(1_usize, vec![c1_for_1, c2_for_1, c3_for_1]),
(3_usize, vec![c1_for_3, c2_for_3]),
];
let result = apply_centroid_splits(&original, &splits).unwrap();
assert_eq!(result.len(), 7);
assert_eq!(centroid_values(&result, 0), [0.0, 0.0]);
assert_eq!(centroid_values(&result, 2), [2.0, 2.0]);
assert_eq!(centroid_values(&result, 1), [1.1, 1.1]);
assert_eq!(centroid_values(&result, 3), [3.1, 3.1]);
assert_eq!(centroid_values(&result, 4), [0.9, 0.9]);
assert_eq!(centroid_values(&result, 5), [1.2, 1.2]);
assert_eq!(centroid_values(&result, 6), [2.9, 2.9]);
}
#[test]
fn select_reassign_candidates_skips_deleted_partition() {
let dim = 4;
let centroid_values = Float32Array::from(vec![0.0_f32; dim * 2]);
let centroids =
FixedSizeListArray::try_new_from_values(centroid_values, dim as i32).unwrap();
let mut ivf = IvfModel::new(centroids, None);
ivf.lengths = vec![10, 20];
ivf.offsets = vec![0, 10];
let c0 = ivf.centroid(1).unwrap();
let (reassign_ids, reassign_centroids) =
select_reassign_candidates_impl(DistanceType::L2, &ivf, 1, &c0, &HashSet::new())
.unwrap();
assert_eq!(reassign_ids.len(), 1);
assert_eq!(reassign_ids.value(0), 0);
assert_eq!(reassign_centroids.len(), 1);
let expected_centroid = ivf.centroid(0).unwrap();
assert_eq!(
reassign_centroids
.value(0)
.as_primitive::<Float32Type>()
.values(),
expected_centroid.as_primitive::<Float32Type>().values()
);
}
#[tokio::test]
async fn optimize_split_after_append_pushes_partition_over_threshold() {
use crate::dataset::{InsertBuilder, WriteMode, WriteParams};
use crate::index::vector::VectorIndexParams;
use crate::index::{DatasetIndexExt, DatasetIndexInternalExt};
use arrow_array::RecordBatchIterator;
use arrow_schema::Schema as ArrowSchema;
use lance_index::optimize::OptimizeOptions;
use lance_linalg::distance::MetricType;
let item_field = Arc::new(Field::new("item", DataType::Float32, true));
let schema = Arc::new(ArrowSchema::new(vec![Field::new(
"vec",
DataType::FixedSizeList(item_field, 4),
false,
)]));
let make_batch = |num_rows: usize, center: f32| -> RecordBatch {
let mut values = Vec::with_capacity(num_rows * 4);
for i in 0..num_rows {
let p = center + (i as f32) * 0.0001;
values.extend_from_slice(&[p, p, p, p]);
}
let fsl =
FixedSizeListArray::try_new_from_values(Float32Array::from(values), 4).unwrap();
RecordBatch::try_new(schema.clone(), vec![Arc::new(fsl)]).unwrap()
};
let tmp = tempfile::tempdir().unwrap();
let uri = tmp.path().to_str().unwrap();
let initial = vec![Ok(make_batch(15_000, 0.0)), Ok(make_batch(1_500, 1000.0))];
let reader = RecordBatchIterator::new(initial, schema.clone());
let mut dataset = crate::Dataset::write(reader, uri, None).await.unwrap();
let params = VectorIndexParams::ivf_flat(2, MetricType::L2);
dataset
.create_index(
&["vec"],
IndexType::Vector,
Some("idx".into()),
¶ms,
false,
)
.await
.unwrap();
let indices = dataset.load_indices_by_name("idx").await.unwrap();
let initial_index = dataset
.open_vector_index("vec", &indices[0].uuid, &NoOpMetricsCollector)
.await
.unwrap();
let initial_ivf = initial_index.ivf_model();
assert_eq!(initial_ivf.num_partitions(), 2);
let max_initial = (0..2).map(|p| initial_ivf.partition_size(p)).max().unwrap();
assert!(
max_initial <= 16_384,
"initial max partition size {max_initial} should be at or under split threshold",
);
let append = make_batch(3_000, 0.0);
let mut dataset = InsertBuilder::new(Arc::new(dataset))
.with_params(&WriteParams {
mode: WriteMode::Append,
..Default::default()
})
.execute(vec![append])
.await
.unwrap();
dataset
.optimize_indices(&OptimizeOptions::default())
.await
.unwrap();
let indices = dataset.load_indices_by_name("idx").await.unwrap();
assert_eq!(indices.len(), 1, "expected merge-all on split");
let optimized = dataset
.open_vector_index("vec", &indices[0].uuid, &NoOpMetricsCollector)
.await
.unwrap();
let ivf = optimized.ivf_model();
assert_eq!(
ivf.num_partitions(),
6,
"expected one 5-way split: 2 original + 4 new = 6 partitions",
);
let total_rows: usize = (0..ivf.num_partitions())
.map(|p| optimized.partition_size(p))
.sum();
assert_eq!(total_rows, 19_500, "all vectors preserved across split");
}
async fn nearest_first_components(dataset: &crate::Dataset, center: f32, k: usize) -> Vec<f32> {
let query = Float32Array::from(vec![center; 4]);
let batch = dataset
.scan()
.nearest("vec", &query, k)
.unwrap()
.minimum_nprobes(1)
.try_into_batch()
.await
.unwrap();
let vectors = batch["vec"].as_fixed_size_list();
(0..vectors.len())
.map(|i| vectors.value(i).as_primitive::<Float32Type>().value(0))
.collect()
}
fn cluster_batch(
schema: &Arc<arrow_schema::Schema>,
num_rows: usize,
center: f32,
) -> RecordBatch {
let mut values = Vec::with_capacity(num_rows * 4);
for i in 0..num_rows {
let p = center + (i as f32) * 0.0001;
values.extend_from_slice(&[p, p, p, p]);
}
let fsl = FixedSizeListArray::try_new_from_values(Float32Array::from(values), 4).unwrap();
RecordBatch::try_new(schema.clone(), vec![Arc::new(fsl)]).unwrap()
}
fn ivf_flat_at_centers(
centers: &[f32],
target_partition_size: Option<usize>,
) -> crate::index::vector::VectorIndexParams {
use crate::index::vector::VectorIndexParams;
use lance_linalg::distance::MetricType;
let values: Vec<f32> = centers.iter().flat_map(|&c| [c; 4]).collect();
let centroids =
FixedSizeListArray::try_new_from_values(Float32Array::from(values), 4).unwrap();
let mut ivf_params =
IvfBuildParams::try_with_centroids(centers.len(), Arc::new(centroids)).unwrap();
ivf_params.target_partition_size = target_partition_size;
VectorIndexParams::with_ivf_flat_params(MetricType::L2, ivf_params)
}
fn cluster_schema() -> Arc<arrow_schema::Schema> {
let item_field = Arc::new(Field::new("item", DataType::Float32, true));
Arc::new(arrow_schema::Schema::new(vec![Field::new(
"vec",
DataType::FixedSizeList(item_field, 4),
false,
)]))
}
async fn write_clusters(uri: &str, clusters: &[(usize, f32)]) -> crate::Dataset {
use arrow_array::RecordBatchIterator;
let schema = cluster_schema();
let batches = clusters
.iter()
.map(|(rows, center)| Ok(cluster_batch(&schema, *rows, *center)))
.collect::<Vec<_>>();
let reader = RecordBatchIterator::new(batches, schema);
crate::Dataset::write(reader, uri, None).await.unwrap()
}
async fn append_cluster(dataset: crate::Dataset, rows: usize, center: f32) -> crate::Dataset {
use crate::dataset::{InsertBuilder, WriteMode, WriteParams};
let batch = cluster_batch(&cluster_schema(), rows, center);
InsertBuilder::new(Arc::new(dataset))
.with_params(&WriteParams {
mode: WriteMode::Append,
..Default::default()
})
.execute(vec![batch])
.await
.unwrap()
}
async fn open_single_segment(dataset: &crate::Dataset) -> Arc<dyn VectorIndex> {
use crate::index::{DatasetIndexExt, DatasetIndexInternalExt};
let indices = dataset.load_indices_by_name("idx").await.unwrap();
assert_eq!(
indices.len(),
1,
"expected the rebalance to merge into one segment"
);
dataset
.open_vector_index("vec", &indices[0].uuid, &NoOpMetricsCollector)
.await
.unwrap()
}
#[tokio::test]
async fn optimize_splits_oversized_partition_to_target_in_one_pass() {
use crate::index::DatasetIndexExt;
use crate::index::vector::VectorIndexParams;
use lance_index::optimize::OptimizeOptions;
use lance_linalg::distance::MetricType;
let tmp = tempfile::tempdir().unwrap();
let uri = tmp.path().to_str().unwrap();
let mut dataset = write_clusters(uri, &[(15_000, 0.0), (1_500, 1000.0)]).await;
let params = VectorIndexParams::ivf_flat(2, MetricType::L2);
dataset
.create_index(
&["vec"],
IndexType::Vector,
Some("idx".into()),
¶ms,
false,
)
.await
.unwrap();
let mut dataset = append_cluster(dataset, 21_000, 0.0).await;
dataset
.optimize_indices(&OptimizeOptions::default())
.await
.unwrap();
let optimized = open_single_segment(&dataset).await;
let ivf = optimized.ivf_model();
assert_eq!(ivf.num_partitions(), 2 + 8, "one 9-way split");
let sizes: Vec<usize> = (0..ivf.num_partitions())
.map(|p| optimized.partition_size(p))
.collect();
assert_eq!(
sizes.iter().sum::<usize>(),
37_500,
"all vectors preserved across split"
);
let max_size = *sizes.iter().max().unwrap();
assert!(
max_size <= 4 * 4096,
"no partition may stay above the split threshold after one pass, got {max_size}"
);
let found = nearest_first_components(&dataset, 3.0, 5).await;
assert_eq!(found.len(), 5);
assert!(found.iter().all(|v| (0.0..=3.6).contains(v)), "{found:?}");
}
#[tokio::test]
async fn optimize_joins_all_undersized_partitions_in_one_pass() {
use crate::index::DatasetIndexExt;
use lance_index::optimize::OptimizeOptions;
let tmp = tempfile::tempdir().unwrap();
let uri = tmp.path().to_str().unwrap();
let mut dataset = write_clusters(
uri,
&[(2_000, 0.0), (100, 1000.0), (100, 2000.0), (100, 3000.0)],
)
.await;
let params = ivf_flat_at_centers(&[0.0, 1000.0, 2000.0, 3000.0], None);
dataset
.create_index(
&["vec"],
IndexType::Vector,
Some("idx".into()),
¶ms,
false,
)
.await
.unwrap();
let mut dataset = append_cluster(dataset, 10, 0.0).await;
dataset
.optimize_indices(&OptimizeOptions::default())
.await
.unwrap();
let optimized = open_single_segment(&dataset).await;
let ivf = optimized.ivf_model();
let sizes: Vec<usize> = (0..ivf.num_partitions())
.map(|p| optimized.partition_size(p))
.collect();
assert_eq!(
sizes,
vec![2_310],
"all three undersized partitions joined in one pass"
);
let found = nearest_first_components(&dataset, 3000.0, 3).await;
assert!(
found.iter().all(|v| (3000.0..=3000.1).contains(v)),
"{found:?}"
);
}
#[tokio::test]
async fn optimize_steady_state_keeps_small_delta_partitions() {
use crate::index::DatasetIndexExt;
use crate::index::vector::VectorIndexParams;
use lance_index::optimize::OptimizeOptions;
use lance_linalg::distance::MetricType;
let tmp = tempfile::tempdir().unwrap();
let uri = tmp.path().to_str().unwrap();
let mut dataset = write_clusters(uri, &[(1_500, 0.0), (1_500, 1000.0)]).await;
let params = VectorIndexParams::ivf_flat(2, MetricType::L2);
dataset
.create_index(
&["vec"],
IndexType::Vector,
Some("idx".into()),
¶ms,
false,
)
.await
.unwrap();
let dataset = append_cluster(dataset, 100, 0.0).await;
let mut dataset = append_cluster(dataset, 100, 1000.0).await;
dataset
.optimize_indices(&OptimizeOptions::default())
.await
.unwrap();
let indices = dataset.load_indices_by_name("idx").await.unwrap();
assert_eq!(indices.len(), 2, "the append is a delta segment");
dataset
.optimize_indices(&OptimizeOptions::default())
.await
.unwrap();
let indices = dataset.load_indices_by_name("idx").await.unwrap();
assert_eq!(indices.len(), 2, "steady state is a no-op");
for index in &indices {
use crate::index::DatasetIndexInternalExt;
let opened = dataset
.open_vector_index("vec", &index.uuid, &NoOpMetricsCollector)
.await
.unwrap();
assert_eq!(
opened.ivf_model().num_partitions(),
2,
"no partition was joined away"
);
}
}
#[tokio::test]
async fn optimize_split_threshold_honors_persisted_target_partition_size() {
use crate::index::DatasetIndexExt;
use crate::index::vector::{StageParams, VectorIndexParams};
use lance_index::optimize::OptimizeOptions;
use lance_linalg::distance::MetricType;
let tmp = tempfile::tempdir().unwrap();
let uri = tmp.path().to_str().unwrap();
let mut dataset = write_clusters(uri, &[(5_000, 0.0), (500, 1000.0)]).await;
let mut params = VectorIndexParams::ivf_flat(2, MetricType::L2);
let StageParams::Ivf(ivf_params) = &mut params.stages[0] else {
panic!("ivf_flat params start with the IVF stage");
};
ivf_params.target_partition_size = Some(1024);
dataset
.create_index(
&["vec"],
IndexType::Vector,
Some("idx".into()),
¶ms,
false,
)
.await
.unwrap();
let mut dataset = append_cluster(dataset, 10, 0.0).await;
dataset
.optimize_indices(&OptimizeOptions::default())
.await
.unwrap();
let optimized = open_single_segment(&dataset).await;
let ivf = optimized.ivf_model();
assert_eq!(
ivf.num_partitions(),
2 + 4,
"5_010 rows split ceil(5010 / 1024) = 5 ways"
);
let total: usize = (0..ivf.num_partitions())
.map(|p| optimized.partition_size(p))
.sum();
assert_eq!(total, 5_510);
}
fn ivf_on_a_line(positions: &[f32]) -> IvfModel {
let centroids =
FixedSizeListArray::try_new_from_values(Float32Array::from(positions.to_vec()), 1)
.unwrap();
IvfModel::new(centroids, None)
}
#[test]
fn plan_partition_adjustment_keeps_room_for_every_joined_row() {
let (_, joins) = plan_partition_adjustment(&[24; 40], 100);
assert_eq!(joins.len(), 37, "{joins:?}");
let (_, joins) = plan_partition_adjustment(&[20; 5], 100);
assert_eq!(joins.len(), 4, "{joins:?}");
}
#[tokio::test]
async fn optimize_join_keeps_every_destination_below_the_split_threshold() {
use crate::index::DatasetIndexExt;
use lance_index::optimize::OptimizeOptions;
let tmp = tempfile::tempdir().unwrap();
let uri = tmp.path().to_str().unwrap();
let mut dataset = write_clusters(uri, &[(24, 4.0), (100, -9.0), (390, 10.0)]).await;
let params = ivf_flat_at_centers(&[0.0, -9.0, 10.0], Some(100));
dataset
.create_index(
&["vec"],
IndexType::Vector,
Some("idx".into()),
¶ms,
false,
)
.await
.unwrap();
let mut dataset = append_cluster(dataset, 1, -9.0).await;
dataset
.optimize_indices(&OptimizeOptions::default())
.await
.unwrap();
let optimized = open_single_segment(&dataset).await;
let ivf = optimized.ivf_model();
let mut sizes: Vec<usize> = (0..ivf.num_partitions())
.map(|p| optimized.partition_size(p))
.collect();
sizes.sort_unstable();
assert_eq!(sizes, vec![115, 400]);
}
fn multivector_batch(rows: &[Vec<[f32; 4]>], first_id: usize) -> RecordBatch {
let mut values = Vec::new();
for (i, row) in rows.iter().enumerate() {
for center in row {
let mut vector = *center;
vector[3] += (first_id + i) as f32 * 0.0001;
values.extend_from_slice(&vector);
}
}
let vectors =
FixedSizeListArray::try_new_from_values(Float32Array::from(values), 4).unwrap();
let item_field = Arc::new(Field::new("item", vectors.data_type().clone(), true));
let offsets = arrow::buffer::OffsetBuffer::from_lengths(rows.iter().map(|row| row.len()));
let list = arrow_array::ListArray::new(item_field, offsets, Arc::new(vectors), None);
let schema = Arc::new(arrow_schema::Schema::new(vec![Field::new(
"vec",
list.data_type().clone(),
false,
)]));
RecordBatch::try_new(schema, vec![Arc::new(list)]).unwrap()
}
#[tokio::test]
async fn optimize_join_counts_room_after_the_reindexed_rows_leave() {
use crate::dataset::{InsertBuilder, WriteMode, WriteParams};
use crate::index::DatasetIndexExt;
use crate::index::vector::VectorIndexParams;
use arrow_array::RecordBatchIterator;
use lance_index::optimize::OptimizeOptions;
use lance_linalg::distance::MetricType;
let tmp = tempfile::tempdir().unwrap();
let uri = tmp.path().to_str().unwrap();
const C0: [f32; 4] = [1.0, 0.0, 0.0, 0.0];
const C1: [f32; 4] = [0.0, 1.0, 0.0, 0.0];
const C2: [f32; 4] = [0.0, 0.0, 1.0, 0.0];
let mut rows: Vec<Vec<[f32; 4]>> = (0..20)
.map(|_| {
std::iter::once(C0)
.chain(std::iter::repeat_n(C1, 18))
.collect()
})
.collect();
rows.extend((0..380).map(|_| vec![C2]));
let batch = multivector_batch(&rows, 0);
let reader = RecordBatchIterator::new(vec![Ok(batch.clone())], batch.schema());
let mut dataset = crate::Dataset::write(reader, uri, None).await.unwrap();
let centroids =
FixedSizeListArray::try_new_from_values(Float32Array::from([C0, C1, C2].concat()), 4)
.unwrap();
let mut ivf_params = IvfBuildParams::try_with_centroids(3, Arc::new(centroids)).unwrap();
ivf_params.target_partition_size = Some(100);
let params = VectorIndexParams::with_ivf_flat_params(MetricType::Cosine, ivf_params);
dataset
.create_index(
&["vec"],
IndexType::Vector,
Some("idx".into()),
¶ms,
false,
)
.await
.unwrap();
let append = multivector_batch(&[vec![C1]], rows.len());
let mut dataset = InsertBuilder::new(Arc::new(dataset))
.with_params(&WriteParams {
mode: WriteMode::Append,
..Default::default()
})
.execute(vec![append])
.await
.unwrap();
dataset
.optimize_indices(&OptimizeOptions::default())
.await
.unwrap();
let optimized = open_single_segment(&dataset).await;
let ivf = optimized.ivf_model();
let sizes: Vec<usize> = (0..ivf.num_partitions())
.map(|p| optimized.partition_size(p))
.collect();
assert_eq!(ivf.num_partitions(), 2, "{sizes:?}");
assert_eq!(
sizes.iter().sum::<usize>(),
761,
"every vector indexed once"
);
assert!(
sizes.iter().all(|&rows| rows <= 400),
"a join destination exceeds the split threshold: {sizes:?}"
);
}
#[tokio::test]
async fn regroup_by_bytes_bounds_each_batch_to_the_budget() {
let one_row = |value: f32| {
let vectors =
FixedSizeListArray::try_new_from_values(Float32Array::from(vec![value; 4]), 4)
.unwrap();
let schema = Arc::new(arrow_schema::Schema::new(vec![Field::new(
"vec",
vectors.data_type().clone(),
false,
)]));
RecordBatch::try_new(schema, vec![Arc::new(vectors)]).unwrap()
};
let rows: Vec<RecordBatch> = (0..5).map(|i| one_row(i as f32)).collect();
let row_bytes = rows[0].get_array_memory_size();
let grouped: Vec<RecordBatch> = regroup_by_bytes(
stream::iter(rows.clone().into_iter().map(Ok)),
2 * row_bytes,
)
.try_collect()
.await
.unwrap();
assert_eq!(
grouped.iter().map(|b| b.num_rows()).collect::<Vec<_>>(),
vec![2, 2, 1]
);
let grouped: Vec<RecordBatch> =
regroup_by_bytes(stream::iter(rows.into_iter().map(Ok)), row_bytes / 2)
.try_collect()
.await
.unwrap();
assert_eq!(grouped.len(), 5);
}
#[test]
fn join_window_follows_each_vector_of_a_multivector_row() {
let mut positions: Vec<f32> = (0..65).map(|i| i as f32 * 0.001).collect();
positions.extend([100.0, 100.0]);
let ivf = ivf_on_a_line(&positions);
let removed: HashSet<usize> = [0, 65].into_iter().collect();
let vector: ArrayRef = Arc::new(Float32Array::from(vec![100.0_f32]));
let (window_ids, window_centroids) =
window_for_vector(DistanceType::L2, &ivf, &vector, &removed).unwrap();
assert_eq!(window_ids.value(0), 66);
let room = vec![1; positions.len()];
let (target, had_room) = choose_join_destination(
DistanceType::L2,
vector.as_primitive::<Float32Type>(),
&window_ids,
&window_centroids,
&ivf,
&removed,
&room,
)
.unwrap();
assert_eq!((target, had_room), (66, true));
let c0 = ivf.centroid(0).unwrap();
let (partition_window, _) =
select_reassign_candidates_impl(DistanceType::L2, &ivf, 0, &c0, &removed).unwrap();
assert!(!partition_window.values().contains(&66));
}
#[test]
fn plan_partition_adjustment_joins_all_undersized_partitions_but_one() {
let (_, joins) = plan_partition_adjustment(&[66, 65, 66, 66], 4096);
assert_eq!(joins, vec![0, 1, 2]);
assert!(plan_partition_adjustment(&[66], 4096).1.is_empty());
}
#[tokio::test]
async fn optimize_join_leaves_no_partition_above_the_split_threshold() {
use crate::index::DatasetIndexExt;
use lance_index::optimize::OptimizeOptions;
let tmp = tempfile::tempdir().unwrap();
let uri = tmp.path().to_str().unwrap();
let centers: Vec<f32> = (0..20).map(|i| i as f32 * 1000.0).collect();
let clusters: Vec<(usize, f32)> = centers.iter().map(|&c| (24, c)).collect();
let mut dataset = write_clusters(uri, &clusters).await;
let params = ivf_flat_at_centers(¢ers, Some(100));
dataset
.create_index(
&["vec"],
IndexType::Vector,
Some("idx".into()),
¶ms,
false,
)
.await
.unwrap();
let mut dataset = append_cluster(dataset, 1, 0.0).await;
dataset
.optimize_indices(&OptimizeOptions::default())
.await
.unwrap();
let optimized = open_single_segment(&dataset).await;
let ivf = optimized.ivf_model();
let sizes: Vec<usize> = (0..ivf.num_partitions())
.map(|p| optimized.partition_size(p))
.collect();
assert_eq!(sizes.iter().sum::<usize>(), 481, "all rows preserved");
assert!(
ivf.num_partitions() < 20,
"some partitions were joined: {sizes:?}"
);
let max_size = *sizes.iter().max().unwrap();
assert!(
max_size <= 4 * 100,
"a join must not create a partition above the split threshold, got {sizes:?}"
);
}
fn duplicate_batch(
schema: &Arc<arrow_schema::Schema>,
distinct: usize,
repeats: usize,
center: f32,
) -> RecordBatch {
let mut values = Vec::with_capacity(distinct * repeats * 4);
for i in 0..distinct {
let p = center + i as f32;
for _ in 0..repeats {
values.extend_from_slice(&[p, p, p, p]);
}
}
let fsl = FixedSizeListArray::try_new_from_values(Float32Array::from(values), 4).unwrap();
RecordBatch::try_new(schema.clone(), vec![Arc::new(fsl)]).unwrap()
}
#[tokio::test]
async fn optimize_splits_duplicate_heavy_partition_as_far_as_its_distinct_rows_allow() {
use crate::dataset::{InsertBuilder, WriteMode, WriteParams};
use crate::index::DatasetIndexExt;
use arrow_array::RecordBatchIterator;
use lance_index::optimize::OptimizeOptions;
let tmp = tempfile::tempdir().unwrap();
let uri = tmp.path().to_str().unwrap();
let schema = cluster_schema();
let batches = vec![
Ok(duplicate_batch(&schema, 5, 240, 0.0)),
Ok(duplicate_batch(&schema, 1, 10, 1000.0)),
];
let reader = RecordBatchIterator::new(batches, schema.clone());
let mut dataset = crate::Dataset::write(reader, uri, None).await.unwrap();
let params = ivf_flat_at_centers(&[2.0, 1000.0], Some(4));
dataset
.create_index(
&["vec"],
IndexType::Vector,
Some("idx".into()),
¶ms,
false,
)
.await
.unwrap();
let append = duplicate_batch(&schema, 1, 10, 0.0);
let mut dataset = InsertBuilder::new(Arc::new(dataset))
.with_params(&WriteParams {
mode: WriteMode::Append,
..Default::default()
})
.execute(vec![append])
.await
.unwrap();
dataset
.optimize_indices(&OptimizeOptions::default())
.await
.unwrap();
let optimized = open_single_segment(&dataset).await;
let ivf = optimized.ivf_model();
let mut sizes: Vec<usize> = (0..ivf.num_partitions())
.map(|p| optimized.partition_size(p))
.collect();
sizes.sort_unstable();
assert_eq!(sizes, vec![10, 240, 240, 240, 240, 250]);
let found = nearest_first_components(&dataset, 3.0, 5).await;
assert!(found.iter().all(|v| *v == 3.0), "{found:?}");
}
#[tokio::test]
async fn take_partition_batches_preserves_partition_order_for_large_fixed_size_list() {
let value_length = 1_073_741_824i32;
let num_rows = 5usize;
let row_ids = UInt64Array::from(vec![4_u64, 3, 2, 1, 0]);
let part_ids = UInt32Array::from(vec![0_u32; num_rows]);
let values = Arc::new(NullArray::new(num_rows * value_length as usize));
let item_field = Arc::new(Field::new("item", DataType::Null, true));
let codes = FixedSizeListArray::try_new(item_field, value_length, values, None).unwrap();
let batch = RecordBatch::try_new(
Arc::new(arrow_schema::Schema::new(vec![
ROW_ID_FIELD.clone(),
PART_ID_FIELD.clone(),
Field::new(PQ_CODE_COLUMN, codes.data_type().clone(), true),
])),
vec![Arc::new(row_ids), Arc::new(part_ids), Arc::new(codes)],
)
.unwrap();
let reader = SingleBatchReader {
batch,
partition_id: 0,
};
let (batches, loss) = IvfIndexBuilder::<FlatIndex, FlatQuantizer>::take_partition_batches(
0,
&[],
Some(&reader),
)
.await
.unwrap();
assert_eq!(loss, 0.0);
assert_eq!(batches.len(), 1);
assert!(batches[0].column_by_name(PART_ID_COLUMN).is_none());
let row_ids = batches[0][ROW_ID].as_primitive::<UInt64Type>();
assert_eq!(row_ids.values(), &[4, 3, 2, 1, 0]);
}
}