Skip to content
Open
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
10 changes: 7 additions & 3 deletions core/src/distance.rs
Original file line number Diff line number Diff line change
Expand Up @@ -972,8 +972,12 @@ pub fn scan_4bit_simd(sim_table: &[f32], codes: &[u8], count: usize, m: usize, d

let cs = m / 2; // code_size = m/2 bytes per vector

// Step 1: Compute first FLAT_NUM vectors with f32 precision
let flat_end = count.min(FLAT_NUM);
// Step 1: Keep large-M configurations out of the u16 accumulation path.
let flat_end = if m > crate::fastscan::MAX_U16_SUBQUANTIZERS {
count
} else {
count.min(FLAT_NUM)
};
for i in 0..flat_end {
let base = i * cs;
let mut d = 0.0f32;
Expand All @@ -987,7 +991,7 @@ pub fn scan_4bit_simd(sim_table: &[f32], codes: &[u8], count: usize, m: usize, d
dists[i] = d;
}

if count <= FLAT_NUM {
if flat_end == count {
return;
}

Expand Down
14 changes: 11 additions & 3 deletions core/src/fastscan.rs
Original file line number Diff line number Diff line change
Expand Up @@ -27,6 +27,10 @@
/// Block size: 32 vectors per block (matches AVX2 register width).
pub const BBS: usize = 32;

// Every subquantizer can contribute 255 to a u16 accumulator. Larger
// configurations must use the exact f32 path instead of wrapping the sum.
pub(crate) const MAX_U16_SUBQUANTIZERS: usize = u16::MAX as usize / u8::MAX as usize;

/// Pack 4-bit codes from row-major [n][cs] into block layout.
/// Output layout: [num_blocks][cs][BBS] where cs = M/2.
/// Pads the last block with zeros if n is not a multiple of BBS.
Expand Down Expand Up @@ -97,9 +101,13 @@ pub fn quantize_distance_table(table: &[f32], qmax_hint: f32) -> (f32, f32, Vec<
pub fn fastscan_4bit(sim_table: &[f32], codes: &[u8], n: usize, m: usize, dists: &mut [f32]) {
let cs = m / 2;

// Step 1: f32 exact for first min(200, n) vectors as qmax calibration
// Step 1: Scan the prefix exactly, or all rows if u16 accumulation is unsafe.
const FLAT_NUM: usize = 200;
let flat_end = n.min(FLAT_NUM);
let flat_end = if m > MAX_U16_SUBQUANTIZERS {
n
} else {
n.min(FLAT_NUM)
};
let block_size = cs * BBS;

for i in 0..flat_end {
Expand All @@ -116,7 +124,7 @@ pub fn fastscan_4bit(sim_table: &[f32], codes: &[u8], n: usize, m: usize, dists:
dists[i] = d;
}

if n <= FLAT_NUM {
if flat_end == n {
return;
}

Expand Down
8 changes: 6 additions & 2 deletions core/src/ivfpq.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1181,7 +1181,11 @@ fn scan_codes_4bit_transposed(
let cs = m / 2;

const FLAT_NUM: usize = 200;
let flat_end = count.min(FLAT_NUM);
let flat_end = if m > crate::fastscan::MAX_U16_SUBQUANTIZERS {
count
} else {
count.min(FLAT_NUM)
};

let mut dists = vec![0.0f32; count];

Expand All @@ -1197,7 +1201,7 @@ fn scan_codes_4bit_transposed(
dists[i] = d;
}

if count > FLAT_NUM {
if flat_end < count {
let qmin = sim_table.iter().cloned().fold(f32::INFINITY, f32::min);
let qmax = dists[..flat_end].iter().cloned().fold(f32::MIN, f32::max);
let range = (qmax - qmin).max(1e-10);
Expand Down
107 changes: 107 additions & 0 deletions core/tests/ivfpq_4bit_overflow.rs
Original file line number Diff line number Diff line change
@@ -0,0 +1,107 @@
// Licensed to the Apache Software Foundation (ASF) under one
// or more contributor license agreements. See the NOTICE file
// distributed with this work for additional information
// regarding copyright ownership. The ASF licenses this file
// to you under the Apache License, Version 2.0 (the
// "License"); you may not use this file except in compliance
// with the License. You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing,
// software distributed under the License is distributed on an
// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
// KIND, either express or implied. See the License for the
// specific language governing permissions and limitations
// under the License.

use paimon_vindex_core::distance::MetricType;
use paimon_vindex_core::io::{write_index, IVFPQIndexReader, PosWriter};
use paimon_vindex_core::ivfpq::{search_batch_reader, IVFPQIndex};
use std::io::Cursor;

const ROWS: usize = 225;

fn build_index(m: usize) -> IVFPQIndex {
let mut index = IVFPQIndex::with_nbits(m, 1, m, 4, MetricType::L2, false);
index.set_quantizer_centroids(vec![0.0; m]);
index.pq.centroids = vec![0.0; m * 16];
for sub in 0..m {
index.pq.centroids[sub * 16 + 1] = 1.0;
}
index.pq.rebuild_norms_cache();

// The first 200 distances equal 1, matching the full LUT maximum.
// This isolates accumulation overflow from LUT-range calibration.
let mut data = vec![0.0; ROWS * m];
for row in 0..200 {
data[row * m] = 1.0;
}
data[200 * m..].fill(1.0);
let ids = (0..ROWS as i64).collect::<Vec<_>>();
index.add(&data, &ids, ROWS);
index
}

fn assert_distances(ids: &[i64], distances: &[f32], m: usize) {
assert_eq!(ids.len(), ROWS);
assert_eq!(distances.len(), ROWS);
let mut seen = vec![false; ROWS];
for (&id, &distance) in ids.iter().zip(distances) {
assert!((0..ROWS as i64).contains(&id));
assert!(!seen[id as usize], "duplicate row ID {id}");
seen[id as usize] = true;
let expected = if id < 200 { 1.0 } else { m as f32 };
assert!(
(distance - expected).abs() < 1e-4,
"m={m}, row={id}: expected {expected}, got {distance}"
);
}
}

#[test]
fn four_bit_in_memory_scans_avoid_accumulator_overflow() {
for m in [256, 258, 512] {
let mut index = build_index(m);
let query = vec![0.0; m];
for fastscan in [false, true] {
if fastscan {
index.build_search_structures();
}
let mut ids = vec![-1; ROWS];
let mut distances = vec![0.0; ROWS];
index.search(&query, 1, ROWS, 1, &mut distances, &mut ids);
assert_distances(&ids, &distances, m);
index.search_with_max_codes(&query, 1, ROWS, 1, ROWS, &mut distances, &mut ids);
assert_distances(&ids, &distances, m);
}
}
}

#[test]
fn four_bit_reader_scans_avoid_accumulator_overflow() {
for m in [256, 258, 512] {
let index = build_index(m);
let query = vec![0.0; m];
let mut bytes = Vec::new();
write_index(&index, &mut PosWriter::new(&mut bytes)).unwrap();
let mut reader = IVFPQIndexReader::open(Cursor::new(bytes)).unwrap();
for precomputed in [false, true] {
if precomputed {
reader.optimize_for_search().unwrap();
}
let (ids, distances) = reader.search(&query, ROWS, 1).unwrap();
assert_distances(&ids, &distances, m);
let (ids, distances) =
search_batch_reader(&mut reader, &query.repeat(2), 2, ROWS, 1).unwrap();
for row in 0..2 {
let start = row * ROWS;
assert_distances(
&ids[start..start + ROWS],
&distances[start..start + ROWS],
m,
);
}
}
}
}