ESPResSo
Extensible Simulation Package for Research on Soft Matter Systems
Loading...
Searching...
No Matches
matrix.hpp
Go to the documentation of this file.
1/*
2 * Copyright (C) 2010-2026 The ESPResSo project
3 *
4 * This file is part of ESPResSo.
5 *
6 * ESPResSo is free software: you can redistribute it and/or modify
7 * it under the terms of the GNU General Public License as published by
8 * the Free Software Foundation, either version 3 of the License, or
9 * (at your option) any later version.
10 *
11 * ESPResSo is distributed in the hope that it will be useful,
12 * but WITHOUT ANY WARRANTY; without even the implied warranty of
13 * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
14 * GNU General Public License for more details.
15 *
16 * You should have received a copy of the GNU General Public License
17 * along with this program. If not, see <http://www.gnu.org/licenses/>.
18 */
19
20#pragma once
21
22/**
23 * @file
24 *
25 * @brief Matrix implementation and trait types
26 * for boost qvm interoperability.
27 */
28
29#include "utils/Array.hpp"
30#include "utils/Vector.hpp"
31#include "utils/flatten.hpp"
32
33#include <algorithm>
34#include <array>
35#include <cassert>
36#include <cstddef>
37#include <functional>
38#include <initializer_list>
39#include <numeric>
40#include <type_traits>
41#include <utility>
42
43// These includes need to come first due to ADL reasons.
44// clang-format off
45#include <boost/qvm/mat_operations.hpp>
46#include <boost/qvm/vec_mat_operations.hpp>
47#include <boost/qvm/vec_operations.hpp>
48// clang-format on
49
50#include <boost/qvm/deduce_mat.hpp>
51#include <boost/qvm/deduce_scalar.hpp>
52#include <boost/qvm/deduce_vec.hpp>
53#include <boost/qvm/map_mat_mat.hpp>
54#include <boost/qvm/map_mat_vec.hpp>
55#include <boost/qvm/map_vec_mat.hpp>
56#include <boost/qvm/mat.hpp>
57#include <boost/qvm/mat_access.hpp>
58#include <boost/qvm/mat_traits.hpp>
59
60namespace Utils {
61
62/**
63 * @brief Matrix representation with static size.
64 * @tparam T The data type.
65 * @tparam Rows Number of rows.
66 * @tparam Cols Number of columns.
67 */
68template <typename T, std::size_t Rows, std::size_t Cols> struct Matrix {
77
79
80private:
82 template <class Archive> void serialize(Archive &ar, const unsigned int) {
83 ar & m_data;
84 }
85
86public:
87 Matrix() = default;
88 Matrix(std::initializer_list<T> init_list) {
89 assert(init_list.size() == Rows * Cols);
90 std::ranges::copy(init_list, begin());
91 }
92 Matrix(std::initializer_list<std::initializer_list<T>> init_list) {
93 assert(init_list.size() == Rows);
95 }
96
97 /**
98 * @brief Element access (const).
99 * @param row The row used for access.
100 * @param col The column used for access.
101 * @return The matrix element at row @p row and column @p col.
102 */
103 constexpr value_type operator()(std::size_t row, std::size_t col) const {
104 assert(row < Rows);
105 assert(col < Cols);
106 return m_data[Cols * row + col];
107 }
108 /**
109 * @brief Element access (non const).
110 * @param row The row used for access.
111 * @param col The column used for access.
112 * @return The matrix element at row @p row and column @p col.
113 */
114 constexpr reference operator()(std::size_t row, std::size_t col) {
115 assert(row < Rows);
116 assert(col < Cols);
117 return m_data[Cols * row + col];
118 }
119
120 /**
121 * @brief Access to the underlying data pointer (non const).
122 * @return Pointer to first element of the data.
123 */
124 constexpr pointer data() { return m_data.data(); }
125 /**
126 * @brief Access to the underlying data pointer (non const).
127 * @return Pointer to first element of the data.
128 */
129 constexpr const_pointer data() const noexcept { return m_data.data(); }
130 /**
131 * @brief Iterator access (non const).
132 * @return Returns an iterator to the first element of the matrix.
133 */
134 constexpr iterator begin() noexcept { return m_data.begin(); }
135 /**
136 * @brief Iterator access (const).
137 * @return Returns an iterator to the first element of the matrix.
138 */
139 constexpr const_iterator begin() const noexcept { return m_data.begin(); }
140 /**
141 * @brief Iterator access (non const).
142 * @return Returns an iterator to the element following the last element of
143 * the matrix.
144 */
145 constexpr iterator end() noexcept { return m_data.end(); }
146 /**
147 * @brief Iterator access (non const).
148 * @return Returns an iterator to the element following the last element of
149 * the matrix.
150 */
151 constexpr const_iterator end() const noexcept { return m_data.end(); }
152 /**
153 * @brief Retrieve an entire matrix row.
154 * @tparam R The row index.
155 * @return A vector containing the elements of row @p R.
156 */
157 template <std::size_t R> Vector<T, Cols> row() const {
158 static_assert(R < Rows, "Invalid row index.");
159 return boost::qvm::row<R>(*this);
160 }
161 /**
162 * @brief Retrieve an entire matrix column.
163 * @tparam C The column index.
164 * @return A vector containing the elements of column @p C.
165 */
166 template <std::size_t C> Vector<T, Rows> col() const {
167 static_assert(C < Cols, "Invalid column index.");
168 return boost::qvm::col<C>(*this);
169 }
170 /**
171 * @brief Retrieve the diagonal.
172 * @return Vector containing the diagonal elements of the matrix.
173 */
175 static_assert(Rows == Cols,
176 "Diagonal can only be retrieved from square matrices.");
177 return boost::qvm::diag(*this);
178 }
179 /**
180 * @brief Retrieve the trace.
181 * @return Vector containing the sum of diagonal matrix elements.
182 */
183 T trace() const {
184 auto const d = diagonal();
185 return std::accumulate(d.begin(), d.end(), T{}, std::plus<T>{});
186 }
187
188 /**
189 * @brief Retrieve a transposed copy of the matrix.
190 * @return Transposed matrix.
191 */
193 return boost::qvm::transposed(*this);
194 }
195
196 /**
197 * @brief Retrieve an inverted copy of the matrix.
198 * @return Inverted matrix.
199 */
201 static_assert(Rows == Cols,
202 "Inversion of a non-square matrix not implemented.");
203 return boost::qvm::inverse(*this);
204 }
205 /**
206 * @brief Retrieve the shape of the matrix.
207 * @return Pair containing number of rows and number of columns of the matrix.
208 */
209 constexpr std::pair<std::size_t, std::size_t> shape() const noexcept {
210 return {Rows, Cols};
211 }
212
216};
217
218using boost::qvm::operator+;
219using boost::qvm::operator+=;
220using boost::qvm::operator-;
221using boost::qvm::operator-=;
222using boost::qvm::operator*;
223using boost::qvm::operator*=;
224using boost::qvm::operator==;
225
226template <typename T, std::size_t M, std::size_t N>
228 return m.flatten();
229}
230
231template <typename T, std::size_t Rows, std::size_t Cols>
233 static_assert(Rows == Cols, "Diagonal matrix has to be a square matrix.");
234 return boost::qvm::diag_mat(v);
235}
236
237template <typename T, std::size_t Rows, std::size_t Cols>
239 static_assert(Rows == Cols,
240 "Identity matrix only defined for square matrices.");
241 return boost::qvm::identity_mat<T, Rows>();
242}
243
244} // namespace Utils
245
246namespace boost::qvm {
247
248template <typename T, std::size_t Rows, std::size_t Cols>
249struct mat_traits<Utils::Matrix<T, Rows, Cols>> {
251 static int const rows = Rows;
252 static int const cols = Cols;
253 using scalar_type = T;
254
255 template <std::size_t R, std::size_t C>
256 static inline scalar_type read_element(mat_type const &m) {
257 static_assert(R < Rows, "Invalid row index.");
258 static_assert(C < Cols, "Invalid column index.");
259 return m(R, C);
260 }
261
262 template <std::size_t R, std::size_t C>
263 static inline scalar_type &write_element(mat_type &m) {
264 static_assert(R < Rows, "Invalid row index.");
265 static_assert(C < Cols, "Invalid column index.");
266 return m(R, C);
267 }
268
269 static inline scalar_type read_element_idx(std::size_t r, std::size_t c,
270 mat_type const &m) {
271 assert(r < Rows);
272 assert(c < Cols);
273 return m(r, c);
274 }
275 static inline scalar_type &write_element_idx(std::size_t r, std::size_t c,
276 mat_type &m) {
277 assert(r < Rows);
278 assert(c < Cols);
279 return m(r, c);
280 }
281};
282
283template <typename T, typename U>
284struct deduce_vec2<Utils::Matrix<T, 2, 2>, Utils::Vector<U, 2>, 2> {
286};
287
288template <typename T, typename U>
289struct deduce_vec2<Utils::Matrix<T, 3, 3>, Utils::Vector<U, 3>, 3> {
291};
292
293template <typename T, typename U>
294struct deduce_vec2<Utils::Matrix<T, 4, 4>, Utils::Vector<U, 4>, 4> {
296};
297
298template <typename T, typename U>
299struct deduce_vec2<Utils::Matrix<T, 2, 3>, Utils::Vector<U, 3>, 2> {
301};
302
303template <typename T, typename U>
304struct deduce_mat2<Utils::Matrix<T, 3, 3>, Utils::Matrix<U, 3, 3>, 3, 3> {
306};
307
308} // namespace boost::qvm
Array implementation with CUDA support.
Vector implementation and trait types for boost qvm interoperability.
cudaStream_t stream[1]
CUDA streams for parallel computing on CPU and GPU.
void flatten(Range const &v, OutputIterator out)
Flatten a range of ranges.
Definition flatten.hpp:56
Matrix< T, Rows, Cols > identity_mat()
Definition matrix.hpp:238
Matrix< T, Rows, Cols > diagonal_mat(Utils::Vector< T, Rows > const &v)
Definition matrix.hpp:232
DEVICE_QUALIFIER constexpr pointer data() noexcept
Definition Array.hpp:133
DEVICE_QUALIFIER constexpr iterator begin() noexcept
Definition Array.hpp:141
const value_type & const_reference
Definition Array.hpp:90
const value_type * const_pointer
Definition Array.hpp:94
DEVICE_QUALIFIER constexpr iterator end() noexcept
Definition Array.hpp:153
const value_type * const_iterator
Definition Array.hpp:92
Matrix representation with static size.
Definition matrix.hpp:68
constexpr iterator begin() noexcept
Iterator access (non const).
Definition matrix.hpp:134
Matrix< T, Cols, Rows > transposed() const
Retrieve a transposed copy of the matrix.
Definition matrix.hpp:192
container::value_type value_type
Definition matrix.hpp:74
constexpr const_pointer data() const noexcept
Access to the underlying data pointer (non const).
Definition matrix.hpp:129
container m_data
Definition matrix.hpp:78
container::const_reference const_reference
Definition matrix.hpp:76
Vector< T, Cols > diagonal() const
Retrieve the diagonal.
Definition matrix.hpp:174
Matrix(std::initializer_list< std::initializer_list< T > > init_list)
Definition matrix.hpp:92
container::reference reference
Definition matrix.hpp:75
container::pointer pointer
Definition matrix.hpp:70
constexpr const_iterator end() const noexcept
Iterator access (non const).
Definition matrix.hpp:151
T trace() const
Retrieve the trace.
Definition matrix.hpp:183
Vector< T, Cols > row() const
Retrieve an entire matrix row.
Definition matrix.hpp:157
Matrix()=default
Matrix< T, Rows, Cols > inversed() const
Retrieve an inverted copy of the matrix.
Definition matrix.hpp:200
Vector< T, Rows > col() const
Retrieve an entire matrix column.
Definition matrix.hpp:166
constexpr std::pair< std::size_t, std::size_t > shape() const noexcept
Retrieve the shape of the matrix.
Definition matrix.hpp:209
container::const_iterator const_iterator
Definition matrix.hpp:73
container::const_pointer const_pointer
Definition matrix.hpp:71
auto flatten() const noexcept
Definition matrix.hpp:213
friend class boost::serialization::access
Definition matrix.hpp:81
constexpr iterator end() noexcept
Iterator access (non const).
Definition matrix.hpp:145
container::iterator iterator
Definition matrix.hpp:72
constexpr pointer data()
Access to the underlying data pointer (non const).
Definition matrix.hpp:124
constexpr value_type operator()(std::size_t row, std::size_t col) const
Element access (const).
Definition matrix.hpp:103
constexpr const_iterator begin() const noexcept
Iterator access (const).
Definition matrix.hpp:139
constexpr reference operator()(std::size_t row, std::size_t col)
Element access (non const).
Definition matrix.hpp:114
Matrix(std::initializer_list< T > init_list)
Definition matrix.hpp:88
static scalar_type read_element_idx(std::size_t r, std::size_t c, mat_type const &m)
Definition matrix.hpp:269
static scalar_type & write_element_idx(std::size_t r, std::size_t c, mat_type &m)
Definition matrix.hpp:275
static scalar_type read_element(mat_type const &m)
Definition matrix.hpp:256
static scalar_type & write_element(mat_type &m)
Definition matrix.hpp:263