commit 3232fe1fd1bc2e63adc6aed5292def5825ada1b5
parent a362378bd7aa6718c464da24a60bd0b00db500e4
Author: David Freifeld <freifeld.david@gmail.com>
Date: Fri, 3 Jul 2020 22:26:55 -0700
Strassen's algorithm is functional
Diffstat:
1 file changed, 7 insertions(+), 14 deletions(-)
diff --git a/tests/strassen.cpp b/tests/strassen.cpp
@@ -16,22 +16,13 @@ Eigen::MatrixXd strassen_mul(Eigen::MatrixXd a, Eigen::MatrixXd b)
int block_len = a.rows()/2;
Eigen::MatrixXd result (a.rows(), a.cols());
- // (A+D)(E+H)
- Eigen::MatrixXd m1 = (a.block(0,0, block_len, block_len)) + a.block(a.rows()-block_len,a.cols()-block_len, block_len, block_len) * (b.block(0,0, block_len, block_len) + b.block(b.rows()-block_len,b.cols()-block_len, block_len, block_len));
- // (C+D)E
+ Eigen::MatrixXd m1 = ((a.block(0,0, block_len, block_len)) + a.block(a.rows()-block_len,a.cols()-block_len, block_len, block_len)) * (b.block(0,0, block_len, block_len) + b.block(b.rows()-block_len,b.cols()-block_len, block_len, block_len));
Eigen::MatrixXd m2 = (a.block(a.rows()-block_len, 0, block_len, block_len) + a.block(a.rows()-block_len,a.cols()-block_len, block_len, block_len)) * (b.block(0,0, block_len, block_len));
- // A(F-H)
Eigen::MatrixXd m3 = a.block(0,0, block_len, block_len) * (b.block(0,b.cols()-block_len, block_len, block_len) - b.block(b.rows()-block_len,b.cols()-block_len, block_len, block_len));
- // D(G-E)
Eigen::MatrixXd m4 = a.block(a.rows()-block_len,a.cols()-block_len, block_len, block_len) * (b.block(b.rows()-block_len,0, block_len, block_len) - b.block(0,0, block_len, block_len));
- // (A+B)H
Eigen::MatrixXd m5 = (a.block(0, 0, block_len, block_len) + a.block(0,a.cols()-block_len, block_len, block_len)) * (b.block(b.rows()-block_len,b.cols()-block_len, block_len, block_len));
- // (A-C)(E+F)
Eigen::MatrixXd m6 = (a.block(a.rows()-block_len,0, block_len, block_len) - a.block(0,0, block_len, block_len)) * (b.block(0,0, block_len, block_len) + b.block(0,b.cols()-block_len, block_len, block_len));
- // (B-D)(G+H)
- Eigen::MatrixXd m7 = (a.block(0,a.cols()-block_len, block_len, block_len) - a.block(a.rows()-block_len,a.cols()-block_len, block_len, block_len)) * (b.block(a.rows()-block_len,0, block_len, block_len) + b.block(b.rows()-block_len,b.cols()-block_len, block_len, block_len));
-
- std::cout << m1 << "\n\n" << m2 << "\n\n" << m3 << "\n\n" << m4 << "\n\n" << m5 << "\n\n" << m6 << "\n\n" << m7 << "\n\n";
+ Eigen::MatrixXd m7 = (a.block(0,a.cols()-block_len, block_len, block_len) - a.block(a.rows()-block_len,a.cols()-block_len, block_len, block_len)) * (b.block(a.rows()-block_len,0, block_len, block_len) + b.block(b.rows()-block_len,b.cols()-block_len, block_len, block_len))
result.block(0,0, block_len, block_len) = m1 + m4 - m5 + m7;
result.block(0,result.cols()-block_len, block_len, block_len) = m3 + m5;
@@ -44,8 +35,10 @@ Eigen::MatrixXd strassen_mul(Eigen::MatrixXd a, Eigen::MatrixXd b)
int main()
{
srand((unsigned int) time(0));
- Eigen::MatrixXd a = Eigen::MatrixXd::Random(4, 4);
- Eigen::MatrixXd b = Eigen::MatrixXd::Random(4, 4);
+ int sz;
+ std::cin >> sz;
+ Eigen::MatrixXd a = Eigen::MatrixXd::Random(sz, sz);
+ Eigen::MatrixXd b = Eigen::MatrixXd::Random(sz, sz);
auto eigen_begin = std::chrono::high_resolution_clock::now();
Eigen::MatrixXd product = a * b;
@@ -56,5 +49,5 @@ int main()
auto strassen_end = std::chrono::high_resolution_clock::now();
std::cout << "EIGEN: " << std::chrono::duration_cast<std::chrono::nanoseconds>(eigen_end - eigen_begin).count() / pow(10,9) << " STRASSEN: " << std::chrono::duration_cast<std::chrono::nanoseconds>(strassen_end - strassen_begin).count() / pow(10,9) << "\n";
- std::cout << "A:\n" << a << "\nB:\n" << b << "\nEigen:\n" << product << "\nStrassen:\n" << sproduct << "\n";
+ //std::cout << "A:\n" << a << "\nB:\n" << b << "\nEigen:\n" << product << "\nStrassen:\n" << sproduct << "\n";
}