Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions Cargo.lock

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

5 changes: 5 additions & 0 deletions diskann-utils/Cargo.toml
Original file line number Diff line number Diff line change
Expand Up @@ -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"]
Expand Down
81 changes: 81 additions & 0 deletions diskann-utils/benches/bench_main.rs
Original file line number Diff line number Diff line change
@@ -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<M>(
matrix: &mut M,
) -> impl IndexedParallelIterator<Item = &mut [M::Element]>
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<Item = &'a mut [u32]>,
{
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);
80 changes: 23 additions & 57 deletions diskann-utils/src/views/rowmajor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -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<Item = &mut [Self::Element]>
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) }
})
}
Comment thread
arrayka marked this conversation as resolved.

/// Return a parallel iterator that divides the matrix into mutable sub-matrices with
Expand All @@ -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,
Expand All @@ -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) }
})
}
}
Expand Down Expand Up @@ -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);
Expand Down Expand Up @@ -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)] {
Expand Down
149 changes: 149 additions & 0 deletions diskann-utils/src/views/rowmajor/iter.rs
Original file line number Diff line number Diff line change
Expand Up @@ -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};

//------//
Expand Down Expand Up @@ -187,3 +189,150 @@ impl<'a, T> Iterator for Windows<'a, T> {

impl<T> ExactSizeIterator for Windows<'_, T> {}
impl<T> 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<T>,
layout: Layout<T>,
_lifetime: PhantomData<&'a mut [T]>,
}

// SAFETY: `ParMut` owns an exclusive slice borrow, so sending it requires `T: Send`.
#[cfg(feature = "rayon")]
unsafe impl<T: Send> 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<T: Send> Sync for ParMut<'_, T> {}

#[cfg(feature = "rayon")]
impl<'a, T> ParMut<'a, T> {
pub(super) fn new<M>(matrix: &'a mut M) -> Self
where
M: MatrixMut<Element = T> + ?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<usize>,
) -> 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]
);
}
}
Loading