1#![doc = include_str!("../README.md")]
2#![no_std]
3
4extern crate alloc;
5
6use alloc::vec::Vec;
7use core::fmt::{Debug, Display, Formatter};
8use core::ops::Deref;
9
10use itertools::Itertools;
11use p3_field::{
12 BasedVectorSpace, ExtensionField, Field, FieldArray, PackedField, PackedFieldExtension,
13 PackedValue, PrimeCharacteristicRing,
14};
15use p3_maybe_rayon::prelude::*;
16use strided::{VerticallyStridedMatrixView, VerticallyStridedRowIndexMap};
17use tracing::instrument;
18
19use crate::dense::RowMajorMatrix;
20
21pub mod bitrev;
22pub mod dense;
23pub mod extension;
24pub mod horizontally_truncated;
25pub mod interpolation;
26pub mod row_index_mapped;
27pub mod stack;
28pub mod strided;
29pub mod util;
30
31#[derive(Copy, Clone, PartialEq, Eq)]
36pub struct Dimensions {
37 pub width: usize,
39 pub height: usize,
41}
42
43impl Debug for Dimensions {
44 fn fmt(&self, f: &mut Formatter<'_>) -> core::fmt::Result {
45 write!(f, "{}x{}", self.width, self.height)
46 }
47}
48
49impl Display for Dimensions {
50 fn fmt(&self, f: &mut Formatter<'_>) -> core::fmt::Result {
51 write!(f, "{}x{}", self.width, self.height)
52 }
53}
54
55pub trait Matrix<T: Send + Sync + Clone>: Send + Sync {
61 fn width(&self) -> usize;
63
64 fn height(&self) -> usize;
66
67 fn dimensions(&self) -> Dimensions {
69 Dimensions {
70 width: self.width(),
71 height: self.height(),
72 }
73 }
74
75 #[inline]
86 fn get(&self, r: usize, c: usize) -> Option<T> {
87 (r < self.height() && c < self.width()).then(|| unsafe {
88 self.get_unchecked(r, c)
90 })
91 }
92
93 #[inline]
101 unsafe fn get_unchecked(&self, r: usize, c: usize) -> T {
102 unsafe { self.row_slice_unchecked(r)[c].clone() }
103 }
104
105 #[inline]
111 fn row(
112 &self,
113 r: usize,
114 ) -> Option<impl IntoIterator<Item = T, IntoIter = impl Iterator<Item = T> + Send + Sync>> {
115 (r < self.height()).then(|| unsafe {
116 self.row_unchecked(r)
118 })
119 }
120
121 #[inline]
131 unsafe fn row_unchecked(
132 &self,
133 r: usize,
134 ) -> impl IntoIterator<Item = T, IntoIter = impl Iterator<Item = T> + Send + Sync> {
135 unsafe { self.row_subseq_unchecked(r, 0, self.width()) }
136 }
137
138 #[inline]
148 unsafe fn row_subseq_unchecked(
149 &self,
150 r: usize,
151 start: usize,
152 end: usize,
153 ) -> impl IntoIterator<Item = T, IntoIter = impl Iterator<Item = T> + Send + Sync> {
154 unsafe {
155 self.row_unchecked(r)
156 .into_iter()
157 .skip(start)
158 .take(end - start)
159 }
160 }
161
162 #[inline]
166 fn row_slice(&self, r: usize) -> Option<impl Deref<Target = [T]>> {
167 (r < self.height()).then(|| unsafe {
168 self.row_slice_unchecked(r)
170 })
171 }
172
173 #[inline]
181 unsafe fn row_slice_unchecked(&self, r: usize) -> impl Deref<Target = [T]> {
182 unsafe { self.row_subslice_unchecked(r, 0, self.width()) }
183 }
184
185 #[inline]
195 unsafe fn row_subslice_unchecked(
196 &self,
197 r: usize,
198 start: usize,
199 end: usize,
200 ) -> impl Deref<Target = [T]> {
201 unsafe {
202 self.row_subseq_unchecked(r, start, end)
203 .into_iter()
204 .collect_vec()
205 }
206 }
207
208 #[inline]
210 fn rows(&self) -> impl Iterator<Item = impl Iterator<Item = T>> + Send + Sync {
211 unsafe {
212 (0..self.height()).map(move |r| self.row_unchecked(r).into_iter())
214 }
215 }
216
217 #[inline]
219 fn par_rows(
220 &self,
221 ) -> impl IndexedParallelIterator<Item = impl Iterator<Item = T>> + Send + Sync {
222 unsafe {
223 (0..self.height())
225 .into_par_iter()
226 .map(move |r| self.row_unchecked(r).into_iter())
227 }
228 }
229
230 fn wrapping_row_slices(&self, r: usize, c: usize) -> Vec<impl Deref<Target = [T]>> {
233 unsafe {
234 (0..c)
236 .map(|i| self.row_slice_unchecked((r + i) % self.height()))
237 .collect_vec()
238 }
239 }
240
241 #[inline]
245 fn first_row(
246 &self,
247 ) -> Option<impl IntoIterator<Item = T, IntoIter = impl Iterator<Item = T> + Send + Sync>> {
248 self.row(0)
249 }
250
251 #[inline]
255 fn last_row(
256 &self,
257 ) -> Option<impl IntoIterator<Item = T, IntoIter = impl Iterator<Item = T> + Send + Sync>> {
258 if self.height() == 0 {
259 None
260 } else {
261 unsafe { Some(self.row_unchecked(self.height() - 1)) }
263 }
264 }
265
266 fn to_row_major_matrix(self) -> RowMajorMatrix<T>
268 where
269 Self: Sized,
270 T: Clone,
271 {
272 RowMajorMatrix::new(self.rows().flatten().collect(), self.width())
273 }
274
275 fn horizontally_packed_row<'a, P>(
283 &'a self,
284 r: usize,
285 ) -> (
286 impl Iterator<Item = P> + Send + Sync,
287 impl Iterator<Item = T> + Send + Sync,
288 )
289 where
290 P: PackedValue<Value = T>,
291 T: Clone + 'a,
292 {
293 assert!(r < self.height(), "Row index out of bounds.");
294 let num_packed = self.width() / P::WIDTH;
295 unsafe {
296 let mut iter = self
298 .row_subseq_unchecked(r, 0, num_packed * P::WIDTH)
299 .into_iter();
300
301 let packed =
303 (0..num_packed).map(move |_| P::from_fn(|_| iter.next().unwrap_unchecked()));
304
305 let sfx = self
306 .row_subseq_unchecked(r, num_packed * P::WIDTH, self.width())
307 .into_iter();
308 (packed, sfx)
309 }
310 }
311
312 fn padded_horizontally_packed_row<'a, P>(
319 &'a self,
320 r: usize,
321 ) -> impl Iterator<Item = P> + Send + Sync
322 where
323 P: PackedValue<Value = T>,
324 T: Clone + Default + 'a,
325 {
326 let mut row_iter = self.row(r).expect("Row index out of bounds.").into_iter();
327 let num_elems = self.width().div_ceil(P::WIDTH);
328 (0..num_elems).map(move |_| P::from_fn(|_| row_iter.next().unwrap_or_default()))
330 }
331
332 fn par_horizontally_packed_rows<'a, P>(
337 &'a self,
338 ) -> impl IndexedParallelIterator<
339 Item = (
340 impl Iterator<Item = P> + Send + Sync,
341 impl Iterator<Item = T> + Send + Sync,
342 ),
343 >
344 where
345 P: PackedValue<Value = T>,
346 T: Clone + 'a,
347 {
348 (0..self.height())
349 .into_par_iter()
350 .map(|r| self.horizontally_packed_row(r))
351 }
352
353 fn par_padded_horizontally_packed_rows<'a, P>(
357 &'a self,
358 ) -> impl IndexedParallelIterator<Item = impl Iterator<Item = P> + Send + Sync>
359 where
360 P: PackedValue<Value = T>,
361 T: Clone + Default + 'a,
362 {
363 (0..self.height())
364 .into_par_iter()
365 .map(|r| self.padded_horizontally_packed_row(r))
366 }
367
368 #[inline]
374 fn vertically_packed_row<P>(&self, r: usize) -> impl Iterator<Item = P>
375 where
376 T: Copy,
377 P: PackedValue<Value = T>,
378 {
379 let rows = self.wrapping_row_slices(r, P::WIDTH);
381
382 (0..self.width()).map(move |c| P::from_fn(|i| rows[i][c]))
384 }
385
386 #[inline]
394 fn vertically_packed_row_pair<P>(&self, r: usize, step: usize) -> Vec<P>
395 where
396 T: Copy,
397 P: PackedValue<Value = T>,
398 {
399 let rows = self.wrapping_row_slices(r, P::WIDTH);
404 let next_rows = self.wrapping_row_slices(r + step, P::WIDTH);
405
406 (0..self.width())
407 .map(|c| P::from_fn(|i| rows[i][c]))
408 .chain((0..self.width()).map(|c| P::from_fn(|i| next_rows[i][c])))
409 .collect_vec()
410 }
411
412 fn vertically_strided(self, stride: usize, offset: usize) -> VerticallyStridedMatrixView<Self>
416 where
417 Self: Sized,
418 {
419 VerticallyStridedRowIndexMap::new_view(self, stride, offset)
420 }
421
422 #[instrument(level = "debug", skip_all, fields(dims = %self.dimensions()))]
426 fn columnwise_dot_product<EF>(&self, v: &[EF]) -> Vec<EF>
427 where
428 T: Field,
429 EF: ExtensionField<T>,
430 {
431 assert_eq!(v.len(), self.height());
432
433 const SMALL_ELEMS: usize = 256;
438 if self.height().saturating_mul(self.width()) <= SMALL_ELEMS {
439 let mut acc = EF::zero_vec(self.width());
440 for (row, &scale) in self.rows().zip(v) {
441 for (l, r) in acc.iter_mut().zip(row) {
442 *l += scale * r;
443 }
444 }
445 return acc;
446 }
447
448 let packed_width = self.width().div_ceil(T::Packing::WIDTH);
449
450 let packed_result = self
451 .par_padded_horizontally_packed_rows::<T::Packing>()
452 .zip(v)
453 .par_fold_reduce(
454 || EF::ExtensionPacking::zero_vec(packed_width),
455 |mut acc, (row, &scale)| {
456 let scale: EF::ExtensionPacking = scale.into();
457 acc.iter_mut().zip(row).for_each(|(l, r)| *l += scale * r);
458 acc
459 },
460 |mut acc_l, acc_r| {
461 acc_l.iter_mut().zip(&acc_r).for_each(|(l, r)| *l += *r);
462 acc_l
463 },
464 );
465
466 EF::ExtensionPacking::to_ext_iter(packed_result)
467 .take(self.width())
468 .collect()
469 }
470
471 #[instrument(level = "debug", skip_all, fields(dims = %self.dimensions()))]
478 fn columnwise_dot_product_batched<EF, const N: usize>(
479 &self,
480 vs: &[FieldArray<EF, N>],
481 ) -> Vec<FieldArray<EF, N>>
482 where
483 T: Field,
484 EF: ExtensionField<T>,
485 {
486 assert_eq!(vs.len(), self.height());
487
488 let packed_width = self.width().div_ceil(T::Packing::WIDTH);
489 let height = self.height();
490
491 let num_chunks = (4 * current_num_threads()).clamp(1, height.max(1));
495 let chunk_rows = height.div_ceil(num_chunks);
496
497 let packed_results: Vec<EF::ExtensionPacking> =
498 (0..num_chunks).into_par_iter().par_fold_reduce(
499 || EF::ExtensionPacking::zero_vec(packed_width * N),
500 |mut acc, chunk| {
501 let rows = chunk * chunk_rows..((chunk + 1) * chunk_rows).min(height);
502 T::batched_columnwise_dot_product::<EF, _, _, N>(
503 &mut acc,
504 rows.map(|r| {
505 (
506 self.padded_horizontally_packed_row::<T::Packing>(r),
507 vs[r].0,
508 )
509 }),
510 );
511 acc
512 },
513 |mut acc_l, acc_r| {
514 acc_l.iter_mut().zip(&acc_r).for_each(|(lj, rj)| *lj += *rj);
515 acc_l
516 },
517 );
518
519 packed_results
521 .chunks(N)
522 .flat_map(|chunk| {
523 (0..T::Packing::WIDTH)
524 .map(move |lane| FieldArray::from_fn(|j| chunk[j].extract(lane)))
525 })
526 .take(self.width())
527 .collect()
528 }
529
530 fn rowwise_packed_dot_product<EF>(
540 &self,
541 vec: &[EF::ExtensionPacking],
542 ) -> impl IndexedParallelIterator<Item = EF>
543 where
544 T: Field,
545 EF: ExtensionField<T>,
546 {
547 assert!(vec.len() >= self.width().div_ceil(T::Packing::WIDTH));
549
550 self.par_padded_horizontally_packed_rows::<T::Packing>()
553 .map(move |row_packed| {
554 let d = <EF::ExtensionPacking as BasedVectorSpace<T::Packing>>::DIMENSION;
556
557 let coeff_accs = T::Packing::coeffwise_dot_product(
559 d,
560 vec.iter()
561 .zip(row_packed)
562 .map(|(v, r)| (v.as_basis_coefficients_slice(), r)),
563 );
564
565 let packed_result =
567 EF::ExtensionPacking::from_basis_coefficients_fn(|i| coeff_accs[i]);
568 EF::ExtensionPacking::to_ext_iter([packed_result]).sum()
569 })
570 }
571}
572
573#[cfg(test)]
574mod tests {
575 use alloc::vec::Vec;
576 use alloc::{format, vec};
577
578 use itertools::izip;
579 use p3_baby_bear::BabyBear;
580 use p3_field::PrimeCharacteristicRing;
581 use p3_field::extension::BinomialExtensionField;
582 use rand::SeedableRng;
583 use rand::rngs::SmallRng;
584
585 use super::*;
586
587 #[test]
588 fn test_columnwise_dot_product() {
589 type F = BabyBear;
590 type EF = BinomialExtensionField<BabyBear, 4>;
591
592 let mut rng = SmallRng::seed_from_u64(1);
593 let m = RowMajorMatrix::<F>::rand(&mut rng, 1 << 8, 1 << 4);
594 let v = RowMajorMatrix::<EF>::rand(&mut rng, 1 << 8, 1).values;
595
596 let mut expected = EF::zero_vec(m.width());
597 for (row, &scale) in izip!(m.rows(), &v) {
598 for (l, r) in izip!(&mut expected, row) {
599 *l += scale * r;
600 }
601 }
602
603 assert_eq!(m.columnwise_dot_product(&v), expected);
604 }
605
606 #[test]
607 fn test_columnwise_dot_product_small_height() {
608 type F = BabyBear;
609 type EF = BinomialExtensionField<BabyBear, 4>;
610
611 let mut rng = SmallRng::seed_from_u64(2);
612
613 for height in [0, 1, 3, 16, 17] {
615 let m = RowMajorMatrix::<F>::rand(&mut rng, height, 1 << 4);
616 let v = RowMajorMatrix::<EF>::rand(&mut rng, height, 1).values;
617
618 let mut expected = EF::zero_vec(m.width());
619 for (row, &scale) in izip!(m.rows(), &v) {
620 for (l, r) in izip!(&mut expected, row) {
621 *l += scale * r;
622 }
623 }
624
625 assert_eq!(m.columnwise_dot_product(&v), expected, "height = {height}");
626 }
627 }
628
629 #[test]
630 fn test_columnwise_dot_product_batched() {
631 type F = BabyBear;
632 type EF = BinomialExtensionField<BabyBear, 4>;
633
634 let mut rng = SmallRng::seed_from_u64(1);
635 let m = RowMajorMatrix::<F>::rand(&mut rng, 1 << 8, 1 << 4);
636 let v1 = RowMajorMatrix::<EF>::rand(&mut rng, 1 << 8, 1).values;
637 let v2 = RowMajorMatrix::<EF>::rand(&mut rng, 1 << 8, 1).values;
638
639 let expected1 = m.columnwise_dot_product(&v1);
641 let expected2 = m.columnwise_dot_product(&v2);
642
643 let vs: Vec<FieldArray<EF, 2>> = v1
645 .into_iter()
646 .zip(v2)
647 .map(|(a, b)| FieldArray([a, b]))
648 .collect();
649 let results = m.columnwise_dot_product_batched::<EF, 2>(&vs);
650
651 let result1: Vec<EF> = results.iter().map(|r| r[0]).collect();
653 let result2: Vec<EF> = results.iter().map(|r| r[1]).collect();
654
655 assert_eq!(result1, expected1);
656 assert_eq!(result2, expected2);
657 }
658
659 struct MockMatrix {
661 data: Vec<Vec<u32>>,
662 width: usize,
663 height: usize,
664 }
665
666 impl Matrix<u32> for MockMatrix {
667 fn width(&self) -> usize {
668 self.width
669 }
670
671 fn height(&self) -> usize {
672 self.height
673 }
674
675 unsafe fn row_unchecked(
676 &self,
677 r: usize,
678 ) -> impl IntoIterator<Item = u32, IntoIter = impl Iterator<Item = u32> + Send + Sync>
679 {
680 self.data[r].clone()
682 }
683 }
684
685 #[test]
686 fn test_dimensions() {
687 let dims = Dimensions {
688 width: 3,
689 height: 5,
690 };
691 assert_eq!(dims.width, 3);
692 assert_eq!(dims.height, 5);
693 assert_eq!(format!("{dims:?}"), "3x5");
694 assert_eq!(format!("{dims}"), "3x5");
695 }
696
697 #[test]
698 fn test_mock_matrix_dimensions() {
699 let matrix = MockMatrix {
700 data: vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]],
701 width: 3,
702 height: 3,
703 };
704 assert_eq!(matrix.width(), 3);
705 assert_eq!(matrix.height(), 3);
706 assert_eq!(
707 matrix.dimensions(),
708 Dimensions {
709 width: 3,
710 height: 3
711 }
712 );
713 }
714
715 #[test]
716 fn test_first_row() {
717 let matrix = MockMatrix {
718 data: vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]],
719 width: 3,
720 height: 3,
721 };
722 let mut first_row = matrix.first_row().unwrap().into_iter();
723 assert_eq!(first_row.next(), Some(1));
724 assert_eq!(first_row.next(), Some(2));
725 assert_eq!(first_row.next(), Some(3));
726 }
727
728 #[test]
729 fn test_last_row() {
730 let matrix = MockMatrix {
731 data: vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]],
732 width: 3,
733 height: 3,
734 };
735 let mut last_row = matrix.last_row().unwrap().into_iter();
736 assert_eq!(last_row.next(), Some(7));
737 assert_eq!(last_row.next(), Some(8));
738 assert_eq!(last_row.next(), Some(9));
739 }
740
741 #[test]
742 fn test_first_last_row_empty_matrix() {
743 let matrix = MockMatrix {
744 data: vec![],
745 width: 3,
746 height: 0,
747 };
748 let first_row = matrix.first_row();
749 let last_row = matrix.last_row();
750 assert!(first_row.is_none());
751 assert!(last_row.is_none());
752 }
753
754 #[test]
755 fn test_to_row_major_matrix() {
756 let matrix = MockMatrix {
757 data: vec![vec![1, 2], vec![3, 4]],
758 width: 2,
759 height: 2,
760 };
761 let row_major = matrix.to_row_major_matrix();
762 assert_eq!(row_major.values, vec![1, 2, 3, 4]);
763 assert_eq!(row_major.width, 2);
764 }
765
766 #[test]
767 fn test_matrix_get_methods() {
768 let matrix = MockMatrix {
769 data: vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]],
770 width: 3,
771 height: 3,
772 };
773 assert_eq!(matrix.get(0, 0), Some(1));
774 assert_eq!(matrix.get(1, 2), Some(6));
775 assert_eq!(matrix.get(2, 1), Some(8));
776
777 unsafe {
778 assert_eq!(matrix.get_unchecked(0, 1), 2);
779 assert_eq!(matrix.get_unchecked(1, 0), 4);
780 assert_eq!(matrix.get_unchecked(2, 2), 9);
781 }
782
783 assert_eq!(matrix.get(3, 0), None); assert_eq!(matrix.get(0, 3), None); }
786
787 #[test]
788 fn test_matrix_row_methods_iteration() {
789 let matrix = MockMatrix {
790 data: vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]],
791 width: 3,
792 height: 3,
793 };
794
795 let mut row_iter = matrix.row(1).unwrap().into_iter();
796 assert_eq!(row_iter.next(), Some(4));
797 assert_eq!(row_iter.next(), Some(5));
798 assert_eq!(row_iter.next(), Some(6));
799 assert_eq!(row_iter.next(), None);
800
801 unsafe {
802 let mut row_iter_unchecked = matrix.row_unchecked(2).into_iter();
803 assert_eq!(row_iter_unchecked.next(), Some(7));
804 assert_eq!(row_iter_unchecked.next(), Some(8));
805 assert_eq!(row_iter_unchecked.next(), Some(9));
806 assert_eq!(row_iter_unchecked.next(), None);
807
808 let mut row_iter_subset = matrix.row_subseq_unchecked(0, 1, 3).into_iter();
809 assert_eq!(row_iter_subset.next(), Some(2));
810 assert_eq!(row_iter_subset.next(), Some(3));
811 assert_eq!(row_iter_subset.next(), None);
812 }
813
814 assert!(matrix.row(3).is_none()); }
816
817 #[test]
818 fn test_row_slice_methods() {
819 let matrix = MockMatrix {
820 data: vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]],
821 width: 3,
822 height: 3,
823 };
824 let row_slice = matrix.row_slice(1).unwrap();
825 assert_eq!(*row_slice, [4, 5, 6]);
826 unsafe {
827 let row_slice_unchecked = matrix.row_slice_unchecked(2);
828 assert_eq!(*row_slice_unchecked, [7, 8, 9]);
829
830 let row_subslice = matrix.row_subslice_unchecked(0, 1, 2);
831 assert_eq!(*row_subslice, [2]);
832 }
833
834 assert!(matrix.row_slice(3).is_none()); }
836
837 #[test]
838 fn test_matrix_rows() {
839 let matrix = MockMatrix {
840 data: vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]],
841 width: 3,
842 height: 3,
843 };
844
845 let all_rows: Vec<Vec<u32>> = matrix.rows().map(|row| row.collect()).collect();
846 assert_eq!(all_rows, vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]]);
847 }
848
849 #[test]
850 fn test_rowwise_packed_dot_product() {
851 use p3_field::PackedFieldExtension;
852
853 type F = BabyBear;
854 type EF = BinomialExtensionField<BabyBear, 4>;
855 type PF = <F as p3_field::Field>::Packing;
856 type EFPacked = <EF as p3_field::ExtensionField<F>>::ExtensionPacking;
857
858 let mut rng = SmallRng::seed_from_u64(42);
859
860 for (height, width) in [(32, 16), (64, 128), (128, 17), (256, 255)] {
862 let m = RowMajorMatrix::<F>::rand(&mut rng, height, width);
863 let v = RowMajorMatrix::<EF>::rand(&mut rng, width, 1).values;
864
865 let expected: Vec<EF> = m
867 .rows()
868 .map(|row| {
869 row.into_iter()
870 .zip(v.iter())
871 .map(|(r, &ve)| ve * r)
872 .sum::<EF>()
873 })
874 .collect();
875
876 let packed_v: Vec<EFPacked> = v
878 .chunks(<PF as PackedValue>::WIDTH)
879 .map(|chunk| {
880 let mut padded = EF::zero_vec(<PF as PackedValue>::WIDTH);
881 padded[..chunk.len()].copy_from_slice(chunk);
882 EFPacked::from_ext_slice(&padded)
883 })
884 .collect();
885
886 let result: Vec<EF> = m.rowwise_packed_dot_product::<EF>(&packed_v).collect();
888
889 assert_eq!(result, expected, "Mismatch for matrix {}x{}", height, width);
890 }
891 }
892}