From b3c8438bf1430ff001da2f8492d010b110cf1546 Mon Sep 17 00:00:00 2001 From: Orange Date: Thu, 16 Jul 2026 01:10:55 +0300 Subject: [PATCH] improved matrix multiplication --- benchmark/benchmark_mat.cpp | 32 ++-- include/omath/linear_algebra/mat.hpp | 219 +++++++++++++++++++-------- tests/general/unit_test_mat.cpp | 54 +++++++ 3 files changed, 233 insertions(+), 72 deletions(-) diff --git a/benchmark/benchmark_mat.cpp b/benchmark/benchmark_mat.cpp index 2e54382..a8da8dd 100644 --- a/benchmark/benchmark_mat.cpp +++ b/benchmark/benchmark_mat.cpp @@ -2,11 +2,9 @@ // Created by Vlad on 9/17/2025. // #include - #include using namespace omath; - void mat_float_multiplication_col_major(benchmark::State& state) { using MatType = Mat<128, 128, float, MatStoreType::COLUMN_MAJOR>; @@ -15,9 +13,12 @@ void mat_float_multiplication_col_major(benchmark::State& state) a.set(3.f); b.set(7.f); - for ([[maybe_unused]] const auto _ : state) - std::ignore = a * b; + { + benchmark::DoNotOptimize(a); + benchmark::DoNotOptimize(b); + benchmark::DoNotOptimize(a * b); + } } void mat_float_multiplication_row_major(benchmark::State& state) { @@ -27,9 +28,12 @@ void mat_float_multiplication_row_major(benchmark::State& state) a.set(3.f); b.set(7.f); - for ([[maybe_unused]] const auto _ : state) - std::ignore = a * b; + { + benchmark::DoNotOptimize(a); + benchmark::DoNotOptimize(b); + benchmark::DoNotOptimize(a * b); + } } void mat_double_multiplication_row_major(benchmark::State& state) @@ -40,9 +44,12 @@ void mat_double_multiplication_row_major(benchmark::State& state) a.set(3.f); b.set(7.f); - for ([[maybe_unused]] const auto _ : state) - std::ignore = a * b; + { + benchmark::DoNotOptimize(a); + benchmark::DoNotOptimize(b); + benchmark::DoNotOptimize(a * b); + } } void mat_double_multiplication_col_major(benchmark::State& state) @@ -53,13 +60,16 @@ void mat_double_multiplication_col_major(benchmark::State& state) a.set(3.f); b.set(7.f); - for ([[maybe_unused]] const auto _ : state) - std::ignore = a * b; + { + benchmark::DoNotOptimize(a); + benchmark::DoNotOptimize(b); + benchmark::DoNotOptimize(a * b); + } } BENCHMARK(mat_float_multiplication_col_major)->Iterations(5000); BENCHMARK(mat_float_multiplication_row_major)->Iterations(5000); BENCHMARK(mat_double_multiplication_col_major)->Iterations(5000); -BENCHMARK(mat_double_multiplication_row_major)->Iterations(5000); \ No newline at end of file +BENCHMARK(mat_double_multiplication_row_major)->Iterations(5000); diff --git a/include/omath/linear_algebra/mat.hpp b/include/omath/linear_algebra/mat.hpp index 6ac96cf..aba4be6 100644 --- a/include/omath/linear_algebra/mat.hpp +++ b/include/omath/linear_algebra/mat.hpp @@ -186,7 +186,14 @@ namespace omath else if constexpr (StoreType == MatStoreType::COLUMN_MAJOR) return cache_friendly_multiply_col_major(other); } - if constexpr (StoreType == MatStoreType::ROW_MAJOR) + if constexpr (!std::is_same_v && !std::is_same_v) + { + if constexpr (StoreType == MatStoreType::ROW_MAJOR) + return cache_friendly_multiply_row_major(other); + else if constexpr (StoreType == MatStoreType::COLUMN_MAJOR) + return cache_friendly_multiply_col_major(other); + } + else if constexpr (StoreType == MatStoreType::ROW_MAJOR) return avx_multiply_row_major(other); else if constexpr (StoreType == MatStoreType::COLUMN_MAJOR) return avx_multiply_col_major(other); @@ -429,13 +436,22 @@ namespace omath cache_friendly_multiply_row_major(const Mat& other) const { Mat result; + const Type* left_data = m_data.data(); + const Type* right_data = other.raw_array().data(); + Type* result_data = result.raw_array().data(); + for (std::size_t row_index = 0; row_index < Rows; ++row_index) + { + const Type* left_row = left_data + row_index * Columns; + Type* result_row = result_data + row_index * OtherColumns; for (std::size_t column_index = 0; column_index < Columns; ++column_index) { - const Type& current_number = at(row_index, column_index); + const Type current_number = left_row[column_index]; + const Type* right_row = right_data + column_index * OtherColumns; for (std::size_t other_column = 0; other_column < OtherColumns; ++other_column) - result.at(row_index, other_column) += current_number * other.at(column_index, other_column); + result_row[other_column] += current_number * right_row[other_column]; } + } return result; } @@ -444,13 +460,22 @@ namespace omath const Mat& other) const { Mat result; + const Type* left_data = m_data.data(); + const Type* right_data = other.raw_array().data(); + Type* result_data = result.raw_array().data(); + for (std::size_t other_column = 0; other_column < OtherColumns; ++other_column) + { + const Type* right_column = right_data + other_column * Columns; + Type* result_column = result_data + other_column * Rows; for (std::size_t column_index = 0; column_index < Columns; ++column_index) { - const Type& current_number = other.at(column_index, other_column); + const Type current_number = right_column[column_index]; + const Type* left_column = left_data + column_index * Rows; for (std::size_t row_index = 0; row_index < Rows; ++row_index) - result.at(row_index, other_column) += at(row_index, column_index) * current_number; + result_column[row_index] += left_column[row_index] * current_number; } + } return result; } #ifdef OMATH_USE_AVX2 @@ -466,56 +491,92 @@ namespace omath if constexpr (std::is_same_v) { - // ReSharper disable once CppTooWideScopeInitStatement constexpr std::size_t vector_size = 8; + constexpr std::size_t block_size = vector_size * 4; for (std::size_t j = 0; j < OtherColumns; ++j) { auto* c_col = reinterpret_cast(result_mat_data + j * Rows); - for (std::size_t k = 0; k < Columns; ++k) + std::size_t i = 0; + for (; i + block_size <= Rows; i += block_size) { - const float bkj = reinterpret_cast(other_mat_data)[k + j * Columns]; - const __m256 bkj_vec = _mm256_set1_ps(bkj); - - const auto* a_col_k = reinterpret_cast(this_mat_data + k * Rows); - - std::size_t i = 0; - for (; i + vector_size <= Rows; i += vector_size) + __m256 cvec0 = _mm256_setzero_ps(); + __m256 cvec1 = _mm256_setzero_ps(); + __m256 cvec2 = _mm256_setzero_ps(); + __m256 cvec3 = _mm256_setzero_ps(); + for (std::size_t k = 0; k < Columns; ++k) { - __m256 cvec = _mm256_loadu_ps(c_col + i); + const __m256 bkj_vec = _mm256_set1_ps(other_mat_data[k + j * Columns]); + const auto* a_col_k = this_mat_data + k * Rows + i; + cvec0 = _mm256_fmadd_ps(_mm256_loadu_ps(a_col_k), bkj_vec, cvec0); + cvec1 = _mm256_fmadd_ps(_mm256_loadu_ps(a_col_k + vector_size), bkj_vec, cvec1); + cvec2 = _mm256_fmadd_ps(_mm256_loadu_ps(a_col_k + vector_size * 2), bkj_vec, cvec2); + cvec3 = _mm256_fmadd_ps(_mm256_loadu_ps(a_col_k + vector_size * 3), bkj_vec, cvec3); + } + _mm256_storeu_ps(c_col + i, cvec0); + _mm256_storeu_ps(c_col + i + vector_size, cvec1); + _mm256_storeu_ps(c_col + i + vector_size * 2, cvec2); + _mm256_storeu_ps(c_col + i + vector_size * 3, cvec3); + } + for (; i + vector_size <= Rows; i += vector_size) + { + __m256 cvec = _mm256_setzero_ps(); + for (std::size_t k = 0; k < Columns; ++k) + { + const __m256 bkj_vec = _mm256_set1_ps(other_mat_data[k + j * Columns]); + const auto* a_col_k = this_mat_data + k * Rows; const __m256 a_vec = _mm256_loadu_ps(a_col_k + i); cvec = _mm256_fmadd_ps(a_vec, bkj_vec, cvec); - _mm256_storeu_ps(c_col + i, cvec); } - for (; i < Rows; ++i) - c_col[i] += a_col_k[i] * bkj; + _mm256_storeu_ps(c_col + i, cvec); } + for (; i < Rows; ++i) + for (std::size_t k = 0; k < Columns; ++k) + c_col[i] += this_mat_data[i + k * Rows] * other_mat_data[k + j * Columns]; } } else if (std::is_same_v) - { // double - // ReSharper disable once CppTooWideScopeInitStatement + { constexpr std::size_t vector_size = 4; + constexpr std::size_t block_size = vector_size * 4; for (std::size_t j = 0; j < OtherColumns; ++j) { auto* c_col = reinterpret_cast(result_mat_data + j * Rows); - for (std::size_t k = 0; k < Columns; ++k) + std::size_t i = 0; + for (; i + block_size <= Rows; i += block_size) { - const double bkj = reinterpret_cast(other_mat_data)[k + j * Columns]; - const __m256d bkj_vec = _mm256_set1_pd(bkj); - - const auto* a_col_k = reinterpret_cast(this_mat_data + k * Rows); - - std::size_t i = 0; - for (; i + vector_size <= Rows; i += vector_size) + __m256d cvec0 = _mm256_setzero_pd(); + __m256d cvec1 = _mm256_setzero_pd(); + __m256d cvec2 = _mm256_setzero_pd(); + __m256d cvec3 = _mm256_setzero_pd(); + for (std::size_t k = 0; k < Columns; ++k) { - __m256d cvec = _mm256_loadu_pd(c_col + i); + const __m256d bkj_vec = _mm256_set1_pd(other_mat_data[k + j * Columns]); + const auto* a_col_k = this_mat_data + k * Rows + i; + cvec0 = _mm256_fmadd_pd(_mm256_loadu_pd(a_col_k), bkj_vec, cvec0); + cvec1 = _mm256_fmadd_pd(_mm256_loadu_pd(a_col_k + vector_size), bkj_vec, cvec1); + cvec2 = _mm256_fmadd_pd(_mm256_loadu_pd(a_col_k + vector_size * 2), bkj_vec, cvec2); + cvec3 = _mm256_fmadd_pd(_mm256_loadu_pd(a_col_k + vector_size * 3), bkj_vec, cvec3); + } + _mm256_storeu_pd(c_col + i, cvec0); + _mm256_storeu_pd(c_col + i + vector_size, cvec1); + _mm256_storeu_pd(c_col + i + vector_size * 2, cvec2); + _mm256_storeu_pd(c_col + i + vector_size * 3, cvec3); + } + for (; i + vector_size <= Rows; i += vector_size) + { + __m256d cvec = _mm256_setzero_pd(); + for (std::size_t k = 0; k < Columns; ++k) + { + const __m256d bkj_vec = _mm256_set1_pd(other_mat_data[k + j * Columns]); + const auto* a_col_k = this_mat_data + k * Rows; const __m256d a_vec = _mm256_loadu_pd(a_col_k + i); cvec = _mm256_fmadd_pd(a_vec, bkj_vec, cvec); - _mm256_storeu_pd(c_col + i, cvec); } - for (; i < Rows; ++i) - c_col[i] += a_col_k[i] * bkj; + _mm256_storeu_pd(c_col + i, cvec); } + for (; i < Rows; ++i) + for (std::size_t k = 0; k < Columns; ++k) + c_col[i] += this_mat_data[i + k * Rows] * other_mat_data[k + j * Columns]; } } else @@ -536,56 +597,92 @@ namespace omath if constexpr (std::is_same_v) { - // ReSharper disable once CppTooWideScopeInitStatement constexpr std::size_t vector_size = 8; + constexpr std::size_t block_size = vector_size * 4; for (std::size_t i = 0; i < Rows; ++i) { - Type* c_row = result_mat_data + i * OtherColumns; - for (std::size_t k = 0; k < Columns; ++k) + auto* c_row = reinterpret_cast(result_mat_data + i * OtherColumns); + std::size_t j = 0; + for (; j + block_size <= OtherColumns; j += block_size) { - const auto aik = static_cast(this_mat_data[i * Columns + k]); - const __m256 aik_vec = _mm256_set1_ps(aik); - const auto* b_row = reinterpret_cast(other_mat_data + k * OtherColumns); - - std::size_t j = 0; - for (; j + vector_size <= OtherColumns; j += vector_size) + __m256 cvec0 = _mm256_setzero_ps(); + __m256 cvec1 = _mm256_setzero_ps(); + __m256 cvec2 = _mm256_setzero_ps(); + __m256 cvec3 = _mm256_setzero_ps(); + for (std::size_t k = 0; k < Columns; ++k) { - __m256 cvec = _mm256_loadu_ps(c_row + j); + const __m256 aik_vec = _mm256_set1_ps(this_mat_data[i * Columns + k]); + const auto* b_row = other_mat_data + k * OtherColumns + j; + cvec0 = _mm256_fmadd_ps(_mm256_loadu_ps(b_row), aik_vec, cvec0); + cvec1 = _mm256_fmadd_ps(_mm256_loadu_ps(b_row + vector_size), aik_vec, cvec1); + cvec2 = _mm256_fmadd_ps(_mm256_loadu_ps(b_row + vector_size * 2), aik_vec, cvec2); + cvec3 = _mm256_fmadd_ps(_mm256_loadu_ps(b_row + vector_size * 3), aik_vec, cvec3); + } + _mm256_storeu_ps(c_row + j, cvec0); + _mm256_storeu_ps(c_row + j + vector_size, cvec1); + _mm256_storeu_ps(c_row + j + vector_size * 2, cvec2); + _mm256_storeu_ps(c_row + j + vector_size * 3, cvec3); + } + for (; j + vector_size <= OtherColumns; j += vector_size) + { + __m256 cvec = _mm256_setzero_ps(); + for (std::size_t k = 0; k < Columns; ++k) + { + const __m256 aik_vec = _mm256_set1_ps(this_mat_data[i * Columns + k]); + const auto* b_row = other_mat_data + k * OtherColumns; const __m256 b_vec = _mm256_loadu_ps(b_row + j); cvec = _mm256_fmadd_ps(b_vec, aik_vec, cvec); - - _mm256_storeu_ps(c_row + j, cvec); } - for (; j < OtherColumns; ++j) - c_row[j] += aik * b_row[j]; + _mm256_storeu_ps(c_row + j, cvec); } + for (; j < OtherColumns; ++j) + for (std::size_t k = 0; k < Columns; ++k) + c_row[j] += this_mat_data[i * Columns + k] * other_mat_data[k * OtherColumns + j]; } } else if (std::is_same_v) - { // double - // ReSharper disable once CppTooWideScopeInitStatement + { constexpr std::size_t vector_size = 4; + constexpr std::size_t block_size = vector_size * 4; for (std::size_t i = 0; i < Rows; ++i) { - Type* c_row = result_mat_data + i * OtherColumns; - for (std::size_t k = 0; k < Columns; ++k) + auto* c_row = reinterpret_cast(result_mat_data + i * OtherColumns); + std::size_t j = 0; + for (; j + block_size <= OtherColumns; j += block_size) { - const auto aik = static_cast(this_mat_data[i * Columns + k]); - const __m256d aik_vec = _mm256_set1_pd(aik); - const auto* b_row = reinterpret_cast(other_mat_data + k * OtherColumns); - - std::size_t j = 0; - for (; j + vector_size <= OtherColumns; j += vector_size) + __m256d cvec0 = _mm256_setzero_pd(); + __m256d cvec1 = _mm256_setzero_pd(); + __m256d cvec2 = _mm256_setzero_pd(); + __m256d cvec3 = _mm256_setzero_pd(); + for (std::size_t k = 0; k < Columns; ++k) { - __m256d cvec = _mm256_loadu_pd(c_row + j); + const __m256d aik_vec = _mm256_set1_pd(this_mat_data[i * Columns + k]); + const auto* b_row = other_mat_data + k * OtherColumns + j; + cvec0 = _mm256_fmadd_pd(_mm256_loadu_pd(b_row), aik_vec, cvec0); + cvec1 = _mm256_fmadd_pd(_mm256_loadu_pd(b_row + vector_size), aik_vec, cvec1); + cvec2 = _mm256_fmadd_pd(_mm256_loadu_pd(b_row + vector_size * 2), aik_vec, cvec2); + cvec3 = _mm256_fmadd_pd(_mm256_loadu_pd(b_row + vector_size * 3), aik_vec, cvec3); + } + _mm256_storeu_pd(c_row + j, cvec0); + _mm256_storeu_pd(c_row + j + vector_size, cvec1); + _mm256_storeu_pd(c_row + j + vector_size * 2, cvec2); + _mm256_storeu_pd(c_row + j + vector_size * 3, cvec3); + } + for (; j + vector_size <= OtherColumns; j += vector_size) + { + __m256d cvec = _mm256_setzero_pd(); + for (std::size_t k = 0; k < Columns; ++k) + { + const __m256d aik_vec = _mm256_set1_pd(this_mat_data[i * Columns + k]); + const auto* b_row = other_mat_data + k * OtherColumns; const __m256d b_vec = _mm256_loadu_pd(b_row + j); cvec = _mm256_fmadd_pd(b_vec, aik_vec, cvec); - - _mm256_storeu_pd(c_row + j, cvec); } - for (; j < OtherColumns; ++j) - c_row[j] += aik * b_row[j]; + _mm256_storeu_pd(c_row + j, cvec); } + for (; j < OtherColumns; ++j) + for (std::size_t k = 0; k < Columns; ++k) + c_row[j] += this_mat_data[i * Columns + k] * other_mat_data[k * OtherColumns + j]; } } else diff --git a/tests/general/unit_test_mat.cpp b/tests/general/unit_test_mat.cpp index e2e3654..e133313 100644 --- a/tests/general/unit_test_mat.cpp +++ b/tests/general/unit_test_mat.cpp @@ -16,6 +16,21 @@ namespace const float diff = actual - expected; return (diff < 0.0f ? -diff : diff) <= epsilon; } + + template + void expect_multiplication_matches_scalar_reference(const Mat& left, + const Mat& right) + { + const auto result = left * right; + for (size_t row = 0; row < Rows; ++row) + for (size_t column = 0; column < OtherColumns; ++column) + { + Type expected{}; + for (size_t shared_index = 0; shared_index < Columns; ++shared_index) + expected += left.at(row, shared_index) * right.at(shared_index, column); + EXPECT_EQ(result.at(row, column), expected); + } + } } // namespace class UnitTestMat : public ::testing::Test @@ -92,6 +107,45 @@ TEST_F(UnitTestMat, Operator_Multiplication_Matrix) EXPECT_FLOAT_EQ(m3.at(1, 1), 22.0f); } +TEST(UnitTestMatStandalone, Operator_Multiplication_RowMajorSimdAndTail) +{ + Mat<3, 5, float, MatStoreType::ROW_MAJOR> left; + Mat<5, 33, float, MatStoreType::ROW_MAJOR> right; + for (size_t row = 0; row < left.row_count(); ++row) + for (size_t column = 0; column < left.columns_count(); ++column) + left.at(row, column) = static_cast(row * 3 + column + 1); + for (size_t row = 0; row < right.row_count(); ++row) + for (size_t column = 0; column < right.columns_count(); ++column) + right.at(row, column) = static_cast((row + 1) * (column % 5 + 1)); + + expect_multiplication_matches_scalar_reference(left, right); +} + +TEST(UnitTestMatStandalone, Operator_Multiplication_ColumnMajorSimdAndTail) +{ + Mat<17, 5, double, MatStoreType::COLUMN_MAJOR> left; + Mat<5, 3, double, MatStoreType::COLUMN_MAJOR> right; + for (size_t row = 0; row < left.row_count(); ++row) + for (size_t column = 0; column < left.columns_count(); ++column) + left.at(row, column) = static_cast(row * 3 + column + 1); + for (size_t row = 0; row < right.row_count(); ++row) + for (size_t column = 0; column < right.columns_count(); ++column) + right.at(row, column) = static_cast((row + 1) * (column + 1)); + + expect_multiplication_matches_scalar_reference(left, right); +} + +TEST(UnitTestMatStandalone, Operator_Multiplication_IntegerFallsBackFromAvx) +{ + constexpr Mat<2, 3, int> left{{1, 2, 3}, {4, 5, 6}}; + constexpr Mat<3, 2, int> right{{7, 8}, {9, 10}, {11, 12}}; + constexpr auto result = left * right; + static_assert(result.at(0, 0) == 58); + static_assert(result.at(1, 1) == 154); + + expect_multiplication_matches_scalar_reference(left, right); +} + TEST_F(UnitTestMat, Operator_Multiplication_Scalar) { Mat<2, 2> m3 = m2 * 2.0f;