diff --git a/Cargo.lock b/Cargo.lock index bc9903c01c..686b732ce8 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -788,6 +788,7 @@ version = "0.59.0" dependencies = [ "bytemuck", "cfg-if", + "criterion", "diskann-vector", "diskann-wide", "half", diff --git a/diskann-utils/Cargo.toml b/diskann-utils/Cargo.toml index d0dbb9b873..f76b971cb0 100644 --- a/diskann-utils/Cargo.toml +++ b/diskann-utils/Cargo.toml @@ -30,9 +30,14 @@ workspace = true [dev-dependencies] cfg-if.workspace = true +criterion.workspace = true rand.workspace = true rstest.workspace = true +[[bench]] +name = "bench_main" +harness = false +required-features = ["rayon"] [features] default = ["rayon"] diff --git a/diskann-utils/benches/bench_main.rs b/diskann-utils/benches/bench_main.rs new file mode 100644 index 0000000000..a33b636b9e --- /dev/null +++ b/diskann-utils/benches/bench_main.rs @@ -0,0 +1,81 @@ +/* + * Copyright (c) Microsoft Corporation. + * Licensed under the MIT license. + */ + +use std::{hint::black_box, time::Duration}; + +use criterion::{criterion_group, criterion_main, BenchmarkId, Criterion, Throughput}; +use diskann_utils::views::rowmajor::{Matrix, MatrixMut, Owned}; +use rayon::prelude::{IndexedParallelIterator, ParallelIterator, ParallelSliceMut}; + +const NCOLS: usize = 100; + +fn par_rows_mut_baseline( + matrix: &mut M, +) -> impl IndexedParallelIterator +where + M: MatrixMut, + M::Element: Send, +{ + let ncols = matrix.ncols(); + assert!( + ncols != 0 || matrix.nrows() == 0, + "`MatrixMut::par_rows_mut` does not support matrices with rows and zero columns" + ); + matrix.as_mut_slice().par_chunks_exact_mut(ncols.max(1)) +} + +fn update_rows<'a, I>(rows: I) +where + I: ParallelIterator, +{ + rows.for_each(|row| { + let (first, rest) = row.split_first_mut().unwrap(); + *first = first.wrapping_add(1); + + let mut previous = *first; + for value in rest { + *value = value.wrapping_add(previous).wrapping_add(1); + previous = *value; + } + }); +} + +fn benchmark_shape(c: &mut Criterion, nrows: usize) { + let mut group = c.benchmark_group(format!("par_rows_mut/{nrows}x{NCOLS}")); + group.throughput(Throughput::Elements((nrows * NCOLS) as u64)); + + group.bench_function(BenchmarkId::new("baseline", nrows), |b| { + let mut matrix = Owned::from_element(nrows, NCOLS, 0u32); + b.iter(|| { + update_rows(par_rows_mut_baseline(black_box(&mut matrix))); + black_box(matrix.as_slice()); + }); + }); + + group.bench_function(BenchmarkId::new("current", nrows), |b| { + let mut matrix = Owned::from_element(nrows, NCOLS, 0u32); + b.iter(|| { + update_rows(black_box(&mut matrix).par_rows_mut()); + black_box(matrix.as_slice()); + }); + }); + + group.finish(); +} + +fn benchmark_par_rows_mut(c: &mut Criterion) { + benchmark_shape(c, 5); + benchmark_shape(c, 1_000_000); +} + +criterion_group! { + name = benches; + config = Criterion::default() + .sample_size(10) + .warm_up_time(Duration::from_secs(2)) + .measurement_time(Duration::from_secs(5)); + targets = benchmark_par_rows_mut +} +criterion_main!(benches); diff --git a/diskann-utils/src/views/rowmajor.rs b/diskann-utils/src/views/rowmajor.rs index ffd569394d..25be0b880b 100644 --- a/diskann-utils/src/views/rowmajor.rs +++ b/diskann-utils/src/views/rowmajor.rs @@ -6,9 +6,7 @@ use std::{marker::PhantomData, mem::ManuallyDrop, num::NonZeroUsize, ptr::NonNull}; #[cfg(feature = "rayon")] -use rayon::prelude::{ - IndexedParallelIterator, IntoParallelIterator, ParallelIterator, ParallelSliceMut, -}; +use rayon::prelude::{IndexedParallelIterator, IntoParallelIterator, ParallelIterator}; use thiserror::Error; pub mod iter; @@ -525,20 +523,18 @@ pub unsafe trait MatrixMut: Matrix { /// Return a parallel iterator over the rows of the matrix. /// - /// # Panics - /// - /// Panics if `self.ncols() == 0 && self.nrows() != 0`. + /// A matrix with zero columns yields one empty mutable slice per row. #[cfg(feature = "rayon")] fn par_rows_mut(&mut self) -> impl IndexedParallelIterator where Self::Element: Send, { - let ncols = self.ncols(); - assert!( - ncols != 0 || self.nrows() == 0, - "`MatrixMut::par_rows_mut` does not support matrices with rows and zero columns" - ); - self.as_mut_slice().par_chunks_exact_mut(ncols.max(1)) + let matrix = iter::ParMut::new(self); + (0..matrix.nrows()).into_par_iter().map(move |row| { + // SAFETY: The range produces each in-bounds row exactly once. Distinct rows + // have disjoint element ranges; zero-column rows contain no elements. + unsafe { matrix.row_disjoint_unchecked(row) } + }) } /// Return a parallel iterator that divides the matrix into mutable sub-matrices with @@ -550,9 +546,12 @@ pub unsafe trait MatrixMut: Matrix { /// It is possible for yielded sub-matrices to have fewer than `batchsize` rows if the /// number of rows in the parent matrix is not evenly divisible by `batchsize`. /// + /// A matrix with zero columns yields zero-column sub-matrices containing up to + /// `batchsize` rows. + /// /// # Panics /// - /// Panics if `batchsize = 0` or `self.ncols() == 0 && self.nrows() != 0`. + /// Panics if `batchsize = 0`. #[cfg(feature = "rayon")] fn par_window_iter_mut( &mut self, @@ -566,30 +565,17 @@ pub unsafe trait MatrixMut: Matrix { "par_window_iter_mut batchsize cannot be zero" ); - let ncols = self.ncols(); - assert!( - ncols != 0 || self.nrows() == 0, - "`MatrixMut::par_window_iter_mut` does not support matrices with rows and zero columns" - ); + let matrix = iter::ParMut::new(self); + let nrows = matrix.nrows(); + (0..nrows) + .into_par_iter() + .step_by(batchsize) + .map(move |start| { + let end = start.saturating_add(batchsize).min(nrows); - // Ensure that `batchsize * ncols` does not overflow. - let batchsize = batchsize.min(self.nrows()); - self.as_mut_slice() - .par_chunks_mut((ncols * batchsize).max(1)) - .map(move |data| { - let blobsize = data.len(); - let nrows = blobsize / ncols; - assert_eq!(blobsize % ncols, 0); - - // SAFETY: - // - // * `Layout::new_unchecked` is safe because `ncols` is the parent column - // count and `nrows <= self.nrows()`, so this layout cannot exceed the - // validated parent layout. - // - // * `Mut::from_data_unchecked` is safe because by construction, - // `data.len() == ncols * nrows`. - unsafe { Mut::from_data_unchecked(data, Layout::new_unchecked(nrows, ncols)) } + // SAFETY: `start` comes from an in-bounds range and `end` is clamped to + // `nrows`. Stepping by `batchsize` makes the yielded ranges disjoint. + unsafe { matrix.window_disjoint_unchecked(start..end) } }) } } @@ -2325,7 +2311,7 @@ mod tests { fn parallel_mutable_iterators_match_scalar_indexing() { use rayon::prelude::*; - for (nrows, ncols) in [(0, 0), (0, 4), (1, 1), (1, 4), (4, 1), (5, 3)] { + for (nrows, ncols) in [(0, 0), (0, 4), (3, 0), (1, 1), (1, 4), (4, 1), (5, 3)] { let context = lazy_format!("nrows = {nrows}, ncols = {ncols}"); let original = striped_matrix(nrows, ncols); @@ -2392,26 +2378,6 @@ mod tests { } } - #[test] - #[cfg(feature = "rayon")] - #[should_panic( - expected = "`MatrixMut::par_rows_mut` does not support matrices with rows and zero columns" - )] - fn par_rows_mut_rejects_nonempty_zero_column_matrix() { - let mut m = striped_matrix(3, 0); - let _ = m.par_rows_mut(); - } - - #[test] - #[cfg(feature = "rayon")] - #[should_panic( - expected = "`MatrixMut::par_window_iter_mut` does not support matrices with rows and zero columns" - )] - fn par_window_iter_mut_rejects_nonempty_zero_column_matrix() { - let mut m = striped_matrix(3, 0); - let _ = m.par_window_iter_mut(2); - } - #[test] fn matrix_iterators_match_scalar_indexing() { for (nrows, ncols) in [(0, 0), (0, 4), (3, 0), (1, 1), (1, 4), (4, 1), (5, 3)] { diff --git a/diskann-utils/src/views/rowmajor/iter.rs b/diskann-utils/src/views/rowmajor/iter.rs index 91a765e551..a7a5f7a148 100644 --- a/diskann-utils/src/views/rowmajor/iter.rs +++ b/diskann-utils/src/views/rowmajor/iter.rs @@ -5,6 +5,8 @@ use std::{marker::PhantomData, num::NonZeroUsize, ptr::NonNull}; +#[cfg(feature = "rayon")] +use crate::views::rowmajor::MatrixMut; use crate::views::rowmajor::{Layout, Matrix, Mut, Ref}; //------// @@ -187,3 +189,150 @@ impl<'a, T> Iterator for Windows<'a, T> { impl ExactSizeIterator for Windows<'_, T> {} impl std::iter::FusedIterator for Windows<'_, T> {} + +//--------// +// ParMut // +//--------// + +#[cfg(feature = "rayon")] +/// Carries an exclusive matrix borrow across Rayon workers. +pub(super) struct ParMut<'a, T> { + ptr: NonNull, + layout: Layout, + _lifetime: PhantomData<&'a mut [T]>, +} + +// SAFETY: `ParMut` owns an exclusive slice borrow, so sending it requires `T: Send`. +#[cfg(feature = "rayon")] +unsafe impl Send for ParMut<'_, T> {} +// SAFETY: The only methods that produce mutable views are unsafe and require callers to +// ensure disjointness. Sending those views between workers requires `T: Send`. +#[cfg(feature = "rayon")] +unsafe impl Sync for ParMut<'_, T> {} + +#[cfg(feature = "rayon")] +impl<'a, T> ParMut<'a, T> { + pub(super) fn new(matrix: &'a mut M) -> Self + where + M: MatrixMut + ?Sized, + { + let layout = matrix.layout(); + let ptr = matrix.as_nonnull_mut(); + Self { + ptr, + layout, + _lifetime: PhantomData, + } + } + + pub(super) fn nrows(&self) -> usize { + self.layout.nrows() + } + + /// # Safety + /// + /// * `row < self.nrows()`. + /// * No other live reference derived from this `ParMut` may overlap any element of + /// row `row`. Zero-column rows never overlap, even when their addresses match. + pub(super) unsafe fn row_disjoint_unchecked(&self, row: usize) -> &'a mut [T] { + debug_assert!(row < self.layout.nrows()); + let ncols = self.layout.ncols(); + + // SAFETY: The caller guarantees that `row` is in-bounds and does not overlap any + // other live view. The validated parent layout makes the offset representable and + // places the row within the initialized matrix span. + unsafe { std::slice::from_raw_parts_mut(self.ptr.as_ptr().add(row * ncols), ncols) } + } + + /// # Safety + /// + /// * `rows.start <= rows.end <= self.nrows()`. + /// * No other live reference derived from this `ParMut` may overlap any element in + /// `rows`. Zero-column windows never overlap. + pub(super) unsafe fn window_disjoint_unchecked( + &self, + rows: std::ops::Range, + ) -> Mut<'a, T> { + debug_assert!(rows.start <= rows.end); + debug_assert!(rows.end <= self.layout.nrows()); + + let ncols = self.layout.ncols(); + let nrows = rows.end - rows.start; + + // SAFETY: The caller guarantees an ordered, in-bounds range. The validated parent + // layout makes the offset representable and places it within or one past the matrix. + let ptr = unsafe { self.ptr.add(rows.start * ncols) }; + + Mut { + ptr, + // SAFETY: This window has no more rows than the validated parent and keeps its + // column count, so its element count and byte span cannot exceed the parent. + layout: unsafe { Layout::new_unchecked(nrows, ncols) }, + _lifetime: PhantomData, + } + } +} + +#[cfg(all(test, feature = "rayon"))] +mod tests { + use super::ParMut; + use crate::views::rowmajor::{Matrix, MatrixMut, Owned}; + + #[test] + fn par_mut_zero_column_views_can_coexist() { + let mut matrix = Owned::from_element(usize::MAX, 0, 0); + let ptr = matrix.as_ptr(); + let matrix = ParMut::new(&mut matrix); + + // SAFETY: Empty views do not overlap any elements, even when their addresses match. + let rows = unsafe { + [ + matrix.row_disjoint_unchecked(0), + matrix.row_disjoint_unchecked(1), + matrix.row_disjoint_unchecked(usize::MAX - 1), + ] + }; + // SAFETY: Empty windows do not overlap any elements, including the live row views + // whose logical rows fall within these windows. + let windows = unsafe { + [ + matrix.window_disjoint_unchecked(0..2), + matrix.window_disjoint_unchecked(2..usize::MAX), + ] + }; + + assert!(rows.iter().all(|row| row.is_empty() && row.as_ptr() == ptr)); + assert!(windows.iter().all(|window| { + window.ncols() == 0 && window.as_slice().is_empty() && window.as_ptr() == ptr + })); + } + + #[test] + fn par_mut_disjoint_nonempty_views_can_coexist() { + let mut matrix = Owned::from_fn(6, 2, |rc| rc.row * 100 + rc.col); + { + let matrix = ParMut::new(&mut matrix); + + // SAFETY: These rows and windows are in-bounds and pairwise disjoint. + let (row0, row1, mut window2, mut window4) = unsafe { + ( + matrix.row_disjoint_unchecked(0), + matrix.row_disjoint_unchecked(1), + matrix.window_disjoint_unchecked(2..4), + matrix.window_disjoint_unchecked(4..6), + ) + }; + row0[0] = 10; + row1[1] = 11; + *window2.element_mut(0, 0) = 20; + *window2.element_mut(1, 1) = 31; + *window4.element_mut(0, 0) = 40; + *window4.element_mut(1, 1) = 51; + } + + assert_eq!( + matrix.as_slice(), + &[10, 1, 100, 11, 20, 201, 300, 31, 40, 401, 500, 51] + ); + } +}