1use 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#[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 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 pub fn nrows(&self) -> usize {
55 self.nrows
56 }
57
58 pub fn ncols(&self) -> usize {
60 self.ncols
61 }
62
63 pub fn cstride(&self) -> usize {
65 self.cstride
66 }
67
68 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#[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#[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 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 Ok(unsafe { Self::from_data_unchecked(data, layout) })
237 }
238 }
239
240 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 pub fn layout(&self) -> Layout<T> {
261 self.layout
262 }
263
264 pub fn as_ptr(&self) -> *const T {
266 self.as_nonnull().as_ptr().cast_const()
267 }
268
269 pub fn ncols(&self) -> usize {
271 self.layout().ncols()
272 }
273
274 pub fn nrows(&self) -> usize {
276 self.layout().nrows()
277 }
278
279 pub fn cstride(&self) -> usize {
281 self.layout().cstride()
282 }
283
284 pub fn as_slice(&self) -> &[T] {
290 let layout = self.layout();
291
292 unsafe { std::slice::from_raw_parts(self.as_ptr(), layout.linear_length()) }
295 }
296
297 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 unsafe { &*self.as_ptr().add(layout.cstride() * row + col) }
314 }
315
316 pub fn get_element(&self, row: usize, col: usize) -> Option<&T> {
318 if row < self.nrows() && col < self.ncols() {
319 Some(unsafe { self.element_unchecked(row, col) })
321 } else {
322 None
323 }
324 }
325
326 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 unsafe { self.element_unchecked(row, col) }
347 }
348
349 pub unsafe fn row_unchecked(&self, row: usize) -> &[T] {
357 let layout = self.layout();
358 debug_assert!(row < layout.nrows());
359
360 unsafe {
364 std::slice::from_raw_parts(self.as_ptr().add(layout.cstride() * row), layout.ncols())
365 }
366 }
367
368 pub fn get_row(&self, row: usize) -> Option<&[T]> {
370 if row < self.nrows() {
371 Some(unsafe { self.row_unchecked(row) })
373 } else {
374 None
375 }
376 }
377
378 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 unsafe { self.row_unchecked(row) }
393 }
394
395 pub fn rows(&self) -> Rows<'_, T> {
399 Rows::new(*self)
400 }
401}
402
403#[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 unsafe { Self::from_data_unchecked(matrix.into_slice(), layout) }
423 }
424}
425
426unsafe impl<T> Send for Strided<'_, T> where T: Sync {}
429unsafe 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#[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
470unsafe impl<T> Send for Rows<'_, T> where T: Sync {}
473unsafe 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 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 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#[cfg(test)]
509mod tests {
510 use super::*;
511
512 use crate::views::rowmajor::MatrixMut;
513
514 #[test]
515 fn test_linear_length() {
516 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 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 assert_eq!(linear_length(1, 5, 10).unwrap(), 5);
531 assert_eq!(linear_length(1, 7, 99).unwrap(), 7);
532
533 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 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 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 assert!(Layout::<usize>::new(0, 0, 0).is_ok());
566
567 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 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 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 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 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 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 assert!(v.get_element(2, 3).is_some());
634 assert!(v.get_row(2).is_some());
635
636 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 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 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 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 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 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 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 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 assert!(dut.rows().eq(expected.rows()));
808 }
809
810 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 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 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 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 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 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 let s = Strided::try_from_data(m.as_slice(), nrows, 5, ncols).unwrap();
912 assert_eq!(s.as_ptr(), m.as_ptr());
913
914 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 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}