22 void ssyevd_(
char* jobz,
char* uplo,
int* n,
float* a,
int* lda,
float* w,
23 float* work,
int* lwork,
int* iwork,
int* liwork,
int* info);
24 void dsyevd_(
char* jobz,
char* uplo,
int* n,
double* a,
int* lda,
double* w,
25 double* work,
int* lwork,
int* iwork,
int* liwork,
int* info);
27 void sgesv_(
int* N,
int* NRHS,
float* A,
int* LDA,
int* IPIV,
float* B,
29 void dgesv_(
int* N,
int* NRHS,
double* A,
int* LDA,
int* IPIV,
double* B,
32 void sgemm_(
char* transa,
char* transb,
int* m,
int* n,
int* k,
float* alpha,
33 float* a,
int* lda,
float* b,
int* ldb,
float* beta,
float* c,
35 void dgemm_(
char* transa,
char* transb,
int* m,
int* n,
int* k,
double* alpha,
36 double* a,
int* lda,
double* b,
int* ldb,
double* beta,
double* c,
39 int sgetrf_(
const int* m,
const int* n,
float* a,
const int* lda,
int* lpiv,
41 int dgetrf_(
const int* m,
const int* n,
double* a,
const int* lda,
int* lpiv,
57template <std::
floating_po
int T>
58void dot_blas(std::span<const T>
A, std::array<std::size_t, 2>
Ashape,
59 std::span<const T>
B, std::array<std::size_t, 2>
Bshape,
62 static_assert(std::is_same_v<T, float>
or std::is_same_v<T, double>);
77 if constexpr (std::is_same_v<T, float>)
82 else if constexpr (std::is_same_v<T, double>)
95template <
typename U,
typename V>
96std::pair<std::vector<typename U::value_type>, std::array<std::size_t, 2>>
99 std::vector<typename U::value_type>
result(
u.size() *
v.size());
100 for (std::size_t
i = 0;
i <
u.size(); ++
i)
101 for (std::size_t
j = 0;
j <
v.size(); ++
j)
103 return {std::move(
result), {
u.size(),
v.size()}};
110template <
typename U,
typename V>
111std::array<typename U::value_type, 3>
cross(
const U&
u,
const V&
v)
115 return {
u[1] *
v[2] -
u[2] *
v[1],
u[2] *
v[0] -
u[0] *
v[2],
116 u[0] *
v[1] -
u[1] *
v[0]};
125template <std::
floating_po
int T>
126std::pair<std::vector<T>, std::vector<T>>
eigh(std::span<const T>
A,
130 std::vector<T> M(
A.begin(),
A.end());
133 std::vector<T>
w(
n, 0);
142 std::vector<T>
work(1);
143 std::vector<int>
iwork(1);
146 if constexpr (std::is_same_v<T, float>)
151 else if constexpr (std::is_same_v<T, double>)
158 throw std::runtime_error(
"Could not find workspace size for syevd.");
165 if constexpr (std::is_same_v<T, float>)
170 else if constexpr (std::is_same_v<T, double>)
176 throw std::runtime_error(
"Eigenvalue computation did not converge.");
178 return {std::move(
w), std::move(M)};
185template <std::
floating_po
int T>
187solve(MDSPAN_IMPL_STANDARD_NAMESPACE::mdspan<
188 const T, MDSPAN_IMPL_STANDARD_NAMESPACE::dextents<std::size_t, 2>>
190 MDSPAN_IMPL_STANDARD_NAMESPACE::mdspan<
191 const T, MDSPAN_IMPL_STANDARD_NAMESPACE::dextents<std::size_t, 2>>
195 = MDSPAN_IMPL_STANDARD_NAMESPACE::MDSPAN_IMPL_PROPOSED_NAMESPACE;
198 stdex::mdarray<T, MDSPAN_IMPL_STANDARD_NAMESPACE::dextents<std::size_t, 2>,
199 MDSPAN_IMPL_STANDARD_NAMESPACE::layout_left>
200 _A(
A.extents()),
_B(
B.extents());
201 for (std::size_t
i = 0;
i <
A.extent(0); ++
i)
202 for (std::size_t
j = 0;
j <
A.extent(1); ++
j)
204 for (std::size_t
i = 0;
i <
B.extent(0); ++
i)
205 for (std::size_t
j = 0;
j <
B.extent(1); ++
j)
208 int N =
_A.extent(0);
210 int lda =
_A.extent(0);
211 int ldb =
_B.extent(0);
213 std::vector<int>
piv(
N);
215 if constexpr (std::is_same_v<T, float>)
217 else if constexpr (std::is_same_v<T, double>)
220 throw std::runtime_error(
"Call to dgesv failed: " + std::to_string(
info));
223 std::vector<T>
rb(
_B.extent(0) *
_B.extent(1));
224 MDSPAN_IMPL_STANDARD_NAMESPACE::mdspan<
225 T, MDSPAN_IMPL_STANDARD_NAMESPACE::dextents<std::size_t, 2>>
226 r(
rb.data(),
_B.extents());
227 for (std::size_t
i = 0;
i <
_B.extent(0); ++
i)
228 for (std::size_t
j = 0;
j <
_B.extent(1); ++
j)
237template <std::
floating_po
int T>
239 MDSPAN_IMPL_STANDARD_NAMESPACE::mdspan<
240 const T, MDSPAN_IMPL_STANDARD_NAMESPACE::dextents<std::size_t, 2>>
245 = MDSPAN_IMPL_STANDARD_NAMESPACE::MDSPAN_IMPL_PROPOSED_NAMESPACE;
246 stdex::mdarray<T, MDSPAN_IMPL_STANDARD_NAMESPACE::dextents<std::size_t, 2>,
247 MDSPAN_IMPL_STANDARD_NAMESPACE::layout_left>
249 for (std::size_t
i = 0;
i <
A.extent(0); ++
i)
250 for (std::size_t
j = 0;
j <
A.extent(1); ++
j)
253 std::vector<T>
B(
A.extent(1), 1);
254 int N =
_A.extent(0);
256 int lda =
_A.extent(0);
260 std::vector<int>
piv(
N);
262 if constexpr (std::is_same_v<T, float>)
264 else if constexpr (std::is_same_v<T, double>)
269 throw std::runtime_error(
"dgesv failed due to invalid value: "
270 + std::to_string(
info));
283template <std::
floating_po
int T>
284std::vector<std::size_t>
287 std::size_t dim =
A.second[0];
294 if constexpr (std::is_same_v<T, float>)
296 else if constexpr (std::is_same_v<T, double>)
301 throw std::runtime_error(
"LU decomposition failed: "
302 + std::to_string(
info));
305 std::vector<std::size_t>
perm(dim);
306 for (std::size_t
i = 0;
i < dim; ++
i)
317template <
typename U,
typename V,
typename W>
323 if (
A.extent(0) *
B.extent(1) *
A.extent(1) < 512)
325 std::fill_n(
C.data_handle(),
C.extent(0) *
C.extent(1), 0);
326 for (std::size_t
i = 0;
i <
A.extent(0); ++
i)
327 for (std::size_t
j = 0;
j <
B.extent(1); ++
j)
328 for (std::size_t
k = 0;
k <
A.extent(1); ++
k)
333 using T =
typename std::decay_t<U>::value_type;
335 std::span(
A.data_handle(),
A.size()), {A.extent(0), A.extent(1)},
336 std::span(
B.data_handle(),
B.size()), {B.extent(0), B.extent(1)},
337 std::span(
C.data_handle(),
C.size()));
344template <std::
floating_po
int T>
345std::vector<T>
eye(std::size_t
n)
347 std::vector<T>
I(
n *
n, 0);
349 = MDSPAN_IMPL_STANDARD_NAMESPACE::MDSPAN_IMPL_PROPOSED_NAMESPACE;
350 MDSPAN_IMPL_STANDARD_NAMESPACE::mdspan<
351 T, MDSPAN_IMPL_STANDARD_NAMESPACE::dextents<std::size_t, 2>>
353 for (std::size_t
i = 0;
i <
n; ++
i)
362template <std::
floating_po
int T>
364 MDSPAN_IMPL_STANDARD_NAMESPACE::mdspan<
365 T, MDSPAN_IMPL_STANDARD_NAMESPACE::dextents<std::size_t, 2>>
367 std::size_t
start = 0)
369 for (std::size_t
i =
start;
i < wcoeffs.extent(0); ++
i)
372 for (std::size_t
k = 0;
k < wcoeffs.extent(1); ++
k)
373 norm += wcoeffs(
i,
k) * wcoeffs(
i,
k);
376 if (
norm < 2 * std::numeric_limits<T>::epsilon())
378 throw std::runtime_error(
379 "Cannot orthogonalise the rows of a matrix with incomplete row rank");
382 for (std::size_t
k = 0;
k < wcoeffs.extent(1); ++
k)
385 for (std::size_t
j =
i + 1;
j < wcoeffs.extent(0); ++
j)
388 for (std::size_t
k = 0;
k < wcoeffs.extent(1); ++
k)
389 a += wcoeffs(
i,
k) * wcoeffs(
j,
k);
390 for (std::size_t
k = 0;
k < wcoeffs.extent(1); ++
k)
391 wcoeffs(
j,
k) -=
a * wcoeffs(
i,
k);
A finite element.
Definition finite-element.h:139
Mathematical functions.
Definition math.h:50
void dot(const U &A, const V &B, W &&C)
Compute C = A * B.
Definition math.h:318
std::vector< T > solve(MDSPAN_IMPL_STANDARD_NAMESPACE::mdspan< const T, MDSPAN_IMPL_STANDARD_NAMESPACE::dextents< std::size_t, 2 > > A, MDSPAN_IMPL_STANDARD_NAMESPACE::mdspan< const T, MDSPAN_IMPL_STANDARD_NAMESPACE::dextents< std::size_t, 2 > > B)
Solve A X = B.
Definition math.h:187
std::array< typename U::value_type, 3 > cross(const U &u, const V &v)
Definition math.h:111
bool is_singular(MDSPAN_IMPL_STANDARD_NAMESPACE::mdspan< const T, MDSPAN_IMPL_STANDARD_NAMESPACE::dextents< std::size_t, 2 > > A)
Check if A is a singular matrix,.
Definition math.h:238
std::vector< std::size_t > transpose_lu(std::pair< std::vector< T >, std::array< std::size_t, 2 > > &A)
Compute the LU decomposition of the transpose of a square matrix A.
Definition math.h:285
void orthogonalise(MDSPAN_IMPL_STANDARD_NAMESPACE::mdspan< T, MDSPAN_IMPL_STANDARD_NAMESPACE::dextents< std::size_t, 2 > > wcoeffs, std::size_t start=0)
Orthogonalise the rows of a matrix (in place).
Definition math.h:363
std::pair< std::vector< T >, std::vector< T > > eigh(std::span< const T > A, std::size_t n)
Definition math.h:126
std::vector< T > eye(std::size_t n)
Build an identity matrix.
Definition math.h:345
std::pair< std::vector< typename U::value_type >, std::array< std::size_t, 2 > > outer(const U &u, const V &v)
Compute the outer product of vectors u and v.
Definition math.h:97