Skip to main content

strassen_rec

Function strassen_rec 

Source
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    }