1use std::{marker::PhantomData, mem::ManuallyDrop, num::NonZeroUsize, ptr::NonNull};
7
8#[cfg(feature = "rayon")]
9use rayon::prelude::{
10 IndexedParallelIterator, IntoParallelIterator, ParallelIterator, ParallelSliceMut,
11};
12use thiserror::Error;
13
14pub mod iter;
15
16use crate::{internal, Reborrow, ReborrowMut};
17
18pub unsafe trait Matrix {
62 type Element;
64
65 fn as_nonnull(&self) -> NonNull<Self::Element>;
67
68 fn layout(&self) -> Layout<Self::Element>;
70
71 fn nrows(&self) -> usize {
77 self.layout().nrows()
78 }
79
80 fn ncols(&self) -> usize {
82 self.layout().ncols()
83 }
84
85 unsafe fn row_unchecked(&self, row: usize) -> &[Self::Element] {
93 let layout = self.layout();
94 debug_assert!(row < layout.nrows());
95
96 unsafe {
100 std::slice::from_raw_parts(self.as_ptr().add(layout.ncols() * row), layout.ncols())
101 }
102 }
103
104 fn as_ptr(&self) -> *const Self::Element {
106 self.as_nonnull().as_ptr().cast_const()
107 }
108
109 fn as_slice(&self) -> &[Self::Element] {
111 unsafe { std::slice::from_raw_parts(self.as_ptr(), self.layout().num_elements()) }
114 }
115
116 fn row(&self, row: usize) -> &[Self::Element] {
122 assert!(
123 row < self.nrows(),
124 "tried to access row {row} of a matrix with {} rows",
125 self.nrows()
126 );
127
128 unsafe { self.row_unchecked(row) }
130 }
131
132 fn get_row(&self, row: usize) -> Option<&[Self::Element]> {
134 if row < self.nrows() {
135 Some(unsafe { self.row_unchecked(row) })
137 } else {
138 None
139 }
140 }
141
142 fn rows(&self) -> iter::Rows<'_, Self::Element> {
146 iter::Rows::new(self.as_view())
147 }
148
149 unsafe fn element_unchecked(&self, row: usize, col: usize) -> &Self::Element {
157 let layout = self.layout();
158 debug_assert!(row < layout.nrows());
159 debug_assert!(col < layout.ncols());
160
161 unsafe { &*self.as_ptr().add(row * layout.ncols() + col) }
165 }
166
167 fn get_element(&self, row: usize, col: usize) -> Option<&Self::Element> {
171 if row >= self.nrows() || col >= self.ncols() {
172 None
173 } else {
174 Some(unsafe { self.element_unchecked(row, col) })
176 }
177 }
178
179 fn element(&self, row: usize, col: usize) -> &Self::Element {
185 assert!(
186 row < self.nrows(),
187 "row {row} is out of bounds (max: {})",
188 self.nrows()
189 );
190 assert!(
191 col < self.ncols(),
192 "col {col} is out of bounds (max: {})",
193 self.ncols()
194 );
195
196 unsafe { self.element_unchecked(row, col) }
198 }
199
200 fn as_view(&self) -> Ref<'_, Self::Element> {
202 Ref {
203 ptr: self.as_nonnull(),
204 layout: self.layout(),
205 _lifetime: PhantomData,
206 }
207 }
208
209 fn subview(&self, rows: std::ops::Range<usize>) -> Option<Ref<'_, Self::Element>> {
211 if rows.start > rows.end || rows.end > self.nrows() {
212 return None;
213 }
214
215 let ncols = self.ncols();
216 let ptr =
220 unsafe { NonNull::new_unchecked(self.as_ptr().add(rows.start * ncols).cast_mut()) };
221 let layout = unsafe { Layout::new_unchecked(rows.end - rows.start, ncols) };
224 Some(Ref {
225 ptr,
226 layout,
227 _lifetime: PhantomData,
228 })
229 }
230
231 fn window_iter(&self, batchsize: NonZeroUsize) -> iter::Windows<'_, Self::Element> {
237 iter::Windows::new(self.as_view(), batchsize)
238 }
239
240 fn to_rowmajor_owned(&self) -> Owned<Self::Element>
242 where
243 Self::Element: Clone,
244 {
245 unsafe { Owned::from_data_unchecked(self.as_slice().into(), self.layout()) }
248 }
249
250 fn try_map<F, R>(&self, f: F) -> Result<Owned<R>, LayoutError>
254 where
255 F: FnMut(&Self::Element) -> R,
256 {
257 let layout = self.layout().rebind::<R>()?;
258 let data: Box<[_]> = self.as_slice().iter().map(f).collect();
259
260 Ok(unsafe { Owned::from_data_unchecked(data, layout) })
263 }
264
265 #[track_caller]
273 fn map<F, R>(&self, f: F) -> Owned<R>
274 where
275 F: FnMut(&Self::Element) -> R,
276 {
277 match self.try_map(f) {
278 Ok(owned) => owned,
279 Err(error) => panic!("`Matrix::map` failed: {error}"),
280 }
281 }
282
283 fn transpose(&self) -> Owned<Self::Element>
285 where
286 Self::Element: Clone,
287 {
288 Owned::from_fn_with_layout(self.layout().transpose(), |RowCol { row, col }| {
289 unsafe { self.element_unchecked(col, row).clone() }
291 })
292 }
293
294 #[cfg(feature = "rayon")]
300 fn par_rows(&self) -> impl IndexedParallelIterator<Item = &[Self::Element]>
301 where
302 Self::Element: Sync,
303 {
304 let r = self.as_view();
305
306 (0..r.nrows()).into_par_iter().map(move |row| {
307 unsafe { r.into_row_unchecked(row) }
309 })
310 }
311
312 #[cfg(feature = "rayon")]
325 fn par_window_iter(
326 &self,
327 batchsize: usize,
328 ) -> impl IndexedParallelIterator<Item = Ref<'_, Self::Element>>
329 where
330 Self::Element: Sync,
331 {
332 assert!(batchsize != 0, "par_window_iter batchsize cannot be zero");
333
334 let r = self.as_view();
335 (0..r.nrows())
336 .into_par_iter()
337 .step_by(batchsize)
338 .map(move |start| {
339 let end = start.saturating_add(batchsize).min(r.nrows());
340
341 unsafe { r.into_subview_unchecked(start..end) }
343 })
344 }
345}
346
347pub unsafe trait MatrixMut: Matrix {
381 fn as_nonnull_mut(&mut self) -> NonNull<Self::Element>;
389
390 unsafe fn row_unchecked_mut(&mut self, row: usize) -> &mut [Self::Element] {
402 let layout = self.layout();
403
404 debug_assert!(row < layout.nrows());
405
406 unsafe {
410 std::slice::from_raw_parts_mut(
411 self.as_mut_ptr().add(layout.ncols() * row),
412 layout.ncols(),
413 )
414 }
415 }
416
417 fn as_mut_ptr(&mut self) -> *mut Self::Element {
419 self.as_nonnull_mut().as_ptr()
420 }
421
422 fn as_mut_slice(&mut self) -> &mut [Self::Element] {
424 unsafe { std::slice::from_raw_parts_mut(self.as_mut_ptr(), self.layout().num_elements()) }
427 }
428
429 fn row_mut(&mut self, row: usize) -> &mut [Self::Element] {
435 assert!(
436 row < self.nrows(),
437 "tried to access row {row} of a matrix with {} rows",
438 self.nrows()
439 );
440
441 unsafe { self.row_unchecked_mut(row) }
443 }
444
445 fn get_row_mut(&mut self, row: usize) -> Option<&mut [Self::Element]> {
447 if row < self.nrows() {
448 Some(unsafe { self.row_unchecked_mut(row) })
450 } else {
451 None
452 }
453 }
454
455 fn rows_mut(&mut self) -> iter::RowsMut<'_, Self::Element> {
459 iter::RowsMut::new(self.as_view_mut())
460 }
461
462 unsafe fn element_unchecked_mut(&mut self, row: usize, col: usize) -> &mut Self::Element {
470 let layout = self.layout();
471 debug_assert!(row < layout.nrows());
472 debug_assert!(col < layout.ncols());
473
474 unsafe { &mut *self.as_mut_ptr().add(row * layout.ncols() + col) }
478 }
479
480 fn get_element_mut(&mut self, row: usize, col: usize) -> Option<&mut Self::Element> {
484 if row >= self.nrows() || col >= self.ncols() {
485 None
486 } else {
487 Some(unsafe { self.element_unchecked_mut(row, col) })
489 }
490 }
491
492 fn element_mut(&mut self, row: usize, col: usize) -> &mut Self::Element {
498 assert!(
499 row < self.nrows(),
500 "row {row} is out of bounds (max: {})",
501 self.nrows()
502 );
503 assert!(
504 col < self.ncols(),
505 "col {col} is out of bounds (max: {})",
506 self.ncols()
507 );
508
509 unsafe { self.element_unchecked_mut(row, col) }
511 }
512
513 fn as_view_mut(&mut self) -> Mut<'_, Self::Element> {
515 Mut {
516 ptr: self.as_nonnull_mut(),
517 layout: self.layout(),
518 _lifetime: PhantomData,
519 }
520 }
521
522 #[cfg(feature = "rayon")]
532 fn par_rows_mut(&mut self) -> impl IndexedParallelIterator<Item = &mut [Self::Element]>
533 where
534 Self::Element: Send,
535 {
536 let ncols = self.ncols();
537 assert!(
538 ncols != 0 || self.nrows() == 0,
539 "`MatrixMut::par_rows_mut` does not support matrices with rows and zero columns"
540 );
541 self.as_mut_slice().par_chunks_exact_mut(ncols.max(1))
542 }
543
544 #[cfg(feature = "rayon")]
557 fn par_window_iter_mut(
558 &mut self,
559 batchsize: usize,
560 ) -> impl IndexedParallelIterator<Item = Mut<'_, Self::Element>>
561 where
562 Self::Element: Send,
563 {
564 assert!(
565 batchsize != 0,
566 "par_window_iter_mut batchsize cannot be zero"
567 );
568
569 let ncols = self.ncols();
570 assert!(
571 ncols != 0 || self.nrows() == 0,
572 "`MatrixMut::par_window_iter_mut` does not support matrices with rows and zero columns"
573 );
574
575 let batchsize = batchsize.min(self.nrows());
577 self.as_mut_slice()
578 .par_chunks_mut((ncols * batchsize).max(1))
579 .map(move |data| {
580 let blobsize = data.len();
581 let nrows = blobsize / ncols;
582 assert_eq!(blobsize % ncols, 0);
583
584 unsafe { Mut::from_data_unchecked(data, Layout::new_unchecked(nrows, ncols)) }
593 })
594 }
595}
596
597pub struct Layout<T> {
608 nrows: usize,
609 ncols: usize,
610 _type: PhantomData<fn() -> T>,
611}
612
613impl<T> Layout<T> {
614 pub const fn new(nrows: usize, ncols: usize) -> Result<Self, LayoutError> {
620 match LayoutError::check::<T>(nrows, ncols) {
621 Ok(()) => Ok(Self {
622 nrows,
623 ncols,
624 _type: PhantomData,
625 }),
626 Err(err) => Err(err),
627 }
628 }
629
630 unsafe fn new_unchecked(nrows: usize, ncols: usize) -> Self {
636 debug_assert!(LayoutError::check::<T>(nrows, ncols).is_ok());
637 Self {
638 nrows,
639 ncols,
640 _type: PhantomData,
641 }
642 }
643
644 pub fn num_elements(&self) -> usize {
646 self.nrows() * self.ncols()
647 }
648
649 pub fn nrows(&self) -> usize {
651 self.nrows
652 }
653
654 pub fn ncols(&self) -> usize {
656 self.ncols
657 }
658
659 pub fn rebind<U>(&self) -> Result<Layout<U>, LayoutError> {
665 if std::mem::size_of::<U>() <= std::mem::size_of::<T>() {
666 Ok(unsafe { Layout::new_unchecked(self.nrows(), self.ncols()) })
672 } else {
673 Layout::new(self.nrows(), self.ncols())
674 }
675 }
676
677 pub fn transpose(&self) -> Layout<T> {
679 unsafe { Layout::new_unchecked(self.ncols, self.nrows) }
682 }
683}
684
685impl<T> Clone for Layout<T> {
686 fn clone(&self) -> Self {
687 *self
688 }
689}
690
691impl<T> Copy for Layout<T> {}
692
693impl<T> std::fmt::Debug for Layout<T> {
694 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
695 f.debug_struct("Layout")
696 .field("nrows", &self.nrows)
697 .field("ncols", &self.ncols)
698 .field("elsize", &std::mem::size_of::<T>())
699 .finish()
700 }
701}
702
703impl<T> PartialEq for Layout<T> {
704 fn eq(&self, other: &Self) -> bool {
705 self.nrows == other.nrows && self.ncols == other.ncols
706 }
707}
708
709impl<T> Eq for Layout<T> {}
710
711#[derive(Debug, Clone, Copy)]
713pub struct LayoutError {
714 nrows: usize,
715 ncols: usize,
716 elsize: Option<NonZeroUsize>,
717}
718
719impl LayoutError {
720 pub(crate) const fn check<T>(nrows: usize, ncols: usize) -> Result<(), Self> {
721 let elsize = std::mem::size_of::<T>();
723 let num_elements = match nrows.checked_mul(ncols) {
724 Some(num_elements) => num_elements,
725 None => {
726 return Err(Self {
727 nrows,
728 ncols,
729 elsize: None,
730 })
731 }
732 };
733
734 if let Some(len) = num_elements.checked_mul(std::mem::size_of::<T>()) {
735 if len <= isize::MAX as usize {
736 return Ok(());
737 }
738 }
739
740 Err(Self {
741 nrows,
742 ncols,
743 elsize: NonZeroUsize::new(elsize),
744 })
745 }
746}
747
748impl std::fmt::Display for LayoutError {
749 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
750 match self.elsize {
751 Some(elsize) => {
752 write!(
753 f,
754 "a matrix of size {}x{} with elements of size {} exceeds `isize::MAX` bytes",
755 self.nrows, self.ncols, elsize
756 )
757 }
758 None => {
759 write!(
760 f,
761 "a matrix of size {}x{} has a length exceeding `usize::MAX`",
762 self.nrows, self.ncols,
763 )
764 }
765 }
766 }
767}
768
769impl std::error::Error for LayoutError {}
770
771macro_rules! constructors {
776 ($element:ident, $data:ty) => {
777 pub fn try_from_data(
782 data: $data,
783 nrows: usize,
784 ncols: usize,
785 ) -> Result<Self, TryFromError<$data>> {
786 let layout = match Layout::<$element>::new(nrows, ncols) {
787 Ok(layout) => layout,
788 Err(err) => return Err(TryFromError::layout(data, err)),
789 };
790
791 let len = data.len();
792 if len == layout.num_elements() {
793 Ok(unsafe { Self::from_data_unchecked(data, layout) })
795 } else {
796 Err(TryFromError::mismatch(
797 data,
798 layout.nrows(),
799 layout.ncols(),
800 len,
801 ))
802 }
803 }
804
805 pub fn row_vector(data: $data) -> Self {
807 let layout = unsafe { Layout::new_unchecked(1, data.len()) };
810
811 unsafe { Self::from_data_unchecked(data, layout) }
813 }
814
815 pub fn column_vector(data: $data) -> Self {
817 let layout = unsafe { Layout::new_unchecked(data.len(), 1) };
820
821 unsafe { Self::from_data_unchecked(data, layout) }
823 }
824 };
825}
826
827#[derive(Debug, Clone, Copy, PartialEq, Eq)]
834pub struct RowCol {
835 pub row: usize,
836 pub col: usize,
837}
838
839#[derive(Debug)]
841pub struct Owned<T> {
842 ptr: NonNull<T>,
843 layout: Layout<T>,
844}
845
846impl<T> Owned<T> {
847 constructors!(T, Box<[T]>);
848
849 #[track_caller]
872 pub fn from_fn<F>(nrows: usize, ncols: usize, init: F) -> Self
873 where
874 F: FnMut(RowCol) -> T,
875 {
876 match Self::try_from_fn(nrows, ncols, init) {
877 Ok(matrix) => matrix,
878 Err(error) => panic!("Owned::from_fn failed with: {error}"),
879 }
880 }
881
882 pub fn try_from_fn<F>(nrows: usize, ncols: usize, init: F) -> Result<Self, LayoutError>
900 where
901 F: FnMut(RowCol) -> T,
902 {
903 let layout = Layout::new(nrows, ncols)?;
904 Ok(Self::from_fn_with_layout(layout, init))
905 }
906
907 #[track_caller]
925 pub fn from_element(nrows: usize, ncols: usize, element: T) -> Self
926 where
927 T: Clone,
928 {
929 match Self::try_from_element(nrows, ncols, element) {
930 Ok(matrix) => matrix,
931 Err(error) => panic!("Owned::from_element failed with: {error}"),
932 }
933 }
934
935 pub fn try_from_element(nrows: usize, ncols: usize, element: T) -> Result<Self, LayoutError>
953 where
954 T: Clone,
955 {
956 let layout = Layout::new(nrows, ncols)?;
957 Ok(Self::from_element_with_layout(layout, element))
958 }
959
960 pub fn from_fn_with_layout<F>(layout: Layout<T>, mut init: F) -> Self
966 where
967 F: FnMut(RowCol) -> T,
968 {
969 let mut row = 0;
970 let mut col = 0;
971
972 let data: Box<[T]> = (0..layout.num_elements())
973 .map(|_| {
974 let v = (init)(RowCol { row, col });
975 col += 1;
976 if col == layout.ncols() {
977 col = 0;
978 row += 1;
979 }
980 v
981 })
982 .collect();
983
984 unsafe { Self::from_data_unchecked(data, layout) }
986 }
987
988 pub fn from_element_with_layout(layout: Layout<T>, element: T) -> Self
992 where
993 T: Clone,
994 {
995 let data: Box<[T]> = std::iter::repeat_n(element, layout.num_elements()).collect();
996
997 unsafe { Self::from_data_unchecked(data, layout) }
999 }
1000
1001 unsafe fn from_data_unchecked(b: Box<[T]>, layout: Layout<T>) -> Self {
1005 debug_assert_eq!(b.len(), layout.num_elements());
1006 Self {
1007 ptr: internal::box_to_nonnull(b),
1008 layout,
1009 }
1010 }
1011
1012 pub fn into_inner(self) -> Box<[T]> {
1025 let me = ManuallyDrop::new(self);
1026
1027 unsafe { internal::nonnull_to_box(me.ptr, me.layout.num_elements()) }
1030 }
1031}
1032
1033unsafe impl<T> Send for Owned<T> where T: Send {}
1035unsafe impl<T> Sync for Owned<T> where T: Sync {}
1037
1038impl<T> Drop for Owned<T> {
1039 fn drop(&mut self) {
1040 let _ = unsafe { internal::nonnull_to_box(self.ptr, self.layout.num_elements()) };
1043 }
1044}
1045
1046impl<T> Clone for Owned<T>
1047where
1048 T: Clone,
1049{
1050 fn clone(&self) -> Self {
1051 unsafe { Owned::from_data_unchecked(self.as_slice().into(), self.layout()) }
1053 }
1054}
1055
1056unsafe impl<T> Matrix for Owned<T> {
1058 type Element = T;
1059
1060 fn as_nonnull(&self) -> NonNull<T> {
1061 self.ptr
1062 }
1063
1064 fn layout(&self) -> Layout<T> {
1065 self.layout
1066 }
1067}
1068
1069unsafe impl<T> MatrixMut for Owned<T> {
1071 fn as_nonnull_mut(&mut self) -> NonNull<T> {
1072 self.ptr
1073 }
1074}
1075
1076impl<T> PartialEq for Owned<T>
1077where
1078 T: PartialEq,
1079{
1080 fn eq(&self, other: &Self) -> bool {
1081 Matrix::as_view(self).eq(&Matrix::as_view(other))
1082 }
1083}
1084
1085impl<'a, T> Reborrow<'a> for Owned<T> {
1086 type Target = Ref<'a, T>;
1087 fn reborrow(&'a self) -> Self::Target {
1088 Matrix::as_view(self)
1089 }
1090}
1091
1092impl<'a, T> ReborrowMut<'a> for Owned<T> {
1093 type Target = Mut<'a, T>;
1094 fn reborrow_mut(&'a mut self) -> Self::Target {
1095 MatrixMut::as_view_mut(self)
1096 }
1097}
1098
1099#[derive(Debug)]
1105pub struct Ref<'a, T> {
1106 ptr: NonNull<T>,
1107 layout: Layout<T>,
1108 _lifetime: PhantomData<&'a [T]>,
1109}
1110
1111unsafe impl<T> Send for Ref<'_, T> where T: Sync {}
1114unsafe impl<T> Sync for Ref<'_, T> where T: Sync {}
1116
1117impl<'a, T> Ref<'a, T> {
1118 constructors!(T, &'a [T]);
1119
1120 unsafe fn from_data_unchecked(b: &'a [T], layout: Layout<T>) -> Self {
1124 debug_assert_eq!(b.len(), layout.num_elements());
1125 Self {
1126 ptr: internal::slice_to_nonnull(b),
1127 layout,
1128 _lifetime: PhantomData,
1129 }
1130 }
1131
1132 pub fn into_slice(self) -> &'a [T] {
1136 unsafe { std::slice::from_raw_parts(self.as_ptr(), self.layout().num_elements()) }
1138 }
1139
1140 #[cfg(feature = "rayon")]
1146 unsafe fn into_row_unchecked(self, row: usize) -> &'a [T] {
1147 let layout = self.layout();
1148 debug_assert!(row < layout.nrows());
1149
1150 unsafe {
1153 std::slice::from_raw_parts(self.as_ptr().add(layout.ncols() * row), layout.ncols())
1154 }
1155 }
1156
1157 #[cfg(feature = "rayon")]
1163 unsafe fn into_subview_unchecked(self, rows: std::ops::Range<usize>) -> Ref<'a, T> {
1164 debug_assert!(rows.start <= rows.end);
1165 debug_assert!(rows.end <= self.nrows());
1166
1167 let ncols = self.ncols();
1168 Self {
1169 ptr: unsafe { self.ptr.add(rows.start * ncols) },
1172 layout: unsafe { Layout::new_unchecked(rows.end - rows.start, ncols) },
1174 _lifetime: PhantomData,
1175 }
1176 }
1177}
1178
1179impl<T> Clone for Ref<'_, T> {
1180 fn clone(&self) -> Self {
1181 *self
1182 }
1183}
1184
1185impl<T> Copy for Ref<'_, T> {}
1186
1187unsafe impl<T> Matrix for Ref<'_, T> {
1189 type Element = T;
1190
1191 fn as_nonnull(&self) -> NonNull<T> {
1192 self.ptr
1193 }
1194
1195 fn layout(&self) -> Layout<T> {
1196 self.layout
1197 }
1198}
1199
1200impl<'a, T> Reborrow<'a> for Ref<'_, T> {
1201 type Target = Ref<'a, T>;
1202 fn reborrow(&'a self) -> Self::Target {
1203 Matrix::as_view(self)
1204 }
1205}
1206
1207impl<T> PartialEq for Ref<'_, T>
1208where
1209 T: PartialEq,
1210{
1211 fn eq(&self, other: &Self) -> bool {
1212 self.layout() == other.layout() && self.as_slice() == other.as_slice()
1213 }
1214}
1215
1216#[derive(Debug)]
1222pub struct Mut<'a, T> {
1223 ptr: NonNull<T>,
1224 layout: Layout<T>,
1225 _lifetime: PhantomData<&'a mut [T]>,
1226}
1227
1228unsafe impl<T> Send for Mut<'_, T> where T: Send {}
1231unsafe impl<T> Sync for Mut<'_, T> where T: Sync {}
1233
1234impl<'a, T> Mut<'a, T> {
1235 constructors!(T, &'a mut [T]);
1236
1237 unsafe fn from_data_unchecked(b: &'a mut [T], layout: Layout<T>) -> Self {
1241 debug_assert_eq!(b.len(), layout.num_elements());
1242 Self {
1243 ptr: internal::mut_slice_to_nonnull(b),
1244 layout,
1245 _lifetime: PhantomData,
1246 }
1247 }
1248
1249 pub fn into_mut_slice(self) -> &'a mut [T] {
1251 unsafe { std::slice::from_raw_parts_mut(self.ptr.as_ptr(), self.layout.num_elements()) }
1254 }
1255}
1256
1257unsafe impl<T> Matrix for Mut<'_, T> {
1259 type Element = T;
1260
1261 fn as_nonnull(&self) -> NonNull<T> {
1262 self.ptr
1263 }
1264
1265 fn layout(&self) -> Layout<T> {
1266 self.layout
1267 }
1268}
1269
1270unsafe impl<T> MatrixMut for Mut<'_, T> {
1272 fn as_nonnull_mut(&mut self) -> NonNull<T> {
1273 self.ptr
1274 }
1275}
1276
1277impl<T> PartialEq for Mut<'_, T>
1278where
1279 T: PartialEq,
1280{
1281 fn eq(&self, other: &Self) -> bool {
1282 Matrix::as_view(self).eq(&Matrix::as_view(other))
1283 }
1284}
1285
1286impl<'a, T> Reborrow<'a> for Mut<'_, T> {
1287 type Target = Ref<'a, T>;
1288 fn reborrow(&'a self) -> Self::Target {
1289 Matrix::as_view(self)
1290 }
1291}
1292
1293impl<'a, T> ReborrowMut<'a> for Mut<'_, T> {
1294 type Target = Mut<'a, T>;
1295 fn reborrow_mut(&'a mut self) -> Self::Target {
1296 MatrixMut::as_view_mut(self)
1297 }
1298}
1299
1300pub struct TryFromError<T> {
1306 data: T,
1307 inner: TryFromErrorInner,
1308}
1309
1310impl<T> TryFromError<T> {
1311 pub fn into_inner(self) -> T {
1313 self.data
1314 }
1315
1316 pub fn as_static(&self) -> TryFromErrorLight {
1319 TryFromErrorLight(self.inner)
1320 }
1321
1322 fn layout(data: T, error: LayoutError) -> Self {
1327 Self {
1328 data,
1329 inner: TryFromErrorInner::Layout(error),
1330 }
1331 }
1332
1333 fn mismatch(data: T, nrows: usize, ncols: usize, len: usize) -> Self {
1334 Self {
1335 data,
1336 inner: TryFromErrorInner::Mismatch { nrows, ncols, len },
1337 }
1338 }
1339}
1340
1341impl<T> std::fmt::Debug for TryFromError<T> {
1342 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1343 f.debug_struct("TryFromError")
1344 .field("data", &"<hidden>")
1345 .field("inner", &self.inner)
1346 .finish()
1347 }
1348}
1349
1350impl<T> std::fmt::Display for TryFromError<T> {
1351 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
1352 self.inner.fmt(f)
1353 }
1354}
1355
1356impl<T> std::error::Error for TryFromError<T> {}
1357
1358#[derive(Debug, Error)]
1360#[error(transparent)]
1361pub struct TryFromErrorLight(TryFromErrorInner);
1362
1363#[derive(Debug, Error, Clone, Copy)]
1364enum TryFromErrorInner {
1365 #[error(transparent)]
1366 Layout(LayoutError),
1367 #[error(
1368 "tried to construct a {}x{} matrix over a span of length {}",
1369 nrows,
1370 ncols,
1371 len
1372 )]
1373 Mismatch {
1374 nrows: usize,
1375 ncols: usize,
1376 len: usize,
1377 },
1378}
1379
1380#[cfg(test)]
1385mod tests {
1386 use super::*;
1387 use crate::{assert_contains, lazy_format};
1388
1389 fn is_copyable<T: Copy>(_x: T) -> bool {
1393 true
1394 }
1395
1396 fn _matrix_view_is_covariant<'a, 'b>(m: Ref<'a, f32>) -> Ref<'b, f32>
1398 where
1399 'a: 'b,
1400 {
1401 m
1402 }
1403
1404 fn _matrix_view_is_covariant_in_t<'a, 'b, 'm>(m: Ref<'m, &'a f32>) -> Ref<'m, &'b f32>
1405 where
1406 'a: 'b,
1407 {
1408 m
1409 }
1410
1411 fn _matrix_is_covariant_in_t<'a, 'b, 'm>(m: &'m Owned<&'a f32>) -> &'m Owned<&'b f32>
1412 where
1413 'a: 'b,
1414 {
1415 m
1416 }
1417
1418 #[test]
1423 fn test_layout() {
1424 for rows in 0..5 {
1426 for cols in 0..5 {
1427 let layout = Layout::<String>::new(rows, cols).unwrap();
1428 assert_eq!(layout.nrows(), rows);
1429 assert_eq!(layout.ncols(), cols);
1430 assert_eq!(layout.num_elements(), rows * cols);
1431
1432 let transpose = layout.transpose();
1433 assert_eq!(transpose.nrows(), cols);
1434 assert_eq!(transpose.ncols(), rows);
1435 assert_eq!(transpose.num_elements(), rows * cols);
1436
1437 let rebind = layout.rebind::<u32>().unwrap();
1438 assert_eq!(rebind.nrows(), rows);
1439 assert_eq!(rebind.ncols(), cols);
1440 assert_eq!(rebind.num_elements(), rows * cols);
1441
1442 is_copyable(layout);
1443 }
1444 }
1445
1446 #[expect(unused, reason = "we need this so the size is non-zero")]
1447 struct NotDebugOrEq(u32);
1448
1449 assert_eq!(
1450 Layout::<NotDebugOrEq>::new(10, 20).unwrap(),
1451 Layout::<NotDebugOrEq>::new(10, 20).unwrap(),
1452 );
1453
1454 assert_eq!(
1455 Layout::<NotDebugOrEq>::new(20, 0).unwrap(),
1456 Layout::<NotDebugOrEq>::new(20, 0).unwrap(),
1457 );
1458
1459 assert_ne!(
1460 Layout::<NotDebugOrEq>::new(10, 20).unwrap(),
1461 Layout::<NotDebugOrEq>::new(20, 0).unwrap(),
1462 );
1463
1464 let fmt = format!("{:?}", Layout::<NotDebugOrEq>::new(5, 6).unwrap());
1465 assert_eq!(fmt, "Layout { nrows: 5, ncols: 6, elsize: 4 }");
1466
1467 let error = Layout::<u8>::new(usize::MAX, 2).unwrap_err();
1469 assert_eq!(
1470 error.to_string(),
1471 format!(
1472 "a matrix of size {}x2 has a length exceeding `usize::MAX`",
1473 usize::MAX
1474 )
1475 );
1476
1477 let layout = Layout::<u8>::new(isize::MAX as usize, 1).unwrap();
1479 assert_eq!(layout.num_elements(), isize::MAX as usize);
1480
1481 let transpose = layout.transpose();
1482 assert_eq!(transpose.nrows(), 1);
1483 assert_eq!(transpose.ncols(), isize::MAX as usize);
1484 assert_eq!(transpose.num_elements(), layout.num_elements());
1485
1486 let error = Layout::<u8>::new(isize::MAX as usize + 1, 1).unwrap_err();
1488 assert_eq!(
1489 error.to_string(),
1490 format!(
1491 "a matrix of size {}x1 with elements of size 1 exceeds `isize::MAX` bytes",
1492 isize::MAX as usize + 1
1493 )
1494 );
1495
1496 let rebound = Layout::<u8>::new(3, 4).unwrap().rebind::<u16>().unwrap();
1498 assert_eq!(rebound.nrows(), 3);
1499 assert_eq!(rebound.ncols(), 4);
1500 assert_eq!(rebound.num_elements(), 12);
1501
1502 let error = layout.rebind::<u16>().unwrap_err();
1503 assert_eq!(
1504 error.to_string(),
1505 format!(
1506 "a matrix of size {}x1 with elements of size 2 exceeds `isize::MAX` bytes",
1507 isize::MAX
1508 )
1509 );
1510 }
1511
1512 #[test]
1517 fn test_sizes() {
1518 let expected = 3 * std::mem::size_of::<usize>();
1519 assert_eq!(std::mem::size_of::<Owned<String>>(), expected);
1520 assert_eq!(std::mem::size_of::<Option<Owned<String>>>(), expected);
1521
1522 assert_eq!(std::mem::size_of::<Ref<'_, String>>(), expected);
1523 assert_eq!(std::mem::size_of::<Option<Ref<'_, String>>>(), expected);
1524
1525 assert_eq!(std::mem::size_of::<Mut<'_, String>>(), expected);
1526 assert_eq!(std::mem::size_of::<Option<Mut<'_, String>>>(), expected);
1527 }
1528
1529 #[test]
1530 fn fallible_matrix_constructors() {
1531 let err = Owned::try_from_element(usize::MAX, usize::MAX, 0u32).unwrap_err();
1532 let msg = err.to_string();
1533 assert_contains!(msg, "exceeding `usize::MAX`");
1534
1535 let err = Owned::try_from_element(isize::MAX as usize, 1, 0u32).unwrap_err();
1536 let msg = err.to_string();
1537 assert_contains!(msg, "exceeds `isize::MAX` bytes");
1538
1539 let err = std::panic::catch_unwind(|| {
1541 Owned::from_element(usize::MAX, usize::MAX, 0u32);
1542 })
1543 .unwrap_err()
1544 .downcast::<String>()
1545 .unwrap();
1546
1547 let msg = err.to_string();
1548 assert_contains!(msg, "exceeding `usize::MAX`");
1549
1550 let err = std::panic::catch_unwind(|| {
1551 Owned::from_element(isize::MAX as usize, 1, 0u32);
1552 })
1553 .unwrap_err()
1554 .downcast::<String>()
1555 .unwrap();
1556 let msg = err.to_string();
1557 assert_contains!(msg, "exceeds `isize::MAX` bytes");
1558
1559 let err = Owned::try_from_fn(usize::MAX, usize::MAX, |_| panic!("boom")).unwrap_err();
1561 let msg = err.to_string();
1562 assert_contains!(msg, "exceeding `usize::MAX`");
1563
1564 let err = std::panic::catch_unwind(|| {
1565 Owned::from_fn(usize::MAX, usize::MAX, |_| {
1566 panic!("initializer must not run")
1567 });
1568 })
1569 .unwrap_err()
1570 .downcast::<String>()
1571 .unwrap();
1572 let msg = err.to_string();
1573 assert_contains!(msg, "Owned::from_fn failed");
1574 assert_contains!(msg, "exceeding `usize::MAX`");
1575 }
1576
1577 fn make_test_matrix() -> Vec<usize> {
1578 vec![0, 1, 2, 1, 2, 3, 2, 3, 4, 3, 4, 5]
1587 }
1588
1589 fn striped_matrix(nrows: usize, ncols: usize) -> Owned<usize> {
1590 Owned::from_fn(nrows, ncols, |rc| rc.row * 100 + rc.col)
1591 }
1592
1593 fn assert_exact_fused<I>(mut iter: I, expected: usize)
1594 where
1595 I: ExactSizeIterator + std::iter::FusedIterator,
1596 {
1597 assert_eq!(iter.len(), expected);
1598 assert_eq!(iter.size_hint(), (expected, Some(expected)));
1599
1600 for remaining in (0..expected).rev() {
1601 assert!(iter.next().is_some());
1602 assert_eq!(iter.len(), remaining);
1603 assert_eq!(iter.size_hint(), (remaining, Some(remaining)));
1604 }
1605
1606 assert!(iter.next().is_none());
1607 assert!(iter.next().is_none());
1608 assert_eq!(iter.len(), 0);
1609 assert_eq!(iter.size_hint(), (0, Some(0)));
1610 }
1611
1612 fn assert_rows_match_scalar(m: Ref<'_, usize>, rows: Vec<&[usize]>) {
1614 let context = lazy_format!("nrows = {}, ncols = {}", m.nrows(), m.ncols());
1615
1616 assert_eq!(rows.len(), m.nrows(), "{context}");
1617
1618 for (row_index, row) in rows.into_iter().enumerate() {
1619 assert_eq!(row.len(), m.ncols(), "row = {row_index} -- {context}");
1620 for (col_index, value) in row.iter().enumerate() {
1621 assert_eq!(
1622 value,
1623 m.element(row_index, col_index),
1624 "row = {row_index}, col = {col_index} -- {context}"
1625 );
1626 }
1627 }
1628 }
1629
1630 fn assert_windows_match_scalar(
1631 m: Ref<'_, usize>,
1632 batchsize: usize,
1633 windows: Vec<Ref<'_, usize>>,
1634 ) {
1635 let context = lazy_format!(
1636 "nrows = {}, ncols = {}, batchsize = {batchsize}",
1637 m.nrows(),
1638 m.ncols()
1639 );
1640
1641 assert_eq!(windows.len(), m.nrows().div_ceil(batchsize), "{context}");
1642
1643 for (window_index, window) in windows.into_iter().enumerate() {
1644 let window_context = lazy_format!("window = {window_index} -- {context}");
1645 let first_row = window_index * batchsize;
1646 let expected_rows = batchsize.min(m.nrows() - first_row);
1647 assert_eq!(window.nrows(), expected_rows, "{window_context}");
1648 assert_eq!(window.ncols(), m.ncols(), "{window_context}");
1649
1650 for row_index in 0..window.nrows() {
1651 for col_index in 0..window.ncols() {
1652 assert_eq!(
1653 window.element(row_index, col_index),
1654 m.element(first_row + row_index, col_index),
1655 "row = {row_index}, col = {col_index} -- {window_context}"
1656 );
1657 }
1658 }
1659 }
1660 }
1661
1662 #[cfg(all(not(miri), feature = "rayon"))]
1663 fn assert_parallel_rows_match_scalar(m: Ref<'_, usize>) {
1664 let rows: Vec<_> = m.par_rows().collect();
1665 assert_rows_match_scalar(m, rows);
1666 }
1667
1668 #[cfg(all(not(miri), feature = "rayon"))]
1669 fn assert_parallel_windows_match_scalar(m: Ref<'_, usize>, batchsize: usize) {
1670 let windows: Vec<_> = m.par_window_iter(batchsize).collect();
1671 assert_windows_match_scalar(m, batchsize, windows);
1672 }
1673
1674 fn test_basic_indexing<T>(m: &T)
1676 where
1677 T: Matrix<Element = usize> + Sync,
1678 {
1679 assert_eq!(m.nrows(), 4);
1680 assert_eq!(m.ncols(), 3);
1681
1682 assert_eq!(*m.element(0, 0), 0);
1684 assert_eq!(*m.element(0, 1), 1);
1685 assert_eq!(*m.element(0, 2), 2);
1686
1687 assert_eq!(*m.element(1, 0), 1);
1688 assert_eq!(*m.element(1, 1), 2);
1689 assert_eq!(*m.element(1, 2), 3);
1690
1691 assert_eq!(*m.element(2, 0), 2);
1692 assert_eq!(*m.element(2, 1), 3);
1693 assert_eq!(*m.element(2, 2), 4);
1694
1695 assert_eq!(*m.element(3, 0), 3);
1696 assert_eq!(*m.element(3, 1), 4);
1697 assert_eq!(*m.element(3, 2), 5);
1698
1699 assert_eq!(*m.get_element(0, 0).unwrap(), 0);
1700 assert_eq!(*m.get_element(0, 1).unwrap(), 1);
1701 assert_eq!(*m.get_element(0, 2).unwrap(), 2);
1702
1703 assert_eq!(*m.get_element(1, 0).unwrap(), 1);
1704 assert_eq!(*m.get_element(1, 1).unwrap(), 2);
1705 assert_eq!(*m.get_element(1, 2).unwrap(), 3);
1706
1707 assert_eq!(*m.get_element(2, 0).unwrap(), 2);
1708 assert_eq!(*m.get_element(2, 1).unwrap(), 3);
1709 assert_eq!(*m.get_element(2, 2).unwrap(), 4);
1710
1711 assert_eq!(*m.get_element(3, 0).unwrap(), 3);
1712 assert_eq!(*m.get_element(3, 1).unwrap(), 4);
1713 assert_eq!(*m.get_element(3, 2).unwrap(), 5);
1714
1715 assert_eq!(m.row(0), &[0, 1, 2]);
1717 assert_eq!(m.row(1), &[1, 2, 3]);
1718 assert_eq!(m.row(2), &[2, 3, 4]);
1719 assert_eq!(m.row(3), &[3, 4, 5]);
1720
1721 let rows: Vec<Vec<usize>> = m.rows().map(|x| x.to_vec()).collect();
1722 assert_eq!(m.row(0), &rows[0]);
1723 assert_eq!(m.row(1), &rows[1]);
1724 assert_eq!(m.row(2), &rows[2]);
1725 assert_eq!(m.row(3), &rows[3]);
1726
1727 let batchsize = 2;
1729 m.window_iter(NonZeroUsize::new(batchsize).unwrap())
1730 .enumerate()
1731 .for_each(|(i, submatrix)| {
1732 assert_eq!(submatrix.nrows(), batchsize);
1733 assert_eq!(submatrix.ncols(), m.ncols());
1734
1735 let base = i * batchsize;
1737 assert_eq!(*submatrix.element(0, 0), base);
1738 assert_eq!(*submatrix.element(0, 1), base + 1);
1739 assert_eq!(*submatrix.element(0, 2), base + 2);
1740
1741 assert_eq!(*submatrix.element(1, 0), base + 1);
1742 assert_eq!(*submatrix.element(1, 1), base + 2);
1743 assert_eq!(*submatrix.element(1, 2), base + 3);
1744 });
1745
1746 let batchsize = 3;
1749 m.window_iter(NonZeroUsize::new(batchsize).unwrap())
1750 .enumerate()
1751 .for_each(|(i, submatrix)| {
1752 if i == 0 {
1753 assert_eq!(submatrix.nrows(), batchsize);
1754 assert_eq!(submatrix.ncols(), m.ncols());
1755
1756 assert_eq!(*submatrix.element(0, 0), 0);
1758 assert_eq!(*submatrix.element(0, 1), 1);
1759 assert_eq!(*submatrix.element(0, 2), 2);
1760
1761 assert_eq!(*submatrix.element(1, 0), 1);
1762 assert_eq!(*submatrix.element(1, 1), 2);
1763 assert_eq!(*submatrix.element(1, 2), 3);
1764
1765 assert_eq!(*submatrix.element(2, 0), 2);
1766 assert_eq!(*submatrix.element(2, 1), 3);
1767 assert_eq!(*submatrix.element(2, 2), 4);
1768 } else {
1769 assert_eq!(submatrix.nrows(), 1);
1770 assert_eq!(submatrix.ncols(), m.ncols());
1771
1772 assert_eq!(*submatrix.element(0, 0), 3);
1774 assert_eq!(*submatrix.element(0, 1), 4);
1775 assert_eq!(*submatrix.element(0, 2), 5);
1776 }
1777 });
1778 }
1779
1780 #[test]
1781 fn matrix_happy_path() {
1782 let data = make_test_matrix();
1783 let m = Owned::try_from_data(data.into(), 4, 3).unwrap();
1784 test_basic_indexing(&m);
1785
1786 let ptr = m.as_ptr();
1789 let view = m.as_view();
1790 assert!(is_copyable(view));
1791 assert_eq!(view.as_ptr(), ptr);
1792 assert_eq!(view.nrows(), m.nrows());
1793 assert_eq!(view.ncols(), m.ncols());
1794 test_basic_indexing(&view);
1795 }
1796
1797 #[test]
1798 fn matrix_try_from_construction_error() {
1799 let data = make_test_matrix();
1800 let ptr = data.as_ptr();
1801 let len = data.len();
1802
1803 let m = Owned::try_from_data(data.into(), 5, 4);
1804 assert!(m.is_err());
1805 let err = m.unwrap_err();
1806 assert_eq!(
1807 err.to_string(),
1808 "tried to construct a 5x4 matrix over a span of length 12"
1809 );
1810
1811 let data = err.into_inner();
1813 assert_eq!(data.as_ptr(), ptr);
1814 assert_eq!(data.len(), len);
1815
1816 let m = Ref::try_from_data(&data, 5, 4);
1817 assert!(m.is_err());
1818 assert_eq!(
1819 m.unwrap_err().to_string(),
1820 "tried to construct a 5x4 matrix over a span of length 12"
1821 );
1822 }
1823
1824 #[test]
1825 fn mutable_matrix_direct_construction() {
1826 let mut data = make_test_matrix();
1827
1828 {
1829 let mut m = Mut::try_from_data(data.as_mut_slice(), 4, 3).unwrap();
1830 *m.element_mut(1, 2) = 30;
1831 }
1832 assert_eq!(data[5], 30);
1833
1834 let err = Mut::try_from_data(data.as_mut_slice(), 5, 4).unwrap_err();
1835 assert_eq!(
1836 err.to_string(),
1837 "tried to construct a 5x4 matrix over a span of length 12"
1838 );
1839 let recovered = err.into_inner();
1840 recovered[0] = 10;
1841 assert_eq!(data[0], 10);
1842 }
1843
1844 #[test]
1845 fn matrix_mut_view() {
1846 let mut m = Owned::<usize>::from_element(4, 3, 0);
1847 assert_eq!(m.nrows(), 4);
1848 assert_eq!(m.ncols(), 3);
1849 assert!(m.as_slice().iter().all(|&i| i == 0));
1850 let ptr = m.as_ptr();
1851 let mut_ptr = m.as_mut_ptr();
1852 assert_eq!(ptr, mut_ptr);
1853
1854 let mut view = m.as_view_mut();
1855 assert_eq!(view.nrows(), 4);
1856 assert_eq!(view.ncols(), 3);
1857 assert_eq!(view.as_ptr(), ptr);
1858 assert_eq!(view.as_mut_ptr(), mut_ptr);
1859
1860 for i in 0..view.nrows() {
1862 for j in 0..view.ncols() {
1863 *view.element_mut(i, j) = i + j;
1864 }
1865 }
1866
1867 test_basic_indexing(&m);
1869
1870 let mut m_clone = m.clone();
1872 assert_eq!(m.as_view_mut(), m_clone.as_view_mut());
1873
1874 let inner = m.into_inner();
1875 assert_eq!(inner.as_ptr(), ptr);
1876 assert_eq!(inner.len(), 4 * 3);
1877 }
1878
1879 #[test]
1880 fn matrix_view_zero_sizes() {
1881 let data: Vec<usize> = vec![];
1882 let m = Ref::try_from_data(data.as_slice(), 0, 10).unwrap();
1884 assert_eq!(m.nrows(), 0);
1885 assert_eq!(m.ncols(), 10);
1886
1887 let m = Ref::try_from_data(data.as_slice(), 3, 0).unwrap();
1889 assert_eq!(m.nrows(), 3);
1890 assert_eq!(m.ncols(), 0);
1891 let empty: &[usize] = &[];
1892 assert_eq!(m.row(0), empty);
1893 assert_eq!(m.row(1), empty);
1894 assert_eq!(m.row(2), empty);
1895
1896 let m = Ref::try_from_data(data.as_slice(), 0, 0).unwrap();
1898 assert_eq!(m.nrows(), 0);
1899 assert_eq!(m.ncols(), 0);
1900 }
1901
1902 #[test]
1903 fn matrix_construction_by_row() {
1904 let mut m = Owned::<usize>::from_element(4, 3, 0);
1905 assert!(m.as_slice().iter().all(|i| *i == 0));
1906
1907 let ncols = m.ncols();
1908 for i in 0..m.nrows() {
1909 let row = m.row_mut(i);
1910 assert_eq!(row.len(), ncols);
1911 row[0] = i;
1912 row[1] = i + 1;
1913 row[2] = i + 2;
1914 }
1915 test_basic_indexing(&m);
1916 }
1917
1918 #[test]
1920 #[should_panic(expected = "tried to access row 3 of a matrix with 3 rows")]
1921 fn test_get_row_panics() {
1922 let m = Owned::<usize>::from_element(3, 7, 0);
1923 m.row(3);
1924 }
1925
1926 #[test]
1927 #[should_panic(expected = "tried to access row 3 of a matrix with 3 rows")]
1928 fn test_get_row_mut_panics() {
1929 let mut m = Owned::<usize>::from_element(3, 7, 0);
1930 m.row_mut(3);
1931 }
1932
1933 #[test]
1934 #[should_panic(expected = "row 3 is out of bounds (max: 3)")]
1935 fn test_element_panics_row() {
1936 let m = Owned::<usize>::from_element(3, 7, 0);
1937 assert!(m.get_element(3, 2).is_none());
1938 let _ = m.element(3, 2);
1939 }
1940
1941 #[test]
1942 #[should_panic(expected = "col 7 is out of bounds (max: 7)")]
1943 fn test_element_panics_col() {
1944 let m = Owned::<usize>::from_element(3, 7, 0);
1945 assert!(m.get_element(2, 7).is_none());
1946 let _ = m.element(2, 7);
1947 }
1948
1949 #[test]
1950 #[should_panic(expected = "row 3 is out of bounds (max: 3)")]
1951 fn test_element_mut_panics_row() {
1952 let mut m = Owned::<usize>::from_element(3, 7, 0);
1953 assert!(m.get_element_mut(3, 2).is_none());
1954 *m.element_mut(3, 2) = 1;
1955 }
1956
1957 #[test]
1958 #[should_panic(expected = "col 7 is out of bounds (max: 7)")]
1959 fn test_element_mut_panics_col() {
1960 let mut m = Owned::<usize>::from_element(3, 7, 0);
1961 assert!(m.get_element_mut(2, 7).is_none());
1962 *m.element_mut(2, 7) = 1;
1963 }
1964
1965 #[test]
1966 #[cfg(feature = "rayon")]
1967 #[should_panic(expected = "par_window_iter batchsize cannot be zero")]
1968 fn test_par_window_iter_panics() {
1969 let m = Owned::<usize>::from_element(4, 4, 0);
1970 let _ = m.par_window_iter(0);
1971 }
1972
1973 #[test]
1974 #[cfg(feature = "rayon")]
1975 #[should_panic(expected = "par_window_iter_mut batchsize cannot be zero")]
1976 fn test_par_window_iter_mut_panics() {
1977 let mut m = Owned::<usize>::from_element(4, 4, 0);
1978 let _ = m.par_window_iter_mut(0);
1979 }
1980
1981 #[test]
1984 fn test_try_from_error_light() {
1985 let data = vec![1, 2, 3];
1987 let err = Ref::try_from_data(data.as_slice(), 2, 3).unwrap_err();
1988
1989 let err_static = err.as_static();
1991 let msg = err_static.to_string();
1992 assert_contains!(
1993 msg,
1994 "tried to construct a 2x3 matrix over a span of length 3",
1995 );
1996 let recovered_data = err.into_inner();
1998 assert_eq!(recovered_data, data.as_slice());
1999
2000 let err = Ref::try_from_data(data.as_slice(), 2, usize::MAX).unwrap_err();
2002 let msg = err.to_string();
2003 assert_contains!(msg, "usize::MAX");
2004
2005 assert_eq!(data.as_slice(), err.into_inner());
2006 }
2007
2008 #[test]
2009 fn test_map_errors() {
2010 #[derive(Debug, Clone, Copy)]
2011 struct Zst;
2012
2013 let b = Box::<[Zst]>::new_uninit_slice((isize::MAX as usize) + 1);
2015
2016 let b = unsafe { b.assume_init() };
2018
2019 let m = Owned::column_vector(b);
2020 let err = m.try_map(|_: &Zst| 0u8).unwrap_err();
2021 let msg = err.to_string();
2022 assert!(msg.contains("isize::MAX"), "{msg}");
2023
2024 let err = std::panic::catch_unwind(|| m.map(|_: &Zst| 0u8))
2026 .unwrap_err()
2027 .downcast::<String>()
2028 .unwrap();
2029 let msg = err.to_string();
2030 assert!(msg.contains("isize::MAX"), "{msg}");
2031 }
2032
2033 #[test]
2034 fn test_get_row_optional() {
2035 let data = make_test_matrix();
2036 let mut m = Owned::try_from_data(data.into(), 4, 3).unwrap();
2037
2038 assert_eq!(m.get_row(0), Some(&[0, 1, 2][..]));
2039 assert_eq!(m.get_row(1), Some(&[1, 2, 3][..]));
2040 assert_eq!(m.get_row(3), Some(&[3, 4, 5][..]));
2041 assert_eq!(m.get_row(4), None);
2042 assert_eq!(m.get_row(100), None);
2043
2044 let row = m.get_row_mut(1).unwrap();
2045 assert_eq!(row, &[1, 2, 3]);
2046 row[0] = 10;
2047 assert_eq!(m.row(1), &[10, 2, 3]);
2048 assert!(m.get_row_mut(4).is_none());
2049 assert!(m.get_row_mut(100).is_none());
2050 }
2051
2052 #[test]
2053 fn test_unsafe_get_unchecked_methods() {
2054 let data = make_test_matrix();
2055 let mut m = Owned::try_from_data(data.into(), 4, 3).unwrap();
2056
2057 unsafe {
2059 assert_eq!(*m.element_unchecked(0, 0), 0);
2060 assert_eq!(*m.element_unchecked(1, 2), 3);
2061 assert_eq!(*m.element_unchecked(3, 1), 4);
2062 }
2063
2064 unsafe {
2066 *m.element_unchecked_mut(0, 0) = 100;
2067 *m.element_unchecked_mut(1, 2) = 200;
2068 }
2069
2070 assert_eq!(*m.element(0, 0), 100);
2071 assert_eq!(*m.element(1, 2), 200);
2072
2073 unsafe {
2075 let row0 = m.row_unchecked(0);
2076 assert_eq!(row0[0], 100);
2077 assert_eq!(row0[1], 1);
2078 assert_eq!(row0[2], 2);
2079 }
2080
2081 unsafe {
2083 let row1 = m.row_unchecked_mut(1);
2084 row1[0] = 300;
2085 }
2086
2087 assert_eq!(*m.element(1, 0), 300);
2088 }
2089
2090 #[test]
2091 fn test_to_owned() {
2092 let data = make_test_matrix();
2093 let view = Ref::try_from_data(data.as_slice(), 4, 3).unwrap();
2094
2095 let owned: Owned<_> = view.to_rowmajor_owned();
2097 assert_eq!(owned.nrows(), view.nrows());
2098 assert_eq!(owned.ncols(), view.ncols());
2099 assert_eq!(owned.as_slice(), view.as_slice());
2100
2101 assert_ne!(owned.as_ptr(), view.as_ptr());
2103
2104 test_basic_indexing(&owned);
2106 }
2107
2108 #[test]
2109 fn test_matrix_from_conversions() {
2110 let data = make_test_matrix();
2111 let m = Owned::try_from_data(data.into(), 4, 3).unwrap();
2112
2113 let view = m.as_view();
2115 let slice: &[usize] = view.into_slice();
2116 assert_eq!(slice.len(), 12);
2117 assert_eq!(slice[0], 0);
2118 assert_eq!(slice[11], 5);
2119
2120 let data2 = make_test_matrix();
2122 let mut m2 = Owned::try_from_data(data2.into(), 4, 3).unwrap();
2123 let mut_view = m2.as_view_mut();
2124 let slice2: &mut [usize] = mut_view.into_mut_slice();
2125 assert_eq!(slice2.len(), 12);
2126 assert_eq!(slice2[0], 0);
2127 assert_eq!(slice2[11], 5);
2128 slice2[11] = 6;
2129 assert_eq!(*m2.element(3, 2), 6);
2130 }
2131
2132 #[test]
2133 fn test_row_vector() {
2134 let data = vec![1, 2, 3];
2135 let m = Ref::row_vector(data.as_slice());
2136 assert_eq!(m.nrows(), 1);
2137 assert_eq!(m.ncols(), 3);
2138 assert_eq!(m.as_slice(), &[1, 2, 3]);
2139 assert_eq!(m.row(0), &[1, 2, 3]);
2140
2141 let empty: &[i32] = &[];
2143 let m = Ref::row_vector(empty);
2144 assert_eq!(m.nrows(), 1);
2145 assert_eq!(m.ncols(), 0);
2146
2147 let m = Owned::row_vector(vec![10u64, 20].into_boxed_slice());
2149 assert_eq!(m.nrows(), 1);
2150 assert_eq!(m.ncols(), 2);
2151 assert_eq!(*m.element(0, 0), 10);
2152 assert_eq!(*m.element(0, 1), 20);
2153 }
2154
2155 #[test]
2156 fn test_column_vector() {
2157 let data = vec![1, 2, 3];
2158 let m = Ref::column_vector(data.as_slice());
2159 assert_eq!(m.nrows(), 3);
2160 assert_eq!(m.ncols(), 1);
2161 assert_eq!(m.as_slice(), &[1, 2, 3]);
2162 assert_eq!(*m.element(0, 0), 1);
2163 assert_eq!(*m.element(1, 0), 2);
2164 assert_eq!(*m.element(2, 0), 3);
2165 assert_eq!(m.row(0), &[1]);
2166 assert_eq!(m.row(1), &[2]);
2167 assert_eq!(m.row(2), &[3]);
2168
2169 let empty: &[i32] = &[];
2171 let m = Ref::column_vector(empty);
2172 assert_eq!(m.nrows(), 0);
2173 assert_eq!(m.ncols(), 1);
2174
2175 let m = Owned::column_vector(vec![10u64, 20].into_boxed_slice());
2177 assert_eq!(m.nrows(), 2);
2178 assert_eq!(m.ncols(), 1);
2179 assert_eq!(*m.element(0, 0), 10);
2180 assert_eq!(*m.element(1, 0), 20);
2181 }
2182
2183 #[test]
2184 fn test_map() {
2185 let m = Owned::try_from_data(vec![1u32, 2, 3, 4].into(), 2, 2).unwrap();
2186 let doubled = m.map(|&x| x * 2);
2187 assert_eq!(doubled.as_slice(), &[2, 4, 6, 8]);
2188 assert_eq!(doubled.nrows(), 2);
2189 assert_eq!(doubled.ncols(), 2);
2190
2191 let as_f64 = m.map(|&x| x as f64);
2193 assert_eq!(as_f64.as_slice(), &[1.0, 2.0, 3.0, 4.0]);
2194 }
2195
2196 #[test]
2197 fn test_get_element() {
2198 let mut m = Owned::try_from_data(vec![1, 2, 3, 4, 5, 6].into(), 2, 3).unwrap();
2199 assert_eq!(m.get_element(0, 0), Some(&1));
2200 assert_eq!(m.get_element(1, 2), Some(&6));
2201 assert_eq!(m.get_element(2, 0), None);
2202 assert_eq!(m.get_element(0, 3), None);
2203
2204 *m.get_element_mut(1, 2).unwrap() = 7;
2205 assert_eq!(m.get_element(1, 2), Some(&7));
2206 assert_eq!(m.get_element_mut(2, 0), None);
2207 assert_eq!(m.get_element_mut(0, 3), None);
2208 }
2209
2210 #[test]
2211 fn test_subview() {
2212 let data = make_test_matrix();
2213 let m = Owned::try_from_data(data.into(), 4, 3).unwrap();
2214
2215 {
2217 let subview = m.subview(0..4).unwrap();
2218 assert_eq!(subview.nrows(), 4);
2219 assert_eq!(subview.ncols(), 3);
2220
2221 assert_eq!(subview.row(0), &[0, 1, 2]);
2222 assert_eq!(subview.row(1), &[1, 2, 3]);
2223 assert_eq!(subview.row(2), &[2, 3, 4]);
2224 assert_eq!(subview.row(3), &[3, 4, 5]);
2225 assert!(subview.get_row(4).is_none());
2226 }
2227
2228 {
2230 let subview = m.subview(1..4).unwrap();
2231 assert_eq!(subview.nrows(), 3);
2232 assert_eq!(subview.ncols(), 3);
2233
2234 assert_eq!(subview.row(0), &[1, 2, 3]);
2235 assert_eq!(subview.row(1), &[2, 3, 4]);
2236 assert_eq!(subview.row(2), &[3, 4, 5]);
2237 assert!(subview.get_row(3).is_none());
2238 }
2239
2240 {
2242 let subview = m.subview(2..2).unwrap();
2243 assert_eq!(subview.nrows(), 0);
2244 assert_eq!(subview.ncols(), 3);
2245 }
2246
2247 {
2249 let subview = m.subview(0..0).unwrap();
2250 assert_eq!(subview.nrows(), 0);
2251 assert_eq!(subview.ncols(), 3);
2252
2253 let subview = m.subview(4..4).unwrap();
2254 assert_eq!(subview.nrows(), 0);
2255 assert_eq!(subview.ncols(), 3);
2256 }
2257
2258 assert!(m.subview(5..5).is_none());
2260
2261 assert!(m.subview(2..10).is_none());
2263
2264 #[expect(
2266 clippy::reversed_empty_ranges,
2267 reason = "we want to make sure it doesn't work"
2268 )]
2269 let empty = 3..2;
2270 assert!(m.subview(empty).is_none());
2271
2272 assert!(m.subview(usize::MAX - 1..usize::MAX).is_none());
2274 }
2275
2276 #[expect(
2277 clippy::reversed_empty_ranges,
2278 reason = "we want to make sure it doesn't work"
2279 )]
2280 #[test]
2281 fn test_subview_zero_cols() {
2282 let m = Owned::from_element(10, 0, 0u32);
2283
2284 assert!(m.subview(100..200).is_none());
2286
2287 assert!(m.subview(10..11).is_none());
2289
2290 assert!(m.subview(5..4).is_none());
2292
2293 let v = m.subview(5..10).unwrap();
2295 assert_eq!(v.nrows(), 5);
2296 assert_eq!(v.ncols(), 0);
2297
2298 let v = m.subview(10..10).unwrap();
2300 assert_eq!(v.nrows(), 0);
2301 assert_eq!(v.ncols(), 0);
2302
2303 let v = m.subview(0..10).unwrap();
2305 assert_eq!(v.nrows(), 10);
2306 assert_eq!(v.ncols(), 0);
2307 }
2308
2309 #[test]
2310 #[cfg(all(not(miri), feature = "rayon"))]
2311 fn parallel_immutable_iterators_match_scalar_indexing() {
2312 for (nrows, ncols) in [(0, 0), (0, 4), (3, 0), (1, 1), (1, 4), (4, 1), (5, 3)] {
2313 let m = striped_matrix(nrows, ncols);
2314 let view = m.as_view();
2315
2316 assert_parallel_rows_match_scalar(view);
2317 for batchsize in [1, 2, 3, usize::MAX] {
2318 assert_parallel_windows_match_scalar(view, batchsize);
2319 }
2320 }
2321 }
2322
2323 #[test]
2324 #[cfg(all(not(miri), feature = "rayon"))]
2325 fn parallel_mutable_iterators_match_scalar_indexing() {
2326 use rayon::prelude::*;
2327
2328 for (nrows, ncols) in [(0, 0), (0, 4), (1, 1), (1, 4), (4, 1), (5, 3)] {
2329 let context = lazy_format!("nrows = {nrows}, ncols = {ncols}");
2330
2331 let original = striped_matrix(nrows, ncols);
2332 let mut rows = original.clone();
2333
2334 let row_views: Vec<_> = rows.par_rows_mut().collect();
2335 assert_eq!(row_views.len(), nrows, "{context}");
2336 for (row_index, row) in row_views.into_iter().enumerate() {
2337 assert_eq!(row.len(), ncols, "row = {row_index} -- {context}");
2338 for value in row {
2339 *value = value.wrapping_add(row_index);
2340 }
2341 }
2342
2343 for row in 0..nrows {
2344 for col in 0..ncols {
2345 assert_eq!(
2346 *rows.element(row, col),
2347 original.element(row, col).wrapping_add(row),
2348 "row = {}, col = {} -- {}",
2349 row,
2350 col,
2351 context,
2352 );
2353 }
2354 }
2355
2356 for batchsize in [1, 2, 3, usize::MAX] {
2357 let context = lazy_format!("{context}, batchsize = {batchsize}");
2358
2359 let mut windows = original.clone();
2360 let window_views: Vec<_> = windows.par_window_iter_mut(batchsize).collect();
2361 assert_eq!(window_views.len(), nrows.div_ceil(batchsize), "{context}");
2362 for (window_index, mut window) in window_views.into_iter().enumerate() {
2363 let first_row = window_index * batchsize;
2364 assert_eq!(
2365 window.nrows(),
2366 batchsize.min(nrows - first_row),
2367 "window = {window_index} -- {context}"
2368 );
2369 assert_eq!(
2370 window.ncols(),
2371 ncols,
2372 "window = {window_index} -- {context}"
2373 );
2374 for value in window.as_mut_slice() {
2375 *value = value.wrapping_add(window_index);
2376 }
2377 }
2378
2379 for row in 0..nrows {
2380 for col in 0..ncols {
2381 assert_eq!(
2382 *windows.element(row, col),
2383 original.element(row, col).wrapping_add(row / batchsize),
2384 "row = {}, col = {} -- {}",
2385 row,
2386 col,
2387 context,
2388 );
2389 }
2390 }
2391 }
2392 }
2393 }
2394
2395 #[test]
2396 #[cfg(feature = "rayon")]
2397 #[should_panic(
2398 expected = "`MatrixMut::par_rows_mut` does not support matrices with rows and zero columns"
2399 )]
2400 fn par_rows_mut_rejects_nonempty_zero_column_matrix() {
2401 let mut m = striped_matrix(3, 0);
2402 let _ = m.par_rows_mut();
2403 }
2404
2405 #[test]
2406 #[cfg(feature = "rayon")]
2407 #[should_panic(
2408 expected = "`MatrixMut::par_window_iter_mut` does not support matrices with rows and zero columns"
2409 )]
2410 fn par_window_iter_mut_rejects_nonempty_zero_column_matrix() {
2411 let mut m = striped_matrix(3, 0);
2412 let _ = m.par_window_iter_mut(2);
2413 }
2414
2415 #[test]
2416 fn matrix_iterators_match_scalar_indexing() {
2417 for (nrows, ncols) in [(0, 0), (0, 4), (3, 0), (1, 1), (1, 4), (4, 1), (5, 3)] {
2418 let original = striped_matrix(nrows, ncols);
2419 let view = original.as_view();
2420
2421 let rows = view.rows().collect();
2422 assert_rows_match_scalar(view, rows);
2423 for batchsize in [1, 2, 3, usize::MAX] {
2424 let windows = view
2425 .window_iter(NonZeroUsize::new(batchsize).unwrap())
2426 .collect();
2427 assert_windows_match_scalar(view, batchsize, windows);
2428 }
2429
2430 let context = lazy_format!("nrows = {nrows}, ncols = {ncols}");
2431 let mut mutable = original.clone();
2432 let rows: Vec<_> = mutable.rows_mut().collect();
2433 assert_eq!(rows.len(), nrows, "{context}");
2434
2435 for (row_index, row) in rows.into_iter().enumerate() {
2436 assert_eq!(row.len(), ncols, "row = {row_index} -- {context}");
2437 for (col_index, value) in row.iter_mut().enumerate() {
2438 assert_eq!(
2439 *value,
2440 *original.element(row_index, col_index),
2441 "row = {row_index}, col = {col_index} -- {context}"
2442 );
2443 *value = value.wrapping_add(1);
2444 }
2445 }
2446
2447 for row in 0..nrows {
2448 for col in 0..ncols {
2449 assert_eq!(
2450 *mutable.element(row, col),
2451 original.element(row, col).wrapping_add(1),
2452 "row = {row}, col = {col} -- {context}"
2453 );
2454 }
2455 }
2456 }
2457 }
2458
2459 #[test]
2460 fn matrix_iterators_track_exact_remaining_lengths() {
2461 for (nrows, ncols) in [(0, 0), (0, 4), (3, 0), (5, 3)] {
2462 let mut m = striped_matrix(nrows, ncols);
2463
2464 assert_exact_fused(m.rows(), nrows);
2465 assert_exact_fused(m.rows_mut(), nrows);
2466
2467 for batchsize in [1, 2, usize::MAX] {
2468 assert_exact_fused(
2469 m.window_iter(NonZeroUsize::new(batchsize).unwrap()),
2470 nrows.div_ceil(batchsize),
2471 );
2472 }
2473 }
2474 }
2475
2476 #[test]
2477 fn matrix_transformations_preserve_empty_shapes() {
2478 for (nrows, ncols) in [(0, 0), (0, 4), (3, 0)] {
2479 let m = striped_matrix(nrows, ncols);
2480
2481 let owned = m.to_rowmajor_owned();
2482 assert_eq!(owned.nrows(), nrows);
2483 assert_eq!(owned.ncols(), ncols);
2484
2485 let mapped = m.map(|_| -> u8 { unreachable!("empty matrix has no elements") });
2486 assert_eq!(mapped.nrows(), nrows);
2487 assert_eq!(mapped.ncols(), ncols);
2488 }
2489 }
2490
2491 #[test]
2492 fn owned_from_fn_initializes_in_memory_order() {
2493 let mut value = 0;
2494 let m = Owned::from_fn(2, 3, |_| {
2495 let result = value;
2496 value += 1;
2497 result
2498 });
2499
2500 assert_eq!(*m.element(0, 0), 0);
2501 assert_eq!(*m.element(0, 1), 1);
2502 assert_eq!(*m.element(0, 2), 2);
2503 assert_eq!(*m.element(1, 0), 3);
2504 assert_eq!(*m.element(1, 1), 4);
2505 assert_eq!(*m.element(1, 2), 5);
2506 }
2507
2508 #[test]
2509 fn test_transpose() {
2510 {
2511 let v = Owned::from_element(0, 0, 0);
2512 let t = v.transpose();
2513 assert_eq!(t.nrows(), 0);
2514 assert_eq!(t.ncols(), 0);
2515 }
2516
2517 {
2518 let v = Owned::from_element(0, 10, 0);
2519 let t = v.transpose();
2520 assert_eq!(t.nrows(), 10);
2521 assert_eq!(t.ncols(), 0);
2522 }
2523
2524 {
2525 let v = Owned::from_element(10, 0, 0);
2526 let t = v.transpose();
2527 assert_eq!(t.nrows(), 0);
2528 assert_eq!(t.ncols(), 10);
2529 }
2530
2531 {
2532 let v = Owned::<usize>::try_from_data(Box::new([1, 2, 3, 4, 5, 6]), 2, 3).unwrap();
2533 let t = v.transpose();
2534
2535 assert_eq!(t.row(0), &[1, 4]);
2536 assert_eq!(t.row(1), &[2, 5]);
2537 assert_eq!(t.row(2), &[3, 6]);
2538 }
2539 }
2540
2541 #[test]
2542 fn test_debug_error_formatting() {
2543 let data = vec![1, 2, 3];
2545 let err = Owned::try_from_data(data.into(), 2, 3).unwrap_err();
2546 let debug_str = format!("{:?}", err);
2547 assert_contains!(debug_str, "TryFromError");
2548
2549 #[derive(Clone)]
2551 struct NonDebug(#[expect(dead_code)] i32);
2552
2553 let non_debug_data: Box<[NonDebug]> = vec![NonDebug(1), NonDebug(2)].into();
2554 let non_debug_err = match Owned::try_from_data(non_debug_data, 1, 3) {
2555 Ok(_) => panic!("should not have succeeded!"),
2556 Err(err) => err,
2557 };
2558 let debug_str = format!("{:?}", non_debug_err);
2559 assert_contains!(debug_str, "TryFromError");
2560 }
2561}