Skip to content
Merged
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
4 changes: 2 additions & 2 deletions diskann-benchmark-core/src/recall.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,7 +9,7 @@ use std::{
};

use diskann_utils::{
strided::StridedView,
strided::Strided,
views::{Matrix, MatrixView},
};
use thiserror::Error;
Expand Down Expand Up @@ -145,7 +145,7 @@ pub enum GroundTruthMode {
/// than `recall_k` candidates.
pub fn knn<T>(
groundtruth: &dyn Rows<T>,
groundtruth_distances: Option<StridedView<'_, f32>>,
groundtruth_distances: Option<Strided<'_, f32>>,
results: &dyn Rows<T>,
recall_k: usize,
recall_n: usize,
Expand Down
10 changes: 5 additions & 5 deletions diskann-providers/src/model/pq/strided.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,14 +4,14 @@
*/

use diskann::ANNError;
use diskann_utils::{strided, views};
use diskann_utils::strided;

use crate::utils::Bridge;

impl<T: views::DenseData> From<Bridge<strided::TryFromError<T>>> for ANNError {
impl From<Bridge<strided::TryFromError>> for ANNError {
#[track_caller]
fn from(value: Bridge<strided::TryFromError<T>>) -> Self {
ANNError::new(value.into_inner().as_static())
fn from(value: Bridge<strided::TryFromError>) -> Self {
ANNError::new(value.into_inner())
}
}

Expand All @@ -32,7 +32,7 @@ mod tests {
let x = vec![u8::default(); nrows * ncols];

// Provided the incorrect dimensions.
let err = strided::StridedView::try_from(&x, nrows, ncols + 1, ncols + 1)
let err = strided::Strided::try_from_data(&x, nrows, ncols + 1, ncols + 1)
.bridge_err()
.unwrap_err();
let message = format!("{}", err);
Expand Down
8 changes: 4 additions & 4 deletions diskann-quantization/src/algorithms/kmeans/lloyds.rs
Original file line number Diff line number Diff line change
Expand Up @@ -8,7 +8,7 @@ use diskann_wide::{SIMDMask, SIMDMulAdd, SIMDPartialOrd, SIMDSelect, SIMDSumTree
use super::common::square_norm;
use crate::multi_vector::{BlockTransposed, BlockTransposedRef};
use diskann_utils::{
strided::StridedView,
strided::Strided,
views::{Matrix, MatrixView, MutMatrixView},
};

Expand Down Expand Up @@ -342,10 +342,10 @@ fn update((d0, i0): (f32s, u32s), (d1, i1): (f32s, u32s)) -> (f32s, u32s) {
// Update Step //
/////////////////

fn update_centroids(mut centers: MutMatrixView<'_, f32>, data: StridedView<'_, f32>, map: &[u32]) {
fn update_centroids(mut centers: MutMatrixView<'_, f32>, data: Strided<'_, f32>, map: &[u32]) {
let mut sums = Matrix::<f64>::new(0.0, centers.nrows(), centers.ncols());
let mut counts: Vec<u32> = vec![0; centers.nrows()];
data.row_iter().zip(map.iter()).for_each(|(row, &center)| {
data.rows().zip(map.iter()).for_each(|(row, &center)| {
counts[center as usize] += 1;
let sum = sums.row_mut(center as usize);
std::iter::zip(sum.iter_mut(), row.iter()).for_each(|(s, r)| {
Expand All @@ -370,7 +370,7 @@ fn update_centroids(mut centers: MutMatrixView<'_, f32>, data: StridedView<'_, f
////////////

pub(crate) fn lloyds_inner(
data: StridedView<'_, f32>,
data: Strided<'_, f32>,
square_norms: &[f32],
transpose: BlockTransposedRef<'_, f32, 16>,
mut centers: MutMatrixView<'_, f32>,
Expand Down
4 changes: 2 additions & 2 deletions diskann-quantization/src/algorithms/kmeans/plusplus.rs
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
use std::{collections::HashSet, fmt};

use diskann_utils::{
strided::StridedView,
strided::Strided,
views::{MatrixView, MutMatrixView},
};
use diskann_wide::{SIMDMulAdd, SIMDPartialOrd, SIMDSelect, SIMDVector};
Expand Down Expand Up @@ -380,7 +380,7 @@ impl KMeansPlusPlusError {

pub(crate) fn kmeans_plusplus_into_inner<const N: usize>(
mut points: MutMatrixView<'_, f32>,
data: StridedView<'_, f32>,
data: Strided<'_, f32>,
transpose: BlockTransposedRef<'_, f32, N>,
norms: &[f32],
rng: &mut dyn RngCore,
Expand Down
14 changes: 6 additions & 8 deletions diskann-quantization/src/multi_vector/block_transposed.rs
Original file line number Diff line number Diff line change
Expand Up @@ -81,7 +81,7 @@ use std::{alloc::Layout, marker::PhantomData, ptr::NonNull};

use diskann_utils::{
Reborrow, ReborrowMut,
strided::StridedView,
strided::Strided,
views::{MatrixView, MutMatrixView},
};

Expand Down Expand Up @@ -1142,7 +1142,7 @@ impl<T: Copy + Default, const GROUP: usize, const PACK: usize> BlockTransposed<T
})
}

/// Construct a block-transposed matrix by copying data from a [`StridedView`].
/// Construct a block-transposed matrix by copying data from a [`Strided`].
///
/// Each source element at `(row, col)` is placed at the correct offset in the
/// block-transposed layout. Padding positions (both partial-block rows and
Expand All @@ -1151,9 +1151,9 @@ impl<T: Copy + Default, const GROUP: usize, const PACK: usize> BlockTransposed<T
///
/// The loop iterates in physical (block-transposed) order — block, column-group,
/// row-within-block, pack-lane — so that writes to the backing allocation are
/// sequential. Source reads stride across rows of the [`StridedView`], which is
/// sequential. Source reads stride across rows of the [`Strided`], which is
/// acceptable because read-side prefetch is more effective than write-side.
pub fn from_strided(v: StridedView<'_, T>) -> Self {
pub fn from_strided(v: Strided<'_, T>) -> Self {
let nrows = v.nrows();
let ncols = v.ncols();
let mut mat = Self::new(nrows, ncols);
Expand All @@ -1175,7 +1175,7 @@ impl<T: Copy + Default, const GROUP: usize, const PACK: usize> BlockTransposed<T
let row = row_base + rib;
if row < nrows {
// SAFETY: row < nrows is checked by the enclosing `if` condition.
let src_row = unsafe { v.get_row_unchecked(row) };
let src_row = unsafe { v.row_unchecked(row) };
for p in 0..PACK {
let col = col_base + p;
if col < ncols {
Expand Down Expand Up @@ -2136,8 +2136,6 @@ mod tests {

#[test]
fn test_from_strided_nonunit_stride() {
use diskann_utils::strided::StridedView;

const GROUP: usize = 4;
const PACK: usize = 2;
let nrows = 5;
Expand All @@ -2152,7 +2150,7 @@ mod tests {
}
}

let strided = StridedView::try_shrink_from(&flat, nrows, ncols, cstride)
let strided = Strided::try_from_data(&flat, nrows, ncols, cstride)
.expect("should construct strided view");
let transpose = BlockTransposed::<f32, GROUP, PACK>::from_strided(strided);

Expand Down
73 changes: 30 additions & 43 deletions diskann-quantization/src/product/tables/transposed/pivots.rs
Original file line number Diff line number Diff line change
Expand Up @@ -5,7 +5,7 @@

use std::fmt;

use diskann_utils::strided;
use diskann_utils::strided::Strided;
use diskann_wide::{SIMDMask, SIMDMulAdd, SIMDPartialOrd, SIMDSelect, SIMDVector};

use crate::{
Expand All @@ -21,9 +21,9 @@ diskann_wide::alias!(u32s = u32x8);
/// Error types returned by Chunk construction.
#[derive(Debug, Clone)]
pub enum ChunkConstructionError {
/// A `StridedView` was provided with a dimension of zero.
/// A `Strided` was provided with a dimension of zero.
DimensionCannotBeZero,
/// A `StridedView` was provided with a length of zero.
/// A `Strided` was provided with a length of zero.
LengthCannotBeZero,
}

Expand Down Expand Up @@ -220,7 +220,7 @@ impl Chunk {
///
/// 1. `data.ncols() == 0`
/// 2. `data.nrows() == 0`
pub(super) fn new(data: strided::StridedView<'_, f32>) -> Result<Self, ChunkConstructionError> {
pub(super) fn new(data: Strided<'_, f32>) -> Result<Self, ChunkConstructionError> {
// Error handling.
if data.ncols() == 0 {
return Err(ChunkConstructionError::DimensionCannotBeZero);
Expand All @@ -229,7 +229,7 @@ impl Chunk {
return Err(ChunkConstructionError::LengthCannotBeZero);
}

let square_norms = data.row_iter().map(kmeans::square_norm).collect();
let square_norms = data.rows().map(kmeans::square_norm).collect();
let data = BlockTransposed::<f32, 16>::from_strided(data);
Ok(Self { data, square_norms })
}
Expand Down Expand Up @@ -346,9 +346,9 @@ impl Chunk {
///
/// This method is generally more efficient to call.
///
/// **IMPORTANT**: The provided `StridedView` must have a length of `Self::batchsize()`.
/// **IMPORTANT**: The provided `Strided` must have a length of `Self::batchsize()`.
///
/// Providing a `StridedView` as an argument allows the compiler to infer that each
/// Providing a `Strided` as an argument allows the compiler to infer that each
/// row in `x` has a strided offset from the base, allowing for better code generation.
///
/// If the distances between a row in `x` and all centers is not finite, return
Expand Down Expand Up @@ -385,7 +385,7 @@ impl Chunk {
/// inner products.
pub(super) fn find_closest_batch<T>(
&self,
x: strided::StridedView<'_, T>,
x: Strided<'_, T>,
) -> [CompressionResult; Self::batchsize()]
where
T: Copy + Into<f32>,
Expand All @@ -394,7 +394,7 @@ impl Chunk {
assert_eq!(
x.nrows(),
Self::batchsize(),
"argument StridedView must have a length of {}",
"argument Strided must have a length of {}",
Self::batchsize()
);

Expand Down Expand Up @@ -437,7 +437,7 @@ impl Chunk {
debug_assert!(k < x.ncols());

// SAFETY:
// * `StridedView` indexing: We have checked in Assertion 1 that it is safe
// * `Strided` indexing: We have checked in Assertion 1 that it is safe
// to index `x` at indices `0, 1, 2, and 3`.
// * Inner slice indexing: It is the caller's responsibility to ensure that
// `k < x.dim()` so that indexing the slices return by `x.get_unchecked` are
Expand All @@ -449,19 +449,19 @@ impl Chunk {
(
f32s::splat(
diskann_wide::ARCH,
<T as Into<f32>>::into(*x.get_row_unchecked(0).get_unchecked(k)),
<T as Into<f32>>::into(*x.row_unchecked(0).get_unchecked(k)),
),
f32s::splat(
diskann_wide::ARCH,
<T as Into<f32>>::into(*x.get_row_unchecked(1).get_unchecked(k)),
<T as Into<f32>>::into(*x.row_unchecked(1).get_unchecked(k)),
),
f32s::splat(
diskann_wide::ARCH,
<T as Into<f32>>::into(*x.get_row_unchecked(2).get_unchecked(k)),
<T as Into<f32>>::into(*x.row_unchecked(2).get_unchecked(k)),
),
f32s::splat(
diskann_wide::ARCH,
<T as Into<f32>>::into(*x.get_row_unchecked(3).get_unchecked(k)),
<T as Into<f32>>::into(*x.row_unchecked(3).get_unchecked(k)),
),
)
}
Expand Down Expand Up @@ -1198,19 +1198,13 @@ mod tests {
let query_batch =
flatten(&[copy_query(0), copy_query(1), copy_query(2), copy_query(3)]);

let view = strided::StridedView::try_from(
query_batch.as_slice(),
Chunk::batchsize(),
dim,
dim,
)
.unwrap();
let view = Strided::try_from_data(query_batch.as_slice(), Chunk::batchsize(), dim, dim)
.unwrap();

// Make sure that the query batch was constructed correctly.
assert_eq!(view.nrows(), Chunk::batchsize());
assert_eq!(view.ncols(), dim);
for k in 0..view.nrows() {
let row = view.row(k);
for (k, row) in view.rows().enumerate() {
if k == j {
assert_eq!(
row, query,
Expand Down Expand Up @@ -1261,13 +1255,8 @@ mod tests {
maybe_broadcast(3, values[3]),
]);

let view = strided::StridedView::try_from(
query_batch.as_slice(),
Chunk::batchsize(),
dim,
dim,
)
.unwrap();
let view = Strided::try_from_data(query_batch.as_slice(), Chunk::batchsize(), dim, dim)
.unwrap();

let closest = chunk.find_closest_batch(view);
// Lane `j` should not be okay. All other lanes should return the correct
Expand Down Expand Up @@ -1299,7 +1288,7 @@ mod tests {
let mut data_aggregate = create_test_pattern(dim, total);

let data = flatten(&data_aggregate);
let sliced = strided::StridedView::try_from(data.as_slice(), total, dim, dim).unwrap();
let sliced = Strided::try_from_data(data.as_slice(), total, dim, dim).unwrap();
let chunk = Chunk::new(sliced).unwrap();

assert_eq!(chunk.num_centers(), total);
Expand All @@ -1309,7 +1298,7 @@ mod tests {
for row in 0..sliced.nrows() {
for col in 0..sliced.ncols() {
assert_eq!(
sliced[(row, col)],
*sliced.element(row, col),
chunk.get(row, col),
"failed on row {} and col {}",
row,
Expand Down Expand Up @@ -1383,7 +1372,7 @@ mod tests {
let last = data_aggregate.last().unwrap().clone();
data_aggregate[0].clone_from(&last);
let data = flatten(&data_aggregate);
let sliced = strided::StridedView::try_from(data.as_slice(), total, dim, dim).unwrap();
let sliced = Strided::try_from_data(data.as_slice(), total, dim, dim).unwrap();
let chunk = Chunk::new(sliced).unwrap();

assert_eq!(chunk.num_centers(), total);
Expand Down Expand Up @@ -1447,15 +1436,15 @@ mod tests {
#[test]
fn test_chunk_construction_error() {
// No dimensions
let chunk = Chunk::new(strided::StridedView::try_from(&[], 3, 0, 0).unwrap());
let chunk = Chunk::new(Strided::try_from_data(&[], 3, 0, 0).unwrap());
let err = chunk.unwrap_err();
assert!(
err.to_string()
.contains("cannot construct a Chunk from a source with zero dimensions")
);

// No length
let chunk = Chunk::new(strided::StridedView::try_from(&[], 0, 10, 10).unwrap());
let chunk = Chunk::new(Strided::try_from_data(&[], 0, 10, 10).unwrap());
let err = chunk.unwrap_err();
assert!(
err.to_string()
Expand All @@ -1470,7 +1459,7 @@ mod tests {
let dim = 10;
let total = 13;
let data = flatten(&create_test_pattern(dim, total));
let sliced = strided::StridedView::try_from(data.as_slice(), total, dim, dim).unwrap();
let sliced = Strided::try_from_data(data.as_slice(), total, dim, dim).unwrap();
let chunk = Chunk::new(sliced).unwrap();

let query: Vec<f32> = vec![0.0; total];
Expand All @@ -1486,32 +1475,30 @@ mod tests {
let dim = 10;
let total = 13;
let data = flatten(&create_test_pattern(dim, total));
let sliced = strided::StridedView::try_from(data.as_slice(), total, dim, dim).unwrap();
let sliced = Strided::try_from_data(data.as_slice(), total, dim, dim).unwrap();
let chunk = Chunk::new(sliced).unwrap();

let query: Vec<f32> = vec![0.0; 4 * total];
let query_view =
strided::StridedView::try_from(query.as_slice(), Chunk::batchsize(), total, total)
.unwrap();
Strided::try_from_data(query.as_slice(), Chunk::batchsize(), total, total).unwrap();

// PANICS
chunk.find_closest_batch(query_view);
}

// Make sure `find_closest_batch` panics for an incorrect length
#[test]
#[should_panic(expected = "argument StridedView must have a length of")]
#[should_panic(expected = "argument Strided must have a length of")]
fn test_find_closest_batch_panics_on_non_batch_length() {
let dim = 10;
let total = 13;
let data = flatten(&create_test_pattern(dim, total));
let sliced = strided::StridedView::try_from(data.as_slice(), total, dim, dim).unwrap();
let sliced = Strided::try_from_data(data.as_slice(), total, dim, dim).unwrap();
let chunk = Chunk::new(sliced).unwrap();

let query: Vec<f32> = vec![0.0; (Chunk::batchsize() + 1) * dim];
let query_view =
strided::StridedView::try_from(query.as_slice(), Chunk::batchsize() + 1, dim, dim)
.unwrap();
Strided::try_from_data(query.as_slice(), Chunk::batchsize() + 1, dim, dim).unwrap();

// PANICS
chunk.find_closest_batch(query_view);
Expand Down
Loading
Loading