commit eb279d71b41c4b784abb25f49f33de62b5801ee0
parent cd95f09c03c6f3b7659a30274fbdd60e3c295d2b
Author: David Freifeld <freifeld.david@gmail.com>
Date: Wed, 1 Jul 2020 11:44:49 -0700
Some fixes to errors, one remains
Diffstat:
1 file changed, 28 insertions(+), 23 deletions(-)
diff --git a/tests/strassen.cpp b/tests/strassen.cpp
@@ -7,37 +7,42 @@
//
#include <Eigen/Dense>
+#include <ctime>
+#include <iostream>
// Not remotely extensible, just initial test.
Eigen::MatrixXd strassen_mul(Eigen::MatrixXd a, Eigen::MatrixXd b)
{
- int block_len = a.rows()/4;
+ int block_len = 1250;
Eigen::MatrixXd result (a.rows(), a.cols());
// Nigh unreadable code follows.
- Eigen::MatrixXd m1 = (a.block<block_len, block_len>(0,0) + a.block<block_len, block_len>(a.rows()-block_len,a.cols()-block_len)) * (b.block<block_len, block_len>(0,0) + b.block<block_len, block_len>(b.rows()-block_len,b.cols()-block_len));
- Eigen::MatrixXd m2 = (a.block<block_len, block_len>(a.rows()-block_len, 0) + a.block<block_len, block_len>(a.rows()-block_len,a.cols()-block_len)) * (b.block<block_len, block_len>(0,0));
- Eigen::MatrixXd m3 = a.block<block_len, block_len>(0,0) * (b.block<block_len, block_len>(0,b.cols()-block_len) - b.block<block_len, block_len>(b.rows()-block_len,b.cols()-block_len));
- Eigen::MatrixXd m4 = a.block<block_len, block_len>(a.rows()-block_len,a.cols()-block_len) * (b.block<block_len, block_len>(b.rows()-block_len,0) - b.block<block_len, block_len>(0,0));
- Eigen::MatrixXd m5 = (a.block<block_len, block_len>(0, 0) + a.block<block_len, block_len>(0,a.cols()-block_len)) * (b.block<block_len, block_len>(b.rows()-block_len,b.cols()-block_len));
- Eigen::MatrixXd m6 = (a.block<block_len, block_len>(a.rows()-block_len,0) - a.block<block_len, block_len>(0,0)) * (b.block<block_len, block_len>(0,0) + b.block<block_len, block_len>(0,b.cols()-block_len));
- Eigen::MatrixXd m7 = (a.block<block_len, block_len>(0,a.cols()-block_len) - a.block<block_len, block_len>(a.rows()-block_len,a.cols()-block_len)) * (b.block<block_len, block_len>(a.rows()-block_len,0) + b.block<block_len, block_len>(b.rows()-block_len,b.cols()-block_len));
+ 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));
+ 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));
+ 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));
+ 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));
+ 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));
+ 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<block_len, block_len>(0,0) = m1 + m4 - m5 + m7;
- result.block<block_len, block_len>(0,results.cols()-block_len) = m3 + m5;
- result.block<block_len, block_len>(results.rows()-block_len,0) = m2 + m4;
- result.block<block_len, block_len>(results.rows()-block_len,result.cols()-block_len) = m1 - m2 + m3 + m6;
- return result
+ 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;
+ result.block(result.rows()-block_len,0, block_len, block_len) = m2 + m4;
+ result.block(result.rows()-block_len,result.cols()-block_len, block_len, block_len) = m1 - m2 + m3 + m6;
+ return result;
}
-Eigen::MatrixXd a = Eigen::MatrixXd::Random(5000, 5000);
-Eigen::MatrixXd b = Eigen::MatrixXd::Random(5000, 5000);
-
-auto eigen_begin = std::chrono::high_resolution_clock::now();
-Eigen::MatrixXd product = a * b;
-auto eigen_end = std::chrono::high_resolution_clock::now();
+int main()
+{
+ Eigen::MatrixXd a = Eigen::MatrixXd::Random(5000, 5000);
+ Eigen::MatrixXd b = Eigen::MatrixXd::Random(5000, 5000);
-auto strassen_begin = std::chrono::high_resolution_clock::now();
-Eigen::MatrixXd sproduct = strassen_mul(a, b);
-auto strassen_end = std::chrono::high_resolution_clock::now();
+ auto eigen_begin = std::chrono::high_resolution_clock::now();
+ Eigen::MatrixXd product = a * b;
+ auto eigen_end = std::chrono::high_resolution_clock::now();
+
+ auto strassen_begin = std::chrono::high_resolution_clock::now();
+ Eigen::MatrixXd sproduct = strassen_mul(a, b);
+ auto strassen_end = std::chrono::high_resolution_clock::now();
-std::cout << "EIGEN: " << std::chrono::duration_cast<std::chrono::nanoseconds>(eigen_end - eigen_begin) << " STRASSEN: " << std::chrono::duration_cast<std::chrono::nanoseconds>(strassen_end - strassen_begin);
+ std::cout << "EIGEN: " << std::chrono::duration_cast<std::chrono::nanoseconds>(eigen_end - eigen_begin) << " STRASSEN: " << std::chrono::duration_cast<std::chrono::nanoseconds>(strassen_end - strassen_begin);
+}