斯特拉森矩阵乘法
斯特拉森矩阵乘法是解决矩阵乘法问题的分治法。通常的矩阵乘法方法是将每一行与每一列相乘,得到乘积矩阵。这种方法的时间复杂度为 O(n3),因为它需要两次循环才能完成乘法。引入 Strassen 方法将时间复杂度从 O(n3) 降低到 O(nlog 7)。
朴素方法
首先,我们将讨论朴素方法及其复杂度。这里,我们计算 Z=𝑿X × Y。使用朴素方法,如果两个矩阵(X 和 Y)的阶分别为 p × q 和 q × r,则可以将它们相乘,所得矩阵的阶数为 p × r。以下伪代码描述了朴素乘法−
Algorithm: Matrix-Multiplication (X, Y, Z)
for i = 1 to p do
for j = 1 to r do
Z[i,j] := 0
for k = 1 to q do
Z[i,j] := Z[i,j] + X[i,k] × Y[k,j]
复杂度
这里,我们假设整数运算的复杂度为 O(1)。该算法包含三个 for 循环,其中一个循环嵌套在另一个循环中。因此,该算法的执行时间为 O(n3)。
Strassen 矩阵乘法算法
在这种情况下,使用 Strassen 矩阵乘法算法可以稍微改善时间消耗。
Strassen 矩阵乘法只能对方阵执行,其中 n 是 2 的幂。两个矩阵的阶均为 n × n.
将 X、Y 和 Z 分成四个 (n/2)×(n/2) 个矩阵,如下所示 −
$Z = \begin{bmatrix}I & J \K & L \end{bmatrix}$ $X = \begin{bmatrix}A & B \C & D \end{bmatrix}$ 和 $Y = \begin{bmatrix}E & F \G & H \end{bmatrix}$
使用施特拉森算法计算以下 −
$$M_{1} \: \colon= (A+C) imes (E+F)$$
$$M_{2} \:\colon= (B+D) imes (G+H)$$
$$M_{3} \:\colon= (A-D) imes (E+H)$$
$$M_{4} \:\colon= A imes (F-H)$$
$$M_{5} \:\colon= (C+D) imes (E)$$
$$M_{6} \:\colon= (A+B) imes (H)$$
$$M_{7} \:\colon= D imes (G-E)$$
然后,
$$I \:\colon= M_{2} + M_{3} - M_{6} - M_{7}$$
$$J \:\colon= M_{4} + M_{6}$$
$$K \:\colon= M_{5} + M_{7}$$
$$L \:\colon= M_{1} - M_{3} - M_{4} - M_{5}$$
分析
$$T(n)=\begin{cases}c & if\:n= 1\7\:x\:T(\frac{n}{2})+d\:x\:n^2 & otherwise\end{cases}其中 c 和 d 为常数$$
利用此递归关系,我们得到 $T(n) = O(n^{log7})$
因此,Strassen 矩阵乘法算法的复杂度为 $O(n^{log7})$。
示例
让我们看看 Strassen 矩阵乘法在各种编程语言中的实现:C、C++、Java、Python。
#include<stdio.h>
int main(){
int z[2][2];
int i, j;
int m1, m2, m3, m4 , m5, m6, m7;
int x[2][2] = {
{12, 34},
{22, 10}
};
int y[2][2] = {
{3, 4},
{2, 1}
};
printf("The first matrix is: ");
for(i = 0; i < 2; i++) {
printf("
");
for(j = 0; j < 2; j++)
printf("%d ", x[i][j]);
}
printf("
The second matrix is: ");
for(i = 0; i < 2; i++) {
printf("
");
for(j = 0; j < 2; j++)
printf("%d ", y[i][j]);
}
m1= (x[0][0] + x[1][1]) * (y[0][0] + y[1][1]);
m2= (x[1][0] + x[1][1]) * y[0][0];
m3= x[0][0] * (y[0][1] - y[1][1]);
m4= x[1][1] * (y[1][0] - y[0][0]);
m5= (x[0][0] + x[0][1]) * y[1][1];
m6= (x[1][0] - x[0][0]) * (y[0][0]+y[0][1]);
m7= (x[0][1] - x[1][1]) * (y[1][0]+y[1][1]);
z[0][0] = m1 + m4- m5 + m7;
z[0][1] = m3 + m5;
z[1][0] = m2 + m4;
z[1][1] = m1 - m2 + m3 + m6;
printf("
Product achieved using Strassen's algorithm: ");
for(i = 0; i < 2 ; i++) {
printf("
");
for(j = 0; j < 2; j++)
printf("%d ", z[i][j]);
}
return 0;
}
输出
The first matrix is: 12 34 22 10 The second matrix is: 3 4 2 1 Product achieved using Strassen's algorithm: 104 82 86 98
#include<iostream>
using namespace std;
int main() {
int z[2][2];
int i, j;
int m1, m2, m3, m4 , m5, m6, m7;
int x[2][2] = {
{12, 34},
{22, 10}
};
int y[2][2] = {
{3, 4},
{2, 1}
};
cout<<"The first matrix is: ";
for(i = 0; i < 2; i++) {
cout<<endl;
for(j = 0; j < 2; j++)
cout<<x[i][j]<<" ";
}
cout<<"
The second matrix is: ";
for(i = 0;i < 2; i++){
cout<<endl;
for(j = 0;j < 2; j++)
cout<<y[i][j]<<" ";
}
m1 = (x[0][0] + x[1][1]) * (y[0][0] + y[1][1]);
m2 = (x[1][0] + x[1][1]) * y[0][0];
m3 = x[0][0] * (y[0][1] - y[1][1]);
m4 = x[1][1] * (y[1][0] - y[0][0]);
m5 = (x[0][0] + x[0][1]) * y[1][1];
m6 = (x[1][0] - x[0][0]) * (y[0][0]+y[0][1]);
m7 = (x[0][1] - x[1][1]) * (y[1][0]+y[1][1]);
z[0][0] = m1 + m4- m5 + m7;
z[0][1] = m3 + m5;
z[1][0] = m2 + m4;
z[1][1] = m1 - m2 + m3 + m6;
cout<<"
Product achieved using Strassen's algorithm: ";
for(i = 0; i < 2 ; i++) {
cout<<endl;
for(j = 0; j < 2; j++)
cout<<z[i][j]<<" ";
}
return 0;
}
输出
The first matrix is: 12 34 22 10 The second matrix is: 3 4 2 1 Product achieved using Strassen's algorithm: 104 82 86 98
public class Strassens {
public static void main(String[] args) {
int[][] x = {{12, 34}, {22, 10}};
int[][] y = {{3, 4}, {2, 1}};
int z[][] = new int[2][2];
int m1, m2, m3, m4 , m5, m6, m7;
System.out.print("The first matrix is: ");
for(int i = 0; i<2; i++) {
System.out.println();//new line
for(int j = 0; j<2; j++) {
System.out.print(x[i][j] + " ");
}
}
System.out.print("
The second matrix is: ");
for(int i = 0; i<2; i++) {
System.out.println();//new line
for(int j = 0; j<2; j++) {
System.out.print(y[i][j] + " ");
}
}
m1 = (x[0][0] + x[1][1]) * (y[0][0] + y[1][1]);
m2 = (x[1][0] + x[1][1]) * y[0][0];
m3 = x[0][0] * (y[0][1] - y[1][1]);
m4 = x[1][1] * (y[1][0] - y[0][0]);
m5 = (x[0][0] + x[0][1]) * y[1][1];
m6 = (x[1][0] - x[0][0]) * (y[0][0]+y[0][1]);
m7 = (x[0][1] - x[1][1]) * (y[1][0]+y[1][1]);
z[0][0] = m1 + m4- m5 + m7;
z[0][1] = m3 + m5;
z[1][0] = m2 + m4;
z[1][1] = m1 - m2 + m3 + m6;
System.out.print("
Product achieved using Strassen's algorithm: ");
for(int i = 0; i<2; i++) {
System.out.println();//new line
for(int j = 0; j<2; j++) {
System.out.print(z[i][j] + " ");
}
}
}
}
输出
The first matrix is: 12 34 22 10 The second matrix is: 3 4 2 1 Product achieved using Strassen's algorithm: 104 82 86 98
import numpy as np
x = np.array([[12, 34], [22, 10]])
y = np.array([[3, 4], [2, 1]])
z = np.zeros((2, 2))
m1, m2, m3, m4, m5, m6, m7 = 0, 0, 0, 0, 0, 0, 0
print("The first matrix is: ")
for i in range(2):
print()
for j in range(2):
print(x[i][j], end=" ")
print("
The second matrix is: ")
for i in range(2):
print()
for j in range(2):
print(y[i][j], end=" ")
m1 = (x[0][0] + x[1][1]) * (y[0][0] + y[1][1])
m2 = (x[1][0] + x[1][1]) * y[0][0]
m3 = x[0][0] * (y[0][1] - y[1][1])
m4 = x[1][1] * (y[1][0] - y[0][0])
m5 = (x[0][0] + x[0][1]) * y[1][1]
m6 = (x[1][0] - x[0][0]) * (y[0][0] + y[0][1])
m7 = (x[0][1] - x[1][1]) * (y[1][0] + y[1][1])
z[0][0] = m1 + m4 - m5 + m7
z[0][1] = m3 + m5
z[1][0] = m2 + m4
z[1][1] = m1 - m2 + m3 + m6
print("
Product achieved using Strassen's algorithm: ")
for i in range(2):
print()
for j in range(2):
print(z[i][j], end=" ")
输出
The first matrix is: 12 34 22 10 The second matrix is: 3 4 2 1 Product achieved using Strassen's algorithm: 104.0 82.0 86.0 98.0

