fn strassen_rec<R: Ring>(
a: &[R::T],
b: &[R::T],
c: &mut [R::T],
shape: (usize, usize, usize),
stride_a: usize,
stride_b: usize,
)Examples found in repository?
crates/competitive/src/math/matrix.rs (line 761)
679fn strassen_rec<R: Ring>(
680 a: &[R::T],
681 b: &[R::T],
682 c: &mut [R::T],
683 shape: (usize, usize, usize),
684 stride_a: usize,
685 stride_b: usize,
686) {
687 let (n, m, p) = shape;
688 fn add_block<R: Ring>(
689 a: &[R::T],
690 b: &[R::T],
691 out: &mut [R::T],
692 n: usize,
693 stride_a: usize,
694 stride_b: usize,
695 ) {
696 for ((a, b), c) in a
697 .chunks(stride_a)
698 .zip(b.chunks(stride_b))
699 .zip(out.chunks_exact_mut(n))
700 {
701 for ((a, b), c) in a.iter().zip(b.iter()).zip(c.iter_mut()) {
702 *c = R::add(a, b);
703 }
704 }
705 }
706
707 fn sub_block<R: Ring>(
708 a: &[R::T],
709 b: &[R::T],
710 out: &mut [R::T],
711 n: usize,
712 stride_a: usize,
713 stride_b: usize,
714 ) {
715 for ((a, b), c) in a
716 .chunks(stride_a)
717 .zip(b.chunks(stride_b))
718 .zip(out.chunks_exact_mut(n))
719 {
720 for ((a, b), c) in a.iter().zip(b.iter()).zip(c.iter_mut()) {
721 *c = R::sub(a, b);
722 }
723 }
724 }
725
726 if n.min(m).min(p) <= 128 {
727 let transposed: Vec<_> = (0..p)
728 .flat_map(|j| (0..m).map(move |i| b[i * stride_b + j].clone()))
729 .collect();
730 for (a, c) in a.chunks(stride_a).zip(c.chunks_exact_mut(p)) {
731 for (b, c) in transposed.chunks_exact(m).zip(c) {
732 *c = R::dot_product(&a[..m], b);
733 }
734 }
735 return;
736 }
737 let (h, k, w) = (n / 2, m / 2, p / 2);
738 let a11 = 0;
739 let a12 = k;
740 let a21 = h * stride_a;
741 let a22 = a21 + k;
742 let b11 = 0;
743 let b12 = w;
744 let b21 = k * stride_b;
745 let b22 = b21 + w;
746
747 let block = h * w;
748 let mut buf = vec![R::zero(); h * k + k * w + block * 7];
749 let (s1, rest) = buf.split_at_mut(h * k);
750 let (s2, m_buf) = rest.split_at_mut(k * w);
751 let (m1, rest) = m_buf.split_at_mut(block);
752 let (m2, rest) = rest.split_at_mut(block);
753 let (m3, rest) = rest.split_at_mut(block);
754 let (m4, rest) = rest.split_at_mut(block);
755 let (m5, rest) = rest.split_at_mut(block);
756 let (m6, m7) = rest.split_at_mut(block);
757
758 // (A11 + A22)(B11 + B22)
759 add_block::<R>(&a[a11..], &a[a22..], s1, k, stride_a, stride_a);
760 add_block::<R>(&b[b11..], &b[b22..], s2, w, stride_b, stride_b);
761 strassen_rec::<R>(s1, s2, m1, (h, k, w), k, w);
762
763 // (A21 + A22) B11
764 add_block::<R>(&a[a21..], &a[a22..], s1, k, stride_a, stride_a);
765 strassen_rec::<R>(s1, &b[b11..], m2, (h, k, w), k, stride_b);
766
767 // A11 (B12 - B22)
768 sub_block::<R>(&b[b12..], &b[b22..], s2, w, stride_b, stride_b);
769 strassen_rec::<R>(&a[a11..], s2, m3, (h, k, w), stride_a, w);
770
771 // A22 (B21 - B11)
772 sub_block::<R>(&b[b21..], &b[b11..], s2, w, stride_b, stride_b);
773 strassen_rec::<R>(&a[a22..], s2, m4, (h, k, w), stride_a, w);
774
775 // (A11 + A12) B22
776 add_block::<R>(&a[a11..], &a[a12..], s1, k, stride_a, stride_a);
777 strassen_rec::<R>(s1, &b[b22..], m5, (h, k, w), k, stride_b);
778
779 // (A21 - A11)(B11 + B12)
780 sub_block::<R>(&a[a21..], &a[a11..], s1, k, stride_a, stride_a);
781 add_block::<R>(&b[b11..], &b[b12..], s2, w, stride_b, stride_b);
782 strassen_rec::<R>(s1, s2, m6, (h, k, w), k, w);
783
784 // (A12 - A22)(B21 + B22)
785 sub_block::<R>(&a[a12..], &a[a22..], s1, k, stride_a, stride_a);
786 add_block::<R>(&b[b21..], &b[b22..], s2, w, stride_b, stride_b);
787 strassen_rec::<R>(s1, s2, m7, (h, k, w), k, w);
788
789 let c11 = 0;
790 let c12 = w;
791 let c21 = h * p;
792 let c22 = c21 + w;
793 for ((((m1, m4), m5), m7), c) in m1
794 .iter()
795 .zip(m4.iter())
796 .zip(m5.iter())
797 .zip(m7.iter())
798 .zip(c[c11..].chunks_mut(p).flat_map(|c| c.iter_mut().take(w)))
799 {
800 *c = R::add(m1, m4);
801 R::sub_assign(c, m5);
802 R::add_assign(c, m7);
803 }
804 for ((m3, m5), c) in m3
805 .iter()
806 .zip(m5.iter())
807 .zip(c[c12..].chunks_mut(p).flat_map(|c| c.iter_mut().take(w)))
808 {
809 *c = R::add(m3, m5);
810 }
811 for ((m2, m4), c) in m2
812 .iter()
813 .zip(m4.iter())
814 .zip(c[c21..].chunks_mut(p).flat_map(|c| c.iter_mut().take(w)))
815 {
816 *c = R::add(m2, m4);
817 }
818 for ((((m1, m2), m3), m6), c) in m1
819 .iter()
820 .zip(m2.iter())
821 .zip(m3.iter())
822 .zip(m6.iter())
823 .zip(c[c22..].chunks_mut(p).flat_map(|c| c.iter_mut().take(w)))
824 {
825 *c = R::sub(m1, m2);
826 R::add_assign(c, m3);
827 R::add_assign(c, m6);
828 }
829}
830
831impl<R> Matrix<R>
832where
833 R: Ring,
834{
835 pub fn mul_strassen(&self, rhs: &Matrix<R>) -> Matrix<R> {
836 assert_eq!(self.shape.1, rhs.shape.0);
837 if let Some(data) = R::try_matrix_product(&self.data, &rhs.data) {
838 return Matrix::from_vec(data);
839 }
840 let (n, m) = self.shape;
841 let p = rhs.shape.1;
842 if n == 0 || m == 0 || p == 0 {
843 return Matrix::zeros((n, p));
844 }
845 let split = n.min(m).min(p).div_ceil(128).next_power_of_two();
846 if split <= 2 {
847 return self * rhs;
848 }
849 let rows = n.div_ceil(split) * split;
850 let inner = m.div_ceil(split) * split;
851 let cols = p.div_ceil(split) * split;
852 let mut a = vec![R::zero(); rows * inner];
853 for (a, data) in a.chunks_exact_mut(inner).zip(&self.data) {
854 a[..m].clone_from_slice(data);
855 }
856 let mut b = vec![R::zero(); inner * cols];
857 for (b, data) in b.chunks_exact_mut(cols).zip(&rhs.data) {
858 b[..p].clone_from_slice(data);
859 }
860 let mut c = vec![R::zero(); rows * cols];
861 strassen_rec::<R>(&a, &b, &mut c, (rows, inner, cols), inner, cols);
862 let mut res = Matrix::zeros((n, p));
863 for (data, c) in res.data.iter_mut().zip(c.chunks_exact(cols)) {
864 data.clone_from_slice(&c[..p]);
865 }
866 res
867 }