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 const SMALL_ELEMS: usize = 256;
436 if self.height().saturating_mul(self.width()) <= SMALL_ELEMS {
437 let mut acc = EF::zero_vec(self.width());
438 for (row, &scale) in self.rows().zip(v) {
439 for (l, r) in acc.iter_mut().zip(row) {
440 *l += scale * r;
441 }
442 }
443 return acc;
444 }
445
446 let packed_width = self.width().div_ceil(T::Packing::WIDTH);
447
448 let packed_result = self
449 .par_padded_horizontally_packed_rows::<T::Packing>()
450 .zip(v)
451 .par_fold_reduce(
452 || EF::ExtensionPacking::zero_vec(packed_width),
453 |mut acc, (row, &scale)| {
454 let scale: EF::ExtensionPacking = scale.into();
455 acc.iter_mut().zip(row).for_each(|(l, r)| *l += scale * r);
456 acc
457 },
458 |mut acc_l, acc_r| {
459 acc_l.iter_mut().zip(&acc_r).for_each(|(l, r)| *l += *r);
460 acc_l
461 },
462 );
463
464 EF::ExtensionPacking::to_ext_iter(packed_result)
465 .take(self.width())
466 .collect()
467 }
468
469 #[instrument(level = "debug", skip_all, fields(dims = %self.dimensions()))]
476 fn columnwise_dot_product_batched<EF, const N: usize>(
477 &self,
478 vs: &[FieldArray<EF, N>],
479 ) -> Vec<FieldArray<EF, N>>
480 where
481 T: Field,
482 EF: ExtensionField<T>,
483 {
484 let packed_width = self.width().div_ceil(T::Packing::WIDTH);
485 let height = self.height().min(vs.len());
486
487 let num_chunks = (4 * current_num_threads()).clamp(1, height.max(1));
491 let chunk_rows = height.div_ceil(num_chunks);
492
493 let packed_results: Vec<EF::ExtensionPacking> =
494 (0..num_chunks).into_par_iter().par_fold_reduce(
495 || EF::ExtensionPacking::zero_vec(packed_width * N),
496 |mut acc, chunk| {
497 let rows = chunk * chunk_rows..((chunk + 1) * chunk_rows).min(height);
498 T::batched_columnwise_dot_product::<EF, _, _, N>(
499 &mut acc,
500 rows.map(|r| {
501 (
502 self.padded_horizontally_packed_row::<T::Packing>(r),
503 vs[r].0,
504 )
505 }),
506 );
507 acc
508 },
509 |mut acc_l, acc_r| {
510 acc_l.iter_mut().zip(&acc_r).for_each(|(lj, rj)| *lj += *rj);
511 acc_l
512 },
513 );
514
515 packed_results
517 .chunks(N)
518 .flat_map(|chunk| {
519 (0..T::Packing::WIDTH)
520 .map(move |lane| FieldArray::from_fn(|j| chunk[j].extract(lane)))
521 })
522 .take(self.width())
523 .collect()
524 }
525
526 fn rowwise_packed_dot_product<EF>(
536 &self,
537 vec: &[EF::ExtensionPacking],
538 ) -> impl IndexedParallelIterator<Item = EF>
539 where
540 T: Field,
541 EF: ExtensionField<T>,
542 {
543 assert!(vec.len() >= self.width().div_ceil(T::Packing::WIDTH));
545
546 self.par_padded_horizontally_packed_rows::<T::Packing>()
549 .map(move |row_packed| {
550 let d = <EF::ExtensionPacking as BasedVectorSpace<T::Packing>>::DIMENSION;
552
553 let coeff_accs = T::Packing::coeffwise_dot_product(
555 d,
556 vec.iter()
557 .zip(row_packed)
558 .map(|(v, r)| (v.as_basis_coefficients_slice(), r)),
559 );
560
561 let packed_result =
563 EF::ExtensionPacking::from_basis_coefficients_fn(|i| coeff_accs[i]);
564 EF::ExtensionPacking::to_ext_iter([packed_result]).sum()
565 })
566 }
567}
568
569#[cfg(test)]
570mod tests {
571 use alloc::vec::Vec;
572 use alloc::{format, vec};
573
574 use itertools::izip;
575 use p3_baby_bear::BabyBear;
576 use p3_field::PrimeCharacteristicRing;
577 use p3_field::extension::BinomialExtensionField;
578 use rand::SeedableRng;
579 use rand::rngs::SmallRng;
580
581 use super::*;
582
583 #[test]
584 fn test_columnwise_dot_product() {
585 type F = BabyBear;
586 type EF = BinomialExtensionField<BabyBear, 4>;
587
588 let mut rng = SmallRng::seed_from_u64(1);
589 let m = RowMajorMatrix::<F>::rand(&mut rng, 1 << 8, 1 << 4);
590 let v = RowMajorMatrix::<EF>::rand(&mut rng, 1 << 8, 1).values;
591
592 let mut expected = EF::zero_vec(m.width());
593 for (row, &scale) in izip!(m.rows(), &v) {
594 for (l, r) in izip!(&mut expected, row) {
595 *l += scale * r;
596 }
597 }
598
599 assert_eq!(m.columnwise_dot_product(&v), expected);
600 }
601
602 #[test]
603 fn test_columnwise_dot_product_small_height() {
604 type F = BabyBear;
605 type EF = BinomialExtensionField<BabyBear, 4>;
606
607 let mut rng = SmallRng::seed_from_u64(2);
608
609 for height in [0, 1, 3, 16, 17] {
611 let m = RowMajorMatrix::<F>::rand(&mut rng, height, 1 << 4);
612 let v = RowMajorMatrix::<EF>::rand(&mut rng, height, 1).values;
613
614 let mut expected = EF::zero_vec(m.width());
615 for (row, &scale) in izip!(m.rows(), &v) {
616 for (l, r) in izip!(&mut expected, row) {
617 *l += scale * r;
618 }
619 }
620
621 assert_eq!(m.columnwise_dot_product(&v), expected, "height = {height}");
622 }
623 }
624
625 #[test]
626 fn test_columnwise_dot_product_batched() {
627 type F = BabyBear;
628 type EF = BinomialExtensionField<BabyBear, 4>;
629
630 let mut rng = SmallRng::seed_from_u64(1);
631 let m = RowMajorMatrix::<F>::rand(&mut rng, 1 << 8, 1 << 4);
632 let v1 = RowMajorMatrix::<EF>::rand(&mut rng, 1 << 8, 1).values;
633 let v2 = RowMajorMatrix::<EF>::rand(&mut rng, 1 << 8, 1).values;
634
635 let expected1 = m.columnwise_dot_product(&v1);
637 let expected2 = m.columnwise_dot_product(&v2);
638
639 let vs: Vec<FieldArray<EF, 2>> = v1
641 .into_iter()
642 .zip(v2)
643 .map(|(a, b)| FieldArray([a, b]))
644 .collect();
645 let results = m.columnwise_dot_product_batched::<EF, 2>(&vs);
646
647 let result1: Vec<EF> = results.iter().map(|r| r[0]).collect();
649 let result2: Vec<EF> = results.iter().map(|r| r[1]).collect();
650
651 assert_eq!(result1, expected1);
652 assert_eq!(result2, expected2);
653 }
654
655 struct MockMatrix {
657 data: Vec<Vec<u32>>,
658 width: usize,
659 height: usize,
660 }
661
662 impl Matrix<u32> for MockMatrix {
663 fn width(&self) -> usize {
664 self.width
665 }
666
667 fn height(&self) -> usize {
668 self.height
669 }
670
671 unsafe fn row_unchecked(
672 &self,
673 r: usize,
674 ) -> impl IntoIterator<Item = u32, IntoIter = impl Iterator<Item = u32> + Send + Sync>
675 {
676 self.data[r].clone()
678 }
679 }
680
681 #[test]
682 fn test_dimensions() {
683 let dims = Dimensions {
684 width: 3,
685 height: 5,
686 };
687 assert_eq!(dims.width, 3);
688 assert_eq!(dims.height, 5);
689 assert_eq!(format!("{dims:?}"), "3x5");
690 assert_eq!(format!("{dims}"), "3x5");
691 }
692
693 #[test]
694 fn test_mock_matrix_dimensions() {
695 let matrix = MockMatrix {
696 data: vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]],
697 width: 3,
698 height: 3,
699 };
700 assert_eq!(matrix.width(), 3);
701 assert_eq!(matrix.height(), 3);
702 assert_eq!(
703 matrix.dimensions(),
704 Dimensions {
705 width: 3,
706 height: 3
707 }
708 );
709 }
710
711 #[test]
712 fn test_first_row() {
713 let matrix = MockMatrix {
714 data: vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]],
715 width: 3,
716 height: 3,
717 };
718 let mut first_row = matrix.first_row().unwrap().into_iter();
719 assert_eq!(first_row.next(), Some(1));
720 assert_eq!(first_row.next(), Some(2));
721 assert_eq!(first_row.next(), Some(3));
722 }
723
724 #[test]
725 fn test_last_row() {
726 let matrix = MockMatrix {
727 data: vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]],
728 width: 3,
729 height: 3,
730 };
731 let mut last_row = matrix.last_row().unwrap().into_iter();
732 assert_eq!(last_row.next(), Some(7));
733 assert_eq!(last_row.next(), Some(8));
734 assert_eq!(last_row.next(), Some(9));
735 }
736
737 #[test]
738 fn test_first_last_row_empty_matrix() {
739 let matrix = MockMatrix {
740 data: vec![],
741 width: 3,
742 height: 0,
743 };
744 let first_row = matrix.first_row();
745 let last_row = matrix.last_row();
746 assert!(first_row.is_none());
747 assert!(last_row.is_none());
748 }
749
750 #[test]
751 fn test_to_row_major_matrix() {
752 let matrix = MockMatrix {
753 data: vec![vec![1, 2], vec![3, 4]],
754 width: 2,
755 height: 2,
756 };
757 let row_major = matrix.to_row_major_matrix();
758 assert_eq!(row_major.values, vec![1, 2, 3, 4]);
759 assert_eq!(row_major.width, 2);
760 }
761
762 #[test]
763 fn test_matrix_get_methods() {
764 let matrix = MockMatrix {
765 data: vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]],
766 width: 3,
767 height: 3,
768 };
769 assert_eq!(matrix.get(0, 0), Some(1));
770 assert_eq!(matrix.get(1, 2), Some(6));
771 assert_eq!(matrix.get(2, 1), Some(8));
772
773 unsafe {
774 assert_eq!(matrix.get_unchecked(0, 1), 2);
775 assert_eq!(matrix.get_unchecked(1, 0), 4);
776 assert_eq!(matrix.get_unchecked(2, 2), 9);
777 }
778
779 assert_eq!(matrix.get(3, 0), None); assert_eq!(matrix.get(0, 3), None); }
782
783 #[test]
784 fn test_matrix_row_methods_iteration() {
785 let matrix = MockMatrix {
786 data: vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]],
787 width: 3,
788 height: 3,
789 };
790
791 let mut row_iter = matrix.row(1).unwrap().into_iter();
792 assert_eq!(row_iter.next(), Some(4));
793 assert_eq!(row_iter.next(), Some(5));
794 assert_eq!(row_iter.next(), Some(6));
795 assert_eq!(row_iter.next(), None);
796
797 unsafe {
798 let mut row_iter_unchecked = matrix.row_unchecked(2).into_iter();
799 assert_eq!(row_iter_unchecked.next(), Some(7));
800 assert_eq!(row_iter_unchecked.next(), Some(8));
801 assert_eq!(row_iter_unchecked.next(), Some(9));
802 assert_eq!(row_iter_unchecked.next(), None);
803
804 let mut row_iter_subset = matrix.row_subseq_unchecked(0, 1, 3).into_iter();
805 assert_eq!(row_iter_subset.next(), Some(2));
806 assert_eq!(row_iter_subset.next(), Some(3));
807 assert_eq!(row_iter_subset.next(), None);
808 }
809
810 assert!(matrix.row(3).is_none()); }
812
813 #[test]
814 fn test_row_slice_methods() {
815 let matrix = MockMatrix {
816 data: vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]],
817 width: 3,
818 height: 3,
819 };
820 let row_slice = matrix.row_slice(1).unwrap();
821 assert_eq!(*row_slice, [4, 5, 6]);
822 unsafe {
823 let row_slice_unchecked = matrix.row_slice_unchecked(2);
824 assert_eq!(*row_slice_unchecked, [7, 8, 9]);
825
826 let row_subslice = matrix.row_subslice_unchecked(0, 1, 2);
827 assert_eq!(*row_subslice, [2]);
828 }
829
830 assert!(matrix.row_slice(3).is_none()); }
832
833 #[test]
834 fn test_matrix_rows() {
835 let matrix = MockMatrix {
836 data: vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]],
837 width: 3,
838 height: 3,
839 };
840
841 let all_rows: Vec<Vec<u32>> = matrix.rows().map(|row| row.collect()).collect();
842 assert_eq!(all_rows, vec![vec![1, 2, 3], vec![4, 5, 6], vec![7, 8, 9]]);
843 }
844
845 #[test]
846 fn test_rowwise_packed_dot_product() {
847 use p3_field::PackedFieldExtension;
848
849 type F = BabyBear;
850 type EF = BinomialExtensionField<BabyBear, 4>;
851 type PF = <F as p3_field::Field>::Packing;
852 type EFPacked = <EF as p3_field::ExtensionField<F>>::ExtensionPacking;
853
854 let mut rng = SmallRng::seed_from_u64(42);
855
856 for (height, width) in [(32, 16), (64, 128), (128, 17), (256, 255)] {
858 let m = RowMajorMatrix::<F>::rand(&mut rng, height, width);
859 let v = RowMajorMatrix::<EF>::rand(&mut rng, width, 1).values;
860
861 let expected: Vec<EF> = m
863 .rows()
864 .map(|row| {
865 row.into_iter()
866 .zip(v.iter())
867 .map(|(r, &ve)| ve * r)
868 .sum::<EF>()
869 })
870 .collect();
871
872 let packed_v: Vec<EFPacked> = v
874 .chunks(<PF as PackedValue>::WIDTH)
875 .map(|chunk| {
876 let mut padded = EF::zero_vec(<PF as PackedValue>::WIDTH);
877 padded[..chunk.len()].copy_from_slice(chunk);
878 EFPacked::from_ext_slice(&padded)
879 })
880 .collect();
881
882 let result: Vec<EF> = m.rowwise_packed_dot_product::<EF>(&packed_v).collect();
884
885 assert_eq!(result, expected, "Mismatch for matrix {}x{}", height, width);
886 }
887 }
888}