aboutsummaryrefslogtreecommitdiff
path: root/src/matrix.cpp
diff options
context:
space:
mode:
Diffstat (limited to 'src/matrix.cpp')
-rw-r--r--src/matrix.cpp162
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