1use std::io::Write;
8use std::sync::Arc;
9
10use rustc_hash::FxHashMap;
11
12use crate::directories::DirectoryWriter;
13use crate::dsl::{
14 BinaryDenseVectorConfig, BinaryIndexType, DenseVectorConfig, Field, FieldType, VectorIndexType,
15};
16use crate::error::{Error, Result};
17use crate::segment::{SegmentFiles, SegmentId, SegmentMeta};
18
19use super::IndexWriter;
20
21const MAX_IVF_CLUSTERS: usize = 1_048_576;
24const MIN_TRAINING_POINTS_PER_CENTROID: usize = 39;
27const MAX_SAMPLE_READ_BYTES: usize = 64 * 1024 * 1024;
30const SAMPLE_BLOCK: usize = 256;
31const VECTOR_ARTIFACT_PREFIX: &str = "vector_artifact_";
34
35struct TrainedFieldUpdate {
36 field_id: u32,
37 index_type: super::metadata::VectorFieldIndexType,
38 vector_count: usize,
39 num_clusters: usize,
40 centroids_file: String,
41 codebook_file: Option<String>,
42}
43
44enum TrainedFieldArtifacts {
45 Float {
46 centroids: crate::structures::CoarseCentroids,
47 codebook: crate::structures::PQCodebook,
48 pq_sample_count: usize,
49 },
50 FloatCentroids(crate::structures::CoarseCentroids),
53 Binary(crate::structures::BinaryCoarseQuantizer),
54}
55
56struct TrainedFieldModel {
57 update: TrainedFieldUpdate,
58 artifacts: TrainedFieldArtifacts,
59}
60
61#[derive(Clone)]
62enum IvfFieldConfig {
63 Float(DenseVectorConfig),
64 Binary(BinaryDenseVectorConfig),
65}
66
67impl IvfFieldConfig {
68 fn dim(&self) -> usize {
69 match self {
70 Self::Float(config) => config.dim,
71 Self::Binary(config) => config.dim,
72 }
73 }
74
75 fn index_type(&self) -> super::metadata::VectorFieldIndexType {
76 match self {
77 Self::Float(config) => config.index_type.into(),
78 Self::Binary(config) => config.index_type.into(),
79 }
80 }
81
82 fn num_clusters(&self) -> Option<usize> {
83 match self {
84 Self::Float(config) => config.num_clusters,
85 Self::Binary(config) => config.num_clusters,
86 }
87 }
88
89 fn optimal_num_clusters(&self, vector_count: usize) -> usize {
90 match self {
91 Self::Float(config) => config.optimal_num_clusters(vector_count),
92 Self::Binary(config) => config.optimal_num_clusters(vector_count),
93 }
94 }
95}
96
97enum TrainingSample {
98 Float(Vec<Vec<f32>>),
99 Binary(Vec<u8>),
100}
101
102#[derive(Clone, Copy, Debug, Eq, PartialEq)]
103enum VectorGenerationMode {
104 BuildMissing,
105 RetrainAll,
106}
107
108impl TrainingSample {
109 fn len(&self, dim: usize) -> usize {
110 match self {
111 Self::Float(vectors) => vectors.len(),
112 Self::Binary(codes) => codes.len() / dim.div_ceil(8),
113 }
114 }
115}
116
117struct SizeLimitedWriter<'a, W: Write + ?Sized> {
122 inner: &'a mut W,
123 written: usize,
124 limit: usize,
125}
126
127impl<'a, W: Write + ?Sized> SizeLimitedWriter<'a, W> {
128 fn new(inner: &'a mut W, limit: usize) -> Self {
129 Self {
130 inner,
131 written: 0,
132 limit,
133 }
134 }
135}
136
137impl<W: Write + ?Sized> Write for SizeLimitedWriter<'_, W> {
138 fn write(&mut self, buffer: &[u8]) -> std::io::Result<usize> {
139 let next_size = self
140 .written
141 .checked_add(buffer.len())
142 .ok_or_else(|| std::io::Error::other("trained artifact size overflow"))?;
143 if next_size > self.limit {
144 return Err(std::io::Error::new(
145 std::io::ErrorKind::InvalidData,
146 format!(
147 "trained artifact exceeds the {}-byte safety limit",
148 self.limit
149 ),
150 ));
151 }
152 let written = self.inner.write(buffer)?;
153 self.written += written;
154 Ok(written)
155 }
156
157 fn flush(&mut self) -> std::io::Result<()> {
158 self.inner.flush()
159 }
160}
161
162fn validate_explicit_cluster_count(num_clusters: Option<usize>) -> Result<()> {
163 match num_clusters {
164 Some(0) => Err(Error::Schema(
165 "dense vector num_clusters must be at least 1".to_string(),
166 )),
167 Some(value) if value > MAX_IVF_CLUSTERS => Err(Error::Schema(format!(
168 "dense vector num_clusters must not exceed {MAX_IVF_CLUSTERS}, got {value}"
169 ))),
170 _ => Ok(()),
171 }
172}
173
174fn effective_field_num_clusters(
175 config: &IvfFieldConfig,
176 corpus_count: usize,
177 sample_count: usize,
178) -> Result<usize> {
179 if sample_count == 0 {
180 return Err(Error::Schema(
181 "cannot train an IVF vector index without sample vectors".to_string(),
182 ));
183 }
184 validate_explicit_cluster_count(config.num_clusters())?;
185 let centroid_bytes = match config {
186 IvfFieldConfig::Float(config) => config.dim.saturating_mul(size_of::<f32>()),
187 IvfFieldConfig::Binary(config) => config.dim.div_ceil(8),
188 };
189 let artifact_limit = super::metadata::MAX_TRAINED_ARTIFACT_BYTES
190 .saturating_sub(1024)
191 .checked_div(centroid_bytes.max(1))
192 .unwrap_or(0)
193 .max(1);
194 let quality_limit = if config.num_clusters().is_some() {
195 sample_count
196 } else {
197 (sample_count / MIN_TRAINING_POINTS_PER_CENTROID)
198 .max(16)
199 .min(sample_count)
200 };
201 let requested = config.optimal_num_clusters(corpus_count);
202 if config.num_clusters().is_some() && requested > artifact_limit {
203 return Err(Error::Schema(format!(
204 "configured IVF codebook needs {} bytes for {} centroids, exceeding the {}-byte artifact limit",
205 requested.saturating_mul(centroid_bytes),
206 requested,
207 super::metadata::MAX_TRAINED_ARTIFACT_BYTES,
208 )));
209 }
210 Ok(requested.min(quality_limit).min(artifact_limit))
211}
212
213fn training_sample_limit(
214 max_samples: usize,
215 max_bytes: usize,
216 bytes_per_sample: usize,
217) -> Result<usize> {
218 if max_samples == 0 || max_bytes == 0 || bytes_per_sample == 0 {
219 return Err(Error::Schema(
220 "vector training sample count, memory budget, and vector size must be greater than zero"
221 .into(),
222 ));
223 }
224 let memory_limited = max_bytes / bytes_per_sample;
225 if memory_limited == 0 {
226 return Err(Error::Schema(format!(
227 "vector training memory budget ({max_bytes} bytes) cannot hold one {bytes_per_sample}-byte sample"
228 )));
229 }
230 Ok(max_samples.min(memory_limited))
231}
232
233#[cfg(test)]
240fn effective_ivf_num_clusters(
241 config: &DenseVectorConfig,
242 corpus_count: usize,
243 sample_count: usize,
244) -> Result<usize> {
245 if sample_count == 0 {
246 return Err(Error::Schema(
247 "cannot train an IVF vector index without sample vectors".to_string(),
248 ));
249 }
250
251 effective_field_num_clusters(
252 &IvfFieldConfig::Float(config.clone()),
253 corpus_count,
254 sample_count,
255 )
256}
257
258impl<D: DirectoryWriter + 'static> IndexWriter<D> {
259 pub async fn build_vector_index(&self) -> Result<()> {
266 self.build_vector_generation(VectorGenerationMode::BuildMissing)
267 .await
268 }
269
270 pub async fn retrain_vector_index(&self) -> Result<()> {
275 self.build_vector_generation(VectorGenerationMode::RetrainAll)
276 .await
277 }
278
279 async fn build_vector_generation(&self, mode: VectorGenerationMode) -> Result<()> {
280 let dense_fields = self.get_ivf_vector_fields();
281 if dense_fields.is_empty() {
282 log::info!("No dense vector fields configured for ANN indexing");
283 return Ok(());
284 }
285
286 let artifact_update = self.segment_manager.begin_vector_artifact_update().await?;
287 self.cleanup_unreferenced_vector_artifacts().await;
288
289 let fields_to_train = match mode {
290 VectorGenerationMode::BuildMissing => self.get_fields_to_build(&dense_fields).await,
291 VectorGenerationMode::RetrainAll => dense_fields.clone(),
292 };
293 for (_, config) in &fields_to_train {
294 validate_explicit_cluster_count(config.num_clusters())?;
295 }
296
297 let snapshot = self.segment_manager.acquire_snapshot().await;
298 if snapshot.is_empty() {
299 if mode == VectorGenerationMode::RetrainAll {
300 return Err(Error::Schema(
301 "cannot retrain vector codebooks without committed segments".into(),
302 ));
303 }
304 return Ok(());
305 }
306
307 let mut candidate_metadata = self.segment_manager.read_metadata(Clone::clone).await;
308 if !fields_to_train.is_empty() {
309 let total_vectors = self
310 .count_vectors_for_training(
311 snapshot.segment_ids(),
312 &fields_to_train,
313 mode == VectorGenerationMode::BuildMissing,
314 )
315 .await?;
316 let artifact_generation = SegmentId::new().to_hex();
317 let updates = self
318 .train_fields(
319 snapshot.segment_ids(),
320 &fields_to_train,
321 &total_vectors,
322 &artifact_generation,
323 )
324 .await?;
325 for update in &updates {
326 candidate_metadata.init_field(update.field_id, update.index_type);
327 candidate_metadata.mark_field_built(
328 update.field_id,
329 update.vector_count,
330 update.num_clusters,
331 update.centroids_file.clone(),
332 update.codebook_file.clone(),
333 );
334 }
335 }
336
337 let target_field_ids = dense_fields
338 .iter()
339 .filter_map(|(field, _)| {
340 candidate_metadata
341 .is_field_built(field.0)
342 .then_some(field.0)
343 })
344 .collect::<Vec<_>>();
345 if target_field_ids.is_empty() {
346 return Ok(());
347 }
348
349 let candidate_trained = super::IndexMetadata::try_load_trained_from_fields(
350 &candidate_metadata.vector_fields,
351 self.schema.as_ref(),
352 self.directory.as_ref(),
353 )
354 .await?
355 .map(Arc::new)
356 .ok_or_else(|| Error::Internal("candidate vector generation has no artifacts".into()))?;
357
358 let staged = self
359 .segment_manager
360 .stage_vector_generation(
361 &artifact_update,
362 snapshot.segment_ids(),
363 &target_field_ids,
364 Arc::clone(&candidate_trained),
365 mode == VectorGenerationMode::RetrainAll,
366 )
367 .await?;
368 self.segment_manager
369 .publish_vector_generation(
370 &artifact_update,
371 candidate_metadata.vector_fields,
372 candidate_trained,
373 staged,
374 )
375 .await?;
376
377 drop(snapshot);
381 drop(artifact_update);
382
383 self.segment_manager
387 .rewrite_vector_segments(&target_field_ids)
388 .await?;
389 self.cleanup_unreferenced_vector_artifacts().await;
390 log::info!(
391 "Dense vector ANN generation {:?} complete for {} field(s)",
392 mode,
393 target_field_ids.len(),
394 );
395 Ok(())
396 }
397
398 async fn train_fields(
399 &self,
400 segment_ids: &[String],
401 fields: &[(Field, IvfFieldConfig)],
402 total_vectors: &FxHashMap<u32, usize>,
403 artifact_generation: &str,
404 ) -> Result<Vec<TrainedFieldUpdate>> {
405 let training_pool = self.segment_manager.background_cpu_pool();
406 let mut missing = Vec::new();
407 let mut updates = Vec::with_capacity(fields.len());
408 for (field, config) in fields {
409 let corpus_count = total_vectors.get(&field.0).copied().unwrap_or(0);
413 let Some(sample) = self
414 .collect_training_sample(segment_ids, *field, config, corpus_count)
415 .await?
416 else {
417 missing.push(field.0);
418 continue;
419 };
420 let model = crate::segment::block_in_place_if_multithread(|| {
421 training_pool.install(|| {
422 Self::train_field_model(
423 *field,
424 config,
425 &sample,
426 corpus_count,
427 artifact_generation,
428 )
429 })
430 })?;
431 drop(sample);
434 updates.push(self.save_trained_field(model).await?);
435 }
436 if updates.is_empty() && !fields.is_empty() {
437 return Err(Error::Schema(format!(
438 "cannot train vector codebooks: no committed vectors for field(s) {missing:?}"
439 )));
440 }
441 if !missing.is_empty() {
442 log::info!(
443 "Skipping dense vector field(s) {missing:?}: the current corpus contains no vectors"
444 );
445 }
446 Ok(updates)
447 }
448
449 async fn cleanup_unreferenced_vector_artifacts(&self) {
454 let referenced = self
455 .segment_manager
456 .read_metadata(|metadata| {
457 metadata
458 .vector_fields
459 .values()
460 .flat_map(|field| {
461 field
462 .centroids_file
463 .iter()
464 .chain(field.codebook_file.iter())
465 })
466 .cloned()
467 .collect::<std::collections::HashSet<_>>()
468 })
469 .await;
470 let files = match self.directory.list_files(std::path::Path::new("")).await {
471 Ok(files) => files,
472 Err(error) => {
473 log::warn!("[trained] failed listing abandoned dense vector artifacts: {error}");
474 return;
475 }
476 };
477 for path in files {
478 let Some(name) = path.file_name().and_then(|name| name.to_str()) else {
479 continue;
480 };
481 if !name.starts_with(VECTOR_ARTIFACT_PREFIX)
482 || referenced.contains(path.to_string_lossy().as_ref())
483 {
484 continue;
485 }
486 if let Err(error) = self.directory.delete(&path).await
487 && error.kind() != std::io::ErrorKind::NotFound
488 {
489 log::warn!("[trained] failed deleting abandoned artifact {path:?}: {error}");
490 }
491 }
492 }
493
494 fn reject_ann_fields(ann_fields: &[u32], id_str: &str, field_ids: &[u32]) -> Result<()> {
499 for &field_id in field_ids {
500 if ann_fields.binary_search(&field_id).is_ok() {
501 return Err(Error::Schema(format!(
502 "metadata-flat field {field_id} already has ANN data in segment {id_str}; \
503 recreate the index instead of mixing vector generations"
504 )));
505 }
506 }
507 Ok(())
508 }
509
510 async fn load_training_vectors(
514 &self,
515 segment_id: SegmentId,
516 field_ids: &[u32],
517 ) -> Result<crate::segment::reader::loader::VectorsFileData> {
518 let files = SegmentFiles::new(segment_id.0);
519 let meta_bytes = self
520 .directory
521 .open_read(&files.meta)
522 .await?
523 .read_bytes()
524 .await?;
525 let meta = SegmentMeta::deserialize(meta_bytes.as_slice())?;
526 if meta.id != segment_id.0 {
527 return Err(Error::Corruption(format!(
528 "segment metadata ID {:032x} does not match file ID {}",
529 meta.id,
530 segment_id.to_hex(),
531 )));
532 }
533 crate::segment::reader::loader::load_flat_vectors_file(
534 self.directory.as_ref(),
535 &files,
536 self.schema.as_ref(),
537 meta.num_docs,
538 field_ids,
539 )
540 .await
541 }
542
543 fn get_ivf_vector_fields(&self) -> Vec<(Field, IvfFieldConfig)> {
545 self.schema
546 .fields()
547 .filter_map(|(field, entry)| {
548 if entry.field_type == FieldType::DenseVector && entry.indexed {
549 entry
550 .dense_vector_config
551 .as_ref()
552 .filter(|c| c.uses_ivf())
555 .map(|c| (field, IvfFieldConfig::Float(c.clone())))
556 } else if entry.field_type == FieldType::BinaryDenseVector && entry.indexed {
557 entry
558 .binary_dense_vector_config
559 .as_ref()
560 .filter(|config| config.index_type == BinaryIndexType::Ivf)
561 .map(|config| (field, IvfFieldConfig::Binary(config.clone())))
562 } else {
563 None
564 }
565 })
566 .collect()
567 }
568
569 async fn get_fields_to_build(
571 &self,
572 dense_fields: &[(Field, IvfFieldConfig)],
573 ) -> Vec<(Field, IvfFieldConfig)> {
574 let field_ids: Vec<u32> = dense_fields.iter().map(|(f, _)| f.0).collect();
575 let built: Vec<u32> = self
576 .segment_manager
577 .read_metadata(|meta| {
578 field_ids
579 .iter()
580 .filter(|fid| meta.is_field_built(**fid))
581 .copied()
582 .collect()
583 })
584 .await;
585 dense_fields
586 .iter()
587 .filter(|(field, _)| !built.contains(&field.0))
588 .cloned()
589 .collect()
590 }
591
592 async fn count_vectors_for_training(
594 &self,
595 segment_ids: &[String],
596 fields_to_build: &[(Field, IvfFieldConfig)],
597 require_flat_generation: bool,
598 ) -> Result<FxHashMap<u32, usize>> {
599 let mut total_vectors: FxHashMap<u32, usize> = FxHashMap::default();
600 let field_ids: Vec<u32> = fields_to_build.iter().map(|(field, _)| field.0).collect();
601
602 for id_str in segment_ids {
606 let segment_id = SegmentId::from_hex(id_str)
607 .ok_or_else(|| Error::Corruption(format!("Invalid segment ID: {}", id_str)))?;
608 let vectors = self.load_training_vectors(segment_id, &field_ids).await?;
609
610 if require_flat_generation {
611 Self::reject_ann_fields(&vectors.ann_fields, id_str, &field_ids)?;
612 }
613
614 for (field, _) in fields_to_build {
615 if let Some(flat) = vectors.flat_vectors.get(&field.0) {
616 let total = total_vectors.entry(field.0).or_default();
617 *total = total.checked_add(flat.num_vectors).ok_or_else(|| {
618 Error::Corruption(format!(
619 "vector count overflows usize for field {}",
620 field.0,
621 ))
622 })?;
623 }
624 }
625 }
626 Ok(total_vectors)
627 }
628
629 async fn collect_training_sample(
634 &self,
635 segment_ids: &[String],
636 field: Field,
637 config: &IvfFieldConfig,
638 total: usize,
639 ) -> Result<Option<TrainingSample>> {
640 if total == 0 {
641 return Ok(None);
642 }
643 let bytes_per_sample = match config {
644 IvfFieldConfig::Float(config) => config
645 .dim
646 .checked_mul(size_of::<f32>())
647 .ok_or_else(|| Error::Schema("float training vector size overflows".into()))?,
648 IvfFieldConfig::Binary(config) => config.dim.div_ceil(8),
649 };
650 let limit = training_sample_limit(
651 self.config.vector_training_max_samples,
652 self.config.vector_training_memory_bytes,
653 bytes_per_sample,
654 )?;
655 let take = total.min(limit);
656 let mut rng = <rand::rngs::StdRng as rand::SeedableRng>::seed_from_u64(
657 0x4845_524d_4553_4956 ^ field.0 as u64 ^ total as u64,
658 );
659 let mut ordinals = Vec::with_capacity(take);
660 if take == total {
661 ordinals.extend(0..total);
662 } else {
663 let blocks = take.div_ceil(SAMPLE_BLOCK);
664 for block in 0..blocks {
665 let block_len = SAMPLE_BLOCK.min(take - ordinals.len());
666 let stratum_start = block.saturating_mul(total) / blocks;
667 let stratum_end = (block + 1).saturating_mul(total) / blocks;
668 let latest_start = stratum_end.saturating_sub(block_len);
669 let start = if latest_start > stratum_start {
670 rand::Rng::random_range(&mut rng, stratum_start..=latest_start)
671 } else {
672 stratum_start
673 };
674 ordinals.extend(start..start + block_len);
675 }
676 }
677
678 let mut sample = match config {
679 IvfFieldConfig::Float(_) => TrainingSample::Float(Vec::with_capacity(take)),
680 IvfFieldConfig::Binary(_) => TrainingSample::Binary(Vec::with_capacity(
681 take.checked_mul(bytes_per_sample)
682 .ok_or_else(|| Error::Schema("binary training sample size overflows".into()))?,
683 )),
684 };
685 let max_read_vectors = (MAX_SAMPLE_READ_BYTES / bytes_per_sample.max(1)).max(1);
686 let mut global_offset = 0usize;
687 let mut cursor = 0usize;
688 let field_ids = [field.0];
689
690 for id_str in segment_ids {
691 let segment_id = SegmentId::from_hex(id_str)
692 .ok_or_else(|| Error::Corruption(format!("Invalid segment ID: {id_str}")))?;
693 let vectors = self.load_training_vectors(segment_id, &field_ids).await?;
694
695 let Some(lazy_flat) = vectors.flat_vectors.get(&field.0) else {
696 continue;
697 };
698 let base = global_offset;
699 let end = base.checked_add(lazy_flat.num_vectors).ok_or_else(|| {
700 Error::Corruption(format!("vector offset overflows for field {}", field.0))
701 })?;
702 global_offset = end;
703 let first = cursor;
704 while cursor < ordinals.len() && ordinals[cursor] < end {
705 cursor += 1;
706 }
707 let selected = &ordinals[first..cursor];
708 let mut run_start = 0;
709 while run_start < selected.len() {
710 let mut run_end = run_start + 1;
711 while run_end < selected.len()
712 && run_end - run_start < max_read_vectors
713 && selected[run_end] == selected[run_end - 1] + 1
714 {
715 run_end += 1;
716 }
717 let local_start = selected[run_start] - base;
718 let run_len = run_end - run_start;
719 let bytes = lazy_flat
720 .read_vectors_batch(local_start, run_len)
721 .await
722 .map_err(crate::Error::Io)?;
723 match &mut sample {
724 TrainingSample::Binary(codes) => {
725 let expected = run_len.checked_mul(bytes_per_sample).ok_or_else(|| {
726 Error::Corruption("binary sample read size overflows".into())
727 })?;
728 if bytes.len() != expected {
729 return Err(Error::Corruption(format!(
730 "binary sample read returned {} bytes, expected {expected}",
731 bytes.len(),
732 )));
733 }
734 codes.extend_from_slice(bytes.as_slice());
735 }
736 TrainingSample::Float(vectors) => {
737 let dim = lazy_flat.dim;
738 let float_count = run_len.checked_mul(dim).ok_or_else(|| {
739 Error::Corruption("float sample read size overflows".into())
740 })?;
741 let mut decoded = vec![0.0; float_count];
742 crate::segment::dequantize_raw(
743 bytes.as_slice(),
744 lazy_flat.quantization,
745 decoded.len(),
746 &mut decoded,
747 )
748 .map_err(crate::Error::Io)?;
749 vectors.extend(decoded.chunks_exact(dim).map(<[f32]>::to_vec));
750 }
751 }
752 run_start = run_end;
753 }
754 }
755
756 let collected = sample.len(config.dim());
757 if global_offset != total || cursor != take || collected != take {
758 return Err(Error::Corruption(format!(
759 "training sample coverage mismatch for field {}: counted={total}, traversed={global_offset}, selected={cursor}, collected={collected}",
760 field.0,
761 )));
762 }
763 if collected < total {
764 log::info!(
765 "Sampled {} / {} dense vectors for field {} (max {} vectors / {} resident)",
766 collected,
767 total,
768 field.0,
769 self.config.vector_training_max_samples,
770 crate::format_bytes(self.config.vector_training_memory_bytes as u64),
771 );
772 }
773 Ok(Some(sample))
774 }
775
776 fn train_field_model(
779 field: Field,
780 config: &IvfFieldConfig,
781 sample: &TrainingSample,
782 corpus_count: usize,
783 artifact_generation: &str,
784 ) -> Result<TrainedFieldModel> {
785 let field_id = field.0;
786 let dim = config.dim();
787 let sample_count = sample.len(dim);
788 if sample_count == 0 || corpus_count == 0 {
789 return Err(Error::Internal(format!(
790 "empty training sample for non-empty field {field_id}"
791 )));
792 }
793 let num_clusters = effective_field_num_clusters(config, corpus_count, sample_count)?;
794
795 log::info!(
796 "Training dense vector index for field {} with {} sampled / {} total vectors, {} clusters (dim={})",
797 field_id,
798 sample_count,
799 corpus_count,
800 num_clusters,
801 dim,
802 );
803
804 let centroids_filename =
805 format!("{VECTOR_ARTIFACT_PREFIX}{artifact_generation}_field_{field_id}_centroids.bin");
806 let mut codebook_filename = None;
807
808 let artifacts = match (config, sample) {
809 (IvfFieldConfig::Float(config), TrainingSample::Float(vectors))
810 if config.index_type == VectorIndexType::IvfPq =>
811 {
812 codebook_filename = Some(format!(
813 "{VECTOR_ARTIFACT_PREFIX}{artifact_generation}_field_{field_id}_codebook.bin"
814 ));
815 let (centroids, codebook, pq_sample_count) = Self::train_ivf_pq_model(
816 dim,
817 num_clusters,
818 config.ivf_routing,
819 config.soar.clone(),
820 vectors,
821 );
822 TrainedFieldArtifacts::Float {
823 centroids,
824 codebook,
825 pq_sample_count,
826 }
827 }
828 (IvfFieldConfig::Float(config), TrainingSample::Float(vectors))
829 if config.index_type == VectorIndexType::IvfTq =>
830 {
831 let mut coarse_config = crate::structures::CoarseConfig::new(dim, num_clusters)
832 .with_routing(config.ivf_routing);
833 if let Some(soar) = config.soar.clone() {
834 coarse_config = coarse_config.with_soar(soar);
835 }
836 TrainedFieldArtifacts::FloatCentroids(crate::structures::CoarseCentroids::train(
837 &coarse_config,
838 vectors,
839 ))
840 }
841 (IvfFieldConfig::Binary(config), TrainingSample::Binary(codes)) => {
842 let mut binary_config = crate::structures::BinaryIvfConfig::new(dim, num_clusters);
843 binary_config.max_train_samples = sample_count;
844 binary_config.routing = config.ivf_routing;
845 TrainedFieldArtifacts::Binary(
846 crate::structures::BinaryCoarseQuantizer::train(
847 binary_config,
848 codes,
849 sample_count,
850 )
851 .map_err(Error::Io)?,
852 )
853 }
854 _ => {
855 return Err(Error::Internal(format!(
856 "training sample kind does not match field {field_id}"
857 )));
858 }
859 };
860
861 let actual_num_clusters = match &artifacts {
862 TrainedFieldArtifacts::Float { centroids, .. } => centroids.num_clusters as usize,
863 TrainedFieldArtifacts::FloatCentroids(centroids) => centroids.num_clusters as usize,
864 TrainedFieldArtifacts::Binary(quantizer) => quantizer.num_clusters as usize,
865 };
866 Ok(TrainedFieldModel {
867 update: TrainedFieldUpdate {
868 field_id,
869 index_type: config.index_type(),
870 vector_count: corpus_count,
871 num_clusters: actual_num_clusters,
872 centroids_file: centroids_filename,
873 codebook_file: codebook_filename,
874 },
875 artifacts,
876 })
877 }
878
879 async fn save_trained_field(&self, model: TrainedFieldModel) -> Result<TrainedFieldUpdate> {
880 let TrainedFieldModel { update, artifacts } = model;
881 match artifacts {
882 TrainedFieldArtifacts::Float {
883 centroids,
884 codebook,
885 pq_sample_count,
886 } => {
887 let codebook_file = update.codebook_file.as_deref().ok_or_else(|| {
888 Error::Internal(format!(
889 "trained IVF-PQ field {} has no codebook filename",
890 update.field_id,
891 ))
892 })?;
893 tokio::try_join!(
894 self.save_trained_artifact(¢roids, &update.centroids_file),
895 self.save_trained_artifact(&codebook, codebook_file),
896 )?;
897 log::info!(
898 "Saved IVF-PQ artifacts for field {} ({} clusters, {} PQ samples)",
899 update.field_id,
900 centroids.num_clusters,
901 pq_sample_count,
902 );
903 }
904 TrainedFieldArtifacts::FloatCentroids(centroids) => {
905 self.save_trained_artifact(¢roids, &update.centroids_file)
906 .await?;
907 log::info!(
908 "Saved IVF-TQ coarse artifact for field {} ({} clusters; leaf codec is derived)",
909 update.field_id,
910 centroids.num_clusters,
911 );
912 }
913 TrainedFieldArtifacts::Binary(quantizer) => {
914 self.save_trained_artifact(&quantizer, &update.centroids_file)
915 .await?;
916 log::info!(
917 "Saved binary IVF artifact for field {} ({} clusters)",
918 update.field_id,
919 quantizer.num_clusters,
920 );
921 }
922 }
923 Ok(update)
924 }
925
926 async fn save_trained_artifact(
928 &self,
929 artifact: &impl serde::Serialize,
930 filename: &str,
931 ) -> Result<()> {
932 let temp_filename = format!("{filename}.tmp");
933 let temp_path = std::path::Path::new(&temp_filename);
934 let final_path = std::path::Path::new(filename);
935 let mut writer = self.directory.streaming_writer(temp_path).await?;
936 let encode_result = {
937 let mut limited = SizeLimitedWriter::new(
938 writer.as_mut(),
939 super::metadata::MAX_TRAINED_ARTIFACT_BYTES,
940 );
941 bincode::serde::encode_into_std_write(
942 artifact,
943 &mut limited,
944 bincode::config::standard(),
945 )
946 };
947 if let Err(error) = encode_result {
948 drop(writer);
949 let _ = self.directory.delete(temp_path).await;
950 return Err(Error::Serialization(format!(
951 "failed to serialize trained artifact '{filename}': {error}"
952 )));
953 }
954 if let Err(error) = writer.finish() {
955 let _ = self.directory.delete(temp_path).await;
956 return Err(Error::Io(error));
957 }
958 if let Err(error) = self.directory.rename(temp_path, final_path).await {
959 let _ = self.directory.delete(temp_path).await;
960 return Err(Error::Io(error));
961 }
962 self.directory.sync().await?;
963 Ok(())
964 }
965
966 fn train_ivf_pq_model(
969 dim: usize,
970 num_clusters: usize,
971 routing: crate::dsl::IvfRoutingMode,
972 soar: Option<crate::structures::SoarConfig>,
973 vectors: &[Vec<f32>],
974 ) -> (
975 crate::structures::CoarseCentroids,
976 crate::structures::PQCodebook,
977 usize,
978 ) {
979 let mut coarse_config =
980 crate::structures::CoarseConfig::new(dim, num_clusters).with_routing(routing);
981 if let Some(soar) = soar {
982 coarse_config = coarse_config.with_soar(soar);
983 }
984 let centroids = crate::structures::CoarseCentroids::train(&coarse_config, vectors);
985
986 const PQ_CENTROIDS: usize = 256;
991 const PQ_TRAINING_POINTS_PER_CENTROID: usize = 256;
992 let pq_sample_count = vectors
993 .len()
994 .min(PQ_CENTROIDS * PQ_TRAINING_POINTS_PER_CENTROID);
995 use rayon::prelude::*;
996 let pq_residuals = (0..pq_sample_count)
997 .into_par_iter()
998 .map(|sample_index| {
999 let vector_index = sample_index.saturating_mul(vectors.len()) / pq_sample_count;
1000 let vector = &vectors[vector_index];
1001 let cluster = centroids
1002 .probe(vector, 1, routing)
1003 .cluster_ids
1004 .first()
1005 .copied()
1006 .unwrap_or(0);
1007 centroids.compute_residual(vector, cluster)
1008 })
1009 .collect::<Vec<_>>();
1010 let pq_config = crate::structures::PQConfig::new(dim);
1011 let codebook = crate::structures::PQCodebook::train(pq_config, &pq_residuals, 10);
1012 (centroids, codebook, pq_sample_count)
1013 }
1014}
1015
1016#[cfg(test)]
1017mod tests {
1018 use super::*;
1019
1020 fn ivf_config(num_clusters: Option<usize>) -> DenseVectorConfig {
1021 DenseVectorConfig::with_ivf_pq(8, num_clusters, 4)
1022 }
1023
1024 #[test]
1025 fn effective_clusters_follow_corpus_heuristic_but_fit_sample() {
1026 let config = ivf_config(None);
1027
1028 assert_eq!(
1029 effective_ivf_num_clusters(&config, 1_000_000, 73).unwrap(),
1030 16
1031 );
1032 assert_eq!(
1033 effective_ivf_num_clusters(&config, 10_000, 1_000).unwrap(),
1034 25
1035 );
1036 }
1037
1038 #[test]
1039 fn effective_clusters_clamp_explicit_value_to_sample() {
1040 let config = ivf_config(Some(256));
1041 assert_eq!(
1042 effective_ivf_num_clusters(&config, 1_000_000, 17).unwrap(),
1043 17
1044 );
1045 }
1046
1047 #[test]
1048 fn effective_clusters_reject_invalid_explicit_bounds() {
1049 let zero = effective_ivf_num_clusters(&ivf_config(Some(0)), 10_000, 100)
1050 .unwrap_err()
1051 .to_string();
1052 assert!(zero.contains("at least 1"));
1053
1054 let too_many =
1055 effective_ivf_num_clusters(&ivf_config(Some(MAX_IVF_CLUSTERS + 1)), 10_000, 100)
1056 .unwrap_err()
1057 .to_string();
1058 assert!(too_many.contains("must not exceed 1048576"));
1059 }
1060
1061 #[test]
1062 fn effective_clusters_reject_empty_training_sample() {
1063 let error = effective_ivf_num_clusters(&ivf_config(None), 10_000, 0)
1064 .unwrap_err()
1065 .to_string();
1066 assert!(error.contains("without sample vectors"));
1067 }
1068
1069 #[test]
1070 fn training_sample_limit_honors_both_cli_bounds() {
1071 assert_eq!(training_sample_limit(10_000_000, 4_096, 4).unwrap(), 1_024);
1072 assert_eq!(training_sample_limit(100, 4_096, 4).unwrap(), 100);
1073 let error = training_sample_limit(100, 3, 4).unwrap_err().to_string();
1074 assert!(error.contains("cannot hold one"), "{error}");
1075 }
1076
1077 #[test]
1078 fn artifact_writer_enforces_limit_without_writing_past_it() {
1079 let mut output = Vec::new();
1080 let mut writer = SizeLimitedWriter::new(&mut output, 3);
1081 writer.write_all(&[1, 2]).unwrap();
1082 let error = writer.write_all(&[3, 4]).unwrap_err().to_string();
1083 assert!(error.contains("3-byte safety limit"), "{error}");
1084 assert_eq!(output, vec![1, 2]);
1085 }
1086
1087 use std::path::Path;
1090 use std::sync::atomic::{AtomicBool, Ordering};
1091
1092 use crate::directories::{
1093 Directory, DirectoryWriter as DirectoryWriterTrait, FileHandle, RamDirectory, RangeReadFn,
1094 };
1095 use crate::dsl::{Document, SchemaBuilder};
1096 use crate::index::{IndexConfig, IndexWriter};
1097
1098 const READ_FAIL_DOCS: usize = 5;
1099 const READ_FAIL_DIM: usize = 4;
1100 const VEC_REGION_START: u64 = 16;
1106 const VEC_REGION_END: u64 = VEC_REGION_START + (READ_FAIL_DOCS * READ_FAIL_DIM * 4) as u64;
1107
1108 #[derive(Clone, Default)]
1113 struct VectorReadFailDirectory {
1114 inner: RamDirectory,
1115 fail_vector_reads: Arc<AtomicBool>,
1116 fail_all_vector_reads: Arc<AtomicBool>,
1117 }
1118
1119 #[async_trait::async_trait]
1120 impl Directory for VectorReadFailDirectory {
1121 async fn exists(&self, path: &Path) -> std::io::Result<bool> {
1122 self.inner.exists(path).await
1123 }
1124
1125 async fn file_size(&self, path: &Path) -> std::io::Result<u64> {
1126 self.inner.file_size(path).await
1127 }
1128
1129 async fn open_read(&self, path: &Path) -> std::io::Result<FileHandle> {
1130 self.inner.open_read(path).await
1131 }
1132
1133 async fn read_range(
1134 &self,
1135 path: &Path,
1136 range: std::ops::Range<u64>,
1137 ) -> std::io::Result<crate::directories::OwnedBytes> {
1138 self.inner.read_range(path, range).await
1139 }
1140
1141 async fn list_files(&self, prefix: &Path) -> std::io::Result<Vec<std::path::PathBuf>> {
1142 self.inner.list_files(prefix).await
1143 }
1144
1145 async fn open_lazy(&self, path: &Path) -> std::io::Result<FileHandle> {
1146 let handle = self.inner.open_lazy(path).await?;
1147 if path.extension().is_some_and(|ext| ext == "vectors") {
1148 let armed = Arc::clone(&self.fail_vector_reads);
1149 let fail_all = Arc::clone(&self.fail_all_vector_reads);
1150 let len = handle.len();
1151 let read_fn: RangeReadFn = Arc::new(move |range: std::ops::Range<u64>| {
1152 let handle = handle.clone();
1153 let armed = Arc::clone(&armed);
1154 let fail_all = Arc::clone(&fail_all);
1155 Box::pin(async move {
1156 if fail_all.load(Ordering::SeqCst)
1157 || (armed.load(Ordering::SeqCst)
1158 && range.start >= VEC_REGION_START
1159 && range.end <= VEC_REGION_END)
1160 {
1161 return Err(std::io::Error::other("injected vector data read failure"));
1162 }
1163 handle.read_bytes_range(range).await
1164 })
1165 });
1166 return Ok(FileHandle::lazy(len, read_fn));
1167 }
1168 Ok(handle)
1169 }
1170 }
1171
1172 #[async_trait::async_trait]
1173 impl DirectoryWriterTrait for VectorReadFailDirectory {
1174 async fn write(&self, path: &Path, data: &[u8]) -> std::io::Result<()> {
1175 self.inner.write(path, data).await
1176 }
1177
1178 async fn delete(&self, path: &Path) -> std::io::Result<()> {
1179 self.inner.delete(path).await
1180 }
1181
1182 async fn rename(&self, from: &Path, to: &Path) -> std::io::Result<()> {
1183 self.inner.rename(from, to).await
1184 }
1185
1186 async fn sync(&self) -> std::io::Result<()> {
1187 self.inner.sync().await
1188 }
1189
1190 async fn streaming_writer(
1191 &self,
1192 path: &Path,
1193 ) -> std::io::Result<Box<dyn crate::directories::StreamingWriter>> {
1194 self.inner.streaming_writer(path).await
1195 }
1196 }
1197
1198 #[tokio::test]
1201 async fn build_propagates_vector_read_errors_without_publishing_artifacts() {
1202 let mut sb = SchemaBuilder::default();
1203 let embedding = sb.add_dense_vector_field_with_config(
1204 "embedding",
1205 true,
1206 true,
1207 DenseVectorConfig::with_ivf_pq(READ_FAIL_DIM, Some(1), 1),
1208 );
1209 let schema = sb.build();
1210
1211 let dir = VectorReadFailDirectory::default();
1212 let config = IndexConfig {
1213 merge_policy: Box::new(crate::merge::NoMergePolicy),
1214 num_indexing_threads: 1,
1215 ..Default::default()
1216 };
1217 let mut writer = IndexWriter::create(dir.clone(), schema, config)
1218 .await
1219 .unwrap();
1220 for i in 0..READ_FAIL_DOCS {
1221 let mut doc = Document::new();
1222 doc.add_dense_vector(embedding, vec![i as f32 + 1.0; READ_FAIL_DIM]);
1223 writer.add_document(doc).unwrap();
1224 }
1225 writer.commit().await.unwrap();
1226 dir.fail_vector_reads.store(true, Ordering::SeqCst);
1227 let error = writer
1228 .build_vector_index()
1229 .await
1230 .expect_err("failed sample collection must fail the build")
1231 .to_string();
1232 assert!(
1233 error.contains("injected vector data read failure"),
1234 "{error}"
1235 );
1236
1237 assert!(
1238 !writer
1239 .segment_manager
1240 .read_metadata(|meta| meta.is_field_built(embedding.0))
1241 .await,
1242 "a failed build must not publish Built metadata"
1243 );
1244 assert!(
1245 writer.segment_manager.trained().is_none(),
1246 "a failed build must not publish trained artifacts"
1247 );
1248 }
1249
1250 #[tokio::test]
1251 async fn retrain_read_failure_keeps_the_complete_published_generation() {
1252 let mut sb = SchemaBuilder::default();
1253 let embedding = sb.add_dense_vector_field_with_config(
1254 "embedding",
1255 true,
1256 true,
1257 DenseVectorConfig::with_ivf_pq(READ_FAIL_DIM, Some(1), 1),
1258 );
1259 let schema = sb.build();
1260 let dir = VectorReadFailDirectory::default();
1261 let config = IndexConfig {
1262 merge_policy: Box::new(crate::merge::NoMergePolicy),
1263 num_indexing_threads: 1,
1264 ..Default::default()
1265 };
1266 let mut writer = IndexWriter::create(dir.clone(), schema, config)
1267 .await
1268 .unwrap();
1269 for i in 0..READ_FAIL_DOCS {
1270 let mut doc = Document::new();
1271 doc.add_dense_vector(embedding, vec![i as f32 + 1.0; READ_FAIL_DIM]);
1272 writer.add_document(doc).unwrap();
1273 }
1274 writer.commit().await.unwrap();
1275 writer.build_vector_index().await.unwrap();
1276
1277 let old_ids = writer.segment_manager.get_segment_ids().await;
1278 let old_meta = writer
1279 .segment_manager
1280 .read_metadata(|metadata| metadata.get_field_meta(embedding.0).cloned())
1281 .await
1282 .unwrap();
1283 let old_version = writer.segment_manager.trained().unwrap().codebooks[&embedding.0].version;
1284
1285 dir.fail_all_vector_reads.store(true, Ordering::SeqCst);
1286 let error = writer
1287 .retrain_vector_index()
1288 .await
1289 .expect_err("failed sample collection must abort the retrain")
1290 .to_string();
1291 assert!(
1292 error.contains("injected vector data read failure"),
1293 "{error}"
1294 );
1295 assert_eq!(writer.segment_manager.get_segment_ids().await, old_ids);
1296 assert_eq!(
1297 writer
1298 .segment_manager
1299 .read_metadata(|metadata| metadata
1300 .get_field_meta(embedding.0)
1301 .map(|field| (field.centroids_file.clone(), field.codebook_file.clone())))
1302 .await,
1303 Some((old_meta.centroids_file, old_meta.codebook_file)),
1304 );
1305 assert_eq!(
1306 writer.segment_manager.trained().unwrap().codebooks[&embedding.0].version,
1307 old_version,
1308 );
1309 }
1310}