Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
13 changes: 13 additions & 0 deletions .github/workflows/ci.yml
Original file line number Diff line number Diff line change
Expand Up @@ -288,6 +288,19 @@ jobs:
- name: Build FFI library
run: cargo build --release -p paimon-vindex-ffi

- name: Test AArch64 PQ kernels
if: runner.os == 'macOS'
run: cargo test --release -p paimon-vindex-core pq::tests

- name: Test no-AVX2/FMA PQ fallback under Rosetta
if: runner.os == 'macOS'
env:
PAIMON_EXPECT_PQ_SGEMM_FALLBACK: '1'
run: |
rustup target add x86_64-apple-darwin
cargo test --release -p paimon-vindex-core --target x86_64-apple-darwin config_from_options_pq_encoding
cargo test --release -p paimon-vindex-core --target x86_64-apple-darwin test_encode_batch_8bit_sgemm_is_thread_and_split_invariant

- name: Build wheel
working-directory: python
run: |
Expand Down
1 change: 1 addition & 0 deletions core/benches/ann_bench.rs
Original file line number Diff line number Diff line change
Expand Up @@ -667,6 +667,7 @@ fn index_specs(config: &Config) -> Vec<IndexSpec> {
metric: MetricType::L2,
use_opq: false,
use_approximate_coarse_assignment: true,
canonical_pq_encoding: false,
},
searches: vec![ivf_search],
},
Expand Down
17 changes: 3 additions & 14 deletions core/benches/ivfpq_add_bench.rs
Original file line number Diff line number Diff line change
Expand Up @@ -106,16 +106,10 @@ const CASES: [Case; 10] = [
},
];

fn new_index(
case: Case,
quantizer_centroids: &[f32],
centroids: &[f32],
norms: &[f32],
) -> IVFPQIndex {
fn new_index(case: Case, quantizer_centroids: &[f32], centroids: &[f32]) -> IVFPQIndex {
let mut index = IVFPQIndex::new(case.d, case.nlist, case.m, MetricType::L2, false);
index.set_quantizer_centroids(quantizer_centroids.to_vec());
index.pq.centroids = centroids.to_vec();
index.pq.centroid_norms_cache = norms.to_vec();
index.pq.set_centroids(centroids.to_vec());
index
}

Expand All @@ -138,11 +132,6 @@ fn bench_ivfpq_add(c: &mut Criterion) {
let centroids = (0..case.m * 256 * dsub)
.map(|_| rng.gen_range(-1.0f32..1.0))
.collect::<Vec<_>>();
let norms = centroids
.chunks_exact(dsub)
.map(|centroid| centroid.iter().map(|value| value * value).sum())
.collect::<Vec<_>>();

group.throughput(Throughput::Elements(case.rows as u64));
group.bench_with_input(
BenchmarkId::new(
Expand All @@ -155,7 +144,7 @@ fn bench_ivfpq_add(c: &mut Criterion) {
&case,
|b, &case| {
b.iter_batched(
|| new_index(case, &quantizer_centroids, &centroids, &norms),
|| new_index(case, &quantizer_centroids, &centroids),
|mut index| index.add(black_box(&data), black_box(&ids), case.rows),
BatchSize::LargeInput,
);
Expand Down
8 changes: 5 additions & 3 deletions core/benches/ivfpq_batch_reuse_bench.rs
Original file line number Diff line number Diff line change
Expand Up @@ -103,9 +103,11 @@ fn main() {
.map(|_| rng.gen_range(-1.0f32..1.0))
.collect(),
);
index.pq.centroids = (0..M * index.pq.ksub * index.pq.dsub)
.map(|_| rng.gen_range(-1.0f32..1.0))
.collect();
index.pq.set_centroids(
(0..M * index.pq.ksub() * index.pq.dsub())
.map(|_| rng.gen_range(-1.0f32..1.0))
.collect(),
);
for list_id in 0..NLIST {
let first_id = list_id * ROWS_PER_LIST;
index.ids[list_id] = (first_id..first_id + ROWS_PER_LIST)
Expand Down
8 changes: 5 additions & 3 deletions core/benches/ivfpq_filter_scan_bench.rs
Original file line number Diff line number Diff line change
Expand Up @@ -104,9 +104,11 @@ fn main() {
.map(|_| rng.gen_range(-1.0f32..1.0))
.collect(),
);
index.pq.centroids = (0..M * index.pq.ksub * index.pq.dsub)
.map(|_| rng.gen_range(-1.0f32..1.0))
.collect();
index.pq.set_centroids(
(0..M * index.pq.ksub() * index.pq.dsub())
.map(|_| rng.gen_range(-1.0f32..1.0))
.collect(),
);
for list_id in 0..NLIST {
let first_id = list_id * ROWS_PER_LIST;
index.ids[list_id] = (first_id..first_id + ROWS_PER_LIST)
Expand Down
2 changes: 1 addition & 1 deletion core/benches/ivfpq_train_bench.rs
Original file line number Diff line number Diff line change
Expand Up @@ -103,7 +103,7 @@ fn run_scenario(s: &Scenario) {

// Keep results observable so nothing is optimized away.
let checksum: f32 =
centroids.iter().take(8).sum::<f32>() + pq.centroids.iter().take(8).sum::<f32>();
centroids.iter().take(8).sum::<f32>() + pq.centroids().iter().take(8).sum::<f32>();

println!(
"{:<11} {:>8} {:>5} {:>6} {:>5} {:>8.3} {:>9.3} {:>8.3} {:>14.3}",
Expand Down
34 changes: 20 additions & 14 deletions core/src/diskann.rs
Original file line number Diff line number Diff line change
Expand Up @@ -317,8 +317,8 @@ impl DiskAnnIndex {
fn training_plan(&self, n: usize) -> io::Result<PqTrainingPlan> {
pq_training_plan_with_sample_buffers(
self.d,
self.pq.m,
self.pq.ksub,
self.pq.m(),
self.pq.ksub(),
n,
self.build_params.memory_budget_bytes,
usize::from(self.metric == MetricType::Cosine) + 1,
Expand All @@ -338,14 +338,15 @@ impl DiskAnnIndex {
let row_ids = checked_bytes(n, size_of::<i64>(), "row IDs")?;
let row_id_encoding_scratch = row_id_encoding_scratch_bytes(n)?;
let pq_codes = checked_bytes(n, self.pq.code_size(), "PQ codes")?;
let pq_codebook = checked_bytes(self.pq.centroids.len(), size_of::<f32>(), "PQ codebook")?;
let pq_codebook =
checked_bytes(self.pq.centroids().len(), size_of::<f32>(), "PQ codebook")?;
let pq_build_distances = if self.build_params.build_distance
== DiskAnnBuildDistance::ProductQuantized
{
self.pq
.m
.checked_mul(self.pq.ksub)
.and_then(|value| value.checked_mul(self.pq.ksub))
.m()
.checked_mul(self.pq.ksub())
.and_then(|value| value.checked_mul(self.pq.ksub()))
.and_then(|value| value.checked_mul(size_of::<f32>()))
.ok_or_else(|| invalid_input("DiskANN PQ build-distance table size overflows"))?
} else {
Expand Down Expand Up @@ -438,12 +439,17 @@ impl DiskAnnIndex {
}

pub(crate) fn validate_for_write(&self) -> io::Result<()> {
validate_diskann_format_configuration(self.d, self.pq.m, self.pq.nbits, self.build_params)?;
validate_diskann_format_configuration(
self.d,
self.pq.m(),
self.pq.nbits(),
self.build_params,
)?;
validate_diskann_training_budget(
self.d,
self.metric,
self.pq.m,
self.pq.nbits,
self.pq.m(),
self.pq.nbits(),
self.build_params.memory_budget_bytes,
)?;
if self.build_params.memory_budget_bytes == 0 {
Expand Down Expand Up @@ -489,22 +495,22 @@ impl DiskAnnIndex {
}

let expected_ksub = 1usize
.checked_shl(self.pq.nbits as u32)
.checked_shl(self.pq.nbits() as u32)
.ok_or_else(|| invalid_input("DiskANN PQ centroid count overflows usize"))?;
let expected_centroids = self
.d
.checked_mul(expected_ksub)
.ok_or_else(|| invalid_input("DiskANN PQ codebook shape overflows usize"))?;
if self.pq.d != self.d
|| self.pq.ksub != expected_ksub
|| self.pq.centroids.len() != expected_centroids
if self.pq.d() != self.d
|| self.pq.ksub() != expected_ksub
|| self.pq.centroids().len() != expected_centroids
|| !self.pq.has_valid_layout()
{
return Err(invalid_input("DiskANN PQ codebook shape is invalid"));
}
if let Some(offset) = self
.pq
.centroids
.centroids()
.iter()
.position(|value| !value.is_finite())
{
Expand Down
40 changes: 21 additions & 19 deletions core/src/diskann_io.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1433,10 +1433,8 @@ impl<R: SeekRead> DiskAnnIndexReader<R> {
)));
}

let (mut pq, row_ids, pq_codes, adjacency_index) =
let (pq, row_ids, pq_codes, adjacency_index) =
read_resident_sections(&mut self.reader, &self.header)?;
pq.try_rebuild_norms_cache()
.map_err(|_| invalid_data("DiskANN PQ norms allocation failed"))?;
validate_pq_code_padding(&self.header, &pq_codes)?;
let adjacency_validation =
AdjacencyValidationCache::new(adjacency_page_count(&self.header)?)?;
Expand Down Expand Up @@ -2312,8 +2310,8 @@ pub fn write_diskann_index_with_stats(
index.d,
index.ids.len(),
prepared.graph.entry_node,
index.pq.m,
index.pq.nbits,
index.pq.m(),
index.pq.nbits(),
index.metric,
index.build_params,
row_ids_len,
Expand Down Expand Up @@ -2588,38 +2586,38 @@ fn write_pq_codebook(
put_u32(
&mut header,
8,
u32::try_from(pq.d).map_err(|_| invalid_input("DiskANN PQ dimension exceeds u32"))?,
u32::try_from(pq.d()).map_err(|_| invalid_input("DiskANN PQ dimension exceeds u32"))?,
);
put_u32(
&mut header,
12,
u32::try_from(pq.m).map_err(|_| invalid_input("DiskANN PQ m exceeds u32"))?,
u32::try_from(pq.m()).map_err(|_| invalid_input("DiskANN PQ m exceeds u32"))?,
);
put_u32(
&mut header,
16,
u32::try_from(pq.nbits).map_err(|_| invalid_input("DiskANN PQ bits exceeds u32"))?,
u32::try_from(pq.nbits()).map_err(|_| invalid_input("DiskANN PQ bits exceeds u32"))?,
);
put_u32(
&mut header,
20,
u32::try_from(pq.ksub).map_err(|_| invalid_input("DiskANN PQ ksub exceeds u32"))?,
u32::try_from(pq.ksub()).map_err(|_| invalid_input("DiskANN PQ ksub exceeds u32"))?,
);
put_u32(
&mut header,
24,
u32::try_from(pq.chunk_offsets.len())
u32::try_from(pq.chunk_offsets().len())
.map_err(|_| invalid_input("DiskANN PQ chunk-offset count exceeds u32"))?,
);
writer.write_bytes(&header)?;
for &offset in &pq.chunk_offsets {
for &offset in pq.chunk_offsets() {
writer.write_bytes(
&u32::try_from(offset)
.map_err(|_| invalid_input("DiskANN PQ chunk offset exceeds u32"))?
.to_le_bytes(),
)?;
}
for &value in &pq.centroids {
for &value in pq.centroids() {
writer.write_bytes(&value.to_le_bytes())?;
}
Ok(())
Expand Down Expand Up @@ -3050,7 +3048,7 @@ fn decode_pq_codebook(bytes: &[u8], header: &DiskAnnHeader) -> io::Result<Produc
.map_err(invalid_data)?;
let mut centroids = Vec::new();
centroids
.try_reserve_exact(header.dimension as usize * pq.ksub)
.try_reserve_exact(header.dimension as usize * pq.ksub())
.map_err(|_| invalid_data("DiskANN PQ centroid allocation failed"))?;
for encoded in bytes[centroids_offset..].chunks_exact(4) {
let value =
Expand All @@ -3060,7 +3058,8 @@ fn decode_pq_codebook(bytes: &[u8], header: &DiskAnnHeader) -> io::Result<Produc
}
centroids.push(value);
}
pq.centroids = centroids;
pq.try_set_centroids(centroids)
.map_err(|_| invalid_data("DiskANN PQ norms allocation failed"))?;
if !pq.has_valid_layout() {
return Err(invalid_data("invalid DiskANN PQ codebook layout"));
}
Expand Down Expand Up @@ -4120,15 +4119,18 @@ mod tests {
..DiskAnnBuildParams::default()
},
);
index.pq.centroids = (0..256).map(|code| code as f32).collect();
index.pq.rebuild_norms_cache();
index
.pq
.set_centroids((0..256).map(|code| code as f32).collect());
index.ids = vec![7];
index.vectors = vec![0.0];
index
}

let mut invalid_codebook = one_vector_index();
invalid_codebook.pq.centroids[0] = f32::NAN;
let mut centroids = invalid_codebook.pq.centroids().to_vec();
centroids[0] = f32::NAN;
invalid_codebook.pq.set_centroids(centroids);
let mut codebook_output = Vec::new();
assert!(
write_diskann_index(&invalid_codebook, &mut PosWriter::new(&mut codebook_output))
Expand All @@ -4154,7 +4156,7 @@ mod tests {
assert!(f16_output.is_empty());

let mut invalid_pq_shape = one_vector_index();
invalid_pq_shape.pq.chunk_offsets[1] = 0;
invalid_pq_shape.pq = ProductQuantizer::new(invalid_pq_shape.d + 1, 1);
let mut pq_shape_output = Vec::new();
assert!(
write_diskann_index(&invalid_pq_shape, &mut PosWriter::new(&mut pq_shape_output))
Expand Down Expand Up @@ -4840,7 +4842,7 @@ mod tests {
);
assert_eq!(reader.row_id_count().unwrap(), indexed_count);
assert_eq!(reader.pq_codes().unwrap().len(), indexed_count * 2);
assert_eq!(reader.pq().unwrap().centroids, index.pq.centroids);
assert_eq!(reader.pq().unwrap().centroids(), index.pq.centroids());

let limited_rounds = Arc::new(Mutex::new(Vec::new()));
let limited_recording = RoundRecordingReader {
Expand Down
8 changes: 4 additions & 4 deletions core/src/diskann_search.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1460,7 +1460,7 @@ impl<R: SeekRead> crate::diskann_io::DiskAnnIndexReader<R> {
let query_count = queries.len() / dimension;
let pq_m = self.header.pq_m as usize;
let pq = self.pq()?;
let pq_ksub = pq.ksub;
let pq_ksub = pq.ksub();
let pq_code_size = pq.code_size();
let pq_codes = self.pq_codes()?;
let metric = self.header.metric_type();
Expand Down Expand Up @@ -1881,7 +1881,7 @@ impl<R: SeekRead> crate::diskann_io::DiskAnnIndexReader<R> {
self.header.max_degree as usize,
)?;
let pq = self.pq()?;
let distance_table_len = pq.m * pq.ksub;
let distance_table_len = pq.m() * pq.ksub();
pq.compute_distance_table(
query,
self.header.metric_type(),
Expand Down Expand Up @@ -2024,7 +2024,7 @@ impl<R: SeekRead> crate::diskann_io::DiskAnnIndexReader<R> {
let result = (|| {
scratch.begin_rerank();
let pq = self.pq()?;
let distance_table_len = pq.m * pq.ksub;
let distance_table_len = pq.m() * pq.ksub();
pq.compute_distance_table(
query,
self.header.metric_type(),
Expand Down Expand Up @@ -5432,7 +5432,7 @@ mod tests {

reader.ensure_resident().unwrap();
assert_eq!(reader.header.pq_bits, 4);
assert_eq!(reader.pq().unwrap().ksub, 16);
assert_eq!(reader.pq().unwrap().ksub(), 16);
assert_eq!(reader.pq_codes().unwrap().len(), indexed_count);

let (result_ids, distances) = reader.search(query, 5, 100).unwrap();
Expand Down
Loading
Loading