1use alloc::borrow::Cow;
2use alloc::vec;
3use alloc::vec::Vec;
4use core::borrow::{Borrow, BorrowMut};
5use core::marker::PhantomData;
6use core::ops::Deref;
7
8use p3_field::{
9 ExtensionField, Field, PackedValue, par_scale_slice_in_place, scale_slice_in_place_single_core,
10};
11use p3_maybe_rayon::prelude::*;
12use rand::distr::{Distribution, StandardUniform};
13use rand::{Rng, RngExt};
14use serde::{Deserialize, Serialize};
15use tracing::instrument;
16
17use crate::Matrix;
18
19#[derive(Copy, Clone, Debug, PartialEq, Eq, Serialize, Deserialize)]
23pub struct DenseMatrix<T, V = Vec<T>> {
24 pub values: V,
26 pub width: usize,
30 _phantom: PhantomData<T>,
34}
35
36pub type RowMajorMatrix<T> = DenseMatrix<T>;
37pub type RowMajorMatrixView<'a, T> = DenseMatrix<T, &'a [T]>;
38pub type RowMajorMatrixViewMut<'a, T> = DenseMatrix<T, &'a mut [T]>;
39pub type RowMajorMatrixCow<'a, T> = DenseMatrix<T, Cow<'a, [T]>>;
40
41pub trait DenseStorage<T>: Borrow<[T]> + Send + Sync {
42 fn to_vec(self) -> Vec<T>;
43}
44
45impl<T: Clone + Send + Sync> DenseStorage<T> for Vec<T> {
47 fn to_vec(self) -> Self {
48 self
49 }
50}
51
52impl<T: Clone + Send + Sync> DenseStorage<T> for &[T] {
53 fn to_vec(self) -> Vec<T> {
54 <[T]>::to_vec(self)
55 }
56}
57
58impl<T: Clone + Send + Sync> DenseStorage<T> for &mut [T] {
59 fn to_vec(self) -> Vec<T> {
60 <[T]>::to_vec(self)
61 }
62}
63
64impl<T: Clone + Send + Sync> DenseStorage<T> for Cow<'_, [T]> {
65 fn to_vec(self) -> Vec<T> {
66 self.into_owned()
67 }
68}
69
70impl<T: Clone + Send + Sync + Default> DenseMatrix<T> {
71 #[must_use]
74 pub fn default(width: usize, height: usize) -> Self {
75 Self::new(vec![T::default(); width * height], width)
76 }
77}
78
79impl<T: Clone + Send + Sync, S: DenseStorage<T>> DenseMatrix<T, S> {
80 #[must_use]
87 pub fn new(values: S, width: usize) -> Self {
88 debug_assert!(values.borrow().len().is_multiple_of(width));
89 Self {
90 values,
91 width,
92 _phantom: PhantomData,
93 }
94 }
95
96 #[must_use]
98 pub fn new_row(values: S) -> Self {
99 let width = values.borrow().len();
100 Self::new(values, width)
101 }
102
103 #[must_use]
105 pub fn new_col(values: S) -> Self {
106 Self::new(values, 1)
107 }
108
109 pub fn as_view(&self) -> RowMajorMatrixView<'_, T> {
111 RowMajorMatrixView::new(self.values.borrow(), self.width)
112 }
113
114 pub fn as_view_mut(&mut self) -> RowMajorMatrixViewMut<'_, T>
116 where
117 S: BorrowMut<[T]>,
118 {
119 RowMajorMatrixViewMut::new(self.values.borrow_mut(), self.width)
120 }
121
122 pub fn copy_from<S2>(&mut self, source: &DenseMatrix<T, S2>)
124 where
125 T: Copy,
126 S: BorrowMut<[T]>,
127 S2: DenseStorage<T>,
128 {
129 assert_eq!(self.dimensions(), source.dimensions());
130 self.par_rows_mut()
133 .zip(source.par_row_slices())
134 .for_each(|(dst, src)| {
135 dst.copy_from_slice(src);
136 });
137 }
138
139 pub fn flatten_to_base<F: Field>(self) -> RowMajorMatrix<F>
141 where
142 T: ExtensionField<F>,
143 {
144 let width = self.width * T::DIMENSION;
145 let values = T::flatten_to_base(self.values.to_vec());
146 RowMajorMatrix::new(values, width)
147 }
148
149 pub fn row_slices(&self) -> impl DoubleEndedIterator<Item = &[T]> {
151 self.values.borrow().chunks_exact(self.width)
152 }
153
154 pub fn par_row_slices(&self) -> impl IndexedParallelIterator<Item = &[T]>
156 where
157 T: Sync,
158 {
159 self.values.borrow().par_chunks_exact(self.width)
160 }
161
162 pub fn row_mut(&mut self, r: usize) -> &mut [T]
167 where
168 S: BorrowMut<[T]>,
169 {
170 &mut self.values.borrow_mut()[r * self.width..(r + 1) * self.width]
171 }
172
173 pub fn rows_mut(&mut self) -> impl Iterator<Item = &mut [T]>
175 where
176 S: BorrowMut<[T]>,
177 {
178 self.values.borrow_mut().chunks_exact_mut(self.width)
179 }
180
181 pub fn par_rows_mut<'a>(&'a mut self) -> impl IndexedParallelIterator<Item = &'a mut [T]>
183 where
184 T: 'a + Send,
185 S: BorrowMut<[T]>,
186 {
187 self.values.borrow_mut().par_chunks_exact_mut(self.width)
188 }
189
190 pub fn horizontally_packed_row_mut<P>(&mut self, r: usize) -> (&mut [P], &mut [T])
195 where
196 P: PackedValue<Value = T>,
197 S: BorrowMut<[T]>,
198 {
199 P::pack_slice_with_suffix_mut(self.row_mut(r))
200 }
201
202 pub fn scale_row(&mut self, r: usize, scale: T)
207 where
208 T: Field,
209 S: BorrowMut<[T]>,
210 {
211 scale_slice_in_place_single_core(self.row_mut(r), scale);
212 }
213
214 pub fn par_scale_row(&mut self, r: usize, scale: T)
223 where
224 T: Field,
225 S: BorrowMut<[T]>,
226 {
227 par_scale_slice_in_place(self.row_mut(r), scale);
228 }
229
230 pub fn scale(&mut self, scale: T)
232 where
233 T: Field,
234 S: BorrowMut<[T]>,
235 {
236 par_scale_slice_in_place(self.values.borrow_mut(), scale);
237 }
238
239 pub fn split_rows(&self, r: usize) -> (RowMajorMatrixView<'_, T>, RowMajorMatrixView<'_, T>) {
244 let (lo, hi) = self.values.borrow().split_at(r * self.width);
245 (
246 DenseMatrix::new(lo, self.width),
247 DenseMatrix::new(hi, self.width),
248 )
249 }
250
251 pub fn split_rows_mut(
256 &mut self,
257 r: usize,
258 ) -> (RowMajorMatrixViewMut<'_, T>, RowMajorMatrixViewMut<'_, T>)
259 where
260 S: BorrowMut<[T]>,
261 {
262 let (lo, hi) = self.values.borrow_mut().split_at_mut(r * self.width);
263 (
264 DenseMatrix::new(lo, self.width),
265 DenseMatrix::new(hi, self.width),
266 )
267 }
268
269 pub fn par_row_chunks(
273 &self,
274 chunk_rows: usize,
275 ) -> impl IndexedParallelIterator<Item = RowMajorMatrixView<'_, T>>
276 where
277 T: Send,
278 {
279 self.values
280 .borrow()
281 .par_chunks(self.width * chunk_rows)
282 .map(|slice| RowMajorMatrixView::new(slice, self.width))
283 }
284
285 pub fn par_row_chunks_exact(
289 &self,
290 chunk_rows: usize,
291 ) -> impl IndexedParallelIterator<Item = RowMajorMatrixView<'_, T>>
292 where
293 T: Send,
294 {
295 self.values
296 .borrow()
297 .par_chunks_exact(self.width * chunk_rows)
298 .map(|slice| RowMajorMatrixView::new(slice, self.width))
299 }
300
301 pub fn par_row_chunks_mut(
305 &mut self,
306 chunk_rows: usize,
307 ) -> impl IndexedParallelIterator<Item = RowMajorMatrixViewMut<'_, T>>
308 where
309 T: Send,
310 S: BorrowMut<[T]>,
311 {
312 self.values
313 .borrow_mut()
314 .par_chunks_mut(self.width * chunk_rows)
315 .map(|slice| RowMajorMatrixViewMut::new(slice, self.width))
316 }
317
318 pub fn row_chunks_exact_mut(
323 &mut self,
324 chunk_rows: usize,
325 ) -> impl Iterator<Item = RowMajorMatrixViewMut<'_, T>>
326 where
327 T: Send,
328 S: BorrowMut<[T]>,
329 {
330 self.values
331 .borrow_mut()
332 .chunks_exact_mut(self.width * chunk_rows)
333 .map(|slice| RowMajorMatrixViewMut::new(slice, self.width))
334 }
335
336 pub fn par_row_chunks_exact_mut(
341 &mut self,
342 chunk_rows: usize,
343 ) -> impl IndexedParallelIterator<Item = RowMajorMatrixViewMut<'_, T>>
344 where
345 T: Send,
346 S: BorrowMut<[T]>,
347 {
348 self.values
349 .borrow_mut()
350 .par_chunks_exact_mut(self.width * chunk_rows)
351 .map(|slice| RowMajorMatrixViewMut::new(slice, self.width))
352 }
353
354 pub fn row_pair_mut(&mut self, row_1: usize, row_2: usize) -> (&mut [T], &mut [T])
359 where
360 S: BorrowMut<[T]>,
361 {
362 debug_assert_ne!(row_1, row_2);
363 let start_1 = row_1 * self.width;
364 let start_2 = row_2 * self.width;
365 let (lo, hi) = self.values.borrow_mut().split_at_mut(start_2);
366 (&mut lo[start_1..][..self.width], &mut hi[..self.width])
367 }
368
369 #[allow(clippy::type_complexity)]
376 pub fn packed_row_pair_mut<P>(
377 &mut self,
378 row_1: usize,
379 row_2: usize,
380 ) -> ((&mut [P], &mut [T]), (&mut [P], &mut [T]))
381 where
382 S: BorrowMut<[T]>,
383 P: PackedValue<Value = T>,
384 {
385 let (slice_1, slice_2) = self.row_pair_mut(row_1, row_2);
386 (
387 P::pack_slice_with_suffix_mut(slice_1),
388 P::pack_slice_with_suffix_mut(slice_2),
389 )
390 }
391
392 #[instrument(level = "debug", skip_all)]
395 pub fn bit_reversed_zero_pad(self, added_bits: usize) -> RowMajorMatrix<T>
396 where
397 T: Field,
398 {
399 if added_bits == 0 {
400 return self.to_row_major_matrix();
401 }
402
403 let w = self.width;
413 let mut padded =
414 RowMajorMatrix::new(T::zero_vec(self.values.borrow().len() << added_bits), w);
415 padded
416 .par_row_chunks_exact_mut(1 << added_bits)
417 .zip(self.par_row_slices())
418 .for_each(|(mut ch, r)| ch.row_mut(0).copy_from_slice(r));
419
420 padded
421 }
422}
423
424impl<T: Clone + Send + Sync, S: DenseStorage<T>> Matrix<T> for DenseMatrix<T, S> {
425 #[inline]
426 fn width(&self) -> usize {
427 self.width
428 }
429
430 #[inline]
431 fn height(&self) -> usize {
432 self.values
433 .borrow()
434 .len()
435 .checked_div(self.width)
436 .unwrap_or(0)
437 }
438
439 #[inline]
440 unsafe fn get_unchecked(&self, r: usize, c: usize) -> T {
441 unsafe {
442 self.values
444 .borrow()
445 .get_unchecked(r * self.width + c)
446 .clone()
447 }
448 }
449
450 #[inline]
451 unsafe fn row_subseq_unchecked(
452 &self,
453 r: usize,
454 start: usize,
455 end: usize,
456 ) -> impl IntoIterator<Item = T, IntoIter = impl Iterator<Item = T> + Send + Sync> {
457 unsafe {
458 self.values
460 .borrow()
461 .get_unchecked(r * self.width + start..r * self.width + end)
462 .iter()
463 .cloned()
464 }
465 }
466
467 #[inline]
468 unsafe fn row_subslice_unchecked(
469 &self,
470 r: usize,
471 start: usize,
472 end: usize,
473 ) -> impl Deref<Target = [T]> {
474 unsafe {
475 self.values
477 .borrow()
478 .get_unchecked(r * self.width + start..r * self.width + end)
479 }
480 }
481
482 fn to_row_major_matrix(self) -> RowMajorMatrix<T>
483 where
484 Self: Sized,
485 T: Clone,
486 {
487 RowMajorMatrix::new(self.values.to_vec(), self.width)
488 }
489
490 #[inline]
491 fn horizontally_packed_row<'a, P>(
492 &'a self,
493 r: usize,
494 ) -> (
495 impl Iterator<Item = P> + Send + Sync,
496 impl Iterator<Item = T> + Send + Sync,
497 )
498 where
499 P: PackedValue<Value = T>,
500 T: Clone + 'a,
501 {
502 let buf = &self.values.borrow()[r * self.width..(r + 1) * self.width];
503 let (packed, sfx) = P::pack_slice_with_suffix(buf);
504 (packed.iter().copied(), sfx.iter().cloned())
505 }
506
507 #[inline]
508 fn padded_horizontally_packed_row<'a, P>(
509 &'a self,
510 r: usize,
511 ) -> impl Iterator<Item = P> + Send + Sync
512 where
513 P: PackedValue<Value = T>,
514 T: Clone + Default + 'a,
515 {
516 let buf = &self.values.borrow()[r * self.width..(r + 1) * self.width];
517 let (packed, sfx) = P::pack_slice_with_suffix(buf);
518 packed.iter().copied().chain(
519 (!sfx.is_empty()).then(|| P::from_fn(|i| sfx.get(i).cloned().unwrap_or_default())),
520 )
521 }
522
523 #[inline]
524 fn vertically_packed_row<P>(&self, r: usize) -> impl Iterator<Item = P>
525 where
526 T: Copy,
527 P: PackedValue<Value = T>,
528 {
529 let values = self.values.borrow();
530 let width = self.width;
531 let height = self.height();
532 let row = r % height;
533 let no_wrap = P::WIDTH != 1 && r + P::WIDTH <= height;
534 let rows = (!no_wrap && P::WIDTH != 1).then(|| self.wrapping_row_slices(r, P::WIDTH));
535
536 (0..width).map(move |c| {
537 if P::WIDTH == 1 {
538 unsafe { P::broadcast(*values.get_unchecked(row * width + c)) }
540 } else if no_wrap {
541 P::from_fn(|i| unsafe { *values.get_unchecked((r + i) * width + c) })
543 } else {
544 let rows = rows.as_ref().unwrap();
545 P::from_fn(|i| rows[i][c])
546 }
547 })
548 }
549
550 #[inline]
551 fn vertically_packed_row_pair<P>(&self, r: usize, step: usize) -> Vec<P>
552 where
553 T: Copy,
554 P: PackedValue<Value = T>,
555 {
556 let values = self.values.borrow();
557 let width = self.width;
558 let height = self.height();
559
560 if P::WIDTH == 1 {
561 let row = r % height;
562 let next_row = (r + step) % height;
563 let mut out = Vec::with_capacity(width * 2);
564 out.extend(
565 (0..width).map(|c| unsafe { P::broadcast(*values.get_unchecked(row * width + c)) }),
567 );
568 out.extend(
569 (0..width)
571 .map(|c| unsafe { P::broadcast(*values.get_unchecked(next_row * width + c)) }),
572 );
573 out
574 } else if r + P::WIDTH <= height && r + step + P::WIDTH <= height {
575 (0..width)
578 .map(|c| P::from_fn(|i| unsafe { *values.get_unchecked((r + i) * width + c) }))
579 .chain((0..width).map(|c| {
580 P::from_fn(|i| unsafe { *values.get_unchecked((r + step + i) * width + c) })
581 }))
582 .collect::<Vec<_>>()
583 } else {
584 let rows = self.wrapping_row_slices(r, P::WIDTH);
585 let next_rows = self.wrapping_row_slices(r + step, P::WIDTH);
586 (0..width)
587 .map(|c| P::from_fn(|i| rows[i][c]))
588 .chain((0..width).map(|c| P::from_fn(|i| next_rows[i][c])))
589 .collect::<Vec<_>>()
590 }
591 }
592}
593
594impl<T: Clone + Default + Send + Sync> DenseMatrix<T> {
595 pub fn as_cow<'a>(self) -> RowMajorMatrixCow<'a, T> {
596 RowMajorMatrixCow::new(Cow::Owned(self.values), self.width)
597 }
598
599 pub fn rand<R: Rng>(rng: &mut R, rows: usize, cols: usize) -> Self
600 where
601 StandardUniform: Distribution<T>,
602 {
603 let values = rng.sample_iter(StandardUniform).take(rows * cols).collect();
604 Self::new(values, cols)
605 }
606
607 pub fn rand_nonzero<R: Rng>(rng: &mut R, rows: usize, cols: usize) -> Self
608 where
609 T: Field,
610 StandardUniform: Distribution<T>,
611 {
612 let values = rng
613 .sample_iter(StandardUniform)
614 .filter(|x| !x.is_zero())
615 .take(rows * cols)
616 .collect();
617 Self::new(values, cols)
618 }
619
620 #[instrument(level = "debug", skip_all)]
645 pub fn with_random_cols<R>(&self, num_cols: usize, mut rng: R) -> Self
646 where
647 T: Field,
648 R: Rng + Send + Sync,
649 StandardUniform: Distribution<T>,
650 {
651 let old_w = self.width();
653 let new_w = old_w + num_cols;
654
655 let new_values = T::zero_vec(new_w * self.height());
657 let mut result = Self::new(new_values, new_w);
658
659 result
662 .par_rows_mut()
663 .zip(self.par_row_slices())
664 .for_each(|(new_row, old_row)| {
665 new_row[..old_w].copy_from_slice(old_row);
666 });
667
668 result.rows_mut().for_each(|new_row| {
671 new_row[old_w..].iter_mut().for_each(|v| *v = rng.random());
672 });
673
674 result
675 }
676
677 #[instrument(level = "debug", skip_all)]
697 pub fn with_zero_cols(&self, num_cols: usize) -> Self
698 where
699 T: Field,
700 {
701 if num_cols == 0 {
702 return self.clone();
703 }
704
705 let old_width = self.width();
706 let new_width = old_width + num_cols;
707 let source_bytes = core::mem::size_of_val(self.values.as_slice());
708
709 if cfg!(target_arch = "aarch64") && current_num_threads() == 1 && source_bytes > 1024 * 1024
712 {
713 let mut result = self.clone();
714 result.widen_right(num_cols, T::ZERO);
715 return result;
716 }
717
718 let mut result = Self::new(T::zero_vec(self.height() * new_width), new_width);
719 if self.values.is_empty() {
720 return result;
721 }
722
723 let serial_copy_bytes = if cfg!(target_arch = "aarch64") {
726 3 * 512 * 1024
727 } else {
728 1024 * 1024
729 };
730 if source_bytes <= serial_copy_bytes {
731 result
732 .values
733 .chunks_exact_mut(new_width)
734 .zip(self.values.chunks_exact(old_width))
735 .for_each(|(destination, source)| {
736 destination[..old_width].copy_from_slice(source);
737 });
738 } else {
739 result
740 .values
741 .par_chunks_exact_mut(new_width)
742 .zip(self.values.par_chunks_exact(old_width))
743 .for_each(|(destination, source)| {
744 destination[..old_width].copy_from_slice(source);
745 });
746 }
747
748 result
749 }
750
751 pub fn pad_to_height(&mut self, new_height: usize, fill: T) {
752 assert!(new_height >= self.height());
753 self.values.resize(self.width * new_height, fill);
754 }
755
756 pub fn pad_to_power_of_two_height(&mut self, fill: T) {
766 let target_height = self.height().next_power_of_two();
768
769 self.values.resize(self.width * target_height, fill);
772 }
773
774 pub fn pad_to_min_power_of_two_height(&mut self, min_height: usize, fill: T) {
792 let target_height = self
794 .height()
795 .next_power_of_two()
796 .max(min_height.next_power_of_two());
797
798 self.values.resize(self.width * target_height, fill);
800 }
801
802 #[must_use]
822 pub fn from_flat_padded(mut values: Vec<T>, width: usize, fill: T) -> Self {
823 assert!(width > 0, "width must be positive");
825
826 let len = values.len();
828 let rem = len % width;
829
830 if rem != 0 {
835 values.resize(len + (width - rem), fill.clone());
836 }
837
838 if values.is_empty() {
840 values.resize(width, fill);
841 }
842
843 Self::new(values, width)
844 }
845
846 pub fn widen_right(&mut self, extra_cols: usize, fill: T)
875 where
876 T: Copy,
877 {
878 if extra_cols == 0 {
880 return;
881 }
882
883 let old_w = self.width;
884 let new_w = old_w + extra_cols;
885 let h = self.height();
886
887 self.values.resize(h * new_w, fill);
892
893 for r in (1..h).rev() {
898 let src_start = r * old_w;
900
901 let dst_start = r * new_w;
903
904 self.values
906 .copy_within(src_start..src_start + old_w, dst_start);
907
908 self.values[dst_start - extra_cols..dst_start].fill(fill);
910 }
911
912 if h == 1 {
916 self.values[old_w..new_w].fill(fill);
917 }
918
919 self.width = new_w;
921 }
922}
923
924impl<T: Copy + Default + Send + Sync, V: DenseStorage<T>> DenseMatrix<T, V> {
925 pub fn transpose(&self) -> RowMajorMatrix<T> {
927 let nelts = self.height() * self.width();
928 let mut values = vec![T::default(); nelts];
929 p3_util::transpose::transpose(
930 self.values.borrow(),
931 &mut values,
932 self.width(),
933 self.height(),
934 );
935 RowMajorMatrix::new(values, self.height())
936 }
937
938 pub fn transpose_into<W: DenseStorage<T> + BorrowMut<[T]>>(
940 &self,
941 other: &mut DenseMatrix<T, W>,
942 ) {
943 assert_eq!(self.height(), other.width());
944 assert_eq!(other.height(), self.width());
945 p3_util::transpose::transpose(
946 self.values.borrow(),
947 other.values.borrow_mut(),
948 self.width(),
949 self.height(),
950 );
951 }
952}
953
954impl<'a, T: Clone + Default + Send + Sync> RowMajorMatrixView<'a, T> {
955 pub fn as_cow(self) -> RowMajorMatrixCow<'a, T> {
956 RowMajorMatrixCow::new(Cow::Borrowed(self.values), self.width)
957 }
958}
959
960#[cfg(test)]
961mod tests {
962 use p3_baby_bear::BabyBear;
963 use p3_field::{FieldArray, PrimeCharacteristicRing};
964 use rand::SeedableRng;
965 use rand::rngs::SmallRng;
966
967 use super::*;
968
969 #[test]
970 fn test_new() {
971 let matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 2);
972 assert_eq!(matrix.width, 2);
973 assert_eq!(matrix.height(), 3);
974 assert_eq!(matrix.values, vec![1, 2, 3, 4, 5, 6]);
975 }
976
977 #[test]
978 fn test_new_row() {
979 let matrix = RowMajorMatrix::new_row(vec![1, 2, 3]);
980 assert_eq!(matrix.width, 3);
981 assert_eq!(matrix.height(), 1);
982 }
983
984 #[test]
985 fn test_new_col() {
986 let matrix = RowMajorMatrix::new_col(vec![1, 2, 3]);
987 assert_eq!(matrix.width, 1);
988 assert_eq!(matrix.height(), 3);
989 }
990
991 #[test]
992 fn test_height_with_zero_width() {
993 let matrix: DenseMatrix<i32> = RowMajorMatrix::new(vec![], 0);
994 assert_eq!(matrix.height(), 0);
995 }
996
997 #[test]
998 fn test_get_methods() {
999 let matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 2); assert_eq!(matrix.get(0, 0), Some(1));
1001 assert_eq!(matrix.get(1, 1), Some(4));
1002 assert_eq!(matrix.get(2, 0), Some(5));
1003 unsafe {
1004 assert_eq!(matrix.get_unchecked(0, 1), 2);
1005 assert_eq!(matrix.get_unchecked(1, 0), 3);
1006 assert_eq!(matrix.get_unchecked(2, 1), 6);
1007 }
1008 assert_eq!(matrix.get(3, 0), None); assert_eq!(matrix.get(0, 2), None); }
1011
1012 #[test]
1013 fn test_row_methods() {
1014 let matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6, 7, 8], 4); let row: Vec<_> = matrix.row(1).unwrap().into_iter().collect();
1016 assert_eq!(row, vec![5, 6, 7, 8]);
1017 unsafe {
1018 let row: Vec<_> = matrix.row_unchecked(0).into_iter().collect();
1019 assert_eq!(row, vec![1, 2, 3, 4]);
1020 let row: Vec<_> = matrix.row_subseq_unchecked(0, 0, 3).into_iter().collect();
1021 assert_eq!(row, vec![1, 2, 3]);
1022 let row: Vec<_> = matrix.row_subseq_unchecked(0, 1, 3).into_iter().collect();
1023 assert_eq!(row, vec![2, 3]);
1024 let row: Vec<_> = matrix.row_subseq_unchecked(0, 2, 4).into_iter().collect();
1025 assert_eq!(row, vec![3, 4]);
1026 }
1027 assert!(matrix.row(2).is_none()); }
1029
1030 #[test]
1031 fn test_row_slice_methods() {
1032 let matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6, 7, 8, 9], 3); let slice0 = matrix.row_slice(0);
1034 let slice2 = matrix.row_slice(2);
1035 assert_eq!(slice0.unwrap().deref(), &[1, 2, 3]);
1036 assert_eq!(slice2.unwrap().deref(), &[7, 8, 9]);
1037 unsafe {
1038 assert_eq!(&[1, 2, 3], matrix.row_slice_unchecked(0).deref());
1039 assert_eq!(&[7, 8, 9], matrix.row_slice_unchecked(2).deref());
1040
1041 assert_eq!(&[1, 2, 3], matrix.row_subslice_unchecked(0, 0, 3).deref());
1042 assert_eq!(&[8], matrix.row_subslice_unchecked(2, 1, 2).deref());
1043 }
1044 assert!(matrix.row_slice(3).is_none()); }
1046
1047 #[test]
1048 fn test_as_view() {
1049 let matrix = RowMajorMatrix::new(vec![1, 2, 3, 4], 2);
1050 let view = matrix.as_view();
1051 assert_eq!(view.values, &[1, 2, 3, 4]);
1052 assert_eq!(view.width, 2);
1053 }
1054
1055 #[test]
1056 fn test_as_view_mut() {
1057 let mut matrix = RowMajorMatrix::new(vec![1, 2, 3, 4], 2);
1058 let view = matrix.as_view_mut();
1059 view.values[0] = 10;
1060 assert_eq!(matrix.values, vec![10, 2, 3, 4]);
1061 }
1062
1063 #[test]
1064 fn test_copy_from() {
1065 let mut matrix1 = RowMajorMatrix::new(vec![0, 0, 0, 0], 2);
1066 let matrix2 = RowMajorMatrix::new(vec![1, 2, 3, 4], 2);
1067 matrix1.copy_from(&matrix2);
1068 assert_eq!(matrix1.values, vec![1, 2, 3, 4]);
1069 }
1070
1071 #[test]
1072 fn test_split_rows() {
1073 let matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 2);
1074 let (top, bottom) = matrix.split_rows(1);
1075 assert_eq!(top.values, vec![1, 2]);
1076 assert_eq!(bottom.values, vec![3, 4, 5, 6]);
1077 }
1078
1079 #[test]
1080 fn test_split_rows_mut() {
1081 let mut matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 2);
1082 let (top, bottom) = matrix.split_rows_mut(1);
1083 assert_eq!(top.values, vec![1, 2]);
1084 assert_eq!(bottom.values, vec![3, 4, 5, 6]);
1085 }
1086
1087 #[test]
1088 fn test_row_mut() {
1089 let mut matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 2);
1090 matrix.row_mut(1)[0] = 10;
1091 assert_eq!(matrix.values, vec![1, 2, 10, 4, 5, 6]);
1092 }
1093
1094 #[test]
1095 fn test_bit_reversed_zero_pad() {
1096 let matrix = RowMajorMatrix::new(
1097 vec![
1098 BabyBear::new(1),
1099 BabyBear::new(2),
1100 BabyBear::new(3),
1101 BabyBear::new(4),
1102 ],
1103 2,
1104 );
1105 let padded = matrix.bit_reversed_zero_pad(1);
1106 assert_eq!(padded.width, 2);
1107 assert_eq!(
1108 padded.values,
1109 vec![
1110 BabyBear::new(1),
1111 BabyBear::new(2),
1112 BabyBear::new(0),
1113 BabyBear::new(0),
1114 BabyBear::new(3),
1115 BabyBear::new(4),
1116 BabyBear::new(0),
1117 BabyBear::new(0)
1118 ]
1119 );
1120 }
1121
1122 #[test]
1123 fn test_bit_reversed_zero_pad_no_change() {
1124 let matrix = RowMajorMatrix::new(
1125 vec![
1126 BabyBear::new(1),
1127 BabyBear::new(2),
1128 BabyBear::new(3),
1129 BabyBear::new(4),
1130 ],
1131 2,
1132 );
1133 let padded = matrix.bit_reversed_zero_pad(0);
1134
1135 assert_eq!(padded.width, 2);
1136 assert_eq!(
1137 padded.values,
1138 vec![
1139 BabyBear::new(1),
1140 BabyBear::new(2),
1141 BabyBear::new(3),
1142 BabyBear::new(4),
1143 ]
1144 );
1145 }
1146
1147 #[test]
1148 fn test_scale() {
1149 let mut matrix = RowMajorMatrix::new(
1150 vec![
1151 BabyBear::new(1),
1152 BabyBear::new(2),
1153 BabyBear::new(3),
1154 BabyBear::new(4),
1155 BabyBear::new(5),
1156 BabyBear::new(6),
1157 ],
1158 2,
1159 );
1160 matrix.scale(BabyBear::new(2));
1161 assert_eq!(
1162 matrix.values,
1163 vec![
1164 BabyBear::new(2),
1165 BabyBear::new(4),
1166 BabyBear::new(6),
1167 BabyBear::new(8),
1168 BabyBear::new(10),
1169 BabyBear::new(12)
1170 ]
1171 );
1172 }
1173
1174 #[test]
1175 fn test_scale_row() {
1176 let mut matrix = RowMajorMatrix::new(
1177 vec![
1178 BabyBear::new(1),
1179 BabyBear::new(2),
1180 BabyBear::new(3),
1181 BabyBear::new(4),
1182 BabyBear::new(5),
1183 BabyBear::new(6),
1184 ],
1185 2,
1186 );
1187 matrix.scale_row(1, BabyBear::new(3));
1188 assert_eq!(
1189 matrix.values,
1190 vec![
1191 BabyBear::new(1),
1192 BabyBear::new(2),
1193 BabyBear::new(9),
1194 BabyBear::new(12),
1195 BabyBear::new(5),
1196 BabyBear::new(6),
1197 ]
1198 );
1199 }
1200
1201 #[test]
1202 fn test_to_row_major_matrix() {
1203 let matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 2);
1204 let converted = matrix.to_row_major_matrix();
1205
1206 assert_eq!(converted.width, 2);
1208 assert_eq!(converted.height(), 3);
1209 assert_eq!(converted.values, vec![1, 2, 3, 4, 5, 6]);
1210 }
1211
1212 #[test]
1213 fn test_horizontally_packed_row() {
1214 type Packed = FieldArray<BabyBear, 2>;
1215
1216 let matrix = RowMajorMatrix::new(
1217 vec![
1218 BabyBear::new(1),
1219 BabyBear::new(2),
1220 BabyBear::new(3),
1221 BabyBear::new(4),
1222 BabyBear::new(5),
1223 BabyBear::new(6),
1224 ],
1225 3,
1226 );
1227
1228 let (packed_iter, suffix_iter) = matrix.horizontally_packed_row::<Packed>(1);
1229
1230 let packed: Vec<_> = packed_iter.collect();
1231 let suffix: Vec<_> = suffix_iter.collect();
1232
1233 assert_eq!(
1234 packed,
1235 vec![Packed::from([BabyBear::new(4), BabyBear::new(5)])]
1236 );
1237 assert_eq!(suffix, vec![BabyBear::new(6)]);
1238 }
1239
1240 #[test]
1241 fn test_padded_horizontally_packed_row() {
1242 use p3_baby_bear::BabyBear;
1243
1244 type Packed = FieldArray<BabyBear, 2>;
1245
1246 let matrix = RowMajorMatrix::new(
1247 vec![
1248 BabyBear::new(1),
1249 BabyBear::new(2),
1250 BabyBear::new(3),
1251 BabyBear::new(4),
1252 BabyBear::new(5),
1253 BabyBear::new(6),
1254 ],
1255 3,
1256 );
1257
1258 let packed_iter = matrix.padded_horizontally_packed_row::<Packed>(1);
1259 let packed: Vec<_> = packed_iter.collect();
1260
1261 assert_eq!(
1262 packed,
1263 vec![
1264 Packed::from([BabyBear::new(4), BabyBear::new(5)]),
1265 Packed::from([BabyBear::new(6), BabyBear::new(0)])
1266 ]
1267 );
1268 }
1269
1270 #[test]
1271 fn test_padded_horizontally_packed_row_exact_width() {
1272 type Packed = FieldArray<BabyBear, 2>;
1273
1274 let matrix = RowMajorMatrix::new(
1278 vec![
1279 BabyBear::new(1),
1280 BabyBear::new(2),
1281 BabyBear::new(3),
1282 BabyBear::new(4),
1283 BabyBear::new(5),
1284 BabyBear::new(6),
1285 BabyBear::new(7),
1286 BabyBear::new(8),
1287 ],
1288 4,
1289 );
1290
1291 let packed: Vec<_> = matrix.padded_horizontally_packed_row::<Packed>(1).collect();
1292
1293 assert_eq!(packed.len(), 2);
1294 assert_eq!(
1295 packed,
1296 vec![
1297 Packed::from([BabyBear::new(5), BabyBear::new(6)]),
1298 Packed::from([BabyBear::new(7), BabyBear::new(8)]),
1299 ]
1300 );
1301 }
1302
1303 #[test]
1304 fn test_pad_to_height() {
1305 let mut matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 3);
1306
1307 matrix.pad_to_height(4, 9);
1312
1313 assert_eq!(matrix.height(), 4);
1320 assert_eq!(matrix.values, vec![1, 2, 3, 4, 5, 6, 9, 9, 9, 9, 9, 9]);
1321 }
1322
1323 #[test]
1324 fn test_pad_to_power_of_two_height() {
1325 let mut matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 2);
1330 assert_eq!(matrix.height(), 3);
1331 matrix.pad_to_power_of_two_height(0);
1332 assert_eq!(matrix.height(), 4);
1333 assert_eq!(matrix.values, vec![1, 2, 3, 4, 5, 6, 0, 0]);
1335
1336 let mut matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6, 7, 8], 2);
1341 assert_eq!(matrix.height(), 4);
1342 matrix.pad_to_power_of_two_height(99);
1343 assert_eq!(matrix.height(), 4);
1344 assert_eq!(matrix.values, vec![1, 2, 3, 4, 5, 6, 7, 8]);
1346
1347 let mut matrix = RowMajorMatrix::new(vec![1, 2, 3], 3);
1351 assert_eq!(matrix.height(), 1);
1352 matrix.pad_to_power_of_two_height(42);
1353 assert_eq!(matrix.height(), 1);
1354 assert_eq!(matrix.values, vec![1, 2, 3]);
1355
1356 let mut matrix = RowMajorMatrix::new(vec![1; 10], 2);
1360 assert_eq!(matrix.height(), 5);
1361 matrix.pad_to_power_of_two_height(-1);
1362 assert_eq!(matrix.height(), 8);
1363 assert_eq!(matrix.values.len(), 16);
1365 assert!(matrix.values[..10].iter().all(|&v| v == 1));
1366 assert!(matrix.values[10..].iter().all(|&v| v == -1));
1367 }
1368
1369 #[test]
1370 fn test_pad_to_power_of_two_height_empty_matrix() {
1371 let mut matrix: RowMajorMatrix<i32> = RowMajorMatrix::new(vec![], 3);
1374 assert_eq!(matrix.height(), 0);
1375 assert_eq!(matrix.width, 3);
1376 matrix.pad_to_power_of_two_height(7);
1377 assert_eq!(matrix.height(), 1);
1379 assert_eq!(matrix.values, vec![7, 7, 7]);
1380 }
1381
1382 #[test]
1383 fn test_pad_to_min_power_of_two_height() {
1384 let mut matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 2);
1390 assert_eq!(matrix.height(), 3);
1391 matrix.pad_to_min_power_of_two_height(5, 0);
1392 assert_eq!(matrix.height(), 8);
1393 assert_eq!(matrix.values[..6], [1, 2, 3, 4, 5, 6]);
1394 assert!(matrix.values[6..].iter().all(|&v| v == 0));
1395
1396 let mut matrix = RowMajorMatrix::new(vec![1; 10], 2);
1402 assert_eq!(matrix.height(), 5);
1403 matrix.pad_to_min_power_of_two_height(2, -1);
1404 assert_eq!(matrix.height(), 8);
1405 assert!(matrix.values[..10].iter().all(|&v| v == 1));
1406 assert!(matrix.values[10..].iter().all(|&v| v == -1));
1407
1408 let mut matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6, 7, 8], 2);
1414 assert_eq!(matrix.height(), 4);
1415 matrix.pad_to_min_power_of_two_height(3, 99);
1416 assert_eq!(matrix.height(), 4);
1417 assert_eq!(matrix.values, vec![1, 2, 3, 4, 5, 6, 7, 8]);
1418
1419 let mut matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 2);
1425 assert_eq!(matrix.height(), 3);
1426 matrix.pad_to_min_power_of_two_height(0, 0);
1427 assert_eq!(matrix.height(), 4);
1428 assert_eq!(matrix.values, vec![1, 2, 3, 4, 5, 6, 0, 0]);
1429
1430 let mut matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 3);
1432 assert_eq!(matrix.height(), 2);
1433 matrix.pad_to_min_power_of_two_height(8, 7);
1434 assert_eq!(matrix.height(), 8);
1435 assert_eq!(matrix.values[..6], [1, 2, 3, 4, 5, 6]);
1436 assert!(matrix.values[6..].iter().all(|&v| v == 7));
1437 assert_eq!(matrix.values.len(), 24); }
1439
1440 #[test]
1441 fn test_pad_to_min_power_of_two_height_empty_matrix() {
1442 let mut matrix: RowMajorMatrix<i32> = RowMajorMatrix::new(vec![], 3);
1444 assert_eq!(matrix.height(), 0);
1445 matrix.pad_to_min_power_of_two_height(5, 7);
1446 assert_eq!(matrix.height(), 8);
1447 assert_eq!(matrix.values.len(), 24);
1448 assert!(matrix.values.iter().all(|&v| v == 7));
1449 }
1450
1451 #[test]
1452 fn test_from_flat_padded() {
1453 let matrix = RowMajorMatrix::from_flat_padded(vec![1, 2, 3, 4, 5, 6], 3, 0);
1457 assert_eq!(matrix.height(), 2);
1458 assert_eq!(matrix.width, 3);
1459 assert_eq!(matrix.values, vec![1, 2, 3, 4, 5, 6]);
1460
1461 let matrix = RowMajorMatrix::from_flat_padded(vec![1, 2, 3, 4, 5], 3, 99);
1465 assert_eq!(matrix.height(), 2);
1466 assert_eq!(matrix.width, 3);
1467 assert_eq!(matrix.values, vec![1, 2, 3, 4, 5, 99]);
1468
1469 let matrix = RowMajorMatrix::from_flat_padded(vec![42], 3, 0);
1471 assert_eq!(matrix.height(), 1);
1472 assert_eq!(matrix.values, vec![42, 0, 0]);
1473
1474 let matrix = RowMajorMatrix::from_flat_padded(vec![], 4, 7);
1476 assert_eq!(matrix.height(), 1);
1477 assert_eq!(matrix.width, 4);
1478 assert_eq!(matrix.values, vec![7, 7, 7, 7]);
1479
1480 let matrix = RowMajorMatrix::from_flat_padded(vec![10, 20, 30], 1, 0);
1482 assert_eq!(matrix.height(), 3);
1483 assert_eq!(matrix.values, vec![10, 20, 30]);
1484 }
1485
1486 #[test]
1487 #[should_panic(expected = "width must be positive")]
1488 fn test_from_flat_padded_zero_width_panics() {
1489 let _ = RowMajorMatrix::from_flat_padded(vec![1, 2, 3], 0, 0);
1490 }
1491
1492 #[test]
1493 fn test_widen_right() {
1494 let mut matrix = RowMajorMatrix::new(vec![1, 2, 3, 4], 2);
1500 matrix.widen_right(1, 0);
1501 assert_eq!(matrix.width, 3);
1502 assert_eq!(matrix.height(), 2);
1503 assert_eq!(matrix.values, vec![1, 2, 0, 3, 4, 0]);
1504
1505 let mut matrix = RowMajorMatrix::new(vec![1, 2, 3, 4], 2);
1511 matrix.widen_right(3, -1);
1512 assert_eq!(matrix.width, 5);
1513 assert_eq!(matrix.height(), 2);
1514 assert_eq!(matrix.values, vec![1, 2, -1, -1, -1, 3, 4, -1, -1, -1]);
1515
1516 let mut matrix = RowMajorMatrix::new(vec![1, 2, 3, 4], 2);
1518 matrix.widen_right(0, 99);
1519 assert_eq!(matrix.width, 2);
1520 assert_eq!(matrix.values, vec![1, 2, 3, 4]);
1521
1522 let mut matrix = RowMajorMatrix::new(vec![10, 20, 30], 3);
1524 matrix.widen_right(2, 0);
1525 assert_eq!(matrix.width, 5);
1526 assert_eq!(matrix.height(), 1);
1527 assert_eq!(matrix.values, vec![10, 20, 30, 0, 0]);
1528
1529 let mut matrix = RowMajorMatrix::new(vec![1, 2, 3], 1);
1536 matrix.widen_right(2, 0);
1537 assert_eq!(matrix.width, 3);
1538 assert_eq!(matrix.height(), 3);
1539 assert_eq!(matrix.values, vec![1, 0, 0, 2, 0, 0, 3, 0, 0]);
1540 }
1541
1542 #[test]
1543 fn test_widen_right_empty_matrix() {
1544 let mut matrix: RowMajorMatrix<i32> = RowMajorMatrix::new(vec![], 3);
1546 matrix.widen_right(2, 0);
1547 assert_eq!(matrix.width, 5);
1548 assert_eq!(matrix.height(), 0);
1549 assert!(matrix.values.is_empty());
1550 }
1551
1552 #[test]
1553 fn test_transpose_into() {
1554 let matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 3);
1555
1556 let mut transposed = RowMajorMatrix::new(vec![0; 6], 2);
1561
1562 matrix.transpose_into(&mut transposed);
1563
1564 assert_eq!(transposed.width, 2);
1570 assert_eq!(transposed.height(), 3);
1571 assert_eq!(transposed.values, vec![1, 4, 2, 5, 3, 6]);
1572 }
1573
1574 #[test]
1575 fn test_flatten_to_base() {
1576 let matrix = RowMajorMatrix::new(
1577 vec![
1578 BabyBear::new(2),
1579 BabyBear::new(3),
1580 BabyBear::new(4),
1581 BabyBear::new(5),
1582 ],
1583 2,
1584 );
1585
1586 let flattened: RowMajorMatrix<BabyBear> = matrix.flatten_to_base();
1587
1588 assert_eq!(flattened.width, 2);
1589 assert_eq!(
1590 flattened.values,
1591 vec![
1592 BabyBear::new(2),
1593 BabyBear::new(3),
1594 BabyBear::new(4),
1595 BabyBear::new(5),
1596 ]
1597 );
1598 }
1599
1600 #[test]
1601 fn test_horizontally_packed_row_mut() {
1602 type Packed = FieldArray<BabyBear, 2>;
1603
1604 let mut matrix = RowMajorMatrix::new(
1605 vec![
1606 BabyBear::new(1),
1607 BabyBear::new(2),
1608 BabyBear::new(3),
1609 BabyBear::new(4),
1610 BabyBear::new(5),
1611 BabyBear::new(6),
1612 ],
1613 3,
1614 );
1615
1616 let (packed, suffix) = matrix.horizontally_packed_row_mut::<Packed>(1);
1617 packed[0] = Packed::from([BabyBear::new(9), BabyBear::new(10)]);
1618 suffix[0] = BabyBear::new(11);
1619
1620 assert_eq!(
1621 matrix.values,
1622 vec![
1623 BabyBear::new(1),
1624 BabyBear::new(2),
1625 BabyBear::new(3),
1626 BabyBear::new(9),
1627 BabyBear::new(10),
1628 BabyBear::new(11),
1629 ]
1630 );
1631 }
1632
1633 #[test]
1634 fn test_par_row_chunks() {
1635 let matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6, 7, 8], 2);
1636
1637 let chunks: Vec<_> = matrix.par_row_chunks(2).collect();
1638
1639 assert_eq!(chunks.len(), 2);
1640 assert_eq!(chunks[0].values, vec![1, 2, 3, 4]);
1641 assert_eq!(chunks[1].values, vec![5, 6, 7, 8]);
1642 }
1643
1644 #[test]
1645 fn test_par_row_chunks_exact() {
1646 let matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 2);
1647
1648 let chunks: Vec<_> = matrix.par_row_chunks_exact(1).collect();
1649
1650 assert_eq!(chunks.len(), 3);
1651 assert_eq!(chunks[0].values, vec![1, 2]);
1652 assert_eq!(chunks[1].values, vec![3, 4]);
1653 assert_eq!(chunks[2].values, vec![5, 6]);
1654 }
1655
1656 #[test]
1657 fn test_par_row_chunks_mut() {
1658 let mut matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6, 7, 8], 2);
1659
1660 matrix
1661 .par_row_chunks_mut(2)
1662 .for_each(|chunk| chunk.values.iter_mut().for_each(|x| *x += 10));
1663
1664 assert_eq!(matrix.values, vec![11, 12, 13, 14, 15, 16, 17, 18]);
1665 }
1666
1667 #[test]
1668 fn test_row_chunks_exact_mut() {
1669 let mut matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 2);
1670
1671 for chunk in matrix.row_chunks_exact_mut(1) {
1672 chunk.values.iter_mut().for_each(|x| *x *= 2);
1673 }
1674
1675 assert_eq!(matrix.values, vec![2, 4, 6, 8, 10, 12]);
1676 }
1677
1678 #[test]
1679 fn test_par_row_chunks_exact_mut() {
1680 let mut matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 2);
1681
1682 matrix
1683 .par_row_chunks_exact_mut(1)
1684 .for_each(|chunk| chunk.values.iter_mut().for_each(|x| *x += 5));
1685
1686 assert_eq!(matrix.values, vec![6, 7, 8, 9, 10, 11]);
1687 }
1688
1689 #[test]
1690 fn test_row_pair_mut() {
1691 let mut matrix = RowMajorMatrix::new(vec![1, 2, 3, 4, 5, 6], 2);
1692
1693 let (row1, row2) = matrix.row_pair_mut(0, 2);
1694 row1[0] = 9;
1695 row2[1] = 10;
1696
1697 assert_eq!(matrix.values, vec![9, 2, 3, 4, 5, 10]);
1698 }
1699
1700 #[test]
1701 fn test_packed_row_pair_mut() {
1702 type Packed = FieldArray<BabyBear, 2>;
1703
1704 let mut matrix = RowMajorMatrix::new(
1705 vec![
1706 BabyBear::new(1),
1707 BabyBear::new(2),
1708 BabyBear::new(3),
1709 BabyBear::new(4),
1710 BabyBear::new(5),
1711 BabyBear::new(6),
1712 ],
1713 3,
1714 );
1715
1716 let ((packed1, sfx1), (packed2, sfx2)) = matrix.packed_row_pair_mut::<Packed>(0, 1);
1717 packed1[0] = Packed::from([BabyBear::new(7), BabyBear::new(8)]);
1718 packed2[0] = Packed::from([BabyBear::new(33), BabyBear::new(44)]);
1719 sfx1[0] = BabyBear::new(99);
1720 sfx2[0] = BabyBear::new(9);
1721
1722 assert_eq!(
1723 matrix.values,
1724 vec![
1725 BabyBear::new(7),
1726 BabyBear::new(8),
1727 BabyBear::new(99),
1728 BabyBear::new(33),
1729 BabyBear::new(44),
1730 BabyBear::new(9),
1731 ]
1732 );
1733 }
1734
1735 #[test]
1736 fn test_transpose_square_matrix() {
1737 const START_INDEX: usize = 1;
1738 const VALUE_LEN: usize = 9;
1739 const WIDTH: usize = 3;
1740 const HEIGHT: usize = 3;
1741
1742 let matrix_values = (START_INDEX..=VALUE_LEN).collect::<Vec<_>>();
1743 let matrix = RowMajorMatrix::new(matrix_values, WIDTH);
1744 let transposed = matrix.transpose();
1745 let should_be_transposed_values = vec![1, 4, 7, 2, 5, 8, 3, 6, 9];
1746 let should_be_transposed = RowMajorMatrix::new(should_be_transposed_values, HEIGHT);
1747 assert_eq!(transposed, should_be_transposed);
1748 }
1749
1750 #[test]
1751 fn test_transpose_row_matrix() {
1752 const START_INDEX: usize = 1;
1753 const VALUE_LEN: usize = 30;
1754 const WIDTH: usize = 1;
1755 const HEIGHT: usize = 30;
1756
1757 let matrix_values = (START_INDEX..=VALUE_LEN).collect::<Vec<_>>();
1758 let matrix = RowMajorMatrix::new(matrix_values.clone(), WIDTH);
1759 let transposed = matrix.transpose();
1760 let should_be_transposed = RowMajorMatrix::new(matrix_values, HEIGHT);
1761 assert_eq!(transposed, should_be_transposed);
1762 }
1763
1764 #[test]
1765 fn test_transpose_rectangular_matrix() {
1766 const START_INDEX: usize = 1;
1767 const VALUE_LEN: usize = 30;
1768 const WIDTH: usize = 5;
1769 const HEIGHT: usize = 6;
1770
1771 let matrix_values = (START_INDEX..=VALUE_LEN).collect::<Vec<_>>();
1772 let matrix = RowMajorMatrix::new(matrix_values, WIDTH);
1773 let transposed = matrix.transpose();
1774 let should_be_transposed_values = vec![
1775 1, 6, 11, 16, 21, 26, 2, 7, 12, 17, 22, 27, 3, 8, 13, 18, 23, 28, 4, 9, 14, 19, 24, 29,
1776 5, 10, 15, 20, 25, 30,
1777 ];
1778 let should_be_transposed = RowMajorMatrix::new(should_be_transposed_values, HEIGHT);
1779 assert_eq!(transposed, should_be_transposed);
1780 }
1781
1782 #[test]
1783 fn test_transpose_larger_rectangular_matrix() {
1784 const START_INDEX: usize = 1;
1785 const VALUE_LEN: usize = 131072; const WIDTH: usize = 256;
1787 const HEIGHT: usize = 512;
1788
1789 let matrix_values = (START_INDEX..=VALUE_LEN).collect::<Vec<_>>();
1790 let matrix = RowMajorMatrix::new(matrix_values, WIDTH);
1791 let transposed = matrix.transpose();
1792
1793 assert_eq!(transposed.width(), HEIGHT);
1794 assert_eq!(transposed.height(), WIDTH);
1795
1796 for col_index in 0..WIDTH {
1797 for row_index in 0..HEIGHT {
1798 assert_eq!(
1799 matrix.values[row_index * WIDTH + col_index],
1800 transposed.values[col_index * HEIGHT + row_index]
1801 );
1802 }
1803 }
1804 }
1805
1806 #[test]
1807 fn test_transpose_very_large_rectangular_matrix() {
1808 const START_INDEX: usize = 1;
1809 const VALUE_LEN: usize = 1048576; const WIDTH: usize = 1024;
1811 const HEIGHT: usize = 1024;
1812
1813 let matrix_values = (START_INDEX..=VALUE_LEN).collect::<Vec<_>>();
1814 let matrix = RowMajorMatrix::new(matrix_values, WIDTH);
1815 let transposed = matrix.transpose();
1816
1817 assert_eq!(transposed.width(), HEIGHT);
1818 assert_eq!(transposed.height(), WIDTH);
1819
1820 for col_index in 0..WIDTH {
1821 for row_index in 0..HEIGHT {
1822 assert_eq!(
1823 matrix.values[row_index * WIDTH + col_index],
1824 transposed.values[col_index * HEIGHT + row_index]
1825 );
1826 }
1827 }
1828 }
1829
1830 #[test]
1831 fn test_vertically_packed_row_scalar_width_1() {
1832 type Packed = BabyBear;
1833
1834 let matrix = RowMajorMatrix::new((1..17).map(BabyBear::new).collect::<Vec<_>>(), 4);
1835 let packed = matrix
1836 .vertically_packed_row::<Packed>(2)
1837 .collect::<Vec<_>>();
1838
1839 assert_eq!(
1840 packed,
1841 vec![
1842 BabyBear::new(9),
1843 BabyBear::new(10),
1844 BabyBear::new(11),
1845 BabyBear::new(12),
1846 ]
1847 );
1848 }
1849
1850 #[test]
1851 fn test_vertically_packed_row_pair() {
1852 type Packed = FieldArray<BabyBear, 2>;
1853
1854 let matrix = RowMajorMatrix::new((1..17).map(BabyBear::new).collect::<Vec<_>>(), 4);
1855
1856 let packed = matrix.vertically_packed_row_pair::<Packed>(0, 2);
1858
1859 assert_eq!(
1875 packed,
1876 (1..5)
1877 .chain(9..13)
1878 .map(|i| [BabyBear::new(i), BabyBear::new(i + 4)].into())
1879 .collect::<Vec<_>>(),
1880 );
1881 }
1882
1883 #[test]
1884 fn test_vertically_packed_row_pair_scalar_width_1() {
1885 type Packed = BabyBear;
1886
1887 let matrix = RowMajorMatrix::new((1..17).map(BabyBear::new).collect::<Vec<_>>(), 4);
1888 let packed = matrix.vertically_packed_row_pair::<Packed>(1, 2);
1889
1890 assert_eq!(
1891 packed,
1892 vec![
1893 BabyBear::new(5),
1894 BabyBear::new(6),
1895 BabyBear::new(7),
1896 BabyBear::new(8),
1897 BabyBear::new(13),
1898 BabyBear::new(14),
1899 BabyBear::new(15),
1900 BabyBear::new(16),
1901 ]
1902 );
1903 }
1904
1905 #[test]
1906 fn test_vertically_packed_row_pair_overlap() {
1907 type Packed = FieldArray<BabyBear, 2>;
1908
1909 let matrix = RowMajorMatrix::new((1..17).map(BabyBear::new).collect::<Vec<_>>(), 4);
1910
1911 let packed = matrix.vertically_packed_row_pair::<Packed>(0, 1);
1928
1929 assert_eq!(
1930 packed,
1931 (1..5)
1932 .chain(5..9)
1933 .map(|i| [BabyBear::new(i), BabyBear::new(i + 4)].into())
1934 .collect::<Vec<_>>(),
1935 );
1936 }
1937
1938 #[test]
1939 fn test_vertically_packed_row_pair_wraparound_start_1() {
1940 use p3_baby_bear::BabyBear;
1941 use p3_field::FieldArray;
1942
1943 type Packed = FieldArray<BabyBear, 2>;
1944
1945 let matrix = RowMajorMatrix::new((1..17).map(BabyBear::new).collect::<Vec<_>>(), 4);
1946
1947 let packed = matrix.vertically_packed_row_pair::<Packed>(1, 2);
1966
1967 assert_eq!(
1968 packed,
1969 vec![
1970 Packed::from([BabyBear::new(5), BabyBear::new(9)]),
1971 Packed::from([BabyBear::new(6), BabyBear::new(10)]),
1972 Packed::from([BabyBear::new(7), BabyBear::new(11)]),
1973 Packed::from([BabyBear::new(8), BabyBear::new(12)]),
1974 Packed::from([BabyBear::new(13), BabyBear::new(1)]),
1975 Packed::from([BabyBear::new(14), BabyBear::new(2)]),
1976 Packed::from([BabyBear::new(15), BabyBear::new(3)]),
1977 Packed::from([BabyBear::new(16), BabyBear::new(4)]),
1978 ]
1979 );
1980 }
1981
1982 #[test]
1983 fn test_with_zero_cols() {
1984 let mat: RowMajorMatrix<BabyBear> =
1990 RowMajorMatrix::new((1..=6).map(BabyBear::new).collect(), 3);
1991 let widened = mat.with_zero_cols(2);
1992
1993 assert_eq!(widened.width(), 5);
1995 assert_eq!(widened.height(), 2);
1996
1997 assert_eq!(
1999 widened.row_slices().next().unwrap(),
2000 &[
2001 BabyBear::new(1),
2002 BabyBear::new(2),
2003 BabyBear::new(3),
2004 BabyBear::ZERO,
2005 BabyBear::ZERO,
2006 ]
2007 );
2008
2009 assert_eq!(
2011 widened.row_slices().nth(1).unwrap(),
2012 &[
2013 BabyBear::new(4),
2014 BabyBear::new(5),
2015 BabyBear::new(6),
2016 BabyBear::ZERO,
2017 BabyBear::ZERO,
2018 ]
2019 );
2020
2021 let same = mat.with_zero_cols(0);
2023 assert_eq!(same.width(), mat.width());
2024 assert_eq!(same.values, mat.values);
2025
2026 let single_row: RowMajorMatrix<BabyBear> =
2030 RowMajorMatrix::new(vec![BabyBear::new(7), BabyBear::new(8)], 2);
2031 let widened = single_row.with_zero_cols(3);
2032 assert_eq!(widened.width(), 5);
2033 assert_eq!(widened.height(), 1);
2034 assert_eq!(
2035 widened.row_slices().next().unwrap(),
2036 &[
2037 BabyBear::new(7),
2038 BabyBear::new(8),
2039 BabyBear::ZERO,
2040 BabyBear::ZERO,
2041 BabyBear::ZERO,
2042 ]
2043 );
2044
2045 let single_col: RowMajorMatrix<BabyBear> = RowMajorMatrix::new(
2051 vec![BabyBear::new(1), BabyBear::new(2), BabyBear::new(3)],
2052 1,
2053 );
2054 let widened = single_col.with_zero_cols(2);
2055 assert_eq!(widened.width(), 3);
2056 assert_eq!(widened.height(), 3);
2057 for (i, row) in widened.row_slices().enumerate() {
2058 assert_eq!(row[0], BabyBear::new((i + 1) as u32));
2060 assert_eq!(row[1], BabyBear::ZERO);
2061 assert_eq!(row[2], BabyBear::ZERO);
2062 }
2063
2064 let empty: RowMajorMatrix<BabyBear> = RowMajorMatrix::new(vec![], 3);
2066 let widened = empty.with_zero_cols(2);
2067 assert_eq!(widened.width(), 5);
2068 assert_eq!(widened.height(), 0);
2069 assert!(widened.values.is_empty());
2070 }
2071
2072 #[test]
2073 fn test_with_zero_cols_matches_widen_right() {
2074 let mat: RowMajorMatrix<BabyBear> =
2082 RowMajorMatrix::new((1..=12).map(BabyBear::new).collect(), 4);
2083
2084 let via_method = mat.with_zero_cols(3);
2086
2087 let mut via_widen = mat;
2089 via_widen.widen_right(3, BabyBear::ZERO);
2090
2091 assert_eq!(via_method.width(), via_widen.width());
2093 assert_eq!(via_method.height(), via_widen.height());
2094 assert_eq!(via_method.values, via_widen.values);
2095 }
2096
2097 #[test]
2098 fn test_with_zero_cols_preserves_rows_for_edge_shapes() {
2099 for (height, width, extra) in [
2100 (0, 3, 1),
2101 (1, 5, 1),
2102 (3, 5, 4),
2103 (4, 3, 0),
2104 (16_385, 17, 3),
2105 (4_097, 129, 1),
2106 ] {
2107 let matrix: RowMajorMatrix<BabyBear> = RowMajorMatrix::new(
2108 (0..height * width)
2109 .map(|i| BabyBear::new((i + 1) as u32))
2110 .collect(),
2111 width,
2112 );
2113 let original = matrix.clone();
2114 let padded = matrix.with_zero_cols(extra);
2115
2116 assert_eq!(padded.width(), width + extra);
2117 assert_eq!(padded.height(), height);
2118 for row in 0..height {
2119 let source = matrix.row_slice(row).unwrap();
2120 let result = padded.row_slice(row).unwrap();
2121 assert_eq!(&result[..width], &*source);
2122 assert!(result[width..].iter().all(|value| value.is_zero()));
2123 }
2124 assert_eq!(matrix, original);
2125 }
2126 }
2127
2128 #[test]
2129 fn test_with_random_cols() {
2130 let mat: RowMajorMatrix<BabyBear> =
2138 RowMajorMatrix::new((1..=4).map(BabyBear::new).collect(), 2);
2139
2140 let seed = 42u64;
2141 let widened = mat.with_random_cols(3, SmallRng::seed_from_u64(seed));
2142
2143 assert_eq!(widened.width(), 5);
2145 assert_eq!(widened.height(), 2);
2146
2147 let mut reference_rng = SmallRng::seed_from_u64(seed);
2151 for (new_row, old_row) in widened.row_slices().zip(mat.row_slices()) {
2152 assert_eq!(&new_row[..2], old_row);
2154
2155 for val in &new_row[2..] {
2157 let expected: BabyBear = reference_rng.random();
2158 assert_eq!(*val, expected);
2159 }
2160 }
2161 }
2162
2163 #[test]
2164 fn test_with_random_cols_zero_extra() {
2165 let mat: RowMajorMatrix<BabyBear> =
2167 RowMajorMatrix::new((1..=6).map(BabyBear::new).collect(), 3);
2168 let same = mat.with_random_cols(0, SmallRng::seed_from_u64(0));
2169 assert_eq!(same.width(), mat.width());
2170 assert_eq!(same.values, mat.values);
2171 }
2172
2173 #[test]
2174 fn test_with_random_cols_empty_matrix() {
2175 let empty: RowMajorMatrix<BabyBear> = RowMajorMatrix::new(vec![], 3);
2177 let widened = empty.with_random_cols(2, SmallRng::seed_from_u64(0));
2178 assert_eq!(widened.width(), 5);
2179 assert_eq!(widened.height(), 0);
2180 assert!(widened.values.is_empty());
2181 }
2182
2183 #[test]
2184 fn test_with_random_cols_different_seeds() {
2185 let mat: RowMajorMatrix<BabyBear> =
2188 RowMajorMatrix::new((1..=4).map(BabyBear::new).collect(), 2);
2189
2190 let num_random = 4;
2191 let seed_a = 1u64;
2192 let seed_b = 2u64;
2193
2194 let result_a = mat.with_random_cols(num_random, SmallRng::seed_from_u64(seed_a));
2195 let result_b = mat.with_random_cols(num_random, SmallRng::seed_from_u64(seed_b));
2196
2197 for (seed, result) in [(seed_a, &result_a), (seed_b, &result_b)] {
2199 let mut reference_rng = SmallRng::seed_from_u64(seed);
2200 for (new_row, old_row) in result.row_slices().zip(mat.row_slices()) {
2201 assert_eq!(&new_row[..2], old_row);
2203
2204 for val in &new_row[2..] {
2206 let expected: BabyBear = reference_rng.random();
2207 assert_eq!(*val, expected);
2208 }
2209 }
2210 }
2211 }
2212}