From 4289e7482584e7e4694dc5debacf3b346e96d517 Mon Sep 17 00:00:00 2001 From: Daoyuan Wang Date: Wed, 9 Sep 2026 11:28:35 +0800 Subject: [PATCH] fix: avoid 4-bit IVF-PQ accumulator overflow Use exact f32 scanning when the number of subquantizers can overflow a u16 accumulator. Apply the bound to row-major, transposed, and FastScan layouts before integer accumulation begins. Cover m=256, 258, and 512 with exactly representable vectors and a prefix whose distance range matches the LUT, independently of range calibration. Exercise in-memory, budgeted, persisted, and batch searches. Validation: Rust 1.95 core tests (497 passed, 2 ignored), rustfmt, and Clippy. --- core/src/distance.rs | 10 ++- core/src/fastscan.rs | 14 +++- core/src/ivfpq.rs | 8 ++- core/tests/ivfpq_4bit_overflow.rs | 107 ++++++++++++++++++++++++++++++ 4 files changed, 131 insertions(+), 8 deletions(-) create mode 100644 core/tests/ivfpq_4bit_overflow.rs diff --git a/core/src/distance.rs b/core/src/distance.rs index 53a13c24..c6ee8b68 100644 --- a/core/src/distance.rs +++ b/core/src/distance.rs @@ -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; @@ -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; } diff --git a/core/src/fastscan.rs b/core/src/fastscan.rs index 68655901..e837ab89 100644 --- a/core/src/fastscan.rs +++ b/core/src/fastscan.rs @@ -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. @@ -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 { @@ -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; } diff --git a/core/src/ivfpq.rs b/core/src/ivfpq.rs index 15e530c2..c9a164db 100644 --- a/core/src/ivfpq.rs +++ b/core/src/ivfpq.rs @@ -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]; @@ -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); diff --git a/core/tests/ivfpq_4bit_overflow.rs b/core/tests/ivfpq_4bit_overflow.rs new file mode 100644 index 00000000..3b6aa1ee --- /dev/null +++ b/core/tests/ivfpq_4bit_overflow.rs @@ -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::>(); + 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, + ); + } + } + } +}