From 59e22f3d0d86526442fede2487feef69dafa7ca0 Mon Sep 17 00:00:00 2001 From: Suryansh Gupta Date: Wed, 16 Sep 2026 22:05:07 +0530 Subject: [PATCH 1/7] PACK support in packed.rs --- .../src/matrix_kernels/blocks/packed.rs | 168 ++++++++++++++---- 1 file changed, 138 insertions(+), 30 deletions(-) diff --git a/diskann-quantization/src/matrix_kernels/blocks/packed.rs b/diskann-quantization/src/matrix_kernels/blocks/packed.rs index aa8d4b0a29..f8e1e05bd3 100644 --- a/diskann-quantization/src/matrix_kernels/blocks/packed.rs +++ b/diskann-quantization/src/matrix_kernels/blocks/packed.rs @@ -14,25 +14,40 @@ use crate::{ multi_vector::BlockTransposedRef, }; +/// Round the logical contraction dimension up to the extent physically stored per band. +#[inline] +fn padded_k(k: usize) -> usize { + k.next_multiple_of(PACK) +} + /// A view over packed memory. /// /// Elements are gathered into groups of size `SZ`. A collection of `self.k` groups forms /// a "block". `self.blocks` tracks how many such blocks are in the view. /// +/// `PACK` interleaves that many consecutive columns within each group, so one physical row +/// spans `SZ * PACK` elements. `k` stays the logical dimension; a trailing row that `PACK` +/// does not fill is zero-padded. +/// /// This layout requires that no block is partially filled. /// /// # Class Invariants /// -/// * The tracked length `ptr.len()` must be equal to `SZ * blocks * k`. +/// * The tracked length `ptr.len()` must be equal to `SZ * blocks * padded_k::(k)`. /// * `SZ` may not be zero. #[derive(Debug, Clone, Copy)] -pub(crate) struct View<'a, T, const SZ: usize> { +pub(crate) struct View<'a, T, const SZ: usize, const PACK: usize = 1> { ptr: Slice<'a, T>, blocks: NonZeroUsize, k: Bound, } -impl<'a, T, const SZ: usize> View<'a, T, SZ> { +impl<'a, T, const SZ: usize, const PACK: usize> View<'a, T, SZ, PACK> { + const _ASSERTIONS: () = { + assert!(PACK > 0, "packing factor PACK must be positive"); + assert!(SZ.is_multiple_of(PACK), "SZ must be divisible by PACK"); + }; + /// Construct a [`View`] from a [`BlockTransposedRef`]. /// /// The mapping of parameters is as follows: @@ -42,7 +57,7 @@ impl<'a, T, const SZ: usize> View<'a, T, SZ> { /// * The number of blocks is [`BlockTransposedRef::num_blocks`]. /// /// Returns `None` if any of the runtime values is zero. - pub(crate) fn from_block_transposed(v: BlockTransposedRef<'a, T, SZ>) -> Option + pub(crate) fn from_block_transposed(v: BlockTransposedRef<'a, T, SZ, PACK>) -> Option where T: Copy, { @@ -53,22 +68,23 @@ impl<'a, T, const SZ: usize> View<'a, T, SZ> { let blocks = NonZeroUsize::new(v.num_blocks())?; let k = DimK::new(NonZeroUsize::new(v.ncols())?); - // SAFETY: `BlockTransposedRef` ensures the underlying slice has a length of - // exactly `SZ * blocks * k`. + // SAFETY: `BlockTransposedRef` sizes its allocation as `SZ * blocks * padded_ncols`, + // and `padded_ncols` is `padded_k::(ncols)`. Some(unsafe { Self::new(Slice::new(v.as_slice()), blocks, k) }) } /// # Safety /// - /// `ptr.len()` must be exactly equal to `SZ * blocks * k`. + /// `ptr.len()` must be exactly equal to `SZ * blocks * padded_k::(k)`. pub(in crate::matrix_kernels) unsafe fn new( ptr: Slice<'a, T>, blocks: NonZeroUsize, k: DimK, ) -> Self { + let () = Self::_ASSERTIONS; bounds::check_eq!( ptr.len(), - blocks.get() * SZ * k.value().get(), + blocks.get() * SZ * padded_k::(k.value().get()), "invalid block-transposed access", ); bounds::check_lt!(Bound::new(0), SZ, "group size may not be zero.",); @@ -79,13 +95,15 @@ impl<'a, T, const SZ: usize> View<'a, T, SZ> { /// # Safety /// - /// `ptr.len()` must be exactly equal to `SZ * blocks * k`. + /// `ptr.len()` must be exactly equal to `SZ * blocks * padded_k::(k)`. unsafe fn new_inner(ptr: Slice<'a, T>, blocks: NonZeroUsize, k: Bound) -> Self { - bounds::check_eq!( - ptr.len(), - Bound::new(blocks.get()) * Bound::new(SZ) * k, - "invalid block-transposed access", - ); + k.with(|k| { + bounds::check_eq!( + ptr.len(), + blocks.get() * SZ * padded_k::(k), + "invalid block-transposed access", + ); + }); bounds::check_lt!(Bound::new(0), SZ, "group size may not be zero.",); Self { ptr, blocks, k } @@ -108,7 +126,7 @@ impl<'a, T, const SZ: usize> View<'a, T, SZ> { /// `k` must be equal to the contraction dimension tracked by [`Self::k`]. pub(in crate::matrix_kernels) fn block_stride(&self, k: DimK) -> Elements { bounds::check_eq!(self.k, k.value()); - Elements::new(SZ * k.value().get()) + Elements::new(SZ * padded_k::(k.value().get())) } /// Return the number of bands stored in `self`. @@ -135,7 +153,7 @@ impl<'a, T, const SZ: usize> View<'a, T, SZ> { k: DimK, mut f: F, ) where - F: FnMut(View<'_, T, SZ>, usize), + F: FnMut(View<'_, T, SZ, PACK>, usize), { let stride = self.block_stride(k); @@ -181,7 +199,7 @@ impl<'a, T, const SZ: usize> View<'a, T, SZ> { /// The bound [`Self::k`] must be equal to `k`. pub(in crate::matrix_kernels) unsafe fn visit_panels(&self, k: DimK, mut f: F) where - F: FnMut(Panel<'_, T, SZ>, usize), + F: FnMut(Panel<'_, T, SZ, PACK>, usize), { let stride = self.block_stride(k); for b in 0..self.blocks().get() { @@ -202,10 +220,10 @@ impl<'a, T, const SZ: usize> View<'a, T, SZ> { } #[cfg(test)] -impl<'a, T, const SZ: usize> View<'a, T, SZ> { +impl View<'_, T, SZ, PACK> { fn checked_visit_sub_views(&self, sub_blocks: NonZeroUsize, f: F) where - F: FnMut(View<'_, T, SZ>, usize), + F: FnMut(View<'_, T, SZ, PACK>, usize), { let k = DimK::from_bound(self.k()); // SAFETY: Checked in test builds. @@ -214,7 +232,7 @@ impl<'a, T, const SZ: usize> View<'a, T, SZ> { fn checked_visit_panels(&self, f: F) where - F: FnMut(Panel<'_, T, SZ>, usize), + F: FnMut(Panel<'_, T, SZ, PACK>, usize), { let k = DimK::from_bound(self.k()); // SAFETY: Checked in test builds. @@ -228,31 +246,33 @@ impl<'a, T, const SZ: usize> View<'a, T, SZ> { /// A block containing `k` contiguous groups of size `SZ`. /// +/// `PACK` consecutive groups are interleaved into one physical row of `SZ * PACK` elements. +/// /// # Class Invariants /// -/// The bound `ptr.len()` must be equal to `SZ * k`. +/// The bound `ptr.len()` must be equal to `SZ * padded_k::(k)`. #[derive(Debug, Clone, Copy)] -pub(in crate::matrix_kernels) struct Panel<'a, T, const SZ: usize> { +pub(in crate::matrix_kernels) struct Panel<'a, T, const SZ: usize, const PACK: usize = 1> { ptr: Slice<'a, T>, k: Bound, } -impl<'a, T, const SZ: usize> Panel<'a, T, SZ> { +impl<'a, T, const SZ: usize, const PACK: usize> Panel<'a, T, SZ, PACK> { /// # Safety /// - /// `ptr.len()` must be equal to `SZ * k`. + /// `ptr.len()` must be equal to `SZ * padded_k::(k)`. #[cfg(test)] pub(in crate::matrix_kernels) unsafe fn new(ptr: Slice<'a, T>, k: DimK) -> Self { - bounds::check_eq!(ptr.len(), SZ * k.value().get()); + bounds::check_eq!(ptr.len(), SZ * padded_k::(k.value().get())); // SAFETY: Inherited from caller. unsafe { Self::new_inner(ptr, Bound::new(k.value().get())) } } /// # Safety /// - /// `ptr.len()` must be equal to `SZ * k`. + /// `ptr.len()` must be equal to `SZ * padded_k::(k)`. unsafe fn new_inner(ptr: Slice<'a, T>, k: Bound) -> Self { - k.with(|k| bounds::check_eq!(ptr.len(), SZ * k)); + k.with(|k| bounds::check_eq!(ptr.len(), SZ * padded_k::(k))); Self { ptr, k } } @@ -268,12 +288,29 @@ impl<'a, T, const SZ: usize> Panel<'a, T, SZ> { pub(in crate::matrix_kernels) const fn k(&self) -> Bound { self.k } + + /// Return the number of elements spanned by one physical row. + // Exercised by tests; consumed by PACK-aware kernels once they land. + #[allow(dead_code)] + pub(in crate::matrix_kernels) const fn row_stride(&self) -> Elements { + Elements::new(SZ * PACK) + } + + /// Return the number of physical rows in `self`. + /// + /// `k` must be equal to the contraction dimension tracked by [`Self::k`]. + // Exercised by tests; consumed by PACK-aware kernels once they land. + #[allow(dead_code)] + pub(in crate::matrix_kernels) fn rows(&self, k: DimK) -> usize { + bounds::check_eq!(self.k, k.value()); + padded_k::(k.value().get()) / PACK + } } #[cfg(test)] -impl<'a, T, const SZ: usize> Panel<'a, T, SZ> { +impl<'a, T, const SZ: usize, const PACK: usize> Panel<'a, T, SZ, PACK> { fn checked_as_std_slice(self) -> &'a [T] { - let len = SZ * self.k().value(); + let len = SZ * padded_k::(self.k().value()); // SAFETY: Bounds are retained under `cfg(test)`. unsafe { self.ptr.as_std_slice(len) } } @@ -289,7 +326,78 @@ mod tests { use diskann_utils::views::{Init, Matrix, MatrixView}; - use crate::matrix_kernels::test_util::{assert_contains, panic_message_for}; + use crate::{ + matrix_kernels::test_util::{assert_contains, panic_message_for}, + multi_vector::BlockTransposed, + }; + + /// Pin the physical element order against the source matrix, rather than inferring it + /// from the layout documentation. Kernels index panel memory with this offset formula. + #[test] + fn test_pack_layout() { + for ncols in 1..20 { + assert_pack_layout::<1, 1>(3, ncols); + assert_pack_layout::<4, 1>(9, ncols); + assert_pack_layout::<8, 2>(20, ncols); + assert_pack_layout::<8, 4>(20, ncols); + assert_pack_layout::<16, 4>(33, ncols); + } + } + + fn assert_pack_layout(nrows: usize, ncols: usize) { + let ctx = format_args!("SZ = {SZ}, PACK = {PACK}, nrows = {nrows}, ncols = {ncols}"); + + // Values start at one so that zero unambiguously marks a padded slot. + let mut value = 0.0; + let matrix = Matrix::new( + Init(|| { + value += 1.0; + value + }), + nrows, + ncols, + ); + + let bt = BlockTransposed::::from_matrix_view(matrix.as_view()); + let padded = bt.padded_ncols(); + + let view = View::::from_block_transposed(bt.as_view()).unwrap(); + assert_eq!(view.blocks().get(), nrows.div_ceil(SZ), "{ctx}"); + assert_eq!(view.k().value(), ncols, "{ctx}"); + + let dim_k = DimK::from_bound(view.k()); + let mut blocks = 0; + + view.checked_visit_panels(|panel, block| { + assert_eq!(block, blocks, "{ctx}"); + assert_eq!(panel.rows(dim_k), padded / PACK, "{ctx}"); + assert_eq!(panel.row_stride().value(), SZ * PACK, "{ctx}"); + + let flat = panel.checked_as_std_slice(); + assert_eq!(flat.len(), SZ * padded, "{ctx}"); + + for col in 0..padded { + for row in 0..SZ { + let global_row = block * SZ + row; + let expected = if col < ncols && global_row < nrows { + matrix[(global_row, col)] + } else { + 0.0 + }; + + let offset = (col / PACK) * SZ * PACK + row * PACK + (col % PACK); + assert_eq!( + flat[offset], expected, + "{ctx}, block = {block}, row = {row}, col = {col}", + ); + } + } + + blocks += 1; + }); + + assert_eq!(blocks, nrows.div_ceil(SZ), "{ctx}"); + } #[test] fn test_visit_panels() { From 6b5293da2312d3684c8a7a48300d61549cf4bd3c Mon Sep 17 00:00:00 2001 From: Suryansh Gupta Date: Mon, 21 Sep 2026 18:38:59 +0530 Subject: [PATCH 2/7] Add v3 and Scalar kernel --- diskann-benchmark/example/multi-vector.json | 22 + .../perf_test_inputs/multi-vector.json | 90 ++ diskann-benchmark/src/multi_vector/driver.rs | 8 +- diskann-benchmark/src/multi_vector/kernels.rs | 1 + .../src/matrix_kernels/blocks/packed.rs | 4 - .../src/matrix_kernels/maxsim/mod.rs | 1 + .../maxsim/packed_f32_x_unpacked_f16.rs | 2 +- .../maxsim/packed_f32_x_unpacked_f32.rs | 7 +- .../maxsim/packed_i8_x_unpacked_i8.rs | 921 ++++++++++++++++++ .../src/matrix_kernels/maxsim/test.rs | 32 +- .../src/matrix_kernels/test_util.rs | 6 + .../src/matrix_kernels/util.rs | 9 + .../src/multi_vector/distance/factory.rs | 367 ++++++- .../src/multi_vector/distance/fallback.rs | 40 + .../src/multi_vector/distance/kernel.rs | 12 +- 15 files changed, 1465 insertions(+), 57 deletions(-) create mode 100644 diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs diff --git a/diskann-benchmark/example/multi-vector.json b/diskann-benchmark/example/multi-vector.json index af66a886d3..2dd53693d0 100644 --- a/diskann-benchmark/example/multi-vector.json +++ b/diskann-benchmark/example/multi-vector.json @@ -42,6 +42,28 @@ { "num_query_vectors": 8, "num_doc_vectors": 32, "dim": 128, "loops_per_measurement": 2, "num_measurements": 1 } ] } + }, + { + "type": "multi-vector-op", + "content": { + "element_type": "int8", + "isa": "auto", + "runs": [ + { "num_query_vectors": 8, "num_doc_vectors": 32, "dim": 128, "loops_per_measurement": 2, "num_measurements": 1 }, + { "num_query_vectors": 32, "num_doc_vectors": 16, "dim": 256, "loops_per_measurement": 2, "num_measurements": 1 } + ] + } + }, + { + "type": "multi-vector-op", + "content": { + "element_type": "int8", + "isa": "reference", + "runs": [ + { "num_query_vectors": 8, "num_doc_vectors": 32, "dim": 128, "loops_per_measurement": 2, "num_measurements": 1 }, + { "num_query_vectors": 64, "num_doc_vectors": 32, "dim": 264, "loops_per_measurement": 2, "num_measurements": 1 } + ] + } } ] } diff --git a/diskann-benchmark/perf_test_inputs/multi-vector.json b/diskann-benchmark/perf_test_inputs/multi-vector.json index c4ce9bb8bf..65ba11130b 100644 --- a/diskann-benchmark/perf_test_inputs/multi-vector.json +++ b/diskann-benchmark/perf_test_inputs/multi-vector.json @@ -144,6 +144,96 @@ { "num_query_vectors": 32, "num_doc_vectors": 32, "dim": 512, "loops_per_measurement": 50, "num_measurements": 50 } ] } + }, + { + "type": "multi-vector-op", + "content": { + "element_type": "int8", + "isa": "auto", + "runs": [ + { "num_query_vectors": 8, "num_doc_vectors": 32, "dim": 128, "loops_per_measurement": 500, "num_measurements": 50 }, + { "num_query_vectors": 16, "num_doc_vectors": 64, "dim": 256, "loops_per_measurement": 100, "num_measurements": 50 }, + { "num_query_vectors": 32, "num_doc_vectors": 128, "dim": 384, "loops_per_measurement": 20, "num_measurements": 50 }, + { "num_query_vectors": 32, "num_doc_vectors": 16, "dim": 256, "loops_per_measurement": 200, "num_measurements": 50 }, + { "num_query_vectors": 64, "num_doc_vectors": 32, "dim": 264, "loops_per_measurement": 50, "num_measurements": 50 }, + { "num_query_vectors": 32, "num_doc_vectors": 1250, "dim": 128, "loops_per_measurement": 10, "num_measurements": 50 }, + { "num_query_vectors": 64, "num_doc_vectors": 1250, "dim": 512, "loops_per_measurement": 2, "num_measurements": 50 }, + { "num_query_vectors": 64, "num_doc_vectors": 32, "dim": 128, "loops_per_measurement": 200, "num_measurements": 50 }, + { "num_query_vectors": 32, "num_doc_vectors": 32, "dim": 512, "loops_per_measurement": 50, "num_measurements": 50 } + ] + } + }, + { + "type": "multi-vector-op", + "content": { + "element_type": "int8", + "isa": "scalar", + "runs": [ + { "num_query_vectors": 8, "num_doc_vectors": 32, "dim": 128, "loops_per_measurement": 500, "num_measurements": 50 }, + { "num_query_vectors": 16, "num_doc_vectors": 64, "dim": 256, "loops_per_measurement": 100, "num_measurements": 50 }, + { "num_query_vectors": 32, "num_doc_vectors": 128, "dim": 384, "loops_per_measurement": 20, "num_measurements": 50 }, + { "num_query_vectors": 32, "num_doc_vectors": 16, "dim": 256, "loops_per_measurement": 200, "num_measurements": 50 }, + { "num_query_vectors": 64, "num_doc_vectors": 32, "dim": 264, "loops_per_measurement": 50, "num_measurements": 50 }, + { "num_query_vectors": 32, "num_doc_vectors": 1250, "dim": 128, "loops_per_measurement": 10, "num_measurements": 50 }, + { "num_query_vectors": 64, "num_doc_vectors": 1250, "dim": 512, "loops_per_measurement": 2, "num_measurements": 50 }, + { "num_query_vectors": 64, "num_doc_vectors": 32, "dim": 128, "loops_per_measurement": 200, "num_measurements": 50 }, + { "num_query_vectors": 32, "num_doc_vectors": 32, "dim": 512, "loops_per_measurement": 50, "num_measurements": 50 } + ] + } + }, + { + "type": "multi-vector-op", + "content": { + "element_type": "int8", + "isa": "x86-64-v3", + "runs": [ + { "num_query_vectors": 8, "num_doc_vectors": 32, "dim": 128, "loops_per_measurement": 500, "num_measurements": 50 }, + { "num_query_vectors": 16, "num_doc_vectors": 64, "dim": 256, "loops_per_measurement": 100, "num_measurements": 50 }, + { "num_query_vectors": 32, "num_doc_vectors": 128, "dim": 384, "loops_per_measurement": 20, "num_measurements": 50 }, + { "num_query_vectors": 32, "num_doc_vectors": 16, "dim": 256, "loops_per_measurement": 200, "num_measurements": 50 }, + { "num_query_vectors": 64, "num_doc_vectors": 32, "dim": 264, "loops_per_measurement": 50, "num_measurements": 50 }, + { "num_query_vectors": 32, "num_doc_vectors": 1250, "dim": 128, "loops_per_measurement": 10, "num_measurements": 50 }, + { "num_query_vectors": 64, "num_doc_vectors": 1250, "dim": 512, "loops_per_measurement": 2, "num_measurements": 50 }, + { "num_query_vectors": 64, "num_doc_vectors": 32, "dim": 128, "loops_per_measurement": 200, "num_measurements": 50 }, + { "num_query_vectors": 32, "num_doc_vectors": 32, "dim": 512, "loops_per_measurement": 50, "num_measurements": 50 } + ] + } + }, + { + "type": "multi-vector-op", + "content": { + "element_type": "int8", + "isa": "x86-64-v4", + "runs": [ + { "num_query_vectors": 8, "num_doc_vectors": 32, "dim": 128, "loops_per_measurement": 500, "num_measurements": 50 }, + { "num_query_vectors": 16, "num_doc_vectors": 64, "dim": 256, "loops_per_measurement": 100, "num_measurements": 50 }, + { "num_query_vectors": 32, "num_doc_vectors": 128, "dim": 384, "loops_per_measurement": 20, "num_measurements": 50 }, + { "num_query_vectors": 32, "num_doc_vectors": 16, "dim": 256, "loops_per_measurement": 200, "num_measurements": 50 }, + { "num_query_vectors": 64, "num_doc_vectors": 32, "dim": 264, "loops_per_measurement": 50, "num_measurements": 50 }, + { "num_query_vectors": 32, "num_doc_vectors": 1250, "dim": 128, "loops_per_measurement": 10, "num_measurements": 50 }, + { "num_query_vectors": 64, "num_doc_vectors": 1250, "dim": 512, "loops_per_measurement": 2, "num_measurements": 50 }, + { "num_query_vectors": 64, "num_doc_vectors": 32, "dim": 128, "loops_per_measurement": 200, "num_measurements": 50 }, + { "num_query_vectors": 32, "num_doc_vectors": 32, "dim": 512, "loops_per_measurement": 50, "num_measurements": 50 } + ] + } + }, + { + "type": "multi-vector-op", + "content": { + "element_type": "int8", + "isa": "reference", + "runs": [ + { "num_query_vectors": 8, "num_doc_vectors": 32, "dim": 128, "loops_per_measurement": 500, "num_measurements": 50 }, + { "num_query_vectors": 16, "num_doc_vectors": 64, "dim": 256, "loops_per_measurement": 100, "num_measurements": 50 }, + { "num_query_vectors": 32, "num_doc_vectors": 128, "dim": 384, "loops_per_measurement": 20, "num_measurements": 50 }, + { "num_query_vectors": 32, "num_doc_vectors": 16, "dim": 256, "loops_per_measurement": 200, "num_measurements": 50 }, + { "num_query_vectors": 64, "num_doc_vectors": 32, "dim": 264, "loops_per_measurement": 50, "num_measurements": 50 }, + { "num_query_vectors": 32, "num_doc_vectors": 1250, "dim": 128, "loops_per_measurement": 10, "num_measurements": 50 }, + { "num_query_vectors": 64, "num_doc_vectors": 1250, "dim": 512, "loops_per_measurement": 2, "num_measurements": 50 }, + { "num_query_vectors": 64, "num_doc_vectors": 32, "dim": 128, "loops_per_measurement": 200, "num_measurements": 50 }, + { "num_query_vectors": 32, "num_doc_vectors": 32, "dim": 512, "loops_per_measurement": 50, "num_measurements": 50 } + ] + } } ] } diff --git a/diskann-benchmark/src/multi_vector/driver.rs b/diskann-benchmark/src/multi_vector/driver.rs index e69c708451..5cec88bc02 100644 --- a/diskann-benchmark/src/multi_vector/driver.rs +++ b/diskann-benchmark/src/multi_vector/driver.rs @@ -14,7 +14,9 @@ use diskann_benchmark_runner::{ }, Checker, Input, }; -use diskann_quantization::multi_vector::{Mat, MatRef, MaxSimKernel, Overflow, Standard}; +use diskann_quantization::multi_vector::{ + Mat, MatRef, MaxSimElement, MaxSimKernel, Overflow, Standard, +}; use rand::{ distr::{Distribution, StandardUniform}, rngs::StdRng, @@ -94,12 +96,12 @@ where // Timing harness // ////////////////////// -pub(super) fn run_with_kernel( +pub(super) fn run_with_kernel( run: &Run, doc: MatRef<'_, Standard>, kernel: &dyn MaxSimKernel, ) -> RunResult { - let mut scores = vec![0.0f32; run.num_query_vectors.get()]; + let mut scores = vec![T::Score::default(); run.num_query_vectors.get()]; let mut latencies = Vec::with_capacity(run.num_measurements.get()); for _ in 0..run.num_measurements.get() { diff --git a/diskann-benchmark/src/multi_vector/kernels.rs b/diskann-benchmark/src/multi_vector/kernels.rs index d45ac646ba..7e169cbd9f 100644 --- a/diskann-benchmark/src/multi_vector/kernels.rs +++ b/diskann-benchmark/src/multi_vector/kernels.rs @@ -147,5 +147,6 @@ where pub(super) fn register(registry: &mut Registry) -> anyhow::Result<()> { registry.register_regression("multi-vector-op-f32", Kernel::::new())?; registry.register_regression("multi-vector-op-f16", Kernel::::new())?; + registry.register_regression("multi-vector-op-i8", Kernel::::new())?; Ok(()) } diff --git a/diskann-quantization/src/matrix_kernels/blocks/packed.rs b/diskann-quantization/src/matrix_kernels/blocks/packed.rs index f8e1e05bd3..48e40a8059 100644 --- a/diskann-quantization/src/matrix_kernels/blocks/packed.rs +++ b/diskann-quantization/src/matrix_kernels/blocks/packed.rs @@ -290,8 +290,6 @@ impl<'a, T, const SZ: usize, const PACK: usize> Panel<'a, T, SZ, PACK> { } /// Return the number of elements spanned by one physical row. - // Exercised by tests; consumed by PACK-aware kernels once they land. - #[allow(dead_code)] pub(in crate::matrix_kernels) const fn row_stride(&self) -> Elements { Elements::new(SZ * PACK) } @@ -299,8 +297,6 @@ impl<'a, T, const SZ: usize, const PACK: usize> Panel<'a, T, SZ, PACK> { /// Return the number of physical rows in `self`. /// /// `k` must be equal to the contraction dimension tracked by [`Self::k`]. - // Exercised by tests; consumed by PACK-aware kernels once they land. - #[allow(dead_code)] pub(in crate::matrix_kernels) fn rows(&self, k: DimK) -> usize { bounds::check_eq!(self.k, k.value()); padded_k::(k.value().get()) / PACK diff --git a/diskann-quantization/src/matrix_kernels/maxsim/mod.rs b/diskann-quantization/src/matrix_kernels/maxsim/mod.rs index 5d2d7cee22..412f9d98c5 100644 --- a/diskann-quantization/src/matrix_kernels/maxsim/mod.rs +++ b/diskann-quantization/src/matrix_kernels/maxsim/mod.rs @@ -18,6 +18,7 @@ pub(crate) mod packed_f32_x_unpacked_f16; pub(crate) mod packed_f32_x_unpacked_f32; +pub(crate) mod packed_i8_x_unpacked_i8; #[cfg(test)] mod test; diff --git a/diskann-quantization/src/matrix_kernels/maxsim/packed_f32_x_unpacked_f16.rs b/diskann-quantization/src/matrix_kernels/maxsim/packed_f32_x_unpacked_f16.rs index 79d48184fc..bcb7e4eabb 100644 --- a/diskann-quantization/src/matrix_kernels/maxsim/packed_f32_x_unpacked_f16.rs +++ b/diskann-quantization/src/matrix_kernels/maxsim/packed_f32_x_unpacked_f16.rs @@ -266,7 +266,7 @@ mod tests { let k = DimK::new(NonZeroUsize::new(k).unwrap()); let (ref_a, ref_b, ref_c) = - maxsim::test::generate(total_a_rows, k.value().get(), total_b_cols, rng); + maxsim::test::generate_f32(total_a_rows, k.value().get(), total_b_cols, rng); // Massage the input data in the form needed by the kernel. let a_bt = BlockTransposed::::from_matrix_view(ref_a.as_view()); diff --git a/diskann-quantization/src/matrix_kernels/maxsim/packed_f32_x_unpacked_f32.rs b/diskann-quantization/src/matrix_kernels/maxsim/packed_f32_x_unpacked_f32.rs index 3de91f7a56..7267178a69 100644 --- a/diskann-quantization/src/matrix_kernels/maxsim/packed_f32_x_unpacked_f32.rs +++ b/diskann-quantization/src/matrix_kernels/maxsim/packed_f32_x_unpacked_f32.rs @@ -803,7 +803,7 @@ mod tests { ) where for<'a> MicroKernel<'a, A, MR, NR>: driver::MicroKernel, { - let (ref_a, ref_b, ref_c) = maxsim::test::generate(MR, k.value().get(), NR, rng); + let (ref_a, ref_b, ref_c) = maxsim::test::generate_f32(MR, k.value().get(), NR, rng); // From the reference problem, we need to transpose both `ref_a` and `ref_b` to get // them into the desired format. @@ -931,7 +931,8 @@ mod tests { continue; } - let (ref_a, ref_b, ref_c) = maxsim::test::generate(MR, k.value().get(), cols, rng); + let (ref_a, ref_b, ref_c) = + maxsim::test::generate_f32(MR, k.value().get(), cols, rng); // From the reference problem, we need to transpose both `ref_a` and `ref_b` to get // them into the desired format. @@ -1056,7 +1057,7 @@ mod tests { let k = DimK::new(NonZeroUsize::new(k).unwrap()); let (ref_a, ref_b, ref_c) = - maxsim::test::generate(total_a_rows, k.value().get(), total_b_cols, rng); + maxsim::test::generate_f32(total_a_rows, k.value().get(), total_b_cols, rng); // Massage the input data in the form needed by the kernel. let a_bt = BlockTransposed::::from_matrix_view(ref_a.as_view()); diff --git a/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs b/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs new file mode 100644 index 0000000000..11e5041557 --- /dev/null +++ b/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs @@ -0,0 +1,921 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +//! The integer counterpart of [`super::packed_f32_x_unpacked_f32`], accumulating in `i32`. +//! +//! The blocking strategy is identical. The difference is that `a` interleaves `PACK` +//! consecutive contraction indices within each packed row so that a single lane of the +//! widening dot-product instructions consumes `PACK` products at a time. `PACK` is chosen +//! per architecture to match the instruction used. +//! +//! Each product is bounded by `128 * 128`, so the `i32` accumulator is exact and cannot +//! overflow for contraction dimensions up to `131_071`. + +use diskann_wide::arch::{Architecture, Scalar}; +use diskann_wide::{SIMDDotProduct, SIMDMinMax, SIMDReinterpret, SIMDVector}; + +use crate::matrix_kernels::{ + Cache, + blocks::{packed, unpacked}, + bounds, driver, + num::{DimK, Elements}, + ptr::{MutSlice, Slice}, + util::{self, Fold, Folder}, +}; + +use super::packed_f32_x_unpacked_f32::Params; + +diskann_wide::alias!(i8x16 = i8x16); +diskann_wide::alias!(i16x16 = i16x16); +diskann_wide::alias!(i32x8 = i32x8); +diskann_wide::alias!(u32x8 = u32x8); + +/// Widen a `PACK = 2` group into the little-endian `i16` lane pair consumed by the 16-bit +/// dot products. +#[inline(always)] +fn i16_pair([lo, hi]: [i8; 2]) -> u32 { + u32::from(i16::from(lo) as u16) | (u32::from(i16::from(hi) as u16) << 16) +} + +//--------// +// Driver // +//--------// + +/// A driver for prepacked by unpacked integer "maxsim" computations. +/// +/// See [`super::packed_f32_x_unpacked_f32::Driver`] for the blocking strategy and for the +/// rationale behind invariant (2). +/// +/// # Class Invariants +/// +/// 1. `a.k()` and `b.k()` must be equal to `k`. +/// 2. `c.len().div_ceil(MR)` must be equal to `a.blocks()`. +pub(crate) struct Driver<'a, A, const MR: usize, const NR: usize, const PACK: usize> { + arch: A, + a: packed::View<'a, i8, MR, PACK>, + b: unpacked::View<'a, i8>, + c: &'a mut [i32], + k: DimK, + params: Params, +} + +impl<'a, A, const MR: usize, const NR: usize, const PACK: usize> Driver<'a, A, MR, NR, PACK> { + /// Prepare for a maxsim on `a` and `b` with the results stored directly into `c`. + /// + /// `c` does not require any specific initial value. + /// + /// # Safety + /// + /// 1. `a.k()` and `b.k()` must be equal to `k`. + /// 2. `c.len().div_ceil(MR)` must be equal to `a.blocks()`. + pub(crate) unsafe fn new( + arch: A, + a: packed::View<'a, i8, MR, PACK>, + b: unpacked::View<'a, i8>, + c: &'a mut [i32], + k: DimK, + cache: Cache, + ) -> Self { + bounds::check_eq!(a.k(), k, "contraction dimensions do not agree"); + bounds::check_eq!(b.k(), k, "contraction dimensions do not agree"); + bounds::check_eq!( + bounds::Bound::new(a.blocks().get()), + c.len().div_ceil(MR), + "output length must occupy exactly the packed A blocks", + ); + + // SAFETY: Inherited from caller. + unsafe { + Self::new_inner( + arch, + a, + b, + c, + k, + Params::new(cache, a.block_stride(k).bytes(), b.stride(k).bytes(), NR), + ) + } + } + + /// # Safety + /// + /// 1. `a.k()` and `b.k()` must be equal to `k`. + /// 2. `c.len().div_ceil(MR)` must be equal to `a.blocks()`. + unsafe fn new_inner( + arch: A, + a: packed::View<'a, i8, MR, PACK>, + b: unpacked::View<'a, i8>, + c: &'a mut [i32], + k: DimK, + params: Params, + ) -> Self { + bounds::check_eq!(a.k(), k, "contraction dimensions do not agree"); + bounds::check_eq!(b.k(), k, "contraction dimensions do not agree"); + bounds::check_eq!( + bounds::Bound::new(a.blocks().get()), + c.len().div_ceil(MR), + "output length must occupy exactly the packed A blocks", + ); + + Self { + arch, + a, + b, + c, + k, + params, + } + } +} + +impl driver::Drive + for Driver<'_, A, MR, NR, PACK> +where + A: util::LoadStore + Architecture, + for<'a> PanelKernel<'a, A, MR, NR, PACK>: driver::PanelKernel, +{ + fn drive(&mut self) { + self.arch.run( + #[inline] + || { + // Pre-fill `c`. + self.c.fill(i32::MIN); + + // We allow `c` to be slightly under-filled. + // + // These variables track if under-fill is happening. + let remainder = self.c.len() % MR; + let last_a_block = self.a.blocks().get() - 1; + + let mut c = MutSlice::new(self.c); + + let on_a_panels = |a_panels: packed::View<'_, i8, MR, PACK>, a_block_base| { + let on_b_panels = |b_panels: unpacked::View<'_, i8>, _| { + let panel_kernel = + |a_panel: packed::Panel<'_, i8, MR, PACK>, a_block_offset| { + // If we are in the very last block and we need to sub-fill, do + // that. Otherwise, reference the output in place. + let a_block = a_block_base + a_block_offset; + let handling_tail = a_block == last_a_block && remainder != 0; + + let bound = bounds::Bound::from_fn(|| { + if handling_tail { remainder } else { MR } + }); + + // SAFETY: By class invariant, + // + // `MR * (self.a.blocks() - 1) < c.len() <= MR * self.a.blocks()`. + // + // From the visitor, `a_block <= self.a.blocks()`. + let mut region = unsafe { c.subslice(MR * a_block, bound) }; + let c = if handling_tail { + util::LoadStore::::load( + self.arch, + // SAFETY: `region` as length exactly `remainder`. + unsafe { region.as_std_slice(remainder) }, + ) + } else { + // SAFETY: `region` has length exactly `MR`. + unsafe { *region.as_array::() } + }; + + // run the kernel + // + // SAFETY: By class invariant, `a_panel.k()` and `b_panels.k()` + // are both equal to `self.k`. + let mut kernel = unsafe { + PanelKernel::new(self.arch, a_panel, b_panels, c, self.k) + }; + + driver::PanelKernel::panel_kernel(&mut kernel); + + let c_final = kernel.take(); + + // Put back `C`. + if handling_tail { + util::LoadStore::::store( + self.arch, + c_final, + // SAFETY: `region` has length exactly `remainder`. + unsafe { region.as_std_mut_slice(remainder) }, + ); + } else { + // SAFETY: `region` has length exactly `MR`. + unsafe { *region.as_array::() = c_final }; + } + }; + + // SAFETY: By class invariant, `a_panels.k() == self.k`. + unsafe { + a_panels.visit_panels(self.k, panel_kernel); + } + }; + + // SAFETY: By class invariant, `self.b.k() == self.k`. + unsafe { + self.b + .visit_sub_views(self.params.b_cols_in_l1, self.k, on_b_panels); + } + }; + + // SAFETY: By class invariant, `self.a.k() == self.k`. + unsafe { + self.a + .visit_sub_views(self.params.a_panels_in_l2, self.k, on_a_panels) + }; + }, + ); + } +} + +//-------------// +// PanelKernel // +//-------------// + +#[derive(Debug)] +pub(super) struct PanelKernel<'a, A, const MR: usize, const NR: usize, const PACK: usize> { + arch: A, + a: packed::Panel<'a, i8, MR, PACK>, + b: unpacked::View<'a, i8>, + c: [i32; MR], + k: DimK, +} + +impl<'a, A, const MR: usize, const NR: usize, const PACK: usize> PanelKernel<'a, A, MR, NR, PACK> { + /// Construct a new kernel. + /// + /// # Safety + /// + /// Bounds `a.k()` and `b.k()` must both be equal to `k`. + pub(super) unsafe fn new( + arch: A, + a: packed::Panel<'a, i8, MR, PACK>, + b: unpacked::View<'a, i8>, + c: [i32; MR], + k: DimK, + ) -> Self { + bounds::check_eq!(a.k(), k); + bounds::check_eq!(b.k(), k); + + Self { arch, a, b, c, k } + } + + pub(super) fn take(self) -> [i32; MR] { + self.c + } +} + +/// A custom visitor for the [`MicroKernel`]. +/// +/// This is needed to ensure the visitor body is inlined to inherit target features. +#[derive(Debug)] +struct Visitor<'a, A, const MR: usize, const NR: usize, const PACK: usize> { + arch: A, + a: packed::Panel<'a, i8, MR, PACK>, + c: &'a mut [i32; MR], + k: DimK, +} + +impl unpacked::PanelVisitor + for Visitor<'_, A, MR, NR, PACK> +where + A: Copy, + for<'a> MicroKernel<'a, A, MR, NR, PACK>: driver::MicroKernel, +{ + #[inline(always)] + fn visit(&mut self, b: unpacked::Panel<'_, i8, NR>, _: usize) { + // SAFETY: This is only used on contexts where `self.a.k()`, `b.k()`, and `self.k` + // are all equal. + let mut micro = unsafe { MicroKernel::new(self.arch, self.a, b, self.c, self.k) }; + driver::MicroKernel::micro_kernel(&mut micro); + } +} + +macro_rules! panel_kernel { + ($arch:ty, $mr:literal, $nr:literal, $pack:literal, [ $($ns:literal),+ $(,)? ]) => { + impl driver::PanelKernel for PanelKernel<'_, $arch, $mr, $nr, $pack> { + #[inline(always)] + fn panel_kernel(&mut self) { + // NOTE: A `Visitor` is used here instead of a closure because a `Visitor` + // is more reliably inlined, which means that target-features are inherited + // more reliably. + let on_b_panels = Visitor { + arch: self.arch, + a: self.a, + c: &mut self.c, + k: self.k, + }; + + // SAFETY: By class invariant, `self.k` is equal to `self.b.k()`. + let b_tail = unsafe { self.b.visit_panels::<$nr>(self.k, on_b_panels) }; + + if let Some(b_tail) = b_tail { + // Repetition Pattern. + $( + const { assert!($ns < $nr) }; + if let Some(b_panel) = b_tail.try_as_panel::<$ns>() { + // SAFETY: By class invariant, `self.a.k()` and `self.b.k()` + // are equal to `self.k`. + let mut micro = unsafe { + MicroKernel::new( + self.arch, + self.a, + b_panel, + &mut self.c, + self.k, + ) + }; + + driver::MicroKernel::micro_kernel(&mut micro); + } + )+ + } + } + } + } +} + +panel_kernel!(Scalar, 8, 2, 2, [1]); + +//--------------// +// Micro Kernel // +//--------------// + +/// # Class Invariants +/// +/// `a.k()` and `b.k()` are equal to `k`. +struct MicroKernel<'a, A, const MR: usize, const NR: usize, const PACK: usize> { + arch: A, + a: packed::Panel<'a, i8, MR, PACK>, + b: unpacked::Panel<'a, i8, NR>, + c: &'a mut [i32; MR], + k: DimK, +} + +impl<'a, A, const MR: usize, const NR: usize, const PACK: usize> MicroKernel<'a, A, MR, NR, PACK> { + /// # Safety + /// + /// Bounds `a.k()` and `b.k()` must be equal to `k`. + unsafe fn new( + arch: A, + a: packed::Panel<'a, i8, MR, PACK>, + b: unpacked::Panel<'a, i8, NR>, + c: &'a mut [i32; MR], + k: DimK, + ) -> Self { + bounds::check_eq!(a.k(), k); + bounds::check_eq!(b.k(), k); + + Self { arch, a, b, c, k } + } +} + +/// Gather the `PACK` contraction indices starting at `ptr` for a single column of `b`, +/// zero filling the last group when `k` is not a multiple of `PACK`. +/// +/// # Safety +/// +/// `valid` must not exceed `PACK` and the first `valid` elements of `ptr` must be readable. +#[inline(always)] +unsafe fn group(ptr: Slice<'_, i8>, valid: usize) -> [i8; PACK] { + core::array::from_fn(|p| { + if p < valid { + // SAFETY: Since `p < valid`, the pointer offset is valid and readable. + unsafe { *ptr.add(Elements::new(p)).as_unit().as_ref() } + } else { + 0 + } + }) +} + +/// # Safety +/// +/// Bounds `a.k()` and `b.k()` must be equal to `k`. +#[inline(always)] +unsafe fn micro_kernel( + wide: W, + a: packed::Panel<'_, i8, MR, PACK>, + b: unpacked::Panel<'_, i8, NR>, + c: &mut [i32; MR], + k: DimK, +) where + W: ExtraWide, + Folder: Fold, +{ + // Check that everyone agrees. + bounds::check_eq!(a.k(), k); + bounds::check_eq!(b.k(), k); + + let ap = a.as_ptr(); + let bp = b.as_ptr(); + + let mut acc = [wide.default(); NR]; + + let astride = a.row_stride(); + let bstride = b.stride(k); + + let rows = a.rows(k); + let k = k.value().get(); + + for row in 0..rows { + // SAFETY: By preconditions, `ap.len() == astride * rows`. Since `row < rows`: + // + // * The pointer offset is valid. + // * The subsequent truncation is valid. + // * The slice passed to `wide.load` has a length equal to `astride`. + let ai = unsafe { wide.load(ap.add(astride * row).truncate(astride)) }; + + // The trailing row of `a` is zero padded, so zero filling `b` past `k` keeps every + // padded product at zero. + let i = row * PACK; + let valid = PACK.min(k - i); + + for (j, acc) in acc.iter_mut().enumerate() { + // SAFETY: By preconditions, `bp.len() == bstride * NR`. Since `i < k`, `j < NR` + // and `i + valid <= k`: + // + // * The pointer offset is valid and its first `valid` elements are readable. + let bj = + wide.splat(unsafe { group::(bp.add(bstride * j + Elements::new(i)), valid) }); + + *acc = W::dot(ai, bj, *acc); + } + } + + wide.max_into(Folder::fold(acc, W::max), c); +} + +macro_rules! micro_kernel { + ($arch:ty, $mr:literal, $nr:literal, $pack:literal) => { + impl driver::MicroKernel for MicroKernel<'_, $arch, $mr, $nr, $pack> { + #[inline(always)] + fn micro_kernel(&mut self) { + // SAFETY: By class invariant, `self.a.k()` and `self.b.k()` equal `self.k`. + unsafe { micro_kernel(self.arch, self.a, self.b, self.c, self.k) } + } + } + }; + ($arch:ty, $mr:literal, $pack:literal, { $($nr:literal),+ $(,)? }) => { + $(micro_kernel!($arch, $mr, $nr, $pack);)+ + } +} + +micro_kernel!(Scalar, 8, 2, { 2, 1 }); + +trait ExtraWide: Copy { + type Wide: Copy; + type Splat: Copy; + type Acc: Copy; + + /// # Safety + /// + /// `slice.len()` must be exactly `ELEMENTS * PACK`. + unsafe fn load(self, slice: Slice<'_, i8>) -> Self::Wide; + + fn default(self) -> Self::Acc; + fn splat(self, group: [i8; PACK]) -> Self::Splat; + fn dot(a: Self::Wide, b: Self::Splat, acc: Self::Acc) -> Self::Acc; + fn max(lhs: Self::Acc, rhs: Self::Acc) -> Self::Acc; + fn max_into(self, max: Self::Acc, into: &mut [i32; ELEMENTS]); +} + +impl ExtraWide<8, 2> for Scalar { + type Wide = i16x16; + type Splat = i16x16; + type Acc = i32x8; + + #[inline(always)] + fn default(self) -> Self::Acc { + SIMDVector::default(self) + } + + #[inline(always)] + unsafe fn load(self, slice: Slice<'_, i8>) -> Self::Wide { + bounds::check_eq!(slice.len(), 16); + + // SAFETY: Since `slice.len()` must be 16, the 16-wide SIMD load is valid. + let bytes: i8x16 = unsafe { SIMDVector::load_simd(self, slice.as_ptr()) }; + + Self::Wide::from(bytes) + } + + #[inline(always)] + fn splat(self, group: [i8; 2]) -> Self::Splat { + u32x8::::splat(self, i16_pair(group)).reinterpret_simd() + } + + #[inline(always)] + fn dot(a: Self::Wide, b: Self::Splat, acc: Self::Acc) -> Self::Acc { + acc.dot_simd(a, b) + } + + #[inline(always)] + fn max(lhs: Self::Acc, rhs: Self::Acc) -> Self::Acc { + lhs.max_simd(rhs) + } + + #[inline(always)] + fn max_into(self, lhs: Self::Acc, into: &mut [i32; 8]) { + // SAFETY: Since `into.len()` is 8, the 8-wide SIMD load is valid. + let previous: Self::Acc = unsafe { SIMDVector::load_simd(self, into.as_ptr()) }; + + // SAFETY: Since `into.len()` is 8, the 8-wide SIMD store is valid. + unsafe { Self::max(lhs, previous).store_simd(into.as_mut_ptr()) }; + } +} + +#[cfg(target_arch = "x86_64")] +mod x86_64 { + use super::*; + + use diskann_wide::arch::x86_64::V3; + + panel_kernel!(V3, 16, 6, 2, [1, 2, 3, 4, 5]); + + micro_kernel!(V3, 16, 2, { 6, 5, 4, 3, 2, 1 }); + + //-----------// + // ExtraWide // + //-----------// + + impl ExtraWide<16, 2> for V3 { + type Wide = [i16x16; 2]; + type Splat = i16x16; + type Acc = [i32x8; 2]; + + #[inline(always)] + fn default(self) -> Self::Acc { + [SIMDVector::default(self), SIMDVector::default(self)] + } + + #[inline(always)] + unsafe fn load(self, slice: Slice<'_, i8>) -> Self::Wide { + bounds::check_eq!(slice.len(), 32); + + // SAFETY: Since `slice.len()` must be 32, the pointer offset and 16-wide SIMD loads + // are valid. + let bytes: [i8x16; 2] = unsafe { + [ + SIMDVector::load_simd(self, slice.as_ptr()), + SIMDVector::load_simd(self, slice.add(Elements::new(16)).as_ptr()), + ] + }; + + bytes.map(Self::Splat::from) + } + + #[inline(always)] + fn splat(self, group: [i8; 2]) -> Self::Splat { + u32x8::::splat(self, i16_pair(group)).reinterpret_simd() + } + + #[inline(always)] + fn dot(a: Self::Wide, b: Self::Splat, acc: Self::Acc) -> Self::Acc { + core::array::from_fn(|i| acc[i].dot_simd(a[i], b)) + } + + #[inline(always)] + fn max(lhs: Self::Acc, rhs: Self::Acc) -> Self::Acc { + core::array::from_fn(|i| lhs[i].max_simd(rhs[i])) + } + + #[inline(always)] + fn max_into(self, lhs: Self::Acc, into: &mut [i32; 16]) { + // SAFETY: Since `into.len()` is 16, the pointer offset and 8-wide SIMD loads are + // valid. + let previous: Self::Acc = unsafe { + [ + SIMDVector::load_simd(self, into.as_ptr()), + SIMDVector::load_simd(self, into.as_ptr().add(8)), + ] + }; + + let max = Self::max(lhs, previous); + + // SAFETY: Since `into.len()` is 16, the pointer offset and 8-wide SIMD stores are + // valid. + unsafe { + max[0].store_simd(into.as_mut_ptr()); + max[1].store_simd(into.as_mut_ptr().add(8)); + } + } + } +} + +/////////// +// Tests // +/////////// + +#[cfg(test)] +mod tests { + use super::*; + + use std::num::NonZeroUsize; + + use rand::{SeedableRng, rngs::StdRng}; + + #[cfg(target_arch = "x86_64")] + use diskann_wide::arch::x86_64::V3; + + use crate::{matrix_kernels::maxsim, multi_vector::BlockTransposed}; + + ///////////////// + // MicroKernel // + ///////////////// + + fn test_micro_kernel( + arch: A, + k: DimK, + rng: &mut impl rand::Rng, + ctx: std::fmt::Arguments<'_>, + ) where + for<'a> MicroKernel<'a, A, MR, NR, PACK>: driver::MicroKernel, + { + let (ref_a, ref_b, ref_c) = maxsim::test::generate_i8(MR, k.value().get(), NR, rng); + + // From the reference problem, `ref_a` needs to be packed and `ref_b` transposed to + // get them into the desired format. + let a_bt = BlockTransposed::::from_matrix_view(ref_a.as_view()); + let ref_b = ref_b.transpose(); + + let mut c = [i32::MIN; MR]; + + // Run the test kernel. + // + // SAFETY: Test builds will verify the bounds we passed. + let mut kernel = unsafe { + MicroKernel::new( + arch, + packed::Panel::new(Slice::new(a_bt.as_slice()), k), + unpacked::Panel::new(Slice::new(ref_b.as_slice()), k), + &mut c, + k, + ) + }; + + driver::MicroKernel::micro_kernel(&mut kernel); + assert_eq!(&*ref_c, kernel.c, "{ctx}"); + + // Try again - but this time use a value that is much bigger than the what should + // be generated by the test problem. + // + // This checks that we don't just overwrite existing contents. + let new_c = kernel.c.map(|i| i + 1); + *kernel.c = new_c; + + driver::MicroKernel::micro_kernel(&mut kernel); + assert_eq!(new_c, *kernel.c, "{ctx}"); + } + + macro_rules! test_micro_kernel { + ( + $fn:ident, + $arch:expr, + $seed:literal, + $PACK:literal, + $( + $MR:literal => { $($NR:literal),+ $(,)? } + ),+ $(,)? + ) => { + #[test] + fn $fn() { + if let Some(arch) = $arch { + let mut rng = StdRng::seed_from_u64($seed); + + for k in [1, 2, 5, 8] { + let k = DimK::new(NonZeroUsize::new(k).unwrap()); + + $( + $( + test_micro_kernel::<_, $MR, $NR, $PACK>( + arch, + k, + &mut rng, + format_args!("k = {:?}", k), + ); + )+ + )+ + } + } + } + } + } + + test_micro_kernel!( + test_micro_kernel_scalar, + Some(Scalar::new()), + 0x4b1d09c2a77e5310, + 2, + 8 => { 2, 1 }, + ); + + #[cfg(target_arch = "x86_64")] + test_micro_kernel!( + test_micro_kernel_v3, + V3::new_checked(), + 0xe0a5c31f8b62d94a, + 2, + 16 => { 6, 5, 4, 3, 2, 1 }, + ); + + ///////////////// + // PanelKernel // + ///////////////// + + // The panel kernel operates on a single A-panel with multiple B-panels. + // + // This test sweeps over a number of rows for the B-panels to exercise all possible + // corner cases. + fn test_panel_kernel( + arch: A, + k: DimK, + rng: &mut impl rand::Rng, + ctx: std::fmt::Arguments<'_>, + ) where + A: Copy, + for<'a> PanelKernel<'a, A, MR, NR, PACK>: driver::PanelKernel, + { + for blocks in 0..4 { + for remainder in 0..NR { + let cols = NR * blocks + remainder; + if cols == 0 { + continue; + } + + let (ref_a, ref_b, ref_c) = + maxsim::test::generate_i8(MR, k.value().get(), cols, rng); + + let a_bt = BlockTransposed::::from_matrix_view(ref_a.as_view()); + let ref_b = ref_b.transpose(); + + let extent = NonZeroUsize::new(cols).unwrap(); + + let c = [i32::MIN; MR]; + + // SAFETY: Test builds will verify the bounds we passed. + let mut kernel = unsafe { + PanelKernel::new( + arch, + packed::Panel::new(Slice::new(a_bt.as_slice()), k), + unpacked::View::new(Slice::new(ref_b.as_slice()), extent, k), + c, + k, + ) + }; + + driver::PanelKernel::panel_kernel(&mut kernel); + assert_eq!(&*ref_c, kernel.c, "{ctx}"); + + // Try again - but this time use a value that is much bigger than the what + // should be generated by the test problem. + // + // This checks that we don't just overwrite existing contents. + let new_c = kernel.c.map(|i| i + 1); + kernel.c = new_c; + + driver::PanelKernel::panel_kernel(&mut kernel); + assert_eq!(new_c, kernel.c, "{ctx}"); + } + } + } + + macro_rules! test_panel_kernel { + ( + $fn:ident, + $arch:expr, + $seed:literal, + $( + ( + $MR:literal, $NR:literal, $PACK:literal + ) + ),+ $(,)? + ) => { + #[test] + fn $fn() { + if let Some(arch) = $arch { + let mut rng = StdRng::seed_from_u64($seed); + + for k in [1, 2, 5, 8] { + let k = DimK::new(NonZeroUsize::new(k).unwrap()); + + $( + test_panel_kernel::<_, $MR, $NR, $PACK>( + arch, + k, + &mut rng, + format_args!("k = {:?}", k), + ); + )+ + } + } + } + } + } + + test_panel_kernel!( + test_panel_kernel_scalar, + Some(Scalar::new()), + 0x9f3e7ab4c05d1268, + (8, 2, 2), + ); + + #[cfg(target_arch = "x86_64")] + test_panel_kernel!( + test_panel_kernel_v3, + V3::new_checked(), + 0x9f3e7ab4c05d1268, + (16, 6, 2), + ); + + //////////// + // Driver // + //////////// + + fn test_driver( + arch: A, + rng: &mut impl rand::Rng, + ) where + A: Copy, + for<'a> Driver<'a, A, MR, NR, PACK>: driver::Drive, + { + let cases = maxsim::test::packed_x_unpacked_test_dims(MR, NR); + for case in cases { + let maxsim::test::TestDims { + a_panels_per_tile, + total_a_rows, + b_cols_per_tile, + total_b_cols, + k, + } = case.clone(); + + let k = DimK::new(NonZeroUsize::new(k).unwrap()); + + let (ref_a, ref_b, ref_c) = + maxsim::test::generate_i8(total_a_rows, k.value().get(), total_b_cols, rng); + + // Massage the input data in the form needed by the kernel. + let a_bt = BlockTransposed::::from_matrix_view(ref_a.as_view()); + let b = ref_b.transpose(); + + let mut c = vec![i32::MAX; a_bt.nrows()]; + + // SAFETY: Test builds will verify the bounds we passed. + let mut driver = unsafe { + Driver::new_inner( + arch, + packed::View::from_block_transposed(a_bt.as_view()).unwrap(), + unpacked::View::from_matrix_view(b.as_view()).unwrap(), + &mut c, + k, + Params { + a_panels_in_l2: NonZeroUsize::new(a_panels_per_tile).unwrap(), + b_cols_in_l1: NonZeroUsize::new(b_cols_per_tile).unwrap(), + }, + ) + }; + + driver::Drive::drive(&mut driver); + + assert_eq!(ref_c, c, "setup: {:?}", case) + } + } + + macro_rules! test_driver { + ( + $fn:ident, + $arch:expr, + $seed:literal, + $( + ( + $MR:literal, $NR:literal, $PACK:literal + ) + ),+ $(,)? + ) => { + #[test] + fn $fn() { + if let Some(arch) = $arch { + let mut rng = StdRng::seed_from_u64($seed); + + $(test_driver::<_, $MR, $NR, $PACK>(arch, &mut rng);)+ + } + } + } + } + + test_driver!( + test_driver_scalar, + Some(Scalar::new()), + 0x63c8ed19f4720ab5, + (8, 2, 2), + ); + + #[cfg(target_arch = "x86_64")] + test_driver!( + test_driver_v3, + V3::new_checked(), + 0x63c8ed19f4720ab5, + (16, 6, 2), + ); +} diff --git a/diskann-quantization/src/matrix_kernels/maxsim/test.rs b/diskann-quantization/src/matrix_kernels/maxsim/test.rs index 8a05fca4c9..605f8fa6e4 100644 --- a/diskann-quantization/src/matrix_kernels/maxsim/test.rs +++ b/diskann-quantization/src/matrix_kernels/maxsim/test.rs @@ -8,7 +8,7 @@ use diskann_utils::views::Matrix; use crate::matrix_kernels::test_util::TestDistr; /// Generate a test MaxSim problem `[M x K] . [K x N]` where both matrices are row-major. -pub(super) fn generate( +pub(super) fn generate_f32( m: usize, k: usize, n: usize, @@ -36,6 +36,36 @@ pub(super) fn generate( (ref_a, ref_b, ref_c) } +/// Generate a test integer MaxSim problem `[M x K] . [K x N]` where both matrices are +/// row-major. +pub(super) fn generate_i8( + m: usize, + k: usize, + n: usize, + rng: &mut impl rand::Rng, +) -> (Matrix, Matrix, Vec) { + let ref_a = TestDistr::matrix::(m, k, rng); + let ref_b = TestDistr::matrix::(k, n, rng); + + let ref_c: Vec = ref_a + .row_iter() + .map(|a_row| { + let mut max_ip = i32::MIN; + for b_col in 0..n { + let mut ip = 0; + for (k, a) in a_row.iter().enumerate() { + ip += i32::from(*a) * i32::from(ref_b[(k, b_col)]); + } + max_ip = max_ip.max(ip); + } + + max_ip + }) + .collect(); + + (ref_a, ref_b, ref_c) +} + #[derive(Debug, Clone)] pub(super) struct TestDims { pub(super) a_panels_per_tile: usize, diff --git a/diskann-quantization/src/matrix_kernels/test_util.rs b/diskann-quantization/src/matrix_kernels/test_util.rs index ac569df139..dea228a560 100644 --- a/diskann-quantization/src/matrix_kernels/test_util.rs +++ b/diskann-quantization/src/matrix_kernels/test_util.rs @@ -64,6 +64,12 @@ impl Distribution for TestDistr { } } +impl Distribution for TestDistr { + fn sample(&self, rng: &mut R) -> i8 { + rng.random_range(i8::MIN..=i8::MAX) + } +} + impl Distribution for TestDistr { fn sample(&self, rng: &mut R) -> f16 { f16::from_f32(>::sample(self, rng)) diff --git a/diskann-quantization/src/matrix_kernels/util.rs b/diskann-quantization/src/matrix_kernels/util.rs index b63ae9c6cd..772877926c 100644 --- a/diskann-quantization/src/matrix_kernels/util.rs +++ b/diskann-quantization/src/matrix_kernels/util.rs @@ -102,6 +102,7 @@ mod x86_64 { impl_loadstore!(f32, 8, f32x8, V3); impl_loadstore!(f32, 16, f32x16, V3); + impl_loadstore!(i32, 16, i32x16, V3); impl_loadstore!(f32, 8, f32x8, V4); impl_loadstore!(f32, 16, f32x16, V4); @@ -273,6 +274,12 @@ mod test { } } + impl FromUsize for i32 { + fn from_usize(v: usize) -> Self { + v as i32 + } + } + fn double(x: usize) -> T where T: FromUsize, @@ -322,6 +329,7 @@ mod test { test_load_store_scalar, Some(Scalar), f32 => { 4, 8, 16 }, + i32 => { 8 }, ); #[cfg(target_arch = "x86_64")] @@ -329,6 +337,7 @@ mod test { test_load_store_v3, V3::new_checked(), f32 => { 8, 16 }, + i32 => { 16 }, ); #[cfg(target_arch = "x86_64")] diff --git a/diskann-quantization/src/multi_vector/distance/factory.rs b/diskann-quantization/src/multi_vector/distance/factory.rs index fe87959ccd..02753cdbb8 100644 --- a/diskann-quantization/src/multi_vector/distance/factory.rs +++ b/diskann-quantization/src/multi_vector/distance/factory.rs @@ -8,8 +8,8 @@ use std::num::NonZeroUsize; +use diskann_vector::PureDistanceFunction; use diskann_vector::distance::InnerProduct; -use diskann_vector::{DistanceFunctionMut, PureDistanceFunction}; use diskann_wide::Architecture; use diskann_wide::arch::Scalar; #[cfg(target_arch = "aarch64")] @@ -17,9 +17,10 @@ use diskann_wide::arch::aarch64::Neon; #[cfg(target_arch = "x86_64")] use diskann_wide::arch::x86_64::{V3, V4}; +use super::fallback::FallbackKernel; use super::isa::{MaxSimIsa, NotSupported}; use super::kernel::{Erase, MaxSimKernel}; -use super::max_sim::{MaxSim, MaxSimError}; +use super::max_sim::MaxSimError; use crate::matrix_kernels as mk; use crate::multi_vector::distance::QueryMatRef; use crate::multi_vector::{BlockTransposed, Mat, MatRef, Standard}; @@ -175,15 +176,87 @@ where } } +impl MaxSimKernel + for Prepared, NR> +where + A: Architecture, + for<'a> mk::maxsim::packed_i8_x_unpacked_i8::Driver<'a, A, GROUP, NR, PACK>: mk::Drive, +{ + fn nrows(&self) -> usize { + self.prepared.nrows() + } + + fn compute_max_sim( + &self, + doc: MatRef<'_, Standard>, + scores: &mut [i32], + ) -> Result<(), MaxSimError> { + if scores.len() != self.nrows() { + return Err(MaxSimError::InvalidBufferLength(scores.len(), self.nrows())); + } + + if doc.vector_dim() != self.prepared.ncols() { + return Err(MaxSimError::UnequalDim( + doc.vector_dim(), + self.prepared.ncols(), + )); + } + + let Some(k) = NonZeroUsize::new(self.prepared.ncols()).map(mk::DimK::new) else { + scores.fill(if doc.num_vectors() == 0 { i32::MAX } else { 0 }); + return Ok(()); + }; + + let Some(a) = mk::blocks::packed::View::from_block_transposed(self.prepared.as_view()) + else { + return Ok(()); + }; + + let Some(b) = mk::blocks::unpacked::View::from_matrix_view(doc.as_matrix_view()) else { + scores.fill(i32::MAX); + return Ok(()); + }; + + // SAFETY: The dimension check establishes that `a.k() == b.k() == k`. + // The length check establishes that `scores` occupies exactly the + // packed blocks in `a`. + let mut driver = unsafe { + mk::maxsim::packed_i8_x_unpacked_i8::Driver::new( + self.arch, + a, + b, + scores, + k, + mk::Cache::detect(), + ) + }; + + mk::Drive::drive(&mut driver); + + scores.iter_mut().for_each(|s| *s = -*s); + + Ok(()) + } +} + // ───────────────────────────────────────────────────────────────────────── -// ReferenceKernel — non-SIMD fallback that wraps MaxSim::evaluate. +// ReferenceKernel — double loop over the single-vector inner product. // ───────────────────────────────────────────────────────────────────────── -struct ReferenceKernel { +/// Reference MaxSim implementation, selected by [`MaxSimElement::build`]. +/// +/// May assume `scores` and `doc` were validated by +/// [`MaxSimKernel::compute_max_sim`]. +type ReferenceFn = + fn(QueryMatRef<'_, Standard>, MatRef<'_, Standard>, &mut [::Score]); + +/// Unoptimized kernel backing [`MaxSimIsa::Reference`]. +struct ReferenceKernel { query: Mat>, + run: ReferenceFn, } -impl std::fmt::Debug for ReferenceKernel { +impl std::fmt::Debug for ReferenceKernel { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { f.debug_struct("ReferenceKernel") .field("nrows", &self.query.num_vectors()) @@ -191,19 +264,16 @@ impl std::fmt::Debug for ReferenceKernel { } } -impl ReferenceKernel { - fn new(query: MatRef<'_, Standard>) -> Self { +impl ReferenceKernel { + fn new(query: MatRef<'_, Standard>, run: ReferenceFn) -> Self { Self { query: query.to_owned(), + run, } } } -impl MaxSimKernel for ReferenceKernel -where - T: Copy + Send + Sync + std::fmt::Debug + 'static, - InnerProduct: for<'a, 'b> PureDistanceFunction<&'a [T], &'b [T], f32>, -{ +impl MaxSimKernel for ReferenceKernel { fn nrows(&self) -> usize { self.query.num_vectors() } @@ -211,7 +281,7 @@ where fn compute_max_sim( &self, doc: MatRef<'_, Standard>, - scores: &mut [f32], + scores: &mut [T::Score], ) -> Result<(), MaxSimError> { if scores.len() != self.nrows() { return Err(MaxSimError::InvalidBufferLength(scores.len(), self.nrows())); @@ -222,16 +292,31 @@ where self.query.vector_dim(), )); } - if doc.num_vectors() == 0 { - scores.fill(f32::MAX); - return Ok(()); - } - let query: QueryMatRef<'_, Standard> = self.query.as_view().into(); - let mut max_sim = MaxSim::new(scores); - max_sim.evaluate(query, doc) + (self.run)(self.query.as_view().into(), doc, scores); + Ok(()) } } +/// [`ReferenceKernel`] implementation for element types scored in `f32`. +fn reference_scores( + query: QueryMatRef<'_, Standard>, + doc: MatRef<'_, Standard>, + scores: &mut [f32], +) where + InnerProduct: for<'a, 'b> PureDistanceFunction<&'a [T], &'b [T], f32>, +{ + FallbackKernel::max_sim_kernel(query, doc, |i, score| scores[i] = score); +} + +/// [`ReferenceKernel`] implementation for `i8`, which scores in exact `i32`. +fn reference_scores_i8( + query: QueryMatRef<'_, Standard>, + doc: MatRef<'_, Standard>, + scores: &mut [i32], +) { + FallbackKernel::max_sim_kernel_i8(query, doc, |i, score| scores[i] = score); +} + // ───────────────────────────────────────────────────────────────────────── // BuildAndErase — Target1 impls used by `dispatch1_no_features` (Auto). // ───────────────────────────────────────────────────────────────────────── @@ -377,6 +462,55 @@ impl> } } +// ───── i8 Target1 impls ───── + +impl> diskann_wide::arch::Target1>> + for BuildAndErase +{ + fn run(self, arch: Scalar, query: MatRef<'_, Standard>) -> E::Output { + let prepared = BlockTransposed::::from_matrix_view(query.as_matrix_view()); + self.0.erase(Prepared { + arch, + prepared, + _packing: Pack::<2>, + }) + } +} + +#[cfg(target_arch = "x86_64")] +impl> diskann_wide::arch::Target1>> + for BuildAndErase +{ + fn run(self, arch: V3, query: MatRef<'_, Standard>) -> E::Output { + let prepared = BlockTransposed::::from_matrix_view(query.as_matrix_view()); + self.0.erase(Prepared { + arch, + prepared, + _packing: Pack::<6>, + }) + } +} + +#[cfg(target_arch = "x86_64")] +impl> diskann_wide::arch::Target1>> + for BuildAndErase +{ + fn run(self, arch: V4, query: MatRef<'_, Standard>) -> E::Output { + // V4 retargets to V3 until the VNNI kernel lands. + diskann_wide::arch::Target1::::run(self, V3::from(arch), query) + } +} + +#[cfg(target_arch = "aarch64")] +impl> diskann_wide::arch::Target1>> + for BuildAndErase +{ + fn run(self, arch: Neon, query: MatRef<'_, Standard>) -> E::Output { + // Neon retargets to Scalar until the dotprod kernel lands. + diskann_wide::arch::Target1::::run(self, Scalar::from(arch), query) + } +} + // ───────────────────────────────────────────────────────────────────────── // MaxSimElement — sealed trait gating accepted element types. // ───────────────────────────────────────────────────────────────────────── @@ -391,6 +525,13 @@ mod sealed { /// (PQ, SQ, packed sub-byte) are intentionally excluded — they need /// codebook/scale state that [`MatRef<'_, Standard>`] can't carry. pub trait MaxSimElement: sealed::Sealed + Sized + Copy + Send + Sync + 'static { + /// Score produced per query row: `f32` for floating-point elements, `i32` + /// for `i8` where the inner product is exact. + type Score: Copy + Default + PartialEq + std::fmt::Debug; + + /// Score written for every query row when the document set is empty. + const NO_MATCH: Self::Score; + /// Build the concrete kernel for this element type and hand it to /// `erase.erase(...)`. /// @@ -407,8 +548,12 @@ pub trait MaxSimElement: sealed::Sealed + Sized + Copy + Send + Sync + 'static { impl sealed::Sealed for f32 {} impl sealed::Sealed for half::f16 {} +impl sealed::Sealed for i8 {} impl MaxSimElement for f32 { + type Score = f32; + const NO_MATCH: f32 = f32::MAX; + fn build>( isa: MaxSimIsa, query: MatRef<'_, Standard>, @@ -454,12 +599,17 @@ impl MaxSimElement for f32 { isa, reason: "aarch64 target only", }), - MaxSimIsa::Reference => Ok(erase.erase(ReferenceKernel::::new(query))), + MaxSimIsa::Reference => { + Ok(erase.erase(ReferenceKernel::new(query, reference_scores::))) + } } } } impl MaxSimElement for half::f16 { + type Score = f32; + const NO_MATCH: f32 = f32::MAX; + fn build>( isa: MaxSimIsa, query: MatRef<'_, Standard>, @@ -505,7 +655,65 @@ impl MaxSimElement for half::f16 { isa, reason: "aarch64 target only", }), - MaxSimIsa::Reference => Ok(erase.erase(ReferenceKernel::::new(query))), + MaxSimIsa::Reference => { + Ok(erase.erase(ReferenceKernel::new(query, reference_scores::))) + } + } + } +} + +impl MaxSimElement for i8 { + type Score = i32; + const NO_MATCH: i32 = i32::MAX; + + fn build>( + isa: MaxSimIsa, + query: MatRef<'_, Standard>, + erase: E, + ) -> Result { + match isa { + MaxSimIsa::Auto => Ok(diskann_wide::arch::dispatch1_no_features( + BuildAndErase(erase), + query, + )), + MaxSimIsa::Scalar => Ok(Scalar::new().run1(BuildAndErase(erase), query)), + #[cfg(target_arch = "x86_64")] + MaxSimIsa::X86_64_V3 => { + let arch = V3::new_checked().ok_or(NotSupported { + isa, + reason: "AVX2/FMA unavailable on this CPU", + })?; + Ok(arch.run1(BuildAndErase(erase), query)) + } + #[cfg(target_arch = "x86_64")] + MaxSimIsa::X86_64_V4 => { + let arch = V4::new_checked().ok_or(NotSupported { + isa, + reason: "AVX-512 unavailable on this CPU", + })?; + Ok(arch.run1(BuildAndErase(erase), query)) + } + #[cfg(not(target_arch = "x86_64"))] + MaxSimIsa::X86_64_V3 | MaxSimIsa::X86_64_V4 => Err(NotSupported { + isa, + reason: "x86_64 target only", + }), + #[cfg(target_arch = "aarch64")] + MaxSimIsa::Neon => { + let arch = Neon::new_checked().ok_or(NotSupported { + isa, + reason: "Neon unavailable on this CPU", + })?; + Ok(arch.run1(BuildAndErase(erase), query)) + } + #[cfg(not(target_arch = "aarch64"))] + MaxSimIsa::Neon => Err(NotSupported { + isa, + reason: "aarch64 target only", + }), + MaxSimIsa::Reference => { + Ok(erase.erase(ReferenceKernel::new(query, reference_scores_i8))) + } } } } @@ -534,6 +742,7 @@ pub fn build_max_sim>( mod tests { use super::*; use crate::multi_vector::{BoxErase, Chamfer, MaxSim, QueryMatRef}; + use diskann_vector::DistanceFunctionMut; /// Local helper trait — picks a sane test value of `T` from an `f32` /// so both `f32` and `half::f16` parameterizations share the same data @@ -554,6 +763,44 @@ mod tests { } } + impl FromF32 for i8 { + fn from_f32(v: f32) -> Self { + v as i8 + } + } + + /// Projects a kernel score onto the `f32` distance the fallback path + /// produces, so both parameterizations share the same assertions. + trait ScoreAsF32: MaxSimElement { + fn score_as_f32(score: Self::Score) -> f32; + } + + impl ScoreAsF32 for f32 { + fn score_as_f32(score: f32) -> f32 { + score + } + } + + impl ScoreAsF32 for half::f16 { + fn score_as_f32(score: f32) -> f32 { + score + } + } + + impl ScoreAsF32 for i8 { + fn score_as_f32(score: i32) -> f32 { + if score == Self::NO_MATCH { + f32::MAX + } else { + score as f32 + } + } + } + + fn scores_buffer(len: usize) -> Vec { + vec![T::Score::default(); len] + } + fn make_mat(data: &[T], nrows: usize, ncols: usize) -> MatRef<'_, Standard> { MatRef::new(Standard::new(nrows, ncols).unwrap(), data).unwrap() } @@ -583,7 +830,7 @@ mod tests { fn check_chamfer_matches(tol: f32, label: &str) where - T: MaxSimElement + FromF32, + T: ScoreAsF32 + FromF32, InnerProduct: for<'a, 'b> PureDistanceFunction<&'a [T], &'b [T], f32>, { for &(nq, nd, dim) in TEST_CASES { @@ -596,9 +843,9 @@ mod tests { let expected = Chamfer::evaluate(QueryMatRef::from(query), doc); let kernel = build_max_sim::(MaxSimIsa::Auto, query, BoxErase).unwrap(); - let mut scores = vec![0.0f32; nq]; + let mut scores = scores_buffer::(nq); kernel.compute_max_sim(doc, &mut scores).unwrap(); - let actual: f32 = scores.iter().sum(); + let actual: f32 = scores.iter().map(|&s| T::score_as_f32(s)).sum(); assert!( (actual - expected).abs() < tol, @@ -609,7 +856,7 @@ mod tests { fn check_max_sim_matches(tol: f32, label: &str) where - T: MaxSimElement + FromF32, + T: ScoreAsF32 + FromF32, InnerProduct: for<'a, 'b> PureDistanceFunction<&'a [T], &'b [T], f32>, { for &(nq, nd, dim) in TEST_CASES { @@ -623,20 +870,56 @@ mod tests { let _ = MaxSim::new(&mut expected_scores).evaluate(QueryMatRef::from(query), doc); let kernel = build_max_sim::(MaxSimIsa::Auto, query, BoxErase).unwrap(); - let mut actual_scores = vec![0.0f32; nq]; + let mut actual_scores = scores_buffer::(nq); kernel.compute_max_sim(doc, &mut actual_scores).unwrap(); for i in 0..nq { + let actual = T::score_as_f32(actual_scores[i]); assert!( - (actual_scores[i] - expected_scores[i]).abs() < tol, - "{label}MaxSim[{i}] mismatch for ({nq},{nd},{dim}): actual={}, expected={}", - actual_scores[i], + (actual - expected_scores[i]).abs() < tol, + "{label}MaxSim[{i}] mismatch for ({nq},{nd},{dim}): actual={actual}, expected={}", expected_scores[i], ); } } } + /// The `i8` reference path is an independent integer implementation + /// ([`FallbackKernel::max_sim_kernel_i8`]), so it needs its own guard; the + /// `f32`/`f16` reference paths share `max_sim_kernel` with the oracle and + /// would only be testing themselves. + /// + /// Every other ISA is reached via [`MaxSimIsa::Auto`] somewhere in the CI + /// matrix. Widen into a full sweep once V4 and Neon gain native `i8` + /// kernels instead of retargeting to V3 and Scalar. + #[test] + fn i8_reference_matches_oracle() { + for &(nq, nd, dim) in TEST_CASES { + let query_data = make_test_data::(nq * dim, dim, dim / 2); + let doc_data = make_test_data::(nd * dim, dim, dim); + + let query = make_mat(&query_data, nq, dim); + let doc = make_mat(&doc_data, nd, dim); + + let mut expected = vec![0.0f32; nq]; + let _ = MaxSim::new(&mut expected).evaluate(QueryMatRef::from(query), doc); + + let kernel = build_max_sim::(MaxSimIsa::Reference, query, BoxErase).unwrap(); + let mut scores = scores_buffer::(nq); + kernel.compute_max_sim(doc, &mut scores).unwrap(); + + for i in 0..nq { + let actual = ::score_as_f32(scores[i]); + assert!( + (actual - expected[i]).abs() < 1e-10, + "i8 reference MaxSim[{i}] mismatch for ({nq},{nd},{dim}): \ + actual={actual}, expected={}", + expected[i], + ); + } + } + } + #[test] fn dimensions_f32() { let data = vec![1.0f32; 5 * 8]; @@ -653,10 +936,17 @@ mod tests { assert_eq!(kernel.nrows(), 5); } + #[test] + fn dimensions_i8() { + let data = vec![1i8; 5 * 8]; + let query = make_mat(&data, 5, 8); + let kernel = build_max_sim::(MaxSimIsa::Auto, query, BoxErase).unwrap(); + assert_eq!(kernel.nrows(), 5); + } + fn check_size_mismatch(label: &str) where T: MaxSimElement + FromF32, - InnerProduct: for<'a, 'b> PureDistanceFunction<&'a [T], &'b [T], f32>, { let query_data = make_test_data::(3 * 4, 4, 0); let doc_data = make_test_data::(2 * 4, 4, 1); @@ -666,7 +956,7 @@ mod tests { for isa in [MaxSimIsa::Auto, MaxSimIsa::Reference] { let kernel = build_max_sim::(isa, query, BoxErase).unwrap(); - let mut too_short = vec![0.0f32; 2]; + let mut too_short = scores_buffer::(2); match kernel.compute_max_sim(doc, &mut too_short) { Err(MaxSimError::InvalidBufferLength(2, 3)) => {} other => { @@ -674,7 +964,7 @@ mod tests { } } - let mut too_long = vec![0.0f32; 4]; + let mut too_long = scores_buffer::(4); match kernel.compute_max_sim(doc, &mut too_long) { Err(MaxSimError::InvalidBufferLength(4, 3)) => {} other => { @@ -687,7 +977,6 @@ mod tests { fn check_zero_docs_fills_sentinel(label: &str) where T: MaxSimElement + FromF32, - InnerProduct: for<'a, 'b> PureDistanceFunction<&'a [T], &'b [T], f32>, { let query_data = make_test_data::(3 * 4, 4, 0); let doc_data: Vec = Vec::new(); @@ -696,13 +985,13 @@ mod tests { for isa in [MaxSimIsa::Auto, MaxSimIsa::Reference] { let kernel = build_max_sim::(isa, query, BoxErase).unwrap(); - let mut scores = vec![0.0f32; 3]; + let mut scores = scores_buffer::(3); kernel.compute_max_sim(doc, &mut scores).unwrap(); for (i, &s) in scores.iter().enumerate() { assert_eq!( s, - f32::MAX, - "{label}({isa:?}) zero-doc slot {i} should be f32::MAX sentinel", + T::NO_MATCH, + "{label}({isa:?}) zero-doc slot {i} should be the NO_MATCH sentinel", ); } } @@ -711,7 +1000,6 @@ mod tests { fn check_zero_query(label: &str) where T: MaxSimElement + FromF32, - InnerProduct: for<'a, 'b> PureDistanceFunction<&'a [T], &'b [T], f32>, { let query_data: Vec = Vec::new(); let doc_data = make_test_data::(2 * 4, 4, 0); @@ -725,7 +1013,7 @@ mod tests { 0, "{label}({isa:?}) empty query should yield nrows=0", ); - let mut scores: Vec = Vec::new(); + let mut scores = scores_buffer::(0); kernel .compute_max_sim(doc, &mut scores) .unwrap_or_else(|e| panic!("{label}({isa:?}) expected Ok, got {e:?}")); @@ -767,4 +1055,5 @@ mod tests { test_matches_fallback!(f32, f32, 1e-10, "f32 "); test_matches_fallback!(f16, half::f16, 1e-10, "f16 "); + test_matches_fallback!(i8, i8, 1e-10, "i8 "); } diff --git a/diskann-quantization/src/multi_vector/distance/fallback.rs b/diskann-quantization/src/multi_vector/distance/fallback.rs index 9bd134ea08..2fe96675fa 100644 --- a/diskann-quantization/src/multi_vector/distance/fallback.rs +++ b/diskann-quantization/src/multi_vector/distance/fallback.rs @@ -99,6 +99,37 @@ impl FallbackKernel { } } + /// Exact `i32` counterpart of [`FallbackKernel::max_sim_kernel`] for `i8`. + /// + /// The `i8` inner product fits an `i32` accumulator exactly for vector dimensions up + /// to `131_071`, so scores are exact rather than rounded through `f32`. If there are + /// no vectors in the `doc`, the score is `i32::MAX`. + /// + /// # Arguments + /// + /// * `query` - The query multi-vector (wrapped as [`QueryMatRef`]) + /// * `doc` - The document multi-vector + /// * `f` - Callback invoked with `(query_index, similarity)` for each query vector + #[inline] + pub(crate) fn max_sim_kernel_i8( + query: QueryMatRef<'_, Standard>, + doc: MatRef<'_, Standard>, + mut f: F, + ) where + F: FnMut(usize, i32), + { + for (i, q_vec) in query.rows().enumerate() { + // Negate to match the similarity convention of `InnerProduct::evaluate`. + let mut min_dist = i32::MAX; + + for d_vec in doc.rows() { + min_dist = min_dist.min(-dot_i32(q_vec, d_vec)); + } + + f(i, min_dist); + } + } + /// Core kernel for computing per-query-vector projected-eigen scores. /// /// For each `query` vector, sums the negated squared inner product @@ -144,6 +175,15 @@ impl FallbackKernel { } } +/// Exact `i8` inner product, matching the `i32` accumulator width of the SIMD +/// kernels. +fn dot_i32(a: &[i8], b: &[i8]) -> i32 { + a.iter() + .zip(b) + .map(|(&x, &y)| i32::from(x) * i32::from(y)) + .sum() +} + //////////// // MaxSim // //////////// diff --git a/diskann-quantization/src/multi_vector/distance/kernel.rs b/diskann-quantization/src/multi_vector/distance/kernel.rs index b9bd5e1be0..5420c13b46 100644 --- a/diskann-quantization/src/multi_vector/distance/kernel.rs +++ b/diskann-quantization/src/multi_vector/distance/kernel.rs @@ -5,15 +5,15 @@ //! Object-safe kernel boundary trait plus BYOTE visitor trait. -use crate::multi_vector::{MatRef, MaxSimError, Standard}; +use crate::multi_vector::{MatRef, MaxSimElement, MaxSimError, Standard}; /// Object-safe interface for computing per-query MaxSim scores. -pub trait MaxSimKernel: Send + Sync + std::fmt::Debug { +pub trait MaxSimKernel: Send + Sync + std::fmt::Debug { /// Number of query rows whose scores this kernel produces. fn nrows(&self) -> usize; /// Compute per-query MaxSim scores into `scores`. On zero docs, fills - /// every slot with `f32::MAX`. + /// every slot with [`MaxSimElement::NO_MATCH`]. /// /// # Errors /// @@ -23,7 +23,7 @@ pub trait MaxSimKernel: Send + Sync + std::fmt::Debug { fn compute_max_sim( &self, doc: MatRef<'_, Standard>, - scores: &mut [f32], + scores: &mut [T::Score], ) -> Result<(), MaxSimError>; } @@ -31,7 +31,7 @@ pub trait MaxSimKernel: Send + Sync + std::fmt::Debug { /// kernel to [`Erase::erase`], which decides how to package it (e.g. as /// `Box>` via [`BoxErase`], a chamfer-only closure, a /// batched evaluator, …). -pub trait Erase { +pub trait Erase { type Output; /// `K` is generic so the body sees its concrete type and the compiler /// can inline it. @@ -42,7 +42,7 @@ pub trait Erase { #[derive(Debug, Clone, Copy)] pub struct BoxErase; -impl Erase for BoxErase { +impl Erase for BoxErase { type Output = Box>; fn erase + 'static>(self, kernel: K) -> Self::Output { From 12e04bd1864efa381465b1896f820d51e353f863 Mon Sep 17 00:00:00 2001 From: Suryansh Gupta Date: Tue, 22 Sep 2026 00:52:09 +0530 Subject: [PATCH 3/7] Add Neon kernel --- .../maxsim/packed_i8_x_unpacked_i8.rs | 262 ++++++++++++++++-- .../src/matrix_kernels/util.rs | 3 + .../src/multi_vector/distance/factory.rs | 8 +- diskann-wide/src/arch/aarch64/u32x4_.rs | 18 +- 4 files changed, 267 insertions(+), 24 deletions(-) diff --git a/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs b/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs index 11e5041557..46c8afe835 100644 --- a/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs +++ b/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs @@ -29,7 +29,9 @@ use super::packed_f32_x_unpacked_f32::Params; diskann_wide::alias!(i8x16 = i8x16); diskann_wide::alias!(i16x16 = i16x16); +diskann_wide::alias!(i32x4 = i32x4); diskann_wide::alias!(i32x8 = i32x8); +diskann_wide::alias!(u32x4 = u32x4); diskann_wide::alias!(u32x8 = u32x8); /// Widen a `PACK = 2` group into the little-endian `i16` lane pair consumed by the 16-bit @@ -390,6 +392,33 @@ unsafe fn group(ptr: Slice<'_, i8>, valid: usize) -> [i8; PAC }) } +/// Accumulate one packed row of `a`, held in `ai`, against every column of `b`, whose +/// contraction offset `bp` already points at. +/// +/// # Safety +/// +/// `valid` must not exceed `PACK`, and for every `j < NR` the first `valid` elements at +/// `bp.add(bstride * j)` must be readable. +#[inline(always)] +unsafe fn accumulate_row( + wide: W, + ai: W::Wide, + bp: Slice<'_, i8>, + bstride: Elements, + valid: usize, + acc: &mut [W::Acc; NR], +) where + W: ExtraWide, +{ + for (j, acc) in acc.iter_mut().enumerate() { + // SAFETY: By preconditions, the pointer offset is valid and its first `valid` + // elements are readable. + let bj = wide.splat(unsafe { group::(bp.add(bstride * j), valid) }); + + *acc = W::dot(ai, bj, *acc); + } +} + /// # Safety /// /// Bounds `a.k()` and `b.k()` must be equal to `k`. @@ -419,29 +448,53 @@ unsafe fn micro_kernel( let rows = a.rows(k); let k = k.value().get(); - for row in 0..rows { - // SAFETY: By preconditions, `ap.len() == astride * rows`. Since `row < rows`: - // - // * The pointer offset is valid. - // * The subsequent truncation is valid. - // * The slice passed to `wide.load` has a length equal to `astride`. - let ai = unsafe { wide.load(ap.add(astride * row).truncate(astride)) }; + // Loads the packed row `row` of `a`, which callers must keep below `rows`. + // + // SAFETY: By preconditions, `ap.len() == astride * rows`. Since `row < rows`: + // + // * The pointer offset is valid. + // * The subsequent truncation is valid. + // * The slice passed to `wide.load` has a length equal to `astride`. + let load = |row| unsafe { wide.load(ap.add(astride * row).truncate(astride)) }; - // The trailing row of `a` is zero padded, so zero filling `b` past `k` keeps every - // padded product at zero. + // Rows whose group lies entirely within `k`. Peeling the trailing partial group keeps + // `valid` constant here, folding away the zero-fill branch in `group` on the hot path. + let full = k / PACK; + + for row in 0..full { let i = row * PACK; - let valid = PACK.min(k - i); - for (j, acc) in acc.iter_mut().enumerate() { - // SAFETY: By preconditions, `bp.len() == bstride * NR`. Since `i < k`, `j < NR` - // and `i + valid <= k`: - // - // * The pointer offset is valid and its first `valid` elements are readable. - let bj = - wide.splat(unsafe { group::(bp.add(bstride * j + Elements::new(i)), valid) }); + // SAFETY: `row < full <= rows`, and since `i + PACK <= k`, every column of `b` has + // `PACK` readable elements at offset `i`. + unsafe { + accumulate_row( + wide, + load(row), + bp.add(Elements::new(i)), + bstride, + PACK, + &mut acc, + ) + }; + } - *acc = W::dot(ai, bj, *acc); - } + // The trailing row of `a` is zero padded, so zero filling `b` past `k` keeps every + // padded product at zero. + if full < rows { + let i = full * PACK; + + // SAFETY: `full < rows`, and every column of `b` has `k - i` readable elements at + // offset `i`. + unsafe { + accumulate_row( + wide, + load(full), + bp.add(Elements::new(i)), + bstride, + k - i, + &mut acc, + ) + }; } wide.max_into(Folder::fold(acc, W::max), c); @@ -604,6 +657,146 @@ mod x86_64 { } } +#[cfg(target_arch = "aarch64")] +mod aarch64 { + use super::*; + + use diskann_wide::arch::aarch64::Neon; + + panel_kernel!(Neon, 8, 6, 4, [1, 2, 3, 4, 5]); + panel_kernel!(Neon, 16, 6, 4, [1, 2, 3, 4, 5]); + + micro_kernel!(Neon, 8, 4, { 6, 5, 4, 3, 2, 1 }); + micro_kernel!(Neon, 16, 4, { 6, 5, 4, 3, 2, 1 }); + + /// Pack a `PACK = 4` group into the little-endian byte quad consumed by `sdot`. + /// + /// Broadcasting through `u32` lowers to a single `ld1r`. + #[inline(always)] + fn i8_quad(group: [i8; 4]) -> u32 { + u32::from_le_bytes(group.map(|x| x as u8)) + } + + //-----------// + // ExtraWide // + //-----------// + + impl ExtraWide<8, 4> for Neon { + type Wide = [i8x16; 2]; + type Splat = i8x16; + type Acc = [i32x4; 2]; + + #[inline(always)] + fn default(self) -> Self::Acc { + [SIMDVector::default(self), SIMDVector::default(self)] + } + + #[inline(always)] + unsafe fn load(self, slice: Slice<'_, i8>) -> Self::Wide { + bounds::check_eq!(slice.len(), 32); + + // SAFETY: Since `slice.len()` must be 32, the pointer offset and 16-wide SIMD + // loads are valid. + unsafe { + [ + SIMDVector::load_simd(self, slice.as_ptr()), + SIMDVector::load_simd(self, slice.add(Elements::new(16)).as_ptr()), + ] + } + } + + #[inline(always)] + fn splat(self, group: [i8; 4]) -> Self::Splat { + u32x4::::splat(self, i8_quad(group)).reinterpret_simd() + } + + #[inline(always)] + fn dot(a: Self::Wide, b: Self::Splat, acc: Self::Acc) -> Self::Acc { + core::array::from_fn(|i| acc[i].dot_simd(a[i], b)) + } + + #[inline(always)] + fn max(lhs: Self::Acc, rhs: Self::Acc) -> Self::Acc { + core::array::from_fn(|i| lhs[i].max_simd(rhs[i])) + } + + #[inline(always)] + fn max_into(self, lhs: Self::Acc, into: &mut [i32; 8]) { + // SAFETY: Since `into.len()` is 8, the pointer offset and 4-wide SIMD loads are + // valid. + let previous: Self::Acc = unsafe { + [ + SIMDVector::load_simd(self, into.as_ptr()), + SIMDVector::load_simd(self, into.as_ptr().add(4)), + ] + }; + + let max = >::max(lhs, previous); + + // SAFETY: Since `into.len()` is 8, the pointer offset and 4-wide SIMD stores are + // valid. + unsafe { + max[0].store_simd(into.as_mut_ptr()); + max[1].store_simd(into.as_mut_ptr().add(4)); + } + } + } + + impl ExtraWide<16, 4> for Neon { + type Wide = [i8x16; 4]; + type Splat = i8x16; + type Acc = [i32x4; 4]; + + #[inline(always)] + fn default(self) -> Self::Acc { + [SIMDVector::default(self); 4] + } + + #[inline(always)] + unsafe fn load(self, slice: Slice<'_, i8>) -> Self::Wide { + bounds::check_eq!(slice.len(), 64); + + // SAFETY: Since `slice.len()` must be 64, the pointer offsets and 16-wide SIMD + // loads are valid. + core::array::from_fn(|i| unsafe { + SIMDVector::load_simd(self, slice.add(Elements::new(16 * i)).as_ptr()) + }) + } + + #[inline(always)] + fn splat(self, group: [i8; 4]) -> Self::Splat { + u32x4::::splat(self, i8_quad(group)).reinterpret_simd() + } + + #[inline(always)] + fn dot(a: Self::Wide, b: Self::Splat, acc: Self::Acc) -> Self::Acc { + core::array::from_fn(|i| acc[i].dot_simd(a[i], b)) + } + + #[inline(always)] + fn max(lhs: Self::Acc, rhs: Self::Acc) -> Self::Acc { + core::array::from_fn(|i| lhs[i].max_simd(rhs[i])) + } + + #[inline(always)] + fn max_into(self, lhs: Self::Acc, into: &mut [i32; 16]) { + // SAFETY: Since `into.len()` is 16, the pointer offsets and 4-wide SIMD loads + // are valid. + let previous: Self::Acc = core::array::from_fn(|i| unsafe { + SIMDVector::load_simd(self, into.as_ptr().add(4 * i)) + }); + + let max = >::max(lhs, previous); + + for (i, max) in max.into_iter().enumerate() { + // SAFETY: Since `into.len()` is 16, the pointer offsets and 4-wide SIMD + // stores are valid. + unsafe { max.store_simd(into.as_mut_ptr().add(4 * i)) }; + } + } + } +} + /////////// // Tests // /////////// @@ -619,6 +812,9 @@ mod tests { #[cfg(target_arch = "x86_64")] use diskann_wide::arch::x86_64::V3; + #[cfg(target_arch = "aarch64")] + use diskann_wide::arch::aarch64::Neon; + use crate::{matrix_kernels::maxsim, multi_vector::BlockTransposed}; ///////////////// @@ -720,6 +916,16 @@ mod tests { 16 => { 6, 5, 4, 3, 2, 1 }, ); + #[cfg(target_arch = "aarch64")] + test_micro_kernel!( + test_micro_kernel_neon, + Neon::new_checked(), + 0x7d4a1e6c93b0f582, + 4, + 8 => { 6, 5, 4, 3, 2, 1 }, + 16 => { 6, 5, 4, 3, 2, 1 }, + ); + ///////////////// // PanelKernel // ///////////////// @@ -829,6 +1035,15 @@ mod tests { (16, 6, 2), ); + #[cfg(target_arch = "aarch64")] + test_panel_kernel!( + test_panel_kernel_neon, + Neon::new_checked(), + 0x9f3e7ab4c05d1268, + (8, 6, 4), + (16, 6, 4), + ); + //////////// // Driver // //////////// @@ -918,4 +1133,13 @@ mod tests { 0x63c8ed19f4720ab5, (16, 6, 2), ); + + #[cfg(target_arch = "aarch64")] + test_driver!( + test_driver_neon, + Neon::new_checked(), + 0x63c8ed19f4720ab5, + (8, 6, 4), + (16, 6, 4), + ); } diff --git a/diskann-quantization/src/matrix_kernels/util.rs b/diskann-quantization/src/matrix_kernels/util.rs index 772877926c..1e4b4a4c63 100644 --- a/diskann-quantization/src/matrix_kernels/util.rs +++ b/diskann-quantization/src/matrix_kernels/util.rs @@ -158,6 +158,8 @@ mod aarch64 { impl_loadstore!(f32, 4, f32x4, Neon); impl_loadstore!(f32, 8, f32x8, Neon); impl_loadstore!(f32, 16, f32x16, Neon); + impl_loadstore!(i32, 8, i32x8, Neon); + impl_loadstore!(i32, 16, i32x16, Neon); } ////////// @@ -352,6 +354,7 @@ mod test { test_load_store_neon, Neon::new_checked(), f32 => { 4, 8, 16 }, + i32 => { 8, 16 }, ); #[test] diff --git a/diskann-quantization/src/multi_vector/distance/factory.rs b/diskann-quantization/src/multi_vector/distance/factory.rs index 02753cdbb8..20b38412a4 100644 --- a/diskann-quantization/src/multi_vector/distance/factory.rs +++ b/diskann-quantization/src/multi_vector/distance/factory.rs @@ -506,8 +506,12 @@ impl> diskann_wide::arch::Target1 { fn run(self, arch: Neon, query: MatRef<'_, Standard>) -> E::Output { - // Neon retargets to Scalar until the dotprod kernel lands. - diskann_wide::arch::Target1::::run(self, Scalar::from(arch), query) + let prepared = BlockTransposed::::from_matrix_view(query.as_matrix_view()); + self.0.erase(Prepared { + arch, + prepared, + _packing: Pack::<6>, + }) } } diff --git a/diskann-wide/src/arch/aarch64/u32x4_.rs b/diskann-wide/src/arch/aarch64/u32x4_.rs index 42de9fd1aa..27a8175055 100644 --- a/diskann-wide/src/arch/aarch64/u32x4_.rs +++ b/diskann-wide/src/arch/aarch64/u32x4_.rs @@ -4,13 +4,13 @@ */ use crate::{ - Emulated, SIMDDotProduct, SIMDMask, SIMDMulAdd, SIMDPartialEq, SIMDPartialOrd, SIMDSelect, - SIMDSumTree, SIMDVector, constant::Const, helpers, + Emulated, SIMDDotProduct, SIMDMask, SIMDMulAdd, SIMDPartialEq, SIMDPartialOrd, SIMDReinterpret, + SIMDSelect, SIMDSumTree, SIMDVector, constant::Const, helpers, }; // AArch64 masks use super::{ - Neon, internal, + Neon, i8x16, internal, macros::{self, AArchLoadStore, AArchSplat}, masks::mask32x4, u8x16, @@ -114,6 +114,18 @@ impl SIMDDotProduct for u32x4 { } } +////////////////// +// Reinterprets // +////////////////// + +impl SIMDReinterpret for u32x4 { + #[inline(always)] + fn reinterpret_simd(self) -> i8x16 { + // SAFETY: Allowed by the `Neon` architecture. + i8x16(unsafe { vreinterpretq_s8_u32(self.0) }) + } +} + /////////// // Tests // /////////// From 0aef9742d6482b84f50ccc1f4cb78eb18562b5fa Mon Sep 17 00:00:00 2001 From: Suryansh Gupta Date: Fri, 25 Sep 2026 02:06:29 +0530 Subject: [PATCH 4/7] Fix review comments --- .../src/matrix_kernels/blocks/packed.rs | 18 +++++----- .../maxsim/packed_i8_x_unpacked_i8.rs | 36 +++++++++---------- .../src/matrix_kernels/test_util.rs | 7 ++-- 3 files changed, 32 insertions(+), 29 deletions(-) diff --git a/diskann-quantization/src/matrix_kernels/blocks/packed.rs b/diskann-quantization/src/matrix_kernels/blocks/packed.rs index 48e40a8059..3bf1120610 100644 --- a/diskann-quantization/src/matrix_kernels/blocks/packed.rs +++ b/diskann-quantization/src/matrix_kernels/blocks/packed.rs @@ -25,8 +25,8 @@ fn padded_k(k: usize) -> usize { /// Elements are gathered into groups of size `SZ`. A collection of `self.k` groups forms /// a "block". `self.blocks` tracks how many such blocks are in the view. /// -/// `PACK` interleaves that many consecutive columns within each group, so one physical row -/// spans `SZ * PACK` elements. `k` stays the logical dimension; a trailing row that `PACK` +/// `PACK` interleaves that many consecutive columns within each group, so one pack spans +/// `SZ * PACK` elements. `k` stays the logical dimension; a trailing pack that `PACK` /// does not fill is zero-padded. /// /// This layout requires that no block is partially filled. @@ -246,7 +246,7 @@ impl View<'_, T, SZ, PACK> { /// A block containing `k` contiguous groups of size `SZ`. /// -/// `PACK` consecutive groups are interleaved into one physical row of `SZ * PACK` elements. +/// `PACK` consecutive groups are interleaved into one pack of `SZ * PACK` elements. /// /// # Class Invariants /// @@ -289,15 +289,15 @@ impl<'a, T, const SZ: usize, const PACK: usize> Panel<'a, T, SZ, PACK> { self.k } - /// Return the number of elements spanned by one physical row. - pub(in crate::matrix_kernels) const fn row_stride(&self) -> Elements { + /// Return the number of elements spanned by one pack. + pub(in crate::matrix_kernels) const fn pack_stride(&self) -> Elements { Elements::new(SZ * PACK) } - /// Return the number of physical rows in `self`. + /// Return the number of packs in `self`. /// /// `k` must be equal to the contraction dimension tracked by [`Self::k`]. - pub(in crate::matrix_kernels) fn rows(&self, k: DimK) -> usize { + pub(in crate::matrix_kernels) fn packs(&self, k: DimK) -> usize { bounds::check_eq!(self.k, k.value()); padded_k::(k.value().get()) / PACK } @@ -366,8 +366,8 @@ mod tests { view.checked_visit_panels(|panel, block| { assert_eq!(block, blocks, "{ctx}"); - assert_eq!(panel.rows(dim_k), padded / PACK, "{ctx}"); - assert_eq!(panel.row_stride().value(), SZ * PACK, "{ctx}"); + assert_eq!(panel.packs(dim_k), padded / PACK, "{ctx}"); + assert_eq!(panel.pack_stride().value(), SZ * PACK, "{ctx}"); let flat = panel.checked_as_std_slice(); assert_eq!(flat.len(), SZ * padded, "{ctx}"); diff --git a/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs b/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs index 46c8afe835..308d7633c9 100644 --- a/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs +++ b/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs @@ -6,7 +6,7 @@ //! The integer counterpart of [`super::packed_f32_x_unpacked_f32`], accumulating in `i32`. //! //! The blocking strategy is identical. The difference is that `a` interleaves `PACK` -//! consecutive contraction indices within each packed row so that a single lane of the +//! consecutive contraction indices within each pack so that a single lane of the //! widening dot-product instructions consumes `PACK` products at a time. `PACK` is chosen //! per architecture to match the instruction used. //! @@ -392,7 +392,7 @@ unsafe fn group(ptr: Slice<'_, i8>, valid: usize) -> [i8; PAC }) } -/// Accumulate one packed row of `a`, held in `ai`, against every column of `b`, whose +/// Accumulate one pack of `a`, held in `ai`, against every column of `b`, whose /// contraction offset `bp` already points at. /// /// # Safety @@ -400,7 +400,7 @@ unsafe fn group(ptr: Slice<'_, i8>, valid: usize) -> [i8; PAC /// `valid` must not exceed `PACK`, and for every `j < NR` the first `valid` elements at /// `bp.add(bstride * j)` must be readable. #[inline(always)] -unsafe fn accumulate_row( +unsafe fn accumulate_pack( wide: W, ai: W::Wide, bp: Slice<'_, i8>, @@ -442,34 +442,34 @@ unsafe fn micro_kernel( let mut acc = [wide.default(); NR]; - let astride = a.row_stride(); + let astride = a.pack_stride(); let bstride = b.stride(k); - let rows = a.rows(k); + let packs = a.packs(k); let k = k.value().get(); - // Loads the packed row `row` of `a`, which callers must keep below `rows`. + // Loads pack `pack` of `a`, which callers must keep below `packs`. // - // SAFETY: By preconditions, `ap.len() == astride * rows`. Since `row < rows`: + // SAFETY: By preconditions, `ap.len() == astride * packs`. Since `pack < packs`: // // * The pointer offset is valid. // * The subsequent truncation is valid. // * The slice passed to `wide.load` has a length equal to `astride`. - let load = |row| unsafe { wide.load(ap.add(astride * row).truncate(astride)) }; + let load = |pack| unsafe { wide.load(ap.add(astride * pack).truncate(astride)) }; - // Rows whose group lies entirely within `k`. Peeling the trailing partial group keeps + // Packs whose group lies entirely within `k`. Peeling the trailing partial group keeps // `valid` constant here, folding away the zero-fill branch in `group` on the hot path. let full = k / PACK; - for row in 0..full { - let i = row * PACK; + for pack in 0..full { + let i = pack * PACK; - // SAFETY: `row < full <= rows`, and since `i + PACK <= k`, every column of `b` has + // SAFETY: `pack < full <= packs`, and since `i + PACK <= k`, every column of `b` has // `PACK` readable elements at offset `i`. unsafe { - accumulate_row( + accumulate_pack( wide, - load(row), + load(pack), bp.add(Elements::new(i)), bstride, PACK, @@ -478,15 +478,15 @@ unsafe fn micro_kernel( }; } - // The trailing row of `a` is zero padded, so zero filling `b` past `k` keeps every + // The trailing pack of `a` is zero padded, so zero filling `b` past `k` keeps every // padded product at zero. - if full < rows { + if full < packs { let i = full * PACK; - // SAFETY: `full < rows`, and every column of `b` has `k - i` readable elements at + // SAFETY: `full < packs`, and every column of `b` has `k - i` readable elements at // offset `i`. unsafe { - accumulate_row( + accumulate_pack( wide, load(full), bp.add(Elements::new(i)), diff --git a/diskann-quantization/src/matrix_kernels/test_util.rs b/diskann-quantization/src/matrix_kernels/test_util.rs index dea228a560..0e4cfebf7d 100644 --- a/diskann-quantization/src/matrix_kernels/test_util.rs +++ b/diskann-quantization/src/matrix_kernels/test_util.rs @@ -5,7 +5,10 @@ use diskann_utils::views::{Init, Matrix}; use half::f16; -use rand::{Rng, distr::Distribution}; +use rand::{ + Rng, + distr::{Distribution, StandardUniform}, +}; /////////////////////// // panic_message_for // @@ -66,7 +69,7 @@ impl Distribution for TestDistr { impl Distribution for TestDistr { fn sample(&self, rng: &mut R) -> i8 { - rng.random_range(i8::MIN..=i8::MAX) + StandardUniform {}.sample(rng) } } From f043eeeba1cd4961f1896c1a634aa0d451e07c97 Mon Sep 17 00:00:00 2001 From: Suryansh Gupta Date: Fri, 25 Sep 2026 20:39:13 +0530 Subject: [PATCH 5/7] Widen i8 to i16 for v3 kernel --- .../maxsim/packed_i8_x_unpacked_i8.rs | 231 ++++++++++++++---- .../src/matrix_kernels/util.rs | 35 +++ .../src/multi_vector/distance/factory.rs | 2 +- 3 files changed, 220 insertions(+), 48 deletions(-) diff --git a/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs b/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs index 308d7633c9..657f3326fe 100644 --- a/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs +++ b/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs @@ -34,11 +34,11 @@ diskann_wide::alias!(i32x8 = i32x8); diskann_wide::alias!(u32x4 = u32x4); diskann_wide::alias!(u32x8 = u32x8); -/// Widen a `PACK = 2` group into the little-endian `i16` lane pair consumed by the 16-bit -/// dot products. +/// Pack a `PACK = 2` group into the little-endian `i16` lane pair consumed by the 16-bit dot +/// products. #[inline(always)] -fn i16_pair([lo, hi]: [i8; 2]) -> u32 { - u32::from(i16::from(lo) as u16) | (u32::from(i16::from(hi) as u16) << 16) +fn i16_pair([lo, hi]: [i16; 2]) -> u32 { + u32::from(lo as u16) | (u32::from(hi as u16) << 16) } //--------// @@ -63,7 +63,10 @@ pub(crate) struct Driver<'a, A, const MR: usize, const NR: usize, const PACK: us params: Params, } -impl<'a, A, const MR: usize, const NR: usize, const PACK: usize> Driver<'a, A, MR, NR, PACK> { +impl<'a, A, const MR: usize, const NR: usize, const PACK: usize> Driver<'a, A, MR, NR, PACK> +where + A: PrepareB, +{ /// Prepare for a maxsim on `a` and `b` with the results stored directly into `c`. /// /// `c` does not require any specific initial value. @@ -88,17 +91,15 @@ impl<'a, A, const MR: usize, const NR: usize, const PACK: usize> Driver<'a, A, M "output length must occupy exactly the packed A blocks", ); + let params = Params::new( + cache, + a.block_stride(k).bytes(), + b.stride(k).cast::().bytes(), + NR, + ); + // SAFETY: Inherited from caller. - unsafe { - Self::new_inner( - arch, - a, - b, - c, - k, - Params::new(cache, a.block_stride(k).bytes(), b.stride(k).bytes(), NR), - ) - } + unsafe { Self::new_inner(arch, a, b, c, k, params) } } /// # Safety @@ -135,7 +136,7 @@ impl<'a, A, const MR: usize, const NR: usize, const PACK: usize> Driver<'a, A, M impl driver::Drive for Driver<'_, A, MR, NR, PACK> where - A: util::LoadStore + Architecture, + A: util::LoadStore + PrepareB + Architecture, for<'a> PanelKernel<'a, A, MR, NR, PACK>: driver::PanelKernel, { fn drive(&mut self) { @@ -152,9 +153,14 @@ where let last_a_block = self.a.blocks().get() - 1; let mut c = MutSlice::new(self.c); + let mut b_scratch = Vec::new(); let on_a_panels = |a_panels: packed::View<'_, i8, MR, PACK>, a_block_base| { let on_b_panels = |b_panels: unpacked::View<'_, i8>, _| { + // SAFETY: By class invariant, `b_panels.k()` is equal to `self.k`. + let b_panels = + unsafe { self.arch.prepare(b_panels, &mut b_scratch, self.k) }; + let panel_kernel = |a_panel: packed::Panel<'_, i8, MR, PACK>, a_block_offset| { // If we are in the very last block and we need to sub-fill, do @@ -232,20 +238,64 @@ where } } +//----------// +// PrepareB // +//----------// + +/// Converts sub-views of `b` into the element type streamed by the micro-kernels. +/// +/// Architectures whose dot products consume elements wider than `i8` widen `b` here, once +/// per sub-view, instead of in the micro-kernel's inner loop. +pub(crate) trait PrepareB: Copy { + type Elem: Copy; + + /// Return `b` as [`Self::Elem`], using `scratch` as storage if a conversion is needed. + /// + /// # Safety + /// + /// `b.k()` must be equal to `k`. + unsafe fn prepare<'a>( + self, + b: unpacked::View<'a, i8>, + scratch: &'a mut Vec, + k: DimK, + ) -> unpacked::View<'a, Self::Elem>; +} + +impl PrepareB for Scalar { + type Elem = i8; + + #[inline(always)] + unsafe fn prepare<'a>( + self, + b: unpacked::View<'a, i8>, + _: &'a mut Vec, + _: DimK, + ) -> unpacked::View<'a, i8> { + b + } +} + //-------------// // PanelKernel // //-------------// #[derive(Debug)] -pub(super) struct PanelKernel<'a, A, const MR: usize, const NR: usize, const PACK: usize> { +pub(super) struct PanelKernel<'a, A, const MR: usize, const NR: usize, const PACK: usize> +where + A: PrepareB, +{ arch: A, a: packed::Panel<'a, i8, MR, PACK>, - b: unpacked::View<'a, i8>, + b: unpacked::View<'a, A::Elem>, c: [i32; MR], k: DimK, } -impl<'a, A, const MR: usize, const NR: usize, const PACK: usize> PanelKernel<'a, A, MR, NR, PACK> { +impl<'a, A, const MR: usize, const NR: usize, const PACK: usize> PanelKernel<'a, A, MR, NR, PACK> +where + A: PrepareB, +{ /// Construct a new kernel. /// /// # Safety @@ -254,7 +304,7 @@ impl<'a, A, const MR: usize, const NR: usize, const PACK: usize> PanelKernel<'a, pub(super) unsafe fn new( arch: A, a: packed::Panel<'a, i8, MR, PACK>, - b: unpacked::View<'a, i8>, + b: unpacked::View<'a, A::Elem>, c: [i32; MR], k: DimK, ) -> Self { @@ -280,14 +330,14 @@ struct Visitor<'a, A, const MR: usize, const NR: usize, const PACK: usize> { k: DimK, } -impl unpacked::PanelVisitor +impl unpacked::PanelVisitor for Visitor<'_, A, MR, NR, PACK> where - A: Copy, + A: PrepareB, for<'a> MicroKernel<'a, A, MR, NR, PACK>: driver::MicroKernel, { #[inline(always)] - fn visit(&mut self, b: unpacked::Panel<'_, i8, NR>, _: usize) { + fn visit(&mut self, b: unpacked::Panel<'_, A::Elem, NR>, _: usize) { // SAFETY: This is only used on contexts where `self.a.k()`, `b.k()`, and `self.k` // are all equal. let mut micro = unsafe { MicroKernel::new(self.arch, self.a, b, self.c, self.k) }; @@ -348,22 +398,28 @@ panel_kernel!(Scalar, 8, 2, 2, [1]); /// # Class Invariants /// /// `a.k()` and `b.k()` are equal to `k`. -struct MicroKernel<'a, A, const MR: usize, const NR: usize, const PACK: usize> { +struct MicroKernel<'a, A, const MR: usize, const NR: usize, const PACK: usize> +where + A: PrepareB, +{ arch: A, a: packed::Panel<'a, i8, MR, PACK>, - b: unpacked::Panel<'a, i8, NR>, + b: unpacked::Panel<'a, A::Elem, NR>, c: &'a mut [i32; MR], k: DimK, } -impl<'a, A, const MR: usize, const NR: usize, const PACK: usize> MicroKernel<'a, A, MR, NR, PACK> { +impl<'a, A, const MR: usize, const NR: usize, const PACK: usize> MicroKernel<'a, A, MR, NR, PACK> +where + A: PrepareB, +{ /// # Safety /// /// Bounds `a.k()` and `b.k()` must be equal to `k`. unsafe fn new( arch: A, a: packed::Panel<'a, i8, MR, PACK>, - b: unpacked::Panel<'a, i8, NR>, + b: unpacked::Panel<'a, A::Elem, NR>, c: &'a mut [i32; MR], k: DimK, ) -> Self { @@ -381,13 +437,16 @@ impl<'a, A, const MR: usize, const NR: usize, const PACK: usize> MicroKernel<'a, /// /// `valid` must not exceed `PACK` and the first `valid` elements of `ptr` must be readable. #[inline(always)] -unsafe fn group(ptr: Slice<'_, i8>, valid: usize) -> [i8; PACK] { +unsafe fn group(ptr: Slice<'_, T>, valid: usize) -> [T; PACK] +where + T: Copy + Default, +{ core::array::from_fn(|p| { if p < valid { // SAFETY: Since `p < valid`, the pointer offset is valid and readable. unsafe { *ptr.add(Elements::new(p)).as_unit().as_ref() } } else { - 0 + T::default() } }) } @@ -403,8 +462,8 @@ unsafe fn group(ptr: Slice<'_, i8>, valid: usize) -> [i8; PAC unsafe fn accumulate_pack( wide: W, ai: W::Wide, - bp: Slice<'_, i8>, - bstride: Elements, + bp: Slice<'_, W::Elem>, + bstride: Elements, valid: usize, acc: &mut [W::Acc; NR], ) where @@ -413,7 +472,7 @@ unsafe fn accumulate_pack(bp.add(bstride * j), valid) }); + let bj = unsafe { wide.splat(bp.add(bstride * j), valid) }; *acc = W::dot(ai, bj, *acc); } @@ -426,7 +485,7 @@ unsafe fn accumulate_pack( wide: W, a: packed::Panel<'_, i8, MR, PACK>, - b: unpacked::Panel<'_, i8, NR>, + b: unpacked::Panel<'_, W::Elem, NR>, c: &mut [i32; MR], k: DimK, ) where @@ -517,7 +576,7 @@ macro_rules! micro_kernel { micro_kernel!(Scalar, 8, 2, { 2, 1 }); -trait ExtraWide: Copy { +trait ExtraWide: PrepareB { type Wide: Copy; type Splat: Copy; type Acc: Copy; @@ -528,7 +587,14 @@ trait ExtraWide: Copy { unsafe fn load(self, slice: Slice<'_, i8>) -> Self::Wide; fn default(self) -> Self::Acc; - fn splat(self, group: [i8; PACK]) -> Self::Splat; + + /// Broadcast the `PACK` elements at `b`, treating those past `valid` as zero. + /// + /// # Safety + /// + /// `valid` must not exceed `PACK` and the first `valid` elements of `b` must be readable. + unsafe fn splat(self, b: Slice<'_, Self::Elem>, valid: usize) -> Self::Splat; + fn dot(a: Self::Wide, b: Self::Splat, acc: Self::Acc) -> Self::Acc; fn max(lhs: Self::Acc, rhs: Self::Acc) -> Self::Acc; fn max_into(self, max: Self::Acc, into: &mut [i32; ELEMENTS]); @@ -555,8 +621,10 @@ impl ExtraWide<8, 2> for Scalar { } #[inline(always)] - fn splat(self, group: [i8; 2]) -> Self::Splat { - u32x8::::splat(self, i16_pair(group)).reinterpret_simd() + unsafe fn splat(self, b: Slice<'_, i8>, valid: usize) -> Self::Splat { + // SAFETY: Inherited from caller. + let pair = i16_pair(unsafe { group(b, valid) }.map(i16::from)); + u32x8::::splat(self, pair).reinterpret_simd() } #[inline(always)] @@ -585,10 +653,37 @@ mod x86_64 { use diskann_wide::arch::x86_64::V3; + use crate::matrix_kernels::util::{Convert, Converter}; + panel_kernel!(V3, 16, 6, 2, [1, 2, 3, 4, 5]); micro_kernel!(V3, 16, 2, { 6, 5, 4, 3, 2, 1 }); + //----------// + // PrepareB // + //----------// + + impl PrepareB for V3 { + type Elem = i16; + + #[inline(always)] + unsafe fn prepare<'a>( + self, + b: unpacked::View<'a, i8>, + scratch: &'a mut Vec, + k: DimK, + ) -> unpacked::View<'a, i16> { + // SAFETY: Inherited from caller. + let from = unsafe { b.as_std_slice(k) }; + + scratch.resize(from.len(), 0); + Converter::new(self).convert(scratch, from); + + // SAFETY: `scratch` has length `b.extent() * k`. + unsafe { unpacked::View::new(Slice::new(scratch), b.extent(), k) } + } + } + //-----------// // ExtraWide // //-----------// @@ -620,8 +715,10 @@ mod x86_64 { } #[inline(always)] - fn splat(self, group: [i8; 2]) -> Self::Splat { - u32x8::::splat(self, i16_pair(group)).reinterpret_simd() + unsafe fn splat(self, b: Slice<'_, i16>, valid: usize) -> Self::Splat { + // SAFETY: Inherited from caller. + let pair = i16_pair(unsafe { group(b, valid) }); + u32x8::::splat(self, pair).reinterpret_simd() } #[inline(always)] @@ -677,6 +774,24 @@ mod aarch64 { u32::from_le_bytes(group.map(|x| x as u8)) } + //----------// + // PrepareB // + //----------// + + impl PrepareB for Neon { + type Elem = i8; + + #[inline(always)] + unsafe fn prepare<'a>( + self, + b: unpacked::View<'a, i8>, + _: &'a mut Vec, + _: DimK, + ) -> unpacked::View<'a, i8> { + b + } + } + //-----------// // ExtraWide // //-----------// @@ -706,8 +821,10 @@ mod aarch64 { } #[inline(always)] - fn splat(self, group: [i8; 4]) -> Self::Splat { - u32x4::::splat(self, i8_quad(group)).reinterpret_simd() + unsafe fn splat(self, b: Slice<'_, i8>, valid: usize) -> Self::Splat { + // SAFETY: Inherited from caller. + let quad = i8_quad(unsafe { group(b, valid) }); + u32x4::::splat(self, quad).reinterpret_simd() } #[inline(always)] @@ -764,8 +881,10 @@ mod aarch64 { } #[inline(always)] - fn splat(self, group: [i8; 4]) -> Self::Splat { - u32x4::::splat(self, i8_quad(group)).reinterpret_simd() + unsafe fn splat(self, b: Slice<'_, i8>, valid: usize) -> Self::Splat { + // SAFETY: Inherited from caller. + let quad = i8_quad(unsafe { group(b, valid) }); + u32x4::::splat(self, quad).reinterpret_simd() } #[inline(always)] @@ -827,6 +946,7 @@ mod tests { rng: &mut impl rand::Rng, ctx: std::fmt::Arguments<'_>, ) where + A: PrepareB, for<'a> MicroKernel<'a, A, MR, NR, PACK>: driver::MicroKernel, { let (ref_a, ref_b, ref_c) = maxsim::test::generate_i8(MR, k.value().get(), NR, rng); @@ -836,6 +956,17 @@ mod tests { let a_bt = BlockTransposed::::from_matrix_view(ref_a.as_view()); let ref_b = ref_b.transpose(); + let mut scratch = Vec::new(); + + // SAFETY: Test builds will verify the bounds we passed. + let b = unsafe { + arch.prepare( + unpacked::View::from_matrix_view(ref_b.as_view()).unwrap(), + &mut scratch, + k, + ) + }; + let mut c = [i32::MIN; MR]; // Run the test kernel. @@ -845,7 +976,7 @@ mod tests { MicroKernel::new( arch, packed::Panel::new(Slice::new(a_bt.as_slice()), k), - unpacked::Panel::new(Slice::new(ref_b.as_slice()), k), + unpacked::Panel::new(Slice::new(b.as_std_slice(k)), k), &mut c, k, ) @@ -940,7 +1071,7 @@ mod tests { rng: &mut impl rand::Rng, ctx: std::fmt::Arguments<'_>, ) where - A: Copy, + A: PrepareB, for<'a> PanelKernel<'a, A, MR, NR, PACK>: driver::PanelKernel, { for blocks in 0..4 { @@ -958,6 +1089,8 @@ mod tests { let extent = NonZeroUsize::new(cols).unwrap(); + let mut scratch = Vec::new(); + let c = [i32::MIN; MR]; // SAFETY: Test builds will verify the bounds we passed. @@ -965,7 +1098,11 @@ mod tests { PanelKernel::new( arch, packed::Panel::new(Slice::new(a_bt.as_slice()), k), - unpacked::View::new(Slice::new(ref_b.as_slice()), extent, k), + arch.prepare( + unpacked::View::new(Slice::new(ref_b.as_slice()), extent, k), + &mut scratch, + k, + ), c, k, ) @@ -1052,7 +1189,7 @@ mod tests { arch: A, rng: &mut impl rand::Rng, ) where - A: Copy, + A: PrepareB, for<'a> Driver<'a, A, MR, NR, PACK>: driver::Drive, { let cases = maxsim::test::packed_x_unpacked_test_dims(MR, NR); diff --git a/diskann-quantization/src/matrix_kernels/util.rs b/diskann-quantization/src/matrix_kernels/util.rs index 1e4b4a4c63..a5c79b0d5a 100644 --- a/diskann-quantization/src/matrix_kernels/util.rs +++ b/diskann-quantization/src/matrix_kernels/util.rs @@ -38,6 +38,25 @@ where } } +#[cfg(target_arch = "x86_64")] +impl Convert for Converter { + #[inline(always)] + fn convert(self, to: &mut [i16], from: &[i8]) { + use diskann_wide::{SIMDVector, arch::x86_64::V3}; + diskann_wide::alias!(i8s = ::i8x16); + diskann_wide::alias!(i16s = ::i16x16); + + debug_assert_eq!(to.len(), from.len(), "lengths must be equal"); + + let (to_chunks, to_tail) = to.as_chunks_mut::<16>(); + let (from_chunks, from_tail) = from.as_chunks::<16>(); + + std::iter::zip(to_chunks, from_chunks) + .for_each(|(to, from)| *to = i16s::from(i8s::from_array(self.0, *from)).to_array()); + std::iter::zip(to_tail, from_tail).for_each(|(to, from)| *to = (*from).into()); + } +} + ////////// // Load // ////////// @@ -266,6 +285,22 @@ mod test { #[cfg(target_arch = "aarch64")] use diskann_wide::arch::aarch64::Neon; + #[cfg(target_arch = "x86_64")] + #[test] + fn test_convert_i16_i8_v3() { + if let Some(arch) = V3::new_checked() { + for len in 0..40 { + let from: Vec = (0..len).map(|i| (37 * i) as i8).collect(); + let mut to = vec![0; len]; + + Converter::new(arch).convert(&mut to, &from); + + let expected: Vec = from.iter().map(|&x| x.into()).collect(); + assert_eq!(to, expected, "len = {len}"); + } + } + } + trait FromUsize { fn from_usize(v: usize) -> Self; } diff --git a/diskann-quantization/src/multi_vector/distance/factory.rs b/diskann-quantization/src/multi_vector/distance/factory.rs index 20b38412a4..e160406f23 100644 --- a/diskann-quantization/src/multi_vector/distance/factory.rs +++ b/diskann-quantization/src/multi_vector/distance/factory.rs @@ -179,7 +179,7 @@ where impl MaxSimKernel for Prepared, NR> where - A: Architecture, + A: mk::maxsim::packed_i8_x_unpacked_i8::PrepareB + Architecture, for<'a> mk::maxsim::packed_i8_x_unpacked_i8::Driver<'a, A, GROUP, NR, PACK>: mk::Drive, { fn nrows(&self) -> usize { From aa4e1029c76ce7157e3e4608b2feec84cd00ca4e Mon Sep 17 00:00:00 2001 From: Suryansh Gupta Date: Thu, 1 Oct 2026 23:23:46 +0530 Subject: [PATCH 6/7] Implement v4 i8 x i8 kernels --- .../src/matrix_kernels/blocks/unpacked.rs | 10 +- .../maxsim/packed_i8_x_unpacked_i8.rs | 535 +++++++++++++++--- .../src/matrix_kernels/util.rs | 86 +-- .../src/multi_vector/distance/factory.rs | 76 ++- 4 files changed, 546 insertions(+), 161 deletions(-) diff --git a/diskann-quantization/src/matrix_kernels/blocks/unpacked.rs b/diskann-quantization/src/matrix_kernels/blocks/unpacked.rs index d8b4d01f5e..ff4b5b2026 100644 --- a/diskann-quantization/src/matrix_kernels/blocks/unpacked.rs +++ b/diskann-quantization/src/matrix_kernels/blocks/unpacked.rs @@ -350,7 +350,7 @@ impl Panel<'_, T, EXTENT> { #[derive(Debug, Clone, Copy)] pub(in crate::matrix_kernels) struct Remainder<'a, T, const CAPACITY: usize> { ptr: Slice<'a, T>, - _start: usize, + start: usize, extent: NonZeroUsize, k: Bound, } @@ -370,7 +370,7 @@ impl<'a, T, const CAPACITY: usize> Remainder<'a, T, CAPACITY> { Self { ptr, - _start: start, + start, extent, k, } @@ -384,12 +384,8 @@ impl<'a, T, const CAPACITY: usize> Remainder<'a, T, CAPACITY> { } /// Return the index of the first band in `self`'s immediate parent [`View`]. - #[cfg_attr( - not(test), - expect(unused, reason = "this completes an API but is not used yet") - )] pub(in crate::matrix_kernels) fn start(&self) -> usize { - self._start + self.start } /// Return the number of elements in each "band" of `self`. diff --git a/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs b/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs index 657f3326fe..2a2dc50674 100644 --- a/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs +++ b/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs @@ -13,7 +13,7 @@ //! Each product is bounded by `128 * 128`, so the `i32` accumulator is exact and cannot //! overflow for contraction dimensions up to `131_071`. -use diskann_wide::arch::{Architecture, Scalar}; +use diskann_wide::arch::{Architecture, Scalar, Target}; use diskann_wide::{SIMDDotProduct, SIMDMinMax, SIMDReinterpret, SIMDVector}; use crate::matrix_kernels::{ @@ -28,11 +28,15 @@ use crate::matrix_kernels::{ use super::packed_f32_x_unpacked_f32::Params; diskann_wide::alias!(i8x16 = i8x16); +diskann_wide::alias!(i8x64 = i8x64); diskann_wide::alias!(i16x16 = i16x16); diskann_wide::alias!(i32x4 = i32x4); diskann_wide::alias!(i32x8 = i32x8); +diskann_wide::alias!(i32x16 = i32x16); +diskann_wide::alias!(u8x64 = u8x64); diskann_wide::alias!(u32x4 = u32x4); diskann_wide::alias!(u32x8 = u32x8); +diskann_wide::alias!(u32x16 = u32x16); /// Pack a `PACK = 2` group into the little-endian `i16` lane pair consumed by the 16-bit dot /// products. @@ -41,6 +45,16 @@ fn i16_pair([lo, hi]: [i16; 2]) -> u32 { u32::from(lo as u16) | (u32::from(hi as u16) << 16) } +/// Pack a `PACK = 4` group into the little-endian byte quad consumed by the 8-bit dot +/// products. +/// +/// Broadcasting through `u32` lowers to a single `ld1r` on Neon. +#[cfg(any(target_arch = "x86_64", target_arch = "aarch64"))] +#[inline(always)] +fn i8_quad(group: [i8; 4]) -> u32 { + u32::from_le_bytes(group.map(|x| x as u8)) +} + //--------// // Driver // //--------// @@ -54,9 +68,12 @@ fn i16_pair([lo, hi]: [i16; 2]) -> u32 { /// /// 1. `a.k()` and `b.k()` must be equal to `k`. /// 2. `c.len().div_ceil(MR)` must be equal to `a.blocks()`. -pub(crate) struct Driver<'a, A, const MR: usize, const NR: usize, const PACK: usize> { +pub(crate) struct Driver<'a, A, const MR: usize, const NR: usize, const PACK: usize> +where + A: PrepareB, +{ arch: A, - a: packed::View<'a, i8, MR, PACK>, + a: packed::View<'a, A::AElem, MR, PACK>, b: unpacked::View<'a, i8>, c: &'a mut [i32], k: DimK, @@ -77,7 +94,7 @@ where /// 2. `c.len().div_ceil(MR)` must be equal to `a.blocks()`. pub(crate) unsafe fn new( arch: A, - a: packed::View<'a, i8, MR, PACK>, + a: packed::View<'a, A::AElem, MR, PACK>, b: unpacked::View<'a, i8>, c: &'a mut [i32], k: DimK, @@ -108,7 +125,7 @@ where /// 2. `c.len().div_ceil(MR)` must be equal to `a.blocks()`. unsafe fn new_inner( arch: A, - a: packed::View<'a, i8, MR, PACK>, + a: packed::View<'a, A::AElem, MR, PACK>, b: unpacked::View<'a, i8>, c: &'a mut [i32], k: DimK, @@ -153,16 +170,16 @@ where let last_a_block = self.a.blocks().get() - 1; let mut c = MutSlice::new(self.c); - let mut b_scratch = Vec::new(); + let mut b_scratch = A::Scratch::default(); - let on_a_panels = |a_panels: packed::View<'_, i8, MR, PACK>, a_block_base| { + let on_a_panels = |a_panels: packed::View<'_, A::AElem, MR, PACK>, a_block_base| { let on_b_panels = |b_panels: unpacked::View<'_, i8>, _| { // SAFETY: By class invariant, `b_panels.k()` is equal to `self.k`. - let b_panels = + let (b_panels, side) = unsafe { self.arch.prepare(b_panels, &mut b_scratch, self.k) }; let panel_kernel = - |a_panel: packed::Panel<'_, i8, MR, PACK>, a_block_offset| { + |a_panel: packed::Panel<'_, A::AElem, MR, PACK>, a_block_offset| { // If we are in the very last block and we need to sub-fill, do // that. Otherwise, reference the output in place. let a_block = a_block_base + a_block_offset; @@ -192,14 +209,15 @@ where // run the kernel // // SAFETY: By class invariant, `a_panel.k()` and `b_panels.k()` - // are both equal to `self.k`. - let mut kernel = unsafe { - PanelKernel::new(self.arch, a_panel, b_panels, c, self.k) + // are both equal to `self.k`, and `side` was returned with + // `b_panels`. + let kernel = unsafe { + PanelKernel::new(self.arch, a_panel, b_panels, side, c, self.k) }; - driver::PanelKernel::panel_kernel(&mut kernel); - - let c_final = kernel.take(); + // Re-enter `arch` so the kernel keeps the target features of + // `A` even if the closures above are not inlined. + let c_final = self.arch.run(kernel); // Put back `C`. if handling_tail { @@ -249,30 +267,40 @@ where pub(crate) trait PrepareB: Copy { type Elem: Copy; + /// The element type of the packed `a`, which callers convert from `i8` when packing. + type AElem: Copy; + + type Scratch: Default; + /// Return `b` as [`Self::Elem`], using `scratch` as storage if a conversion is needed. /// + /// Also return the per-column data read by [`ExtraWide::init`], which is empty for + /// architectures that do not need any. + /// /// # Safety /// /// `b.k()` must be equal to `k`. unsafe fn prepare<'a>( self, b: unpacked::View<'a, i8>, - scratch: &'a mut Vec, + scratch: &'a mut Self::Scratch, k: DimK, - ) -> unpacked::View<'a, Self::Elem>; + ) -> (unpacked::View<'a, Self::Elem>, Slice<'a, i32>); } impl PrepareB for Scalar { type Elem = i8; + type AElem = i8; + type Scratch = (); #[inline(always)] unsafe fn prepare<'a>( self, b: unpacked::View<'a, i8>, - _: &'a mut Vec, + _: &'a mut (), _: DimK, - ) -> unpacked::View<'a, i8> { - b + ) -> (unpacked::View<'a, i8>, Slice<'a, i32>) { + (b, Slice::new(&[])) } } @@ -286,8 +314,9 @@ where A: PrepareB, { arch: A, - a: packed::Panel<'a, i8, MR, PACK>, + a: packed::Panel<'a, A::AElem, MR, PACK>, b: unpacked::View<'a, A::Elem>, + side: Slice<'a, i32>, c: [i32; MR], k: DimK, } @@ -300,21 +329,41 @@ where /// /// # Safety /// - /// Bounds `a.k()` and `b.k()` must both be equal to `k`. + /// Bounds `a.k()` and `b.k()` must both be equal to `k`, and `side` must be returned with + /// `b` by [`PrepareB::prepare`]. pub(super) unsafe fn new( arch: A, - a: packed::Panel<'a, i8, MR, PACK>, + a: packed::Panel<'a, A::AElem, MR, PACK>, b: unpacked::View<'a, A::Elem>, + side: Slice<'a, i32>, c: [i32; MR], k: DimK, ) -> Self { bounds::check_eq!(a.k(), k); bounds::check_eq!(b.k(), k); - Self { arch, a, b, c, k } + Self { + arch, + a, + b, + side, + c, + k, + } } +} - pub(super) fn take(self) -> [i32; MR] { +/// Unlike the closure impl of [`Target`], this is `#[inline(always)]`, so the kernel +/// reliably inherits the target features of `A`. +impl<'a, A, const MR: usize, const NR: usize, const PACK: usize> Target + for PanelKernel<'a, A, MR, NR, PACK> +where + A: PrepareB + Architecture, + PanelKernel<'a, A, MR, NR, PACK>: driver::PanelKernel, +{ + #[inline(always)] + fn run(mut self, _: A) -> [i32; MR] { + driver::PanelKernel::panel_kernel(&mut self); self.c } } @@ -323,9 +372,13 @@ where /// /// This is needed to ensure the visitor body is inlined to inherit target features. #[derive(Debug)] -struct Visitor<'a, A, const MR: usize, const NR: usize, const PACK: usize> { +struct Visitor<'a, A, const MR: usize, const NR: usize, const PACK: usize> +where + A: PrepareB, +{ arch: A, - a: packed::Panel<'a, i8, MR, PACK>, + a: packed::Panel<'a, A::AElem, MR, PACK>, + side: Slice<'a, i32>, c: &'a mut [i32; MR], k: DimK, } @@ -337,10 +390,12 @@ where for<'a> MicroKernel<'a, A, MR, NR, PACK>: driver::MicroKernel, { #[inline(always)] - fn visit(&mut self, b: unpacked::Panel<'_, A::Elem, NR>, _: usize) { + fn visit(&mut self, b: unpacked::Panel<'_, A::Elem, NR>, start: usize) { // SAFETY: This is only used on contexts where `self.a.k()`, `b.k()`, and `self.k` - // are all equal. - let mut micro = unsafe { MicroKernel::new(self.arch, self.a, b, self.c, self.k) }; + // are all equal, and where `b` is the panel at column `start` of the view returned + // with `self.side`. + let mut micro = + unsafe { MicroKernel::new(self.arch, self.a, b, self.side, start, self.c, self.k) }; driver::MicroKernel::micro_kernel(&mut micro); } } @@ -356,6 +411,7 @@ macro_rules! panel_kernel { let on_b_panels = Visitor { arch: self.arch, a: self.a, + side: self.side, c: &mut self.c, k: self.k, }; @@ -369,12 +425,15 @@ macro_rules! panel_kernel { const { assert!($ns < $nr) }; if let Some(b_panel) = b_tail.try_as_panel::<$ns>() { // SAFETY: By class invariant, `self.a.k()` and `self.b.k()` - // are equal to `self.k`. + // are equal to `self.k`, and `self.side` was returned with + // `self.b`, whose column `b_tail.start()` starts `b_panel`. let mut micro = unsafe { MicroKernel::new( self.arch, self.a, b_panel, + self.side, + b_tail.start(), &mut self.c, self.k, ) @@ -397,14 +456,17 @@ panel_kernel!(Scalar, 8, 2, 2, [1]); /// # Class Invariants /// -/// `a.k()` and `b.k()` are equal to `k`. +/// `a.k()` and `b.k()` are equal to `k`, and `b` is the panel at column `start` of the view +/// returned with `side` by [`PrepareB::prepare`]. struct MicroKernel<'a, A, const MR: usize, const NR: usize, const PACK: usize> where A: PrepareB, { arch: A, - a: packed::Panel<'a, i8, MR, PACK>, + a: packed::Panel<'a, A::AElem, MR, PACK>, b: unpacked::Panel<'a, A::Elem, NR>, + side: Slice<'a, i32>, + start: usize, c: &'a mut [i32; MR], k: DimK, } @@ -415,18 +477,29 @@ where { /// # Safety /// - /// Bounds `a.k()` and `b.k()` must be equal to `k`. + /// Bounds `a.k()` and `b.k()` must be equal to `k`, and `b` must be the panel at column + /// `start` of the view returned with `side` by [`PrepareB::prepare`]. unsafe fn new( arch: A, - a: packed::Panel<'a, i8, MR, PACK>, + a: packed::Panel<'a, A::AElem, MR, PACK>, b: unpacked::Panel<'a, A::Elem, NR>, + side: Slice<'a, i32>, + start: usize, c: &'a mut [i32; MR], k: DimK, ) -> Self { bounds::check_eq!(a.k(), k); bounds::check_eq!(b.k(), k); - Self { arch, a, b, c, k } + Self { + arch, + a, + b, + side, + start, + c, + k, + } } } @@ -480,12 +553,15 @@ unsafe fn accumulate_pack( wide: W, - a: packed::Panel<'_, i8, MR, PACK>, + a: packed::Panel<'_, W::AElem, MR, PACK>, b: unpacked::Panel<'_, W::Elem, NR>, + side: Slice<'_, i32>, + start: usize, c: &mut [i32; MR], k: DimK, ) where @@ -499,7 +575,9 @@ unsafe fn micro_kernel( let ap = a.as_ptr(); let bp = b.as_ptr(); - let mut acc = [wide.default(); NR]; + // SAFETY: By preconditions, `start + j` is a column of the view returned with `side` + // for every `j < NR`. + let mut acc: [W::Acc; NR] = core::array::from_fn(|j| unsafe { wide.init(side, start + j) }); let astride = a.pack_stride(); let bstride = b.stride(k); @@ -564,8 +642,14 @@ macro_rules! micro_kernel { impl driver::MicroKernel for MicroKernel<'_, $arch, $mr, $nr, $pack> { #[inline(always)] fn micro_kernel(&mut self) { - // SAFETY: By class invariant, `self.a.k()` and `self.b.k()` equal `self.k`. - unsafe { micro_kernel(self.arch, self.a, self.b, self.c, self.k) } + // SAFETY: By class invariant, `self.a.k()` and `self.b.k()` equal `self.k`, + // and `self.b` is the panel at column `self.start` of the view returned with + // `self.side`. + unsafe { + micro_kernel( + self.arch, self.a, self.b, self.side, self.start, self.c, self.k, + ) + } } } }; @@ -584,9 +668,14 @@ trait ExtraWide: PrepareB { /// # Safety /// /// `slice.len()` must be exactly `ELEMENTS * PACK`. - unsafe fn load(self, slice: Slice<'_, i8>) -> Self::Wide; + unsafe fn load(self, slice: Slice<'_, Self::AElem>) -> Self::Wide; - fn default(self) -> Self::Acc; + /// Return the starting accumulator for `column`. + /// + /// # Safety + /// + /// `column` must be a column of the view returned with `side` by [`PrepareB::prepare`]. + unsafe fn init(self, side: Slice<'_, i32>, column: usize) -> Self::Acc; /// Broadcast the `PACK` elements at `b`, treating those past `valid` as zero. /// @@ -606,7 +695,7 @@ impl ExtraWide<8, 2> for Scalar { type Acc = i32x8; #[inline(always)] - fn default(self) -> Self::Acc { + unsafe fn init(self, _: Slice<'_, i32>, _: usize) -> Self::Acc { SIMDVector::default(self) } @@ -651,13 +740,18 @@ impl ExtraWide<8, 2> for Scalar { mod x86_64 { use super::*; - use diskann_wide::arch::x86_64::V3; + use diskann_wide::{ + SIMDSumTree, + arch::x86_64::{V3, V4}, + }; use crate::matrix_kernels::util::{Convert, Converter}; panel_kernel!(V3, 16, 6, 2, [1, 2, 3, 4, 5]); + panel_kernel!(V4, 32, 6, 4, [1, 2, 3, 4, 5]); micro_kernel!(V3, 16, 2, { 6, 5, 4, 3, 2, 1 }); + micro_kernel!(V4, 32, 4, { 6, 5, 4, 3, 2, 1 }); //----------// // PrepareB // @@ -665,6 +759,8 @@ mod x86_64 { impl PrepareB for V3 { type Elem = i16; + type AElem = i16; + type Scratch = Vec; #[inline(always)] unsafe fn prepare<'a>( @@ -672,7 +768,7 @@ mod x86_64 { b: unpacked::View<'a, i8>, scratch: &'a mut Vec, k: DimK, - ) -> unpacked::View<'a, i16> { + ) -> (unpacked::View<'a, i16>, Slice<'a, i32>) { // SAFETY: Inherited from caller. let from = unsafe { b.as_std_slice(k) }; @@ -680,8 +776,96 @@ mod x86_64 { Converter::new(self).convert(scratch, from); // SAFETY: `scratch` has length `b.extent() * k`. - unsafe { unpacked::View::new(Slice::new(scratch), b.extent(), k) } + let b = unsafe { unpacked::View::new(Slice::new(scratch), b.extent(), k) }; + + (b, Slice::new(&[])) + } + } + + /// `a` is packed as `x + 128` in `u8` for the unsigned by signed dot product, so each + /// accumulator starts at `-128` times the sum of its column of `b` to cancel the shift. + impl PrepareB for V4 { + type Elem = i8; + type AElem = u8; + type Scratch = Vec; + + #[inline(always)] + unsafe fn prepare<'a>( + self, + b: unpacked::View<'a, i8>, + scratch: &'a mut Vec, + k: DimK, + ) -> (unpacked::View<'a, i8>, Slice<'a, i32>) { + // SAFETY: Inherited from caller. + let from = unsafe { b.as_std_slice(k) }; + let k = k.value().get(); + + scratch.resize(b.extent().get(), 0); + + let (to_blocks, to_tail) = scratch.as_chunks_mut::<16>(); + let (from_blocks, from_tail) = from.split_at(16 * k * to_blocks.len()); + + // Plain loops rather than closures, which may not inherit the target features. + for (to, cols) in std::iter::zip(to_blocks, from_blocks.chunks_exact(16 * k)) { + for (to, sum) in std::iter::zip(to, column_sums::<16>(self, cols)) { + *to = -128 * sum; + } + } + for (to, col) in std::iter::zip(to_tail, from_tail.chunks_exact(k)) { + let [sum] = column_sums::<1>(self, col); + *to = -128 * sum; + } + + (b, Slice::new(scratch)) + } + } + + /// Sum each of the `N` equal length columns in `cols`, walking them together so that + /// their dot products overlap. + #[inline(always)] + fn column_sums(arch: V4, cols: &[i8]) -> [i32; N] { + let k = cols.len() / N; + let cols = Slice::new(cols); + + let ones = u8x64::::splat(arch, 1); + let mut sums: [i32x16; N] = [SIMDVector::default(arch); N]; + + let full = k / 64; + for chunk in 0..full { + for (j, sum) in sums.iter_mut().enumerate() { + // SAFETY: Since `j < N`, `j * k + 64 * (chunk + 1) <= (j + 1) * k` is at most + // `cols.len()`. + let x: i8x64 = unsafe { + SIMDVector::load_simd( + arch, + cols.add(Elements::new(j * k + 64 * chunk)).as_ptr(), + ) + }; + *sum = sum.dot_simd(ones, x); + } + } + + let rest = k - 64 * full; + if rest != 0 { + for (j, sum) in sums.iter_mut().enumerate() { + // SAFETY: Since `j < N`, `j * k + 64 * full + rest == (j + 1) * k` is at most + // `cols.len()`. + let x: i8x64 = unsafe { + SIMDVector::load_simd_first( + arch, + cols.add(Elements::new(j * k + 64 * full)).as_ptr(), + rest, + ) + }; + *sum = sum.dot_simd(ones, x); + } + } + + let mut out = [0; N]; + for (out, sum) in std::iter::zip(&mut out, &sums) { + *out = sum.sum_tree(); } + out } //-----------// @@ -694,24 +878,22 @@ mod x86_64 { type Acc = [i32x8; 2]; #[inline(always)] - fn default(self) -> Self::Acc { + unsafe fn init(self, _: Slice<'_, i32>, _: usize) -> Self::Acc { [SIMDVector::default(self), SIMDVector::default(self)] } #[inline(always)] - unsafe fn load(self, slice: Slice<'_, i8>) -> Self::Wide { + unsafe fn load(self, slice: Slice<'_, i16>) -> Self::Wide { bounds::check_eq!(slice.len(), 32); // SAFETY: Since `slice.len()` must be 32, the pointer offset and 16-wide SIMD loads // are valid. - let bytes: [i8x16; 2] = unsafe { + unsafe { [ SIMDVector::load_simd(self, slice.as_ptr()), SIMDVector::load_simd(self, slice.add(Elements::new(16)).as_ptr()), ] - }; - - bytes.map(Self::Splat::from) + } } #[inline(always)] @@ -752,6 +934,72 @@ mod x86_64 { } } } + + impl ExtraWide<32, 4> for V4 { + type Wide = [u8x64; 2]; + type Splat = i8x64; + type Acc = [i32x16; 2]; + + #[inline(always)] + unsafe fn init(self, side: Slice<'_, i32>, column: usize) -> Self::Acc { + // SAFETY: `prepare` returns one value per column in `side`, and by preconditions + // `column` is one of those columns. + let start = unsafe { *side.add(Elements::new(column)).as_unit().as_ref() }; + [SIMDVector::splat(self, start); 2] + } + + #[inline(always)] + unsafe fn load(self, slice: Slice<'_, u8>) -> Self::Wide { + bounds::check_eq!(slice.len(), 128); + + // SAFETY: Since `slice.len()` must be 128, the pointer offset and 64-wide SIMD + // loads are valid. + unsafe { + [ + SIMDVector::load_simd(self, slice.as_ptr()), + SIMDVector::load_simd(self, slice.add(Elements::new(64)).as_ptr()), + ] + } + } + + #[inline(always)] + unsafe fn splat(self, b: Slice<'_, i8>, valid: usize) -> Self::Splat { + // SAFETY: Inherited from caller. + let quad = i8_quad(unsafe { group(b, valid) }); + u32x16::::splat(self, quad).reinterpret_simd() + } + + #[inline(always)] + fn dot(a: Self::Wide, b: Self::Splat, acc: Self::Acc) -> Self::Acc { + core::array::from_fn(|i| acc[i].dot_simd(a[i], b)) + } + + #[inline(always)] + fn max(lhs: Self::Acc, rhs: Self::Acc) -> Self::Acc { + core::array::from_fn(|i| lhs[i].max_simd(rhs[i])) + } + + #[inline(always)] + fn max_into(self, lhs: Self::Acc, into: &mut [i32; 32]) { + // SAFETY: Since `into.len()` is 32, the pointer offset and 16-wide SIMD loads are + // valid. + let previous: Self::Acc = unsafe { + [ + SIMDVector::load_simd(self, into.as_ptr()), + SIMDVector::load_simd(self, into.as_ptr().add(16)), + ] + }; + + let max = Self::max(lhs, previous); + + // SAFETY: Since `into.len()` is 32, the pointer offset and 16-wide SIMD stores + // are valid. + unsafe { + max[0].store_simd(into.as_mut_ptr()); + max[1].store_simd(into.as_mut_ptr().add(16)); + } + } + } } #[cfg(target_arch = "aarch64")] @@ -766,29 +1014,23 @@ mod aarch64 { micro_kernel!(Neon, 8, 4, { 6, 5, 4, 3, 2, 1 }); micro_kernel!(Neon, 16, 4, { 6, 5, 4, 3, 2, 1 }); - /// Pack a `PACK = 4` group into the little-endian byte quad consumed by `sdot`. - /// - /// Broadcasting through `u32` lowers to a single `ld1r`. - #[inline(always)] - fn i8_quad(group: [i8; 4]) -> u32 { - u32::from_le_bytes(group.map(|x| x as u8)) - } - //----------// // PrepareB // //----------// impl PrepareB for Neon { type Elem = i8; + type AElem = i8; + type Scratch = (); #[inline(always)] unsafe fn prepare<'a>( self, b: unpacked::View<'a, i8>, - _: &'a mut Vec, + _: &'a mut (), _: DimK, - ) -> unpacked::View<'a, i8> { - b + ) -> (unpacked::View<'a, i8>, Slice<'a, i32>) { + (b, Slice::new(&[])) } } @@ -802,7 +1044,7 @@ mod aarch64 { type Acc = [i32x4; 2]; #[inline(always)] - fn default(self) -> Self::Acc { + unsafe fn init(self, _: Slice<'_, i32>, _: usize) -> Self::Acc { [SIMDVector::default(self), SIMDVector::default(self)] } @@ -865,7 +1107,7 @@ mod aarch64 { type Acc = [i32x4; 4]; #[inline(always)] - fn default(self) -> Self::Acc { + unsafe fn init(self, _: Slice<'_, i32>, _: usize) -> Self::Acc { [SIMDVector::default(self); 4] } @@ -929,13 +1171,56 @@ mod tests { use rand::{SeedableRng, rngs::StdRng}; #[cfg(target_arch = "x86_64")] - use diskann_wide::arch::x86_64::V3; + use diskann_wide::arch::x86_64::{V3, V4}; #[cfg(target_arch = "aarch64")] use diskann_wide::arch::aarch64::Neon; + use diskann_utils::views::Matrix; + use crate::{matrix_kernels::maxsim, multi_vector::BlockTransposed}; + /// Convert an element of `a` into [`PrepareB::AElem`] before packing. + trait ConvertA: PrepareB { + fn convert_a(x: i8) -> Self::AElem; + } + + impl ConvertA for Scalar { + fn convert_a(x: i8) -> i8 { + x + } + } + + #[cfg(target_arch = "x86_64")] + impl ConvertA for V3 { + fn convert_a(x: i8) -> i16 { + i16::from(x) + } + } + + #[cfg(target_arch = "x86_64")] + impl ConvertA for V4 { + fn convert_a(x: i8) -> u8 { + (x as u8) ^ 0x80 + } + } + + #[cfg(target_arch = "aarch64")] + impl ConvertA for Neon { + fn convert_a(x: i8) -> i8 { + x + } + } + + fn pack_a( + a: &Matrix, + ) -> BlockTransposed + where + A: ConvertA, + { + BlockTransposed::from_matrix_view(a.map(|x| A::convert_a(*x)).as_view()) + } + ///////////////// // MicroKernel // ///////////////// @@ -946,20 +1231,20 @@ mod tests { rng: &mut impl rand::Rng, ctx: std::fmt::Arguments<'_>, ) where - A: PrepareB, + A: ConvertA, for<'a> MicroKernel<'a, A, MR, NR, PACK>: driver::MicroKernel, { let (ref_a, ref_b, ref_c) = maxsim::test::generate_i8(MR, k.value().get(), NR, rng); // From the reference problem, `ref_a` needs to be packed and `ref_b` transposed to // get them into the desired format. - let a_bt = BlockTransposed::::from_matrix_view(ref_a.as_view()); + let a_bt = pack_a::(&ref_a); let ref_b = ref_b.transpose(); - let mut scratch = Vec::new(); + let mut scratch = A::Scratch::default(); // SAFETY: Test builds will verify the bounds we passed. - let b = unsafe { + let (b, side) = unsafe { arch.prepare( unpacked::View::from_matrix_view(ref_b.as_view()).unwrap(), &mut scratch, @@ -977,6 +1262,8 @@ mod tests { arch, packed::Panel::new(Slice::new(a_bt.as_slice()), k), unpacked::Panel::new(Slice::new(b.as_std_slice(k)), k), + side, + 0, &mut c, k, ) @@ -985,7 +1272,7 @@ mod tests { driver::MicroKernel::micro_kernel(&mut kernel); assert_eq!(&*ref_c, kernel.c, "{ctx}"); - // Try again - but this time use a value that is much bigger than the what should + // Try again - but this time use a value that is much bigger than what should // be generated by the test problem. // // This checks that we don't just overwrite existing contents. @@ -1047,6 +1334,15 @@ mod tests { 16 => { 6, 5, 4, 3, 2, 1 }, ); + #[cfg(target_arch = "x86_64")] + test_micro_kernel!( + test_micro_kernel_v4, + V4::new_checked_miri(), + 0x2c9e5b7a41f0d863, + 4, + 32 => { 6, 5, 4, 3, 2, 1 }, + ); + #[cfg(target_arch = "aarch64")] test_micro_kernel!( test_micro_kernel_neon, @@ -1071,7 +1367,7 @@ mod tests { rng: &mut impl rand::Rng, ctx: std::fmt::Arguments<'_>, ) where - A: PrepareB, + A: ConvertA, for<'a> PanelKernel<'a, A, MR, NR, PACK>: driver::PanelKernel, { for blocks in 0..4 { @@ -1084,12 +1380,21 @@ mod tests { let (ref_a, ref_b, ref_c) = maxsim::test::generate_i8(MR, k.value().get(), cols, rng); - let a_bt = BlockTransposed::::from_matrix_view(ref_a.as_view()); + let a_bt = pack_a::(&ref_a); let ref_b = ref_b.transpose(); let extent = NonZeroUsize::new(cols).unwrap(); - let mut scratch = Vec::new(); + let mut scratch = A::Scratch::default(); + + // SAFETY: Test builds will verify the bounds we passed. + let (b, side) = unsafe { + arch.prepare( + unpacked::View::new(Slice::new(ref_b.as_slice()), extent, k), + &mut scratch, + k, + ) + }; let c = [i32::MIN; MR]; @@ -1098,11 +1403,8 @@ mod tests { PanelKernel::new( arch, packed::Panel::new(Slice::new(a_bt.as_slice()), k), - arch.prepare( - unpacked::View::new(Slice::new(ref_b.as_slice()), extent, k), - &mut scratch, - k, - ), + b, + side, c, k, ) @@ -1111,7 +1413,7 @@ mod tests { driver::PanelKernel::panel_kernel(&mut kernel); assert_eq!(&*ref_c, kernel.c, "{ctx}"); - // Try again - but this time use a value that is much bigger than the what + // Try again - but this time use a value that is much bigger than what // should be generated by the test problem. // // This checks that we don't just overwrite existing contents. @@ -1172,6 +1474,14 @@ mod tests { (16, 6, 2), ); + #[cfg(target_arch = "x86_64")] + test_panel_kernel!( + test_panel_kernel_v4, + V4::new_checked_miri(), + 0x9f3e7ab4c05d1268, + (32, 6, 4), + ); + #[cfg(target_arch = "aarch64")] test_panel_kernel!( test_panel_kernel_neon, @@ -1189,10 +1499,20 @@ mod tests { arch: A, rng: &mut impl rand::Rng, ) where - A: PrepareB, + A: ConvertA, for<'a> Driver<'a, A, MR, NR, PACK>: driver::Drive, { - let cases = maxsim::test::packed_x_unpacked_test_dims(MR, NR); + let mut cases = maxsim::test::packed_x_unpacked_test_dims(MR, NR); + + // Sub-views of 17 and 16 columns reach the 16-column blocks in `PrepareB for V4`. + cases.push(maxsim::test::TestDims { + a_panels_per_tile: 1, + total_a_rows: MR + 1, + b_cols_per_tile: 17, + total_b_cols: 33, + k: 5, + }); + for case in cases { let maxsim::test::TestDims { a_panels_per_tile, @@ -1208,7 +1528,7 @@ mod tests { maxsim::test::generate_i8(total_a_rows, k.value().get(), total_b_cols, rng); // Massage the input data in the form needed by the kernel. - let a_bt = BlockTransposed::::from_matrix_view(ref_a.as_view()); + let a_bt = pack_a::(&ref_a); let b = ref_b.transpose(); let mut c = vec![i32::MAX; a_bt.nrows()]; @@ -1271,6 +1591,14 @@ mod tests { (16, 6, 2), ); + #[cfg(target_arch = "x86_64")] + test_driver!( + test_driver_v4, + V4::new_checked_miri(), + 0x63c8ed19f4720ab5, + (32, 6, 4), + ); + #[cfg(target_arch = "aarch64")] test_driver!( test_driver_neon, @@ -1279,4 +1607,43 @@ mod tests { (8, 6, 4), (16, 6, 4), ); + + ////////////// + // PrepareB // + ////////////// + + #[cfg(target_arch = "x86_64")] + #[test] + fn test_prepare_v4() { + use crate::matrix_kernels::test_util::TestDistr; + + if let Some(arch) = V4::new_checked_miri() { + let mut rng = StdRng::seed_from_u64(0x5d0b8e3f7a1c6942); + let mut scratch = Vec::new(); + + for k in [1, 63, 64, 65, 128, 131] { + for n in [1, 15, 16, 17, 33] { + let b = TestDistr::matrix::(n, k, &mut rng); + let expected: Vec = b + .row_iter() + .map(|col| -128 * col.iter().map(|&x| i32::from(x)).sum::()) + .collect(); + + let dim = DimK::new(NonZeroUsize::new(k).unwrap()); + + // SAFETY: Test builds will verify the bounds we passed. + unsafe { + let (view, side) = arch.prepare( + unpacked::View::from_matrix_view(b.as_view()).unwrap(), + &mut scratch, + dim, + ); + + assert_eq!(view.as_std_slice(dim), b.as_slice(), "k = {k}, n = {n}"); + assert_eq!(side.as_std_slice(n), &*expected, "k = {k}, n = {n}"); + } + } + } + } + } } diff --git a/diskann-quantization/src/matrix_kernels/util.rs b/diskann-quantization/src/matrix_kernels/util.rs index a5c79b0d5a..1c2956bbb0 100644 --- a/diskann-quantization/src/matrix_kernels/util.rs +++ b/diskann-quantization/src/matrix_kernels/util.rs @@ -111,6 +111,48 @@ macro_rules! impl_loadstore { } } }; + ($T:ty, 32, [$wide:ident; 2], $arch:ty) => { + impl LoadStore<$T, 32> for $arch { + #[inline(always)] + fn load(self, src: &[$T]) -> [$T; 32] { + use diskann_wide::{LoHi, SIMDVector}; + diskann_wide::alias!(wide = <$arch>::$wide); + + // SAFETY: Loading the first `src.len().min(16)` elements from `src` is valid. + let lo = unsafe { wide::load_simd_first(self, src.as_ptr(), src.len()) }.to_array(); + + // SAFETY: This only reads `src.len() - 16` values if `src.len()` exceeds 16. + let hi = unsafe { + wide::load_simd_first( + self, + src.as_ptr().wrapping_offset(16), + src.len().saturating_sub(16), + ) + } + .to_array(); + + LoHi::new(lo, hi).join() + } + + #[inline(always)] + fn store(self, v: [$T; 32], dst: &mut [$T]) { + use diskann_wide::{LoHi, SIMDVector, SplitJoin}; + diskann_wide::alias!(wide = <$arch>::$wide); + + let LoHi { lo, hi } = v.split(); + + // SAFETY: Storing the first `dst.len().min(16)` elements to `dst` is valid. + unsafe { wide::from_array(self, lo).store_simd_first(dst.as_mut_ptr(), dst.len()) }; + + if let Some(rest) = dst.len().checked_sub(16) { + // SAFETY: This only writes `dst.len() - 16` values if `dst.len()` exceeds 16. + unsafe { + wide::from_array(self, hi).store_simd_first(dst.as_mut_ptr().add(16), rest) + }; + } + } + } + }; } #[cfg(target_arch = "x86_64")] @@ -125,47 +167,8 @@ mod x86_64 { impl_loadstore!(f32, 8, f32x8, V4); impl_loadstore!(f32, 16, f32x16, V4); - - impl LoadStore for V4 { - #[inline(always)] - fn load(self, src: &[f32]) -> [f32; 32] { - use diskann_wide::{LoHi, SIMDVector}; - diskann_wide::alias!(wide = ::f32x16); - - // SAFETY: Loading the first `src.len().min(16)` elements from `src` is valid. - let lo = unsafe { wide::load_simd_first(self, src.as_ptr(), src.len()) }.to_array(); - - // SAFETY: This only reads `src.len() - 16` values if `src.len()` exceeds 16. - let hi = unsafe { - wide::load_simd_first( - self, - src.as_ptr().wrapping_offset(16), - src.len().saturating_sub(16), - ) - } - .to_array(); - - LoHi::new(lo, hi).join() - } - - #[inline(always)] - fn store(self, v: [f32; 32], dst: &mut [f32]) { - use diskann_wide::{LoHi, SIMDVector, SplitJoin}; - diskann_wide::alias!(wide = ::f32x16); - - let LoHi { lo, hi } = v.split(); - - // SAFETY: Storing the first `dst.len().min(16)` elements to `dst` is valid. - unsafe { wide::from_array(self, lo).store_simd_first(dst.as_mut_ptr(), dst.len()) }; - - if let Some(rest) = dst.len().checked_sub(16) { - // SAFETY: This only writes if `dst.len() - 16` values if `dst.len()` exceeds 16. - unsafe { - wide::from_array(self, hi).store_simd_first(dst.as_mut_ptr().add(16), rest) - }; - } - } - } + impl_loadstore!(f32, 32, [f32x16; 2], V4); + impl_loadstore!(i32, 32, [i32x16; 2], V4); } #[cfg(target_arch = "aarch64")] @@ -382,6 +385,7 @@ mod test { test_load_store_v4, V4::new_checked_miri(), f32 => { 8, 16, 32 }, + i32 => { 32 }, ); #[cfg(target_arch = "aarch64")] diff --git a/diskann-quantization/src/multi_vector/distance/factory.rs b/diskann-quantization/src/multi_vector/distance/factory.rs index e160406f23..20e17fff78 100644 --- a/diskann-quantization/src/multi_vector/distance/factory.rs +++ b/diskann-quantization/src/multi_vector/distance/factory.rs @@ -176,10 +176,11 @@ where } } -impl MaxSimKernel - for Prepared, NR> +impl MaxSimKernel + for Prepared, NR> where - A: mk::maxsim::packed_i8_x_unpacked_i8::PrepareB + Architecture, + Q: Copy + Send + Sync + std::fmt::Debug, + A: mk::maxsim::packed_i8_x_unpacked_i8::PrepareB + Architecture, for<'a> mk::maxsim::packed_i8_x_unpacked_i8::Driver<'a, A, GROUP, NR, PACK>: mk::Drive, { fn nrows(&self) -> usize { @@ -482,7 +483,10 @@ impl> diskann_wide::arch::Target1 { fn run(self, arch: V3, query: MatRef<'_, Standard>) -> E::Output { - let prepared = BlockTransposed::::from_matrix_view(query.as_matrix_view()); + // `PrepareB for V3` expects the query widened to `i16`. + let prepared = BlockTransposed::::from_matrix_view( + query.as_matrix_view().map(|v| i16::from(*v)).as_view(), + ); self.0.erase(Prepared { arch, prepared, @@ -496,8 +500,15 @@ impl> diskann_wide::arch::Target1 { fn run(self, arch: V4, query: MatRef<'_, Standard>) -> E::Output { - // V4 retargets to V3 until the VNNI kernel lands. - diskann_wide::arch::Target1::::run(self, V3::from(arch), query) + // `PrepareB for V4` expects the query packed as `x + 128` in `u8`. + let prepared = BlockTransposed::::from_matrix_view( + query.as_matrix_view().map(|v| (*v as u8) ^ 0x80).as_view(), + ); + self.0.erase(Prepared { + arch, + prepared, + _packing: Pack::<6>, + }) } } @@ -749,8 +760,7 @@ mod tests { use diskann_vector::DistanceFunctionMut; /// Local helper trait — picks a sane test value of `T` from an `f32` - /// so both `f32` and `half::f16` parameterizations share the same data - /// generator. + /// so every element type shares the same data generator. trait FromF32 { fn from_f32(v: f32) -> Self; } @@ -774,7 +784,7 @@ mod tests { } /// Projects a kernel score onto the `f32` distance the fallback path - /// produces, so both parameterizations share the same assertions. + /// produces, so every element type shares the same assertions. trait ScoreAsF32: MaxSimElement { fn score_as_f32(score: Self::Score) -> f32; } @@ -888,27 +898,15 @@ mod tests { } } - /// The `i8` reference path is an independent integer implementation - /// ([`FallbackKernel::max_sim_kernel_i8`]), so it needs its own guard; the - /// `f32`/`f16` reference paths share `max_sim_kernel` with the oracle and - /// would only be testing themselves. - /// - /// Every other ISA is reached via [`MaxSimIsa::Auto`] somewhere in the CI - /// matrix. Widen into a full sweep once V4 and Neon gain native `i8` - /// kernels instead of retargeting to V3 and Scalar. - #[test] - fn i8_reference_matches_oracle() { - for &(nq, nd, dim) in TEST_CASES { - let query_data = make_test_data::(nq * dim, dim, dim / 2); - let doc_data = make_test_data::(nd * dim, dim, dim); - - let query = make_mat(&query_data, nq, dim); - let doc = make_mat(&doc_data, nd, dim); + fn check_i8_isas(nq: usize, nd: usize, dim: usize, query_data: &[i8], doc_data: &[i8]) { + let query = make_mat(query_data, nq, dim); + let doc = make_mat(doc_data, nd, dim); - let mut expected = vec![0.0f32; nq]; - let _ = MaxSim::new(&mut expected).evaluate(QueryMatRef::from(query), doc); + let mut expected = vec![0.0f32; nq]; + let _ = MaxSim::new(&mut expected).evaluate(QueryMatRef::from(query), doc); - let kernel = build_max_sim::(MaxSimIsa::Reference, query, BoxErase).unwrap(); + for isa in [MaxSimIsa::Auto, MaxSimIsa::Reference] { + let kernel = build_max_sim::(isa, query, BoxErase).unwrap(); let mut scores = scores_buffer::(nq); kernel.compute_max_sim(doc, &mut scores).unwrap(); @@ -916,7 +914,7 @@ mod tests { let actual = ::score_as_f32(scores[i]); assert!( (actual - expected[i]).abs() < 1e-10, - "i8 reference MaxSim[{i}] mismatch for ({nq},{nd},{dim}): \ + "i8 {isa} MaxSim[{i}] mismatch for ({nq},{nd},{dim}): \ actual={actual}, expected={}", expected[i], ); @@ -924,6 +922,26 @@ mod tests { } } + #[test] + fn i8_isas_match_oracle() { + for &(nq, nd, dim) in TEST_CASES { + let query_data = make_test_data::(nq * dim, dim, dim / 2); + let doc_data = make_test_data::(nd * dim, dim, dim); + check_i8_isas(nq, nd, dim, &query_data, &doc_data); + } + + let full_range: &[(usize, usize, usize)] = if cfg!(miri) { + &[(5, 4, 64)] + } else { + &[(33, 13, 64), (70, 1000, 131)] + }; + for &(nq, nd, dim) in full_range { + let query_data: Vec = (0..nq * dim).map(|v| (37 * v) as i8).collect(); + let doc_data: Vec = (0..nd * dim).map(|v| (91 * v + 5) as i8).collect(); + check_i8_isas(nq, nd, dim, &query_data, &doc_data); + } + } + #[test] fn dimensions_f32() { let data = vec![1.0f32; 5 * 8]; From c91a7faf05fbe1c621531afb6ff61b897170d62a Mon Sep 17 00:00:00 2001 From: Suryansh Gupta Date: Sat, 3 Oct 2026 06:21:07 +0530 Subject: [PATCH 7/7] Address Review Comments --- .../maxsim/packed_i8_x_unpacked_i8.rs | 7 +- .../src/matrix_kernels/util.rs | 35 -------- .../src/multi_vector/distance/factory.rs | 87 ++++++++++++++++--- .../src/multi_vector/distance/mod.rs | 2 +- diskann-quantization/src/multi_vector/mod.rs | 4 +- diskann-wide/src/arch/aarch64/u32x4_.rs | 3 + 6 files changed, 82 insertions(+), 56 deletions(-) diff --git a/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs b/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs index 2a2dc50674..e5961dedf8 100644 --- a/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs +++ b/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs @@ -49,7 +49,6 @@ fn i16_pair([lo, hi]: [i16; 2]) -> u32 { /// products. /// /// Broadcasting through `u32` lowers to a single `ld1r` on Neon. -#[cfg(any(target_arch = "x86_64", target_arch = "aarch64"))] #[inline(always)] fn i8_quad(group: [i8; 4]) -> u32 { u32::from_le_bytes(group.map(|x| x as u8)) @@ -745,8 +744,6 @@ mod x86_64 { arch::x86_64::{V3, V4}, }; - use crate::matrix_kernels::util::{Convert, Converter}; - panel_kernel!(V3, 16, 6, 2, [1, 2, 3, 4, 5]); panel_kernel!(V4, 32, 6, 4, [1, 2, 3, 4, 5]); @@ -773,7 +770,9 @@ mod x86_64 { let from = unsafe { b.as_std_slice(k) }; scratch.resize(from.len(), 0); - Converter::new(self).convert(scratch, from); + for (to, from) in std::iter::zip(scratch.iter_mut(), from) { + *to = (*from).into(); + } // SAFETY: `scratch` has length `b.extent() * k`. let b = unsafe { unpacked::View::new(Slice::new(scratch), b.extent(), k) }; diff --git a/diskann-quantization/src/matrix_kernels/util.rs b/diskann-quantization/src/matrix_kernels/util.rs index 1c2956bbb0..cdb4ef434e 100644 --- a/diskann-quantization/src/matrix_kernels/util.rs +++ b/diskann-quantization/src/matrix_kernels/util.rs @@ -38,25 +38,6 @@ where } } -#[cfg(target_arch = "x86_64")] -impl Convert for Converter { - #[inline(always)] - fn convert(self, to: &mut [i16], from: &[i8]) { - use diskann_wide::{SIMDVector, arch::x86_64::V3}; - diskann_wide::alias!(i8s = ::i8x16); - diskann_wide::alias!(i16s = ::i16x16); - - debug_assert_eq!(to.len(), from.len(), "lengths must be equal"); - - let (to_chunks, to_tail) = to.as_chunks_mut::<16>(); - let (from_chunks, from_tail) = from.as_chunks::<16>(); - - std::iter::zip(to_chunks, from_chunks) - .for_each(|(to, from)| *to = i16s::from(i8s::from_array(self.0, *from)).to_array()); - std::iter::zip(to_tail, from_tail).for_each(|(to, from)| *to = (*from).into()); - } -} - ////////// // Load // ////////// @@ -288,22 +269,6 @@ mod test { #[cfg(target_arch = "aarch64")] use diskann_wide::arch::aarch64::Neon; - #[cfg(target_arch = "x86_64")] - #[test] - fn test_convert_i16_i8_v3() { - if let Some(arch) = V3::new_checked() { - for len in 0..40 { - let from: Vec = (0..len).map(|i| (37 * i) as i8).collect(); - let mut to = vec![0; len]; - - Converter::new(arch).convert(&mut to, &from); - - let expected: Vec = from.iter().map(|&x| x.into()).collect(); - assert_eq!(to, expected, "len = {len}"); - } - } - } - trait FromUsize { fn from_usize(v: usize) -> Self; } diff --git a/diskann-quantization/src/multi_vector/distance/factory.rs b/diskann-quantization/src/multi_vector/distance/factory.rs index 20e17fff78..064b9bb425 100644 --- a/diskann-quantization/src/multi_vector/distance/factory.rs +++ b/diskann-quantization/src/multi_vector/distance/factory.rs @@ -16,6 +16,7 @@ use diskann_wide::arch::Scalar; use diskann_wide::arch::aarch64::Neon; #[cfg(target_arch = "x86_64")] use diskann_wide::arch::x86_64::{V3, V4}; +use thiserror::Error; use super::fallback::FallbackKernel; use super::isa::{MaxSimIsa, NotSupported}; @@ -552,13 +553,15 @@ pub trait MaxSimElement: sealed::Sealed + Sized + Copy + Send + Sync + 'static { /// /// # Errors /// - /// Returns [`NotSupported`] when the requested ISA cannot run on this - /// build (e.g. AVX-512 unavailable; aarch64 on x86_64). + /// Returns [`BuildMaxSimError::NotSupported`] when the requested ISA cannot + /// run on this build (e.g. AVX-512 unavailable; aarch64 on x86_64), and + /// [`BuildMaxSimError::DimTooLarge`] when an `i8` query has more than + /// `131_071` dimensions, beyond which `i32` scores could overflow. fn build>( isa: MaxSimIsa, query: MatRef<'_, Standard>, erase: E, - ) -> Result; + ) -> Result; } impl sealed::Sealed for f32 {} @@ -573,7 +576,7 @@ impl MaxSimElement for f32 { isa: MaxSimIsa, query: MatRef<'_, Standard>, erase: E, - ) -> Result { + ) -> Result { match isa { MaxSimIsa::Auto => Ok(diskann_wide::arch::dispatch1_no_features( BuildAndErase(erase), @@ -600,7 +603,8 @@ impl MaxSimElement for f32 { MaxSimIsa::X86_64_V3 | MaxSimIsa::X86_64_V4 => Err(NotSupported { isa, reason: "x86_64 target only", - }), + } + .into()), #[cfg(target_arch = "aarch64")] MaxSimIsa::Neon => { let arch = Neon::new_checked().ok_or(NotSupported { @@ -613,7 +617,8 @@ impl MaxSimElement for f32 { MaxSimIsa::Neon => Err(NotSupported { isa, reason: "aarch64 target only", - }), + } + .into()), MaxSimIsa::Reference => { Ok(erase.erase(ReferenceKernel::new(query, reference_scores::))) } @@ -629,7 +634,7 @@ impl MaxSimElement for half::f16 { isa: MaxSimIsa, query: MatRef<'_, Standard>, erase: E, - ) -> Result { + ) -> Result { match isa { MaxSimIsa::Auto => Ok(diskann_wide::arch::dispatch1_no_features( BuildAndErase(erase), @@ -656,7 +661,8 @@ impl MaxSimElement for half::f16 { MaxSimIsa::X86_64_V3 | MaxSimIsa::X86_64_V4 => Err(NotSupported { isa, reason: "x86_64 target only", - }), + } + .into()), #[cfg(target_arch = "aarch64")] MaxSimIsa::Neon => { let arch = Neon::new_checked().ok_or(NotSupported { @@ -669,7 +675,8 @@ impl MaxSimElement for half::f16 { MaxSimIsa::Neon => Err(NotSupported { isa, reason: "aarch64 target only", - }), + } + .into()), MaxSimIsa::Reference => { Ok(erase.erase(ReferenceKernel::new(query, reference_scores::))) } @@ -677,6 +684,9 @@ impl MaxSimElement for half::f16 { } } +/// Largest `i8` dimension for which every inner product fits in `i32`. +const MAX_I8_DIM: usize = (i32::MAX / (128 * 128)) as usize; + impl MaxSimElement for i8 { type Score = i32; const NO_MATCH: i32 = i32::MAX; @@ -685,7 +695,14 @@ impl MaxSimElement for i8 { isa: MaxSimIsa, query: MatRef<'_, Standard>, erase: E, - ) -> Result { + ) -> Result { + if query.vector_dim() > MAX_I8_DIM { + return Err(BuildMaxSimError::DimTooLarge( + query.vector_dim(), + MAX_I8_DIM, + )); + } + match isa { MaxSimIsa::Auto => Ok(diskann_wide::arch::dispatch1_no_features( BuildAndErase(erase), @@ -712,7 +729,8 @@ impl MaxSimElement for i8 { MaxSimIsa::X86_64_V3 | MaxSimIsa::X86_64_V4 => Err(NotSupported { isa, reason: "x86_64 target only", - }), + } + .into()), #[cfg(target_arch = "aarch64")] MaxSimIsa::Neon => { let arch = Neon::new_checked().ok_or(NotSupported { @@ -725,7 +743,8 @@ impl MaxSimElement for i8 { MaxSimIsa::Neon => Err(NotSupported { isa, reason: "aarch64 target only", - }), + } + .into()), MaxSimIsa::Reference => { Ok(erase.erase(ReferenceKernel::new(query, reference_scores_i8))) } @@ -737,6 +756,16 @@ impl MaxSimElement for i8 { // Factory entry point. // ───────────────────────────────────────────────────────────────────────── +/// Error returned by [`build_max_sim`]. +#[derive(Debug, Clone, Copy, Error)] +#[non_exhaustive] +pub enum BuildMaxSimError { + #[error(transparent)] + NotSupported(#[from] NotSupported), + #[error("query-vector dim {0} exceeds the maximum of {1}")] + DimTooLarge(usize, usize), +} + /// Build a multi-vector MaxSim kernel for any [`MaxSimElement`] type. /// /// Thin wrapper over [`MaxSimElement::build`] so callers don't have to name @@ -744,12 +773,14 @@ impl MaxSimElement for i8 { /// /// # Errors /// -/// Returns [`NotSupported`] when the requested ISA cannot run on this build. +/// Returns [`BuildMaxSimError::NotSupported`] when the requested ISA cannot run +/// on this build, and [`BuildMaxSimError::DimTooLarge`] when an `i8` query has +/// more than `131_071` dimensions. pub fn build_max_sim>( isa: MaxSimIsa, query: MatRef<'_, Standard>, erase: E, -) -> Result { +) -> Result { T::build(isa, query, erase) } @@ -966,6 +997,34 @@ mod tests { assert_eq!(kernel.nrows(), 5); } + #[test] + fn i8_rejects_dim_too_large() { + let data = vec![0i8; 131_072]; + let query = make_mat(&data, 1, 131_072); + + for isa in [MaxSimIsa::Auto, MaxSimIsa::Reference] { + let err = build_max_sim::(isa, query, BoxErase).err(); + assert!( + matches!(err, Some(BuildMaxSimError::DimTooLarge(131_072, 131_071))), + "{isa:?}: expected DimTooLarge(131_072, 131_071), got {err:?}", + ); + } + } + + #[test] + #[cfg(not(miri))] + fn i8_max_dim_is_exact() { + let data = vec![i8::MIN; 131_071]; + let query = make_mat(&data, 1, 131_071); + + for isa in [MaxSimIsa::Auto, MaxSimIsa::Reference] { + let kernel = build_max_sim::(isa, query, BoxErase).unwrap(); + let mut scores = scores_buffer::(1); + kernel.compute_max_sim(query, &mut scores).unwrap(); + assert_eq!(scores[0], -131_071 * 16_384, "{isa:?}"); + } + } + fn check_size_mismatch(label: &str) where T: MaxSimElement + FromF32, diff --git a/diskann-quantization/src/multi_vector/distance/mod.rs b/diskann-quantization/src/multi_vector/distance/mod.rs index 24f056f99c..3801240902 100644 --- a/diskann-quantization/src/multi_vector/distance/mod.rs +++ b/diskann-quantization/src/multi_vector/distance/mod.rs @@ -50,7 +50,7 @@ mod kernel; mod max_sim; mod projected_eigen; -pub use factory::{MaxSimElement, build_max_sim}; +pub use factory::{BuildMaxSimError, MaxSimElement, build_max_sim}; pub use fallback::QueryMatRef; pub use isa::{MaxSimIsa, NotSupported}; pub use kernel::{BoxErase, Erase, MaxSimKernel}; diff --git a/diskann-quantization/src/multi_vector/mod.rs b/diskann-quantization/src/multi_vector/mod.rs index d4a684f75e..10a24ff621 100644 --- a/diskann-quantization/src/multi_vector/mod.rs +++ b/diskann-quantization/src/multi_vector/mod.rs @@ -57,8 +57,8 @@ pub(crate) mod matrix; pub use block_transposed::{BlockTransposed, BlockTransposedMut, BlockTransposedRef}; pub use distance::{ - BoxErase, Chamfer, Erase, MaxSim, MaxSimElement, MaxSimError, MaxSimIsa, MaxSimKernel, - NotSupported, ProjectedEigen, QueryMatRef, build_max_sim, + BoxErase, BuildMaxSimError, Chamfer, Erase, MaxSim, MaxSimElement, MaxSimError, MaxSimIsa, + MaxSimKernel, NotSupported, ProjectedEigen, QueryMatRef, build_max_sim, }; pub use matrix::{ Defaulted, LayoutError, Mat, MatMut, MatRef, NewCloned, NewMut, NewOwned, NewRef, Overflow, diff --git a/diskann-wide/src/arch/aarch64/u32x4_.rs b/diskann-wide/src/arch/aarch64/u32x4_.rs index 1e2cf56448..cefa6393c6 100644 --- a/diskann-wide/src/arch/aarch64/u32x4_.rs +++ b/diskann-wide/src/arch/aarch64/u32x4_.rs @@ -182,4 +182,7 @@ mod tests { // Reductions test_utils::ops::test_sumtree!(u32x4, 0xb9ac82ab23a855da, test_neon()); + + // Reinterprets + test_utils::ops::test_reinterpret!(u32x4 => i8x16, 0xbfad755f32d25e5c, test_neon()); }