diff options
Diffstat (limited to 'src/matrix.cpp')
| -rw-r--r-- | src/matrix.cpp | 162 |
1 files changed, 57 insertions, 105 deletions
diff --git a/src/matrix.cpp b/src/matrix.cpp index c27264b..75fcb07 100644 --- a/src/matrix.cpp +++ b/src/matrix.cpp @@ -1,120 +1,76 @@ -#include "matrix.hpp" -#include "linalg_error.hpp" +module; -#include <algorithm> -#include <cstddef> -#include <sstream> -#include <stdexcept> - -#if !defined(LINEAR_ALGEBRA_FORCE_SCALAR_MATMUL) && \ - (defined(__AVX__) || defined(__AVX2__) || defined(__AVX512F__)) -#include <immintrin.h> -#endif - -#if !defined(LINEAR_ALGEBRA_FORCE_SCALAR_MATMUL) && defined(__ARM_NEON) && defined(__aarch64__) && defined(__ARM_FEATURE_FP64_VECTOR_ARITHMETIC) +#if !defined(LINEAR_ALGEBRA_FORCE_SCALAR_MATMUL) && defined(__ARM_NEON) #include <arm_neon.h> #endif -namespace linalg { +export module linalgebra:matrix; +import std; +import :error; +import :vector; -namespace { +export namespace linalgebra { -void check_same_shape(const Matrix& lhs, const Matrix& rhs, const char* operation) { - if (lhs.rows() != rhs.rows() || lhs.cols() != rhs.cols()) { - std::ostringstream oss; - oss << operation << " requires equal matrix shapes, got " << lhs.rows() << "x" << lhs.cols() << " and " << rhs.rows() << "x" << rhs.cols(); - throw DimensionMismatchError(oss.str()); - } -} +class Matrix { +public: + Matrix() = default; + Matrix(std::size_t rows, std::size_t cols); + Matrix(std::size_t rows, std::size_t cols, double value); + Matrix(std::initializer_list<std::initializer_list<double>> values); -double dot_product_scalar(const double* lhs, const double* rhs, std::size_t count) { - double sum = 0.0; - for (std::size_t i = 0; i < count; ++i) { - sum += lhs[i] * rhs[i]; - } - return sum; -} - -#if !defined(LINEAR_ALGEBRA_FORCE_SCALAR_MATMUL) && defined(__AVX512F__) -double horizontal_sum(__m512d values) { - alignas(64) double lanes[8]; - _mm512_store_pd(lanes, values); - double sum = 0.0; - for (double lane : lanes) { - sum += lane; - } - return sum; -} + [[nodiscard]] std::size_t rows() const noexcept; + [[nodiscard]] std::size_t cols() const noexcept; + [[nodiscard]] bool empty() const noexcept; -double dot_product_avx512(const double* lhs, const double* rhs, std::size_t count) { - std::size_t i = 0; - __m512d acc0 = _mm512_setzero_pd(); - __m512d acc1 = _mm512_setzero_pd(); + double& operator()(std::size_t i, std::size_t j); + const double& operator()(std::size_t i, std::size_t j) const; - for (; i + 15 < count; i += 16) { - const __m512d lhs0 = _mm512_loadu_pd(lhs + i); - const __m512d rhs0 = _mm512_loadu_pd(rhs + i); - const __m512d lhs1 = _mm512_loadu_pd(lhs + i + 8); - const __m512d rhs1 = _mm512_loadu_pd(rhs + i + 8); + void fill(double value); - acc0 = _mm512_add_pd(acc0, _mm512_mul_pd(lhs0, rhs0)); - acc1 = _mm512_add_pd(acc1, _mm512_mul_pd(lhs1, rhs1)); - } + double* data() noexcept; + const double* data() const noexcept; - return horizontal_sum(acc0) + horizontal_sum(acc1) + - dot_product_scalar(lhs + i, rhs + i, count - i); -} -#endif + static Matrix identity(std::size_t n); + static Matrix zeros(std::size_t rows, std::size_t cols); -#if !defined(LINEAR_ALGEBRA_FORCE_SCALAR_MATMUL) && defined(__AVX2__) -double horizontal_sum(__m256d values) { - alignas(32) double lanes[4]; - _mm256_store_pd(lanes, values); - return lanes[0] + lanes[1] + lanes[2] + lanes[3]; -} +private: + [[nodiscard]] std::size_t index(std::size_t i, std::size_t j) const; + void check_bounds(std::size_t i, std::size_t j) const; -double dot_product_avx2(const double* lhs, const double* rhs, std::size_t count) { - std::size_t i = 0; - __m256d acc0 = _mm256_setzero_pd(); - __m256d acc1 = _mm256_setzero_pd(); + std::size_t rows_ = 0; + std::size_t cols_ = 0; + std::vector<double> data_; +}; - for (; i + 7 < count; i += 8) { - const __m256d lhs0 = _mm256_loadu_pd(lhs + i); - const __m256d rhs0 = _mm256_loadu_pd(rhs + i); - const __m256d lhs1 = _mm256_loadu_pd(lhs + i + 4); - const __m256d rhs1 = _mm256_loadu_pd(rhs + i + 4); +Matrix transpose(const Matrix& matrix); +Matrix operator+(const Matrix& lhs, const Matrix& rhs); +Matrix operator-(const Matrix& lhs, const Matrix& rhs); +Vector operator*(const Matrix& matrix, const Vector& vector); +Matrix operator*(const Matrix& lhs, const Matrix& rhs); - acc0 = _mm256_add_pd(acc0, _mm256_mul_pd(lhs0, rhs0)); - acc1 = _mm256_add_pd(acc1, _mm256_mul_pd(lhs1, rhs1)); - } +} // namespace linalgebra - return horizontal_sum(acc0) + horizontal_sum(acc1) + - dot_product_scalar(lhs + i, rhs + i, count - i); -} -#endif +namespace { -#if !defined(LINEAR_ALGEBRA_FORCE_SCALAR_MATMUL) && defined(__AVX__) && !defined(__AVX2__) -double horizontal_sum(__m256d values) { - alignas(32) double lanes[4]; - _mm256_store_pd(lanes, values); - return lanes[0] + lanes[1] + lanes[2] + lanes[3]; +void check_same_shape(const linalgebra::Matrix& lhs, const linalgebra::Matrix& rhs, + const char* operation) { + if (lhs.rows() != rhs.rows() || lhs.cols() != rhs.cols()) { + std::ostringstream oss; + oss << operation << " requires equal matrix shapes, got " << lhs.rows() << "x" + << lhs.cols() << " and " << rhs.rows() << "x" << rhs.cols(); + throw linalgebra::DimensionMismatchError(oss.str()); + } } -double dot_product_avx(const double* lhs, const double* rhs, std::size_t count) { - std::size_t i = 0; - __m256d acc = _mm256_setzero_pd(); - - for (; i + 3 < count; i += 4) { - const __m256d lhs_values = _mm256_loadu_pd(lhs + i); - const __m256d rhs_values = _mm256_loadu_pd(rhs + i); - acc = _mm256_add_pd(acc, _mm256_mul_pd(lhs_values, rhs_values)); +double dot_product_scalar(const double* lhs, const double* rhs, std::size_t count) { + double sum = 0.0; + for (std::size_t i = 0; i < count; ++i) { + sum += lhs[i] * rhs[i]; } - - return horizontal_sum(acc) + dot_product_scalar(lhs + i, rhs + i, count - i); + return sum; } -#endif -#if !defined(LINEAR_ALGEBRA_FORCE_SCALAR_MATMUL) && defined(__ARM_NEON) && defined(__aarch64__) && defined(__ARM_FEATURE_FP64_VECTOR_ARITHMETIC) +#if !defined(LINEAR_ALGEBRA_FORCE_SCALAR_MATMUL) && defined(__ARM_NEON) double horizontal_sum(float64x2_t values) { return vgetq_lane_f64(values, 0) + vgetq_lane_f64(values, 1); } @@ -140,13 +96,7 @@ double dot_product_neon(const double* lhs, const double* rhs, std::size_t count) #endif double dot_product_simd(const double* lhs, const double* rhs, std::size_t count) { -#if !defined(LINEAR_ALGEBRA_FORCE_SCALAR_MATMUL) && defined(__AVX512F__) - return dot_product_avx512(lhs, rhs, count); -#elif !defined(LINEAR_ALGEBRA_FORCE_SCALAR_MATMUL) && defined(__AVX2__) - return dot_product_avx2(lhs, rhs, count); -#elif !defined(LINEAR_ALGEBRA_FORCE_SCALAR_MATMUL) && defined(__AVX__) && !defined(__AVX2__) - return dot_product_avx(lhs, rhs, count); -#elif !defined(LINEAR_ALGEBRA_FORCE_SCALAR_MATMUL) && defined(__ARM_NEON) && defined(__aarch64__) && defined(__ARM_FEATURE_FP64_VECTOR_ARITHMETIC) +#if !defined(LINEAR_ALGEBRA_FORCE_SCALAR_MATMUL) && defined(__ARM_NEON) return dot_product_neon(lhs, rhs, count); #else return dot_product_scalar(lhs, rhs, count); @@ -155,13 +105,16 @@ double dot_product_simd(const double* lhs, const double* rhs, std::size_t count) } // namespace +namespace linalgebra { + Matrix::Matrix(std::size_t rows, std::size_t cols) : rows_(rows), cols_(cols), data_(rows * cols) {} Matrix::Matrix(std::size_t rows, std::size_t cols, double value) : rows_(rows), cols_(cols), data_(rows * cols, value) {} -Matrix::Matrix(std::initializer_list<std::initializer_list<double>> values) : rows_(values.size()) { +Matrix::Matrix(std::initializer_list<std::initializer_list<double>> values) + : rows_(values.size()) { if (rows_ == 0) { cols_ = 0; return; @@ -178,7 +131,6 @@ Matrix::Matrix(std::initializer_list<std::initializer_list<double>> values) : ro } } - std::size_t Matrix::rows() const noexcept { return rows_; } std::size_t Matrix::cols() const noexcept { return cols_; } @@ -298,4 +250,4 @@ Matrix operator*(const Matrix& lhs, const Matrix& rhs) { return result; } -} // namespace linalg +} // namespace linalgebra |