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