Skip to main content

diskann_utils/
strided.rs

1/*
2 * Copyright (c) Microsoft Corporation.
3 * Licensed under the MIT license.
4 */
5
6use std::{fmt, marker::PhantomData, num::NonZeroUsize, ptr::NonNull};
7use thiserror::Error;
8
9use crate::{
10    internal,
11    views::rowmajor::{self, Matrix},
12    Reborrow,
13};
14
15/// The layout for [`Strided`].
16///
17/// This struct ensures that the [`Self::cstride`] is greater than or equal to [`Self::ncols`]
18/// and that the addressable span is valid for elements of type `T`. In particular, the
19/// linear length does not overflow `usize::MAX` and its size in bytes does not exceed
20/// `isize::MAX`.
21///
22/// The linear length of [`Strided`] is given by the formula
23/// ```text
24/// self.nrows.saturating_sub(1) * self.cstride + self.nrows.min(1) * self.ncols
25/// ```
26/// This allows the last row to occupy less than a full stride.
27#[derive(Debug)]
28pub struct Layout<T> {
29    nrows: usize,
30    ncols: usize,
31    cstride: usize,
32    _type: PhantomData<fn() -> T>,
33}
34
35impl<T> Layout<T> {
36    /// Construct a new [`Layout`].
37    ///
38    /// Errors if:
39    ///
40    /// * `cstride < ncols`.
41    /// * The computation of the linear length overflows `usize::MAX`.
42    /// * The number of bytes required for the full span exceeds `isize::MAX`.
43    pub fn new(nrows: usize, ncols: usize, cstride: usize) -> Result<Self, LayoutError> {
44        LayoutError::check::<T>(nrows, ncols, cstride)?;
45        Ok(Self {
46            nrows,
47            ncols,
48            cstride,
49            _type: PhantomData,
50        })
51    }
52
53    /// Return the number of rows.
54    pub fn nrows(&self) -> usize {
55        self.nrows
56    }
57
58    /// Return the number of columns.
59    pub fn ncols(&self) -> usize {
60        self.ncols
61    }
62
63    /// Return the stride between subsequent rows.
64    pub fn cstride(&self) -> usize {
65        self.cstride
66    }
67
68    /// Return the length of the addressable span described by this [`Layout`], including
69    /// gaps between rows.
70    pub fn linear_length(&self) -> usize {
71        self.nrows.saturating_sub(1) * self.cstride + self.nrows.min(1) * self.ncols
72    }
73}
74
75impl<T> Clone for Layout<T> {
76    fn clone(&self) -> Self {
77        *self
78    }
79}
80
81impl<T> Copy for Layout<T> {}
82
83impl<T> From<rowmajor::Layout<T>> for Layout<T> {
84    fn from(layout: rowmajor::Layout<T>) -> Self {
85        Self {
86            nrows: layout.nrows(),
87            ncols: layout.ncols(),
88            cstride: layout.ncols(),
89            _type: PhantomData,
90        }
91    }
92}
93
94fn linear_length(nrows: usize, ncols: usize, cstride: usize) -> Option<usize> {
95    nrows
96        .saturating_sub(1)
97        .checked_mul(cstride)
98        .and_then(|main| main.checked_add(nrows.min(1) * ncols))
99}
100
101/// Errors for [`Layout::new`].
102#[derive(Debug)]
103pub struct LayoutError(LayoutErrorInner);
104
105impl LayoutError {
106    fn check<T>(nrows: usize, ncols: usize, cstride: usize) -> Result<usize, Self> {
107        if cstride < ncols {
108            Err(Self(LayoutErrorInner::InvalidStride { ncols, cstride }))
109        } else {
110            let linear_length = match linear_length(nrows, ncols, cstride) {
111                Some(len) => len,
112                None => {
113                    return Err(Self(LayoutErrorInner::Overflow {
114                        nrows,
115                        cstride,
116                        elsize: None,
117                    }));
118                }
119            };
120
121            let elsize = std::mem::size_of::<T>();
122            let bytes = linear_length.saturating_mul(elsize);
123            if bytes > (isize::MAX as usize) {
124                Err(Self(LayoutErrorInner::Overflow {
125                    nrows,
126                    cstride,
127                    elsize: NonZeroUsize::new(elsize),
128                }))
129            } else {
130                Ok(linear_length)
131            }
132        }
133    }
134}
135
136impl fmt::Display for LayoutError {
137    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
138        self.0.fmt(f)
139    }
140}
141
142impl std::error::Error for LayoutError {}
143
144#[derive(Debug)]
145enum LayoutErrorInner {
146    InvalidStride {
147        ncols: usize,
148        cstride: usize,
149    },
150    Overflow {
151        nrows: usize,
152        cstride: usize,
153        elsize: Option<NonZeroUsize>,
154    },
155}
156
157impl fmt::Display for LayoutErrorInner {
158    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
159        match self {
160            Self::InvalidStride { ncols, cstride } => write!(
161                f,
162                "column stride {} must be greater than or equal to number of columns {}",
163                cstride, ncols
164            ),
165            Self::Overflow {
166                nrows,
167                cstride,
168                elsize,
169            } => match elsize {
170                Some(elsize) => write!(
171                    f,
172                    "a {}x{} strided matrix with element size {} exceeds isize::MAX bytes",
173                    nrows, cstride, elsize
174                ),
175                None => write!(
176                    f,
177                    "a {}x{} strided matrix has a length exceeding usize::MAX",
178                    nrows, cstride
179                ),
180            },
181        }
182    }
183}
184
185/// A row-major strided matrix.
186///
187/// This is a generalization of the [`Matrix`] trait as it does not mandate a dense
188/// layout in memory.
189///
190/// ```text
191///            |<------ cstride ----->|
192///                     |<-- ncols -->|
193///                     +-------------+
194/// slice 0 -> a0 a1 a2 | a3 a4 a5 a6 |    ^
195/// slice 1 -> b0 b1 b2 | b3 b4 b4 b5 |    |
196/// slice 2 -> c0 c1 c2 | c3 c4 c5 c6 |  nrows
197/// slice 3 -> d0 d1 d2 | d3 d4 d5 d6 |    |
198/// slice 4 -> e0 e1 e2 | e3 e4 e5 e6 |    |
199/// slice 5 -> f0 f1 f2 | f3 f4 f5 f6 |    v
200///                     +-------------+
201///                           ^
202///                           |
203///                        Strided
204/// ```
205///
206/// This abstraction is useful when performing PQ related operations such as training or
207/// compression as it provides a convenient abstraction for working with columnar subsets
208/// of dense data in-place.
209#[derive(Debug)]
210pub struct Strided<'a, T> {
211    ptr: NonNull<T>,
212    layout: Layout<T>,
213    _lifetime: PhantomData<&'a [T]>,
214}
215
216impl<'a, T> Strided<'a, T> {
217    /// Construct a strided view over data slice, shrinking the slice as needed.
218    ///
219    /// Returns an error if the layout is invalid or `data` is not at least
220    /// [`Layout::linear_length`].
221    pub fn try_from_data(
222        data: &'a [T],
223        nrows: usize,
224        ncols: usize,
225        cstride: usize,
226    ) -> Result<Self, TryFromError> {
227        let layout = Layout::new(nrows, ncols, cstride).map_err(TryFromError::LayoutError)?;
228        let expected = layout.linear_length();
229        if data.len() < expected {
230            Err(TryFromError::InvalidLength {
231                got: data.len(),
232                expected,
233            })
234        } else {
235            // SAFETY: `data.len() >= layout.linear_length()`.
236            Ok(unsafe { Self::from_data_unchecked(data, layout) })
237        }
238    }
239
240    /// Construct a strided view over `data`, shrinking the slice as needed, without
241    /// verifying that `data.len() >= layout.linear_length()`.
242    ///
243    /// # Safety
244    ///
245    /// `data.len()` must be greater-than or equal to `layout.linear_length()`.
246    unsafe fn from_data_unchecked(data: &'a [T], layout: Layout<T>) -> Self {
247        debug_assert!(data.len() >= layout.linear_length());
248        Self {
249            ptr: internal::slice_to_nonnull(data),
250            layout,
251            _lifetime: PhantomData,
252        }
253    }
254
255    fn as_nonnull(&self) -> NonNull<T> {
256        self.ptr
257    }
258
259    /// Return the [`Layout`] for the matrix.
260    pub fn layout(&self) -> Layout<T> {
261        self.layout
262    }
263
264    /// Return a pointer to the base of the matrix.
265    pub fn as_ptr(&self) -> *const T {
266        self.as_nonnull().as_ptr().cast_const()
267    }
268
269    /// Return the number of columns in the matrix.
270    pub fn ncols(&self) -> usize {
271        self.layout().ncols()
272    }
273
274    /// Return the number of rows in the matrix.
275    pub fn nrows(&self) -> usize {
276        self.layout().nrows()
277    }
278
279    /// Return the count of elements between the start of each row.
280    pub fn cstride(&self) -> usize {
281        self.layout().cstride()
282    }
283
284    /// Return the underlying data as a slice.
285    ///
286    /// # Note
287    ///
288    /// The underlying representation for a strided matrix is not necessarily dense.
289    pub fn as_slice(&self) -> &[T] {
290        let layout = self.layout();
291
292        // SAFETY: Constructors verify that the backing memory has a length at least
293        // `layout.linear_length()`.
294        unsafe { std::slice::from_raw_parts(self.as_ptr(), layout.linear_length()) }
295    }
296
297    // Element Access
298
299    /// Return the specified element.
300    ///
301    /// # Safety
302    ///
303    /// * `row < self.nrows()`.
304    /// * `col < self.ncols()`.
305    pub unsafe fn element_unchecked(&self, row: usize, col: usize) -> &T {
306        let layout = self.layout();
307        debug_assert!(row < layout.nrows());
308        debug_assert!(col < layout.ncols());
309
310        // SAFETY: Constructors verify that the backing memory has a length of at least
311        // `layout.linear_length`. Since `row` and `col` are in-bounds, the pointer offset
312        // is valid and it is safe to return a reference.
313        unsafe { &*self.as_ptr().add(layout.cstride() * row + col) }
314    }
315
316    /// Return the specified element if `row < self.nrows()` and `col < self.ncols()`.
317    pub fn get_element(&self, row: usize, col: usize) -> Option<&T> {
318        if row < self.nrows() && col < self.ncols() {
319            // SAFETY: `row` and `col` are in-bounds.
320            Some(unsafe { self.element_unchecked(row, col) })
321        } else {
322            None
323        }
324    }
325
326    /// Return the specified element.
327    ///
328    /// # Panics
329    ///
330    /// Panics if `row >= self.nrows()` or `col >= self.ncols()`.
331    pub fn element(&self, row: usize, col: usize) -> &T {
332        assert!(
333            row < self.nrows(),
334            "row {} is out of bounds for a matrix with {} rows",
335            row,
336            self.nrows()
337        );
338        assert!(
339            col < self.ncols(),
340            "col {} is out of bounds for a matrix with {} cols",
341            col,
342            self.ncols()
343        );
344
345        // SAFETY: `row` and `col` are in-bounds.
346        unsafe { self.element_unchecked(row, col) }
347    }
348
349    // Row Access
350
351    /// Returns the requested row without boundschecking.
352    ///
353    /// # Safety
354    ///
355    /// Caller must ensure `row < self.nrows()`.
356    pub unsafe fn row_unchecked(&self, row: usize) -> &[T] {
357        let layout = self.layout();
358        debug_assert!(row < layout.nrows());
359
360        // SAFETY: Constructors verify that the backing memory has a length of at least
361        // `layout.linear_length`. Since `row` is in-bounds, the pointer offset
362        // is valid and it is safe to form a slice of length `layout.ncols()`.
363        unsafe {
364            std::slice::from_raw_parts(self.as_ptr().add(layout.cstride() * row), layout.ncols())
365        }
366    }
367
368    /// Return the requested row if `row < self.nrows()`.
369    pub fn get_row(&self, row: usize) -> Option<&[T]> {
370        if row < self.nrows() {
371            // SAFETY: `row` is in-bounds.
372            Some(unsafe { self.row_unchecked(row) })
373        } else {
374            None
375        }
376    }
377
378    /// Return row `row` as a slice.
379    ///
380    /// # Panics
381    ///
382    /// Panics if `row >= self.nrows()`.
383    pub fn row(&self, row: usize) -> &[T] {
384        assert!(
385            row < self.nrows(),
386            "row {} is out of bounds for a matrix with {} rows",
387            row,
388            self.nrows()
389        );
390
391        // SAFETY: `row` is in-bounds.
392        unsafe { self.row_unchecked(row) }
393    }
394
395    /// Return a iterator over all rows in the matrix.
396    ///
397    /// Rows are yielded sequentially beginning with row 0.
398    pub fn rows(&self) -> Rows<'_, T> {
399        Rows::new(*self)
400    }
401}
402
403/// Errors for [`Strided::try_from_data`].
404#[derive(Debug, Error)]
405pub enum TryFromError {
406    #[error(transparent)]
407    LayoutError(LayoutError),
408    #[error(
409        "argument of length {} is shorter than the expected length {}",
410        got,
411        expected
412    )]
413    InvalidLength { got: usize, expected: usize },
414}
415
416impl<'a, T> From<rowmajor::Ref<'a, T>> for Strided<'a, T> {
417    fn from(matrix: rowmajor::Ref<'a, T>) -> Self {
418        let layout = Layout::from(matrix.layout());
419
420        // SAFETY: `rowmajor::Ref` guarantees that the length of the base slice for `matrix`
421        // is exactly `layout.linear_length()`.
422        unsafe { Self::from_data_unchecked(matrix.into_slice(), layout) }
423    }
424}
425
426// SAFETY: `Strided` only exposes shared access to its `T` elements (via `&T`/`&[T]`), so
427// sharing or transferring a `Strided<T>` across threads is sound whenever `T: Sync`.
428unsafe impl<T> Send for Strided<'_, T> where T: Sync {}
429// SAFETY: See above.
430unsafe impl<T> Sync for Strided<'_, T> where T: Sync {}
431
432impl<T> Clone for Strided<'_, T> {
433    fn clone(&self) -> Self {
434        *self
435    }
436}
437
438impl<T> Copy for Strided<'_, T> {}
439
440impl<'a, T> Reborrow<'a> for Strided<'_, T> {
441    type Target = Strided<'a, T>;
442    fn reborrow(&'a self) -> Self::Target {
443        *self
444    }
445}
446
447/// Iterator for [`Strided::rows`].
448#[derive(Debug)]
449pub struct Rows<'a, T> {
450    ptr: NonNull<T>,
451    remaining: usize,
452    ncols: usize,
453    cstride: usize,
454    _lifetime: PhantomData<&'a T>,
455}
456
457impl<'a, T> Rows<'a, T> {
458    fn new(strided: Strided<'a, T>) -> Self {
459        let layout = strided.layout();
460        Self {
461            ptr: strided.as_nonnull(),
462            remaining: layout.nrows(),
463            ncols: layout.ncols(),
464            cstride: layout.cstride(),
465            _lifetime: PhantomData,
466        }
467    }
468}
469
470// SAFETY: `Rows` only yields shared `&[T]` slices borrowed from a `Strided`, so sharing or
471// transferring a `Rows<T>` across threads is sound whenever `T: Sync`.
472unsafe impl<T> Send for Rows<'_, T> where T: Sync {}
473// SAFETY: See above.
474unsafe impl<T> Sync for Rows<'_, T> where T: Sync {}
475
476impl<'a, T> Iterator for Rows<'a, T> {
477    type Item = &'a [T];
478    fn next(&mut self) -> Option<&'a [T]> {
479        self.remaining.checked_sub(1).map(|remaining| {
480            // SAFETY: The originating `Strided` guarantees `self.remaining` rows of
481            // `self.ncols` elements each are readable starting from `self.ptr`, so the
482            // current row is valid for `self.ncols` elements.
483            let item =
484                unsafe { std::slice::from_raw_parts(self.ptr.as_ptr().cast_const(), self.ncols) };
485            self.remaining = remaining;
486            if remaining != 0 {
487                // SAFETY: There is at least one more row remaining, so advancing by
488                // `self.cstride` stays within (or one-past-the-end of) the originating
489                // allocation, per `Layout::linear_length`.
490                self.ptr = unsafe { self.ptr.add(self.cstride) };
491            }
492            item
493        })
494    }
495
496    fn size_hint(&self) -> (usize, Option<usize>) {
497        (self.remaining, Some(self.remaining))
498    }
499}
500
501impl<T> ExactSizeIterator for Rows<'_, T> {}
502impl<T> std::iter::FusedIterator for Rows<'_, T> {}
503
504///////////
505// Tests //
506///////////
507
508#[cfg(test)]
509mod tests {
510    use super::*;
511
512    use crate::views::rowmajor::MatrixMut;
513
514    #[test]
515    fn test_linear_length() {
516        // If the number of rows is zero - the output should always be zero.
517        assert_eq!(linear_length(0, 1, 1).unwrap(), 0);
518        assert_eq!(linear_length(0, 2, 2).unwrap(), 0);
519        assert_eq!(linear_length(0, 2, 3).unwrap(), 0);
520        assert_eq!(linear_length(0, 2, 4).unwrap(), 0);
521
522        // If `cstride == ncols`, then the computation should be trivial.
523        for row in 1..10 {
524            for col in 1..10 {
525                assert_eq!(linear_length(row, col, col).unwrap(), row * col);
526            }
527        }
528
529        // If there is only one row, then `cstride` should be ignored.
530        assert_eq!(linear_length(1, 5, 10).unwrap(), 5);
531        assert_eq!(linear_length(1, 7, 99).unwrap(), 7);
532
533        // Otherwise, the computation is a block of `nrows - 1` chunks of `cstride` and then
534        // `ncols`. Yes - this runs a bunch of computations.
535        for row in 2..10 {
536            for col in 0..10 {
537                for cstride in col..12 {
538                    assert_eq!(
539                        linear_length(row, col, cstride).unwrap(),
540                        (row - 1) * cstride + col
541                    );
542                }
543            }
544        }
545
546        // Check OOB detection.
547        assert!(linear_length(usize::MAX, 2, 2).is_none());
548        assert!(linear_length(2, usize::MAX, 2).is_none());
549        assert!(linear_length(2, 2, usize::MAX).is_none());
550    }
551
552    #[test]
553    fn test_layout_new() {
554        // Valid layouts.
555        let layout = Layout::<usize>::new(3, 4, 4).unwrap();
556        assert_eq!(layout.nrows(), 3);
557        assert_eq!(layout.ncols(), 4);
558        assert_eq!(layout.cstride(), 4);
559        assert_eq!(layout.linear_length(), 12);
560
561        let layout = Layout::<usize>::new(3, 4, 6).unwrap();
562        assert_eq!(layout.linear_length(), 2 * 6 + 4);
563
564        // `cstride == ncols` is fine, even at zero.
565        assert!(Layout::<usize>::new(0, 0, 0).is_ok());
566
567        // Invalid stride: `cstride < ncols`.
568        let err = Layout::<usize>::new(3, 4, 3).unwrap_err();
569        assert_eq!(
570            err.to_string(),
571            "column stride 3 must be greater than or equal to number of columns 4"
572        );
573
574        // Overflow: linear length exceeds `usize::MAX`.
575        let err = Layout::<usize>::new(usize::MAX, usize::MAX, usize::MAX).unwrap_err();
576        assert_eq!(
577            err.to_string(),
578            format!(
579                "a {}x{} strided matrix has a length exceeding usize::MAX",
580                usize::MAX,
581                usize::MAX
582            )
583        );
584
585        // Overflow: bytes exceeds `isize::MAX`.
586        let err = Layout::<usize>::new(isize::MAX as usize, 1, 1).unwrap_err();
587        assert_eq!(
588            err.to_string(),
589            format!(
590                "a {}x{} strided matrix with element size {} exceeds isize::MAX bytes",
591                isize::MAX,
592                1,
593                std::mem::size_of::<usize>(),
594            )
595        );
596
597        // The element type affects whether the addressable span is valid.
598        let length = isize::MAX as usize;
599        let layout = Layout::<u8>::new(length, 1, 1).unwrap();
600        assert_eq!(layout.linear_length(), length);
601        assert!(Layout::<u8>::new(length + 1, 1, 1).is_err());
602        assert!(Layout::<u16>::new(length, 1, 1).is_err());
603    }
604
605    #[test]
606    fn test_try_from_data_errors() {
607        let m = rowmajor::Owned::<usize>::from_element(10, 10, 0);
608        let nrows = m.nrows();
609        let ncols = m.ncols();
610
611        // An invalid layout (`cstride < ncols`) is reported as a `LayoutError`, not a panic.
612        let err = Strided::try_from_data(m.as_slice(), 2, 2, 1).unwrap_err();
613        assert_eq!(
614            err.to_string(),
615            "column stride 1 must be greater than or equal to number of columns 2"
616        );
617
618        // A slice shorter than `Layout::linear_length` is reported as `InvalidLength`.
619        let err = Strided::try_from_data(m.as_slice(), nrows, ncols, ncols + 1).unwrap_err();
620
621        assert_eq!(
622            err.to_string(),
623            "argument of length 100 is shorter than the expected length 109",
624        );
625    }
626
627    #[test]
628    fn test_element_and_row_out_of_bounds() {
629        let m = create_test_matrix(3, 4);
630        let v = Strided::try_from_data(m.as_slice(), m.nrows(), m.ncols(), m.ncols()).unwrap();
631
632        // In-bounds accesses succeed.
633        assert!(v.get_element(2, 3).is_some());
634        assert!(v.get_row(2).is_some());
635
636        // Out-of-bounds row and/or col return `None` rather than panicking.
637        assert!(v.get_element(3, 0).is_none(), "row out-of-bounds");
638        assert!(v.get_element(0, 4).is_none(), "col out-of-bounds");
639        assert!(v.get_element(3, 4).is_none(), "both out-of-bounds");
640        assert!(v.get_row(3).is_none());
641    }
642
643    #[test]
644    #[should_panic(expected = "row 3 is out of bounds for a matrix with 3 rows")]
645    fn test_element_panics_on_row() {
646        let m = create_test_matrix(3, 4);
647        let v = Strided::try_from_data(m.as_slice(), m.nrows(), m.ncols(), m.ncols()).unwrap();
648        v.element(3, 0);
649    }
650
651    #[test]
652    #[should_panic(expected = "col 4 is out of bounds for a matrix with 4 cols")]
653    fn test_element_panics_on_col() {
654        let m = create_test_matrix(3, 4);
655        let v = Strided::try_from_data(m.as_slice(), m.nrows(), m.ncols(), m.ncols()).unwrap();
656        v.element(0, 4);
657    }
658
659    #[test]
660    #[should_panic(expected = "row 3 is out of bounds for a matrix with 3 rows")]
661    fn test_row_panics() {
662        let m = create_test_matrix(3, 4);
663        let v = Strided::try_from_data(m.as_slice(), m.nrows(), m.ncols(), m.ncols()).unwrap();
664        v.row(3);
665    }
666
667    #[test]
668    fn test_clone_copy_reborrow() {
669        let m = create_test_matrix(3, 4);
670        let v = Strided::try_from_data(m.as_slice(), m.nrows(), m.ncols(), m.ncols()).unwrap();
671
672        // `Copy`/`Clone` produce an independent handle to the same data.
673        let copied = v;
674        let cloned = Clone::clone(&v);
675        assert_eq!(v.as_ptr(), copied.as_ptr());
676        assert_eq!(v.as_ptr(), cloned.as_ptr());
677
678        // `Reborrow` should yield an equivalent view.
679        let reborrowed = v.reborrow();
680        assert_eq!(reborrowed.as_ptr(), v.as_ptr());
681        assert_eq!(reborrowed.nrows(), v.nrows());
682        assert_eq!(reborrowed.ncols(), v.ncols());
683    }
684
685    #[test]
686    fn test_send_sync() {
687        fn assert_send_sync<T: Send + Sync>() {}
688        assert_send_sync::<Strided<'_, u8>>();
689        assert_send_sync::<Rows<'_, u8>>();
690    }
691
692    #[test]
693    fn test_rows_iterator_properties() {
694        let m = create_test_matrix(4, 3);
695        let v = Strided::try_from_data(&m.as_slice()[1..], m.nrows(), m.ncols() - 1, m.ncols())
696            .unwrap();
697
698        let mut rows = v.rows();
699        assert_eq!(rows.len(), 4);
700        assert_eq!(rows.size_hint(), (4, Some(4)));
701
702        for expected_row in 0..4 {
703            let row = rows.next().unwrap();
704            assert_eq!(row, &m.row(expected_row)[1..]);
705        }
706
707        // Exhausted iterators keep returning `None` (`FusedIterator`).
708        assert_eq!(rows.next(), None);
709        assert_eq!(rows.next(), None);
710        assert_eq!(rows.len(), 0);
711    }
712
713    #[test]
714    fn test_rows_iterator_zero_rows() {
715        let m = create_test_matrix(5, 5);
716        let v = Strided::try_from_data(m.as_slice(), 0, 4, 5).unwrap();
717
718        let mut rows = v.rows();
719        assert_eq!(rows.len(), 0);
720        assert_eq!(rows.next(), None);
721    }
722
723    #[test]
724    fn test_rows_iterator_zero_cols() {
725        let m = create_test_matrix(5, 5);
726        let v = Strided::try_from_data(m.as_slice(), 5, 0, 5).unwrap();
727
728        let rows = v.rows();
729        assert_eq!(rows.len(), 5);
730        assert_eq!(rows.size_hint(), (5, Some(5)));
731
732        let mut count = 0;
733        for r in rows {
734            assert!(r.is_empty());
735            count += 1;
736        }
737
738        assert_eq!(count, 5);
739    }
740
741    #[test]
742    fn test_rows_iterator_zero_cstride() {
743        let m = create_test_matrix(5, 5);
744        let v = Strided::try_from_data(m.as_slice(), 5, 0, 0).unwrap();
745
746        let rows = v.rows();
747        assert_eq!(rows.len(), 5);
748        assert_eq!(rows.size_hint(), (5, Some(5)));
749
750        let mut count = 0;
751        for r in rows {
752            assert!(r.is_empty());
753            count += 1;
754        }
755
756        assert_eq!(count, 5);
757    }
758
759    // Test that the contents of `dut` match those in the dense 2d matrix.
760    fn test_indexing(dut: Strided<'_, usize>, expected: rowmajor::Ref<'_, usize>) {
761        assert_eq!(dut.nrows(), expected.nrows());
762        assert_eq!(dut.ncols(), expected.ncols());
763
764        // Check the underlying data.
765        if dut.cstride() == dut.ncols() {
766            assert_eq!(dut.as_slice(), expected.as_slice());
767        } else {
768            assert_ne!(dut.as_slice(), expected.as_slice());
769        }
770
771        // Compare via linear indexing.
772        for row in 0..dut.nrows() {
773            for col in 0..dut.ncols() {
774                let e = *expected.element(row, col);
775
776                assert_eq!(
777                    *dut.element(row, col),
778                    e,
779                    "failed on (row, col) = ({}, {})",
780                    row,
781                    col
782                );
783
784                assert_eq!(
785                    *dut.get_element(row, col).unwrap(),
786                    e,
787                    "failed on (row, col) = ({}, {})",
788                    row,
789                    col
790                );
791            }
792        }
793
794        // Compare via row.
795        for row in 0..dut.nrows() {
796            assert_eq!(dut.row(row), expected.row(row), "failed on row {}", row);
797
798            assert_eq!(
799                dut.get_row(row).unwrap(),
800                expected.row(row),
801                "failed on row {}",
802                row
803            );
804        }
805
806        // Compare via row iterators.
807        assert!(dut.rows().eq(expected.rows()));
808    }
809
810    // Create a base Matrix with the following pattern:
811    // ```text
812    //       0         1         2 ...   ncols-1
813    //   ncols   ncols+1   ncols+2 ... 2*ncols-1
814    // 2*ncols 2*ncols+1 2*ncols+2 ... 3*ncols-1
815    // ...
816    // ```
817    fn create_test_matrix(nrows: usize, ncols: usize) -> rowmajor::Owned<usize> {
818        let mut i = 0;
819        rowmajor::Owned::from_fn(nrows, ncols, |_| {
820            let v = i;
821            i += 1;
822            v
823        })
824    }
825
826    #[test]
827    fn test_basic_indexing() {
828        let m = create_test_matrix(5, 3);
829
830        // First - test a dense Strided view over the entire matrix.
831        let ptr = m.as_ptr();
832        let v = Strided::try_from_data(m.as_slice(), m.nrows(), m.ncols(), m.ncols()).unwrap();
833        assert_eq!(v.as_ptr(), ptr, "base pointer was not preserved");
834
835        assert_eq!(v.nrows(), m.nrows());
836        assert_eq!(v.ncols(), m.ncols());
837        assert_eq!(v.cstride(), m.ncols());
838        test_indexing(v, m.as_view());
839
840        // Now - create a truly strided view over the first two columns.
841        let v = Strided::try_from_data(
842            &(m.as_slice()[..(4 * m.ncols() + 2)]),
843            m.nrows(),
844            2,
845            m.ncols(),
846        )
847        .unwrap();
848        assert_eq!(v.as_ptr(), ptr, "base pointer was not preserved");
849
850        // Create the expected matrix.
851        let mut expected = rowmajor::Owned::from_element(5, 2, 0);
852        for row in 0..expected.nrows() {
853            for col in 0..expected.ncols() {
854                *expected.element_mut(row, col) = *m.element(row, col);
855            }
856        }
857        test_indexing(v, expected.as_view());
858
859        // Create a strided view over the last two columns.
860        let v = Strided::try_from_data(&(m.as_slice()[1..]), m.nrows(), 2, m.ncols()).unwrap();
861        let mut expected = rowmajor::Owned::from_element(5, 2, 0);
862        for row in 0..expected.nrows() {
863            for col in 0..expected.ncols() {
864                *expected.element_mut(row, col) = *m.element(row, col + 1);
865            }
866        }
867        test_indexing(v, expected.as_view());
868    }
869
870    #[test]
871    fn matrix_conversion() {
872        let m = create_test_matrix(3, 4);
873        let ptr = m.as_ptr();
874        let v: Strided<_> = m.as_view().into();
875        assert_eq!(v.as_ptr(), ptr);
876        assert_eq!(v.cstride(), m.ncols());
877        assert_eq!(v.layout().linear_length(), m.layout().num_elements());
878        test_indexing(v, m.as_view());
879    }
880
881    #[test]
882    fn test_zero_sized() {
883        let m = create_test_matrix(5, 5);
884        let v = Strided::try_from_data(m.as_slice(), 0, 4, 5).unwrap();
885
886        assert_eq!(v.nrows(), 0);
887        assert_eq!(v.ncols(), 4);
888        assert_eq!(v.cstride(), 5);
889
890        let v = Strided::try_from_data(m.as_slice(), 5, 0, 5).unwrap();
891        assert_eq!(v.nrows(), 5);
892        assert_eq!(v.ncols(), 0);
893        assert_eq!(v.cstride(), 5);
894
895        for row in 0..v.nrows() {
896            let empty: &[usize] = &[];
897            assert_eq!(v.get_row(row).unwrap(), empty);
898        }
899    }
900
901    #[test]
902    fn test_try_shrink_from() {
903        // Exact is okay.
904        let m = rowmajor::Owned::<usize>::from_element(10, 10, 0);
905        let nrows = m.nrows();
906        let ncols = m.ncols();
907        let s = Strided::try_from_data(m.as_slice(), nrows, ncols, ncols).unwrap();
908        assert_eq!(s.as_slice(), m.as_slice());
909
910        // Giving a slice that is too large is okay.
911        let s = Strided::try_from_data(m.as_slice(), nrows, 5, ncols).unwrap();
912        assert_eq!(s.as_ptr(), m.as_ptr());
913
914        // Too small is a problem, and is reported as an `Err`, not a panic.
915        let s = Strided::try_from_data(m.as_slice(), nrows, ncols, ncols + 1);
916        assert!(s.is_err());
917    }
918
919    #[test]
920    fn test_invalid_stride_is_an_error_not_a_panic() {
921        // Constructing a `Strided` with an invalid layout (`cstride < ncols`) returns an
922        // `Err` rather than panicking - only unwrapping the result panics.
923        let m = rowmajor::Owned::<usize>::from_element(4, 4, 0);
924        let err = Strided::try_from_data(m.as_slice(), 2, 2, 1).unwrap_err();
925        assert!(matches!(err, TryFromError::LayoutError(_)));
926    }
927}