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 402858a022..57ef92c4cb 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 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. /// /// # 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 pack 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,25 @@ 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 pack. + pub(in crate::matrix_kernels) const fn pack_stride(&self) -> Elements { + Elements::new(SZ * PACK) + } + + /// 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 packs(&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) } } @@ -292,7 +325,71 @@ mod tests { views::{Matrix, MatrixView}, }; - use crate::matrix_kernels::test_util::panic_message_for; + use crate::{matrix_kernels::test_util::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::from_fn(nrows, ncols, |_| { + value += 1.0; + value + }); + + 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.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}"); + + 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.element(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() { diff --git a/diskann-quantization/src/matrix_kernels/blocks/unpacked.rs b/diskann-quantization/src/matrix_kernels/blocks/unpacked.rs index 0ae7a28a66..c4c97b4df4 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/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..e5961dedf8 --- /dev/null +++ b/diskann-quantization/src/matrix_kernels/maxsim/packed_i8_x_unpacked_i8.rs @@ -0,0 +1,1648 @@ +/* + * 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 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. +//! +//! 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, Target}; +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!(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. +#[inline(always)] +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. +#[inline(always)] +fn i8_quad(group: [i8; 4]) -> u32 { + u32::from_le_bytes(group.map(|x| x as u8)) +} + +//--------// +// 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> +where + A: PrepareB, +{ + arch: A, + a: packed::View<'a, A::AElem, 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> +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. + /// + /// # 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, A::AElem, 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", + ); + + 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) } + } + + /// # 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, A::AElem, 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 + PrepareB + 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 mut b_scratch = A::Scratch::default(); + + 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, side) = + unsafe { self.arch.prepare(b_panels, &mut b_scratch, self.k) }; + + let panel_kernel = + |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; + 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`, and `side` was returned with + // `b_panels`. + let kernel = unsafe { + PanelKernel::new(self.arch, a_panel, b_panels, side, c, self.k) + }; + + // 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 { + 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) + }; + }, + ); + } +} + +//----------// +// 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; + + /// 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 Self::Scratch, + k: DimK, + ) -> (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 (), + _: DimK, + ) -> (unpacked::View<'a, i8>, Slice<'a, i32>) { + (b, Slice::new(&[])) + } +} + +//-------------// +// PanelKernel // +//-------------// + +#[derive(Debug)] +pub(super) struct PanelKernel<'a, A, const MR: usize, const NR: usize, const PACK: usize> +where + A: PrepareB, +{ + arch: A, + a: packed::Panel<'a, A::AElem, MR, PACK>, + b: unpacked::View<'a, A::Elem>, + side: Slice<'a, i32>, + c: [i32; MR], + k: DimK, +} + +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 + /// + /// 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, 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, + side, + c, + k, + } + } +} + +/// 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 + } +} + +/// 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> +where + A: PrepareB, +{ + arch: A, + a: packed::Panel<'a, A::AElem, MR, PACK>, + side: Slice<'a, i32>, + c: &'a mut [i32; MR], + k: DimK, +} + +impl unpacked::PanelVisitor + for Visitor<'_, A, MR, NR, PACK> +where + A: PrepareB, + for<'a> MicroKernel<'a, A, MR, NR, PACK>: driver::MicroKernel, +{ + #[inline(always)] + 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, 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); + } +} + +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, + side: self.side, + 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`, 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, + ) + }; + + 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`, 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, A::AElem, MR, PACK>, + b: unpacked::Panel<'a, A::Elem, NR>, + side: Slice<'a, i32>, + start: usize, + 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> +where + A: PrepareB, +{ + /// # Safety + /// + /// 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, 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, + side, + start, + 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<'_, 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 { + T::default() + } + }) +} + +/// Accumulate one pack 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_pack( + wide: W, + ai: W::Wide, + bp: Slice<'_, W::Elem>, + 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 = unsafe { wide.splat(bp.add(bstride * j), valid) }; + + *acc = W::dot(ai, bj, *acc); + } +} + +/// # Safety +/// +/// 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`]. +#[inline(always)] +unsafe fn micro_kernel( + wide: W, + 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 + 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(); + + // 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); + + let packs = a.packs(k); + let k = k.value().get(); + + // Loads pack `pack` of `a`, which callers must keep below `packs`. + // + // 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 = |pack| unsafe { wide.load(ap.add(astride * pack).truncate(astride)) }; + + // 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 pack in 0..full { + let i = pack * PACK; + + // SAFETY: `pack < full <= packs`, and since `i + PACK <= k`, every column of `b` has + // `PACK` readable elements at offset `i`. + unsafe { + accumulate_pack( + wide, + load(pack), + bp.add(Elements::new(i)), + bstride, + PACK, + &mut acc, + ) + }; + } + + // The trailing pack of `a` is zero padded, so zero filling `b` past `k` keeps every + // padded product at zero. + if full < packs { + let i = full * PACK; + + // SAFETY: `full < packs`, and every column of `b` has `k - i` readable elements at + // offset `i`. + unsafe { + accumulate_pack( + wide, + load(full), + bp.add(Elements::new(i)), + bstride, + k - i, + &mut 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`, + // 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, + ) + } + } + } + }; + ($arch:ty, $mr:literal, $pack:literal, { $($nr:literal),+ $(,)? }) => { + $(micro_kernel!($arch, $mr, $nr, $pack);)+ + } +} + +micro_kernel!(Scalar, 8, 2, { 2, 1 }); + +trait ExtraWide: PrepareB { + type Wide: Copy; + type Splat: Copy; + type Acc: Copy; + + /// # Safety + /// + /// `slice.len()` must be exactly `ELEMENTS * PACK`. + unsafe fn load(self, slice: Slice<'_, Self::AElem>) -> Self::Wide; + + /// 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. + /// + /// # 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]); +} + +impl ExtraWide<8, 2> for Scalar { + type Wide = i16x16; + type Splat = i16x16; + type Acc = i32x8; + + #[inline(always)] + unsafe fn init(self, _: Slice<'_, i32>, _: usize) -> 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)] + 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)] + 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::{ + SIMDSumTree, + arch::x86_64::{V3, V4}, + }; + + 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 // + //----------// + + impl PrepareB for V3 { + type Elem = i16; + type AElem = i16; + type Scratch = Vec; + + #[inline(always)] + unsafe fn prepare<'a>( + self, + b: unpacked::View<'a, i8>, + scratch: &'a mut Vec, + k: DimK, + ) -> (unpacked::View<'a, i16>, Slice<'a, i32>) { + // SAFETY: Inherited from caller. + let from = unsafe { b.as_std_slice(k) }; + + scratch.resize(from.len(), 0); + 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) }; + + (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 + } + + //-----------// + // ExtraWide // + //-----------// + + impl ExtraWide<16, 2> for V3 { + type Wide = [i16x16; 2]; + type Splat = i16x16; + type Acc = [i32x8; 2]; + + #[inline(always)] + unsafe fn init(self, _: Slice<'_, i32>, _: usize) -> Self::Acc { + [SIMDVector::default(self), SIMDVector::default(self)] + } + + #[inline(always)] + 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. + unsafe { + [ + SIMDVector::load_simd(self, slice.as_ptr()), + SIMDVector::load_simd(self, slice.add(Elements::new(16)).as_ptr()), + ] + } + } + + #[inline(always)] + 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)] + 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)); + } + } + } + + 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")] +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 }); + + //----------// + // 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 (), + _: DimK, + ) -> (unpacked::View<'a, i8>, Slice<'a, i32>) { + (b, Slice::new(&[])) + } + } + + //-----------// + // ExtraWide // + //-----------// + + impl ExtraWide<8, 4> for Neon { + type Wide = [i8x16; 2]; + type Splat = i8x16; + type Acc = [i32x4; 2]; + + #[inline(always)] + 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 { + 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)] + 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)] + 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)] + unsafe fn init(self, _: Slice<'_, i32>, _: usize) -> 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)] + 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)] + 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 // +/////////// + +#[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, 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 // + ///////////////// + + fn test_micro_kernel( + arch: A, + k: DimK, + rng: &mut impl rand::Rng, + ctx: std::fmt::Arguments<'_>, + ) where + 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 = pack_a::(&ref_a); + let ref_b = ref_b.transpose(); + + let mut scratch = A::Scratch::default(); + + // SAFETY: Test builds will verify the bounds we passed. + let (b, side) = 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. + // + // 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(b.as_std_slice(k)), k), + side, + 0, + &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 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 }, + ); + + #[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, + Neon::new_checked(), + 0x7d4a1e6c93b0f582, + 4, + 8 => { 6, 5, 4, 3, 2, 1 }, + 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: ConvertA, + 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 = pack_a::(&ref_a); + let ref_b = ref_b.transpose(); + + let extent = NonZeroUsize::new(cols).unwrap(); + + 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]; + + // 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), + b, + side, + 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 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), + ); + + #[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, + Neon::new_checked(), + 0x9f3e7ab4c05d1268, + (8, 6, 4), + (16, 6, 4), + ); + + //////////// + // Driver // + //////////// + + fn test_driver( + arch: A, + rng: &mut impl rand::Rng, + ) where + A: ConvertA, + for<'a> Driver<'a, A, MR, NR, PACK>: driver::Drive, + { + 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, + 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 = pack_a::(&ref_a); + 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), + ); + + #[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, + Neon::new_checked(), + 0x63c8ed19f4720ab5, + (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/maxsim/test.rs b/diskann-quantization/src/matrix_kernels/maxsim/test.rs index b087d56df4..9cf3a4798a 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.element(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 3bff0cef05..91103ca4aa 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::Matrix; use half::f16; -use rand::{Rng, distr::Distribution}; +use rand::{ + Rng, + distr::{Distribution, StandardUniform}, +}; /////////////////////// // panic_message_for // @@ -46,6 +49,12 @@ impl Distribution for TestDistr { } } +impl Distribution for TestDistr { + fn sample(&self, rng: &mut R) -> i8 { + StandardUniform {}.sample(rng) + } +} + 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..cdb4ef434e 100644 --- a/diskann-quantization/src/matrix_kernels/util.rs +++ b/diskann-quantization/src/matrix_kernels/util.rs @@ -92,6 +92,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")] @@ -102,50 +144,12 @@ 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); - - 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")] @@ -157,6 +161,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); } ////////// @@ -273,6 +279,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 +334,7 @@ mod test { test_load_store_scalar, Some(Scalar), f32 => { 4, 8, 16 }, + i32 => { 8 }, ); #[cfg(target_arch = "x86_64")] @@ -329,6 +342,7 @@ mod test { test_load_store_v3, V3::new_checked(), f32 => { 8, 16 }, + i32 => { 16 }, ); #[cfg(target_arch = "x86_64")] @@ -336,6 +350,7 @@ mod test { test_load_store_v4, V4::new_checked_miri(), f32 => { 8, 16, 32 }, + i32 => { 32 }, ); #[cfg(target_arch = "aarch64")] @@ -343,6 +358,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 fe87959ccd..064b9bb425 100644 --- a/diskann-quantization/src/multi_vector/distance/factory.rs +++ b/diskann-quantization/src/multi_vector/distance/factory.rs @@ -8,18 +8,20 @@ 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")] 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}; 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 +177,88 @@ where } } +impl MaxSimKernel + for Prepared, NR> +where + 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 { + 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 +266,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 +283,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 +294,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 +464,69 @@ 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 { + // `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, + _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 { + // `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>, + }) + } +} + +#[cfg(target_arch = "aarch64")] +impl> diskann_wide::arch::Target1>> + for BuildAndErase +{ + fn run(self, arch: Neon, query: MatRef<'_, Standard>) -> E::Output { + let prepared = BlockTransposed::::from_matrix_view(query.as_matrix_view()); + self.0.erase(Prepared { + arch, + prepared, + _packing: Pack::<6>, + }) + } +} + // ───────────────────────────────────────────────────────────────────────── // MaxSimElement — sealed trait gating accepted element types. // ───────────────────────────────────────────────────────────────────────── @@ -391,29 +541,42 @@ 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(...)`. /// /// # 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 {} 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>, erase: E, - ) -> Result { + ) -> Result { match isa { MaxSimIsa::Auto => Ok(diskann_wide::arch::dispatch1_no_features( BuildAndErase(erase), @@ -440,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 { @@ -453,18 +617,24 @@ impl MaxSimElement for f32 { MaxSimIsa::Neon => Err(NotSupported { isa, reason: "aarch64 target only", - }), - MaxSimIsa::Reference => Ok(erase.erase(ReferenceKernel::::new(query))), + } + .into()), + 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>, erase: E, - ) -> Result { + ) -> Result { match isa { MaxSimIsa::Auto => Ok(diskann_wide::arch::dispatch1_no_features( BuildAndErase(erase), @@ -491,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 { @@ -504,8 +675,79 @@ impl MaxSimElement for half::f16 { MaxSimIsa::Neon => Err(NotSupported { isa, reason: "aarch64 target only", - }), - MaxSimIsa::Reference => Ok(erase.erase(ReferenceKernel::::new(query))), + } + .into()), + MaxSimIsa::Reference => { + Ok(erase.erase(ReferenceKernel::new(query, reference_scores::))) + } + } + } +} + +/// 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; + + fn build>( + isa: MaxSimIsa, + query: MatRef<'_, Standard>, + erase: E, + ) -> 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), + 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", + } + .into()), + #[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", + } + .into()), + MaxSimIsa::Reference => { + Ok(erase.erase(ReferenceKernel::new(query, reference_scores_i8))) + } } } } @@ -514,6 +756,16 @@ impl MaxSimElement for half::f16 { // 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 @@ -521,12 +773,14 @@ impl MaxSimElement for half::f16 { /// /// # 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) } @@ -534,10 +788,10 @@ 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 - /// generator. + /// so every element type shares the same data generator. trait FromF32 { fn from_f32(v: f32) -> Self; } @@ -554,6 +808,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 every element type shares 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 +875,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 +888,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 +901,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 +915,64 @@ 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], ); } } } + 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); + + 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(); + + for i in 0..nq { + let actual = ::score_as_f32(scores[i]); + assert!( + (actual - expected[i]).abs() < 1e-10, + "i8 {isa} MaxSim[{i}] mismatch for ({nq},{nd},{dim}): \ + actual={actual}, expected={}", + expected[i], + ); + } + } + } + + #[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]; @@ -653,10 +989,45 @@ 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); + } + + #[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, - 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 +1037,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 +1045,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 +1058,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 +1066,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 +1081,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 +1094,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 +1136,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 { 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 cc2a36daec..cefa6393c6 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, u32x2, @@ -116,6 +116,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 // /////////// @@ -170,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()); }