strassen.cpp (4097B)
1 // Experimental benchmarking of Strassen's Algorithm vs plain Eigen matrix multiply. 2 3 #include <Eigen/Dense> 4 #include <ctime> 5 #include <iostream> 6 7 // Not remotely extensible, just initial test. 8 Eigen::MatrixXf strassen_mul(Eigen::MatrixXf x, Eigen::MatrixXf y) 9 { 10 int power = 0; 11 int largest; 12 if (x.rows() >= x.cols() || x.rows() >= y.cols() || x.rows() >= y.rows()) largest = x.rows(); 13 else if (x.cols() >= x.rows() || x.cols() >= y.cols() || x.cols() >= y.rows()) largest = x.cols(); 14 else if (y.cols() >= x.rows() || y.cols() >= x.cols() || y.cols() >= y.rows()) largest = y.cols(); 15 else if (y.rows() >= x.rows() || y.rows() >= x.cols() || y.rows() >= y.cols()) largest = y.rows(); 16 for (; pow(2, power) < largest; power++) { 17 std::cout << pow(2, power) << " " << power <<"\n"; 18 } 19 std::cout << "POWER: " << power << " LARGEST: " << largest << "\n"; 20 Eigen::MatrixXf a ((int)pow(2,power), (int)pow(2,power)); 21 Eigen::MatrixXf b ((int)pow(2,power), (int)pow(2,power)); 22 a.block(0,0,x.rows(), x.cols()) = x; 23 b.block(0,0,y.rows(), y.cols()) = y; 24 std::cout << "\nINIT\n" << a << "\n\n" << b << "\n\n\n"; 25 std::cout << "\nEIGEN_VER\n" << a*b << "\n\n\n"; 26 27 int block_len = largest/2; 28 Eigen::MatrixXf result ((int)pow(2,power), (int)pow(2,power)); 29 30 Eigen::MatrixXf m1 = ((a.block(0,0, block_len, block_len)) + a.block(a.rows()-block_len,a.cols()-block_len, block_len, block_len)) * 31 (b.block(0,0, block_len, block_len) + b.block(b.rows()-block_len,b.cols()-block_len, block_len, block_len)); 32 Eigen::MatrixXf 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)) * 33 (b.block(0,0, block_len, block_len)); 34 Eigen::MatrixXf m3 = a.block(0,0, block_len, block_len) * 35 (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)); 36 Eigen::MatrixXf m4 = a.block(a.rows()-block_len,a.cols()-block_len, block_len, block_len) * 37 (b.block(b.rows()-block_len,0, block_len, block_len) - b.block(0,0, block_len, block_len)); 38 Eigen::MatrixXf m5 = (a.block(0, 0, block_len, block_len) + a.block(0,a.cols()-block_len, block_len, block_len)) * 39 (b.block(b.rows()-block_len,b.cols()-block_len, block_len, block_len)); 40 Eigen::MatrixXf m6 = (a.block(a.rows()-block_len,0, block_len, block_len) - a.block(0,0, block_len, block_len)) * 41 (b.block(0,0, block_len, block_len) + b.block(0,b.cols()-block_len, block_len, block_len)); 42 Eigen::MatrixXf 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)) * 43 (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)); 44 45 result.block(0,0, block_len, block_len) = m1 + m4 - m5 + m7; 46 result.block(0,result.cols()-block_len, block_len, block_len) = m3 + m5; 47 result.block(result.rows()-block_len,0, block_len, block_len) = m2 + m4; 48 result.block(result.rows()-block_len,result.cols()-block_len, block_len, block_len) = m1 -m2 + m3 + m6; 49 std::cout << "\nSTRASSEN_VER\n" << result << "\n\n\n"; 50 return result; 51 } 52 53 int main() 54 { 55 srand((unsigned int) time(0)); 56 int sz; 57 std::cin >> sz; 58 Eigen::MatrixXf a = Eigen::MatrixXf::Random(sz, 3); 59 Eigen::MatrixXf b = Eigen::MatrixXf::Random(3, sz); 60 61 auto eigen_begin = std::chrono::high_resolution_clock::now(); 62 Eigen::MatrixXf product = a * b; 63 auto eigen_end = std::chrono::high_resolution_clock::now(); 64 65 auto strassen_begin = std::chrono::high_resolution_clock::now(); 66 Eigen::MatrixXf sproduct = strassen_mul(a, b); 67 auto strassen_end = std::chrono::high_resolution_clock::now(); 68 69 std::cout << "EIGEN: " << std::chrono::duration_cast<std::chrono::nanoseconds>(eigen_end - eigen_begin).count() / pow(10,9) 70 << " STRASSEN: " << std::chrono::duration_cast<std::chrono::nanoseconds>(strassen_end - strassen_begin).count() / pow(10,9) << "\n"; 71 std::cout << "A:\n" << a << "\nB:\n" << b << "\nEigen:\n" << product << "\nStrassen:\n" << sproduct << "\n"; 72 }