BELFEM 0.9.0
Berkeley Lab Finite Element Framework
Loading...
Searching...
No Matches
fn_gemm.hpp
Go to the documentation of this file.
1/*
2 * BELFEM -- The Berkeley Lab Finite Element Framework
3 * Copyright (c) 2026, The Regents of the University of California,
4 * through Lawrence Berkeley National Laboratory (subject to receipt of any required
5 * approvals from the U.S. Dept. of Energy). All rights reserved.
6 *
7 * Developers: Christian Messe, Gregory Giard
8 *
9 * See the top-level LICENSE file for the complete license and disclaimer.
10 */
11
21
22#ifndef BELFEM_FN_GEMM_HPP
23#define BELFEM_FN_GEMM_HPP
24
25#include "assert.hpp"
26#include "lapacktools.hpp"
27
28//------------------------------------------------------------------------------
29namespace belfem
30{
31 namespace lapack
32 {
33
34#ifdef __cplusplus
35 extern "C"
36 {
37#endif
38//------------------------------------------------------------------------------
39
40 void
41 sgemm_( char * transa,
42 char * transb,
43 int_t * m,
44 int_t * n,
45 int_t * k,
46 float * alpha,
47 float * a,
48 int_t * lda,
49 float * b,
50 int_t * ldb,
51 float * beta,
52 float * c,
53 int_t * ldc,
56 );
57
58//------------------------------------------------------------------------------
59
60 void
61 dgemm_( char * transa,
62 char * transb,
63 int_t * m,
64 int_t * n,
65 int_t * k,
66 double * alpha,
67 double * a,
68 int_t * lda,
69 double * b,
70 int_t * ldb,
71 double * beta,
72 double * c,
73 int_t * ldc,
76//------------------------------------------------------------------------------
77
78 void
79 cgemm_( char * transa,
80 char * transb,
81 int_t * m,
82 int_t * n,
83 int_t * k,
84 cplx_float_t * alpha,
85 cplx_float_t * a,
86 int_t * lda,
87 cplx_float_t * b,
88 int_t * ldb,
89 cplx_float_t * beta,
90 cplx_float_t * c,
91 int_t * ldc,
94
95//------------------------------------------------------------------------------
96
97 void
98 zgemm_( char * transa,
99 char * transb,
100 int_t * m,
101 int_t * n,
102 int_t * k,
103 cplx_double_t * alpha,
104 cplx_double_t * a,
105 int_t * lda,
106 cplx_double_t * b,
107 int_t * ldb,
108 cplx_double_t * beta,
109 cplx_double_t * c,
110 int_t * ldc,
112 fortran_charlen_t ltb );
113
114//------------------------------------------------------------------------------
115
116#ifdef __cplusplus
117}
118#endif
119
120//------------------------------------------------------------------------------
121
122 template< typename T >
123 void
124 gemm( const char * transa,
125 const char * transb,
126 const int_t * m,
127 const int_t * n,
128 const int_t * k,
129 const T * alpha,
130 const T * a,
131 const int_t * lda,
132 const T * b,
133 const int_t * ldb,
134 const T * beta,
135 T * c,
136 const int_t * ldc )
137 {
138 static_assert( dependent_false< T >,
139 "gemm not implemented for selected data type" );
140 }
141
142// - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -
143
144 template<>
145 inline void
146 gemm( const char * transa,
147 const char * transb,
148 const int_t * m,
149 const int_t * n,
150 const int_t * k,
151 const float * alpha,
152 const float * a,
153 const int_t * lda,
154 const float * b,
155 const int_t * ldb,
156 const float * beta,
157 float * c,
158 const int_t * ldc )
159 {
160 sgemm_( const_cast< char * > ( transa ),
161 const_cast< char * > ( transb ),
162 const_cast< int_t * > ( m ),
163 const_cast< int_t * > ( n ),
164 const_cast< int_t * > ( k ),
165 const_cast< float * > ( alpha ),
166 const_cast< float * > ( a ),
167 const_cast< int_t * > ( lda ),
168 const_cast< float * > ( b ),
169 const_cast< int_t * > ( ldb ),
170 const_cast< float * > ( beta ),
171 c,
172 const_cast< int_t * > ( ldc ),
173 1, 1 );
174 }
175
176// - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -
177
178 template<>
179 inline void
180 gemm( const char * transa,
181 const char * transb,
182 const int_t * m,
183 const int_t * n,
184 const int_t * k,
185 const double * alpha,
186 const double * a,
187 const int_t * lda,
188 const double * b,
189 const int_t * ldb,
190 const double * beta,
191 double * c,
192 const int_t * ldc )
193 {
194 dgemm_( const_cast< char * > ( transa ),
195 const_cast< char * > ( transb ),
196 const_cast< int_t * > ( m ),
197 const_cast< int_t * > ( n ),
198 const_cast< int_t * > ( k ),
199 const_cast< double * >( alpha ),
200 const_cast< double * >( a ),
201 const_cast< int_t * > ( lda ),
202 const_cast< double * >( b ),
203 const_cast< int_t * > ( ldb ),
204 const_cast< double * >( beta ),
205 c,
206 const_cast< int_t * > ( ldc ),
207 1, 1 );
208 }
209
210// - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -
211
212 template<>
213 inline void
214 gemm( const char * transa,
215 const char * transb,
216 const int_t * m,
217 const int_t * n,
218 const int_t * k,
219 const std::complex< float > * alpha,
220 const std::complex< float > * a,
221 const int_t * lda,
222 const std::complex< float > * b,
223 const int_t * ldb,
224 const std::complex< float > * beta,
225 std::complex< float > * c,
226 const int_t * ldc )
227 {
228 cgemm_( const_cast< char * > ( transa ),
229 const_cast< char * > ( transb ),
230 const_cast< int_t * > ( m ),
231 const_cast< int_t * > ( n ),
232 const_cast< int_t * > ( k ),
233 const_cast< cplx_float_t * >( reinterpret_cast< const cplx_float_t * >( alpha ) ),
234 const_cast< cplx_float_t * >( reinterpret_cast< const cplx_float_t * >( a ) ),
235 const_cast< int_t * > ( lda ),
236 const_cast< cplx_float_t * >( reinterpret_cast< const cplx_float_t * >( b ) ),
237 const_cast< int_t * > ( ldb ),
238 const_cast< cplx_float_t * >( reinterpret_cast< const cplx_float_t * >( beta ) ),
239 reinterpret_cast< cplx_float_t * >( c ),
240 const_cast< int_t * > ( ldc ),
241 1, 1 );
242 }
243
244// - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - - -
245
246 template<>
247 inline void
248 gemm( const char * transa,
249 const char * transb,
250 const int_t * m,
251 const int_t * n,
252 const int_t * k,
253 const std::complex< double > * alpha,
254 const std::complex< double > * a,
255 const int_t * lda,
256 const std::complex< double > * b,
257 const int_t * ldb,
258 const std::complex< double > * beta,
259 std::complex< double > * c,
260 const int_t * ldc )
261 {
262 zgemm_( const_cast< char * > ( transa ),
263 const_cast< char * > ( transb ),
264 const_cast< int_t * > ( m ),
265 const_cast< int_t * > ( n ),
266 const_cast< int_t * > ( k ),
267 const_cast< cplx_double_t * >( reinterpret_cast< const cplx_double_t * >( alpha ) ),
268 const_cast< cplx_double_t * >( reinterpret_cast< const cplx_double_t * >( a ) ),
269 const_cast< int_t * > ( lda ),
270 const_cast< cplx_double_t * >( reinterpret_cast< const cplx_double_t * >( b ) ),
271 const_cast< int_t * > ( ldb ),
272 const_cast< cplx_double_t * >( reinterpret_cast< const cplx_double_t * >( beta ) ),
273 reinterpret_cast< cplx_double_t * >( c ),
274 const_cast< int_t * > ( ldc ),
275 1, 1 );
276 }
277
278//------------------------------------------------------------------------------
279 } /* end namespace lapack */
280//------------------------------------------------------------------------------
281
296 template< typename T >
297 void
299 const Matrix< T > & A,
300 const Matrix< T > & B,
301 Matrix< T > & C,
302 const T alpha = 1.0,
303 const T beta = 0.0,
304 const char transa = 'N',
305 const char transb = 'N' )
306 {
307 // dimensions of op(A) ( m x k ) and op(B) ( k x n )
308 const bool tNoTransA = ( transa == 'N' || transa == 'n' );
309 const bool tNoTransB = ( transb == 'N' || transb == 'n' );
310
311 BELFEM_ASSERT( transa == 'N' || transa == 'T' || transa == 'C' ||
312 transa == 'n' || transa == 't' || transa == 'c',
313 "unsupported transa flag '%c'", transa );
314
315 BELFEM_ASSERT( transb == 'N' || transb == 'T' || transb == 'C' ||
316 transb == 'n' || transb == 't' || transb == 'c',
317 "unsupported transb flag '%c'", transb );
318
319 int_t m = ( int_t ) ( tNoTransA ? A.n_rows() : A.n_cols() );
320 int_t k = ( int_t ) ( tNoTransA ? A.n_cols() : A.n_rows() );
321 int_t n = ( int_t ) ( tNoTransB ? B.n_cols() : B.n_rows() );
322
323 BELFEM_ASSERT( ( int_t ) ( tNoTransB ? B.n_rows() : B.n_cols() ) == k,
324 "Inner dimensions of op(A) and op(B) do not match ( %u vs %u )",
325 ( unsigned int ) k,
326 ( unsigned int ) ( tNoTransB ? B.n_rows() : B.n_cols() ) );
327
328 if ( beta == static_cast< T >( 0.0 ) )
329 {
330 C.set_size( m, n, 0.0 );
331 }
332
333 BELFEM_ASSERT( ( int_t ) C.n_rows() == m,
334 "Number of rows of op(A) and C do not match ( %u vs %u )",
335 ( unsigned int ) m, ( unsigned int ) C.n_rows() );
336
337 BELFEM_ASSERT( ( int_t ) C.n_cols() == n,
338 "Number of cols of op(B) and C do not match ( %u vs %u )",
339 ( unsigned int ) n, ( unsigned int ) C.n_cols() );
340
341 // leading dimensions describe the stored arrays and are
342 // independent of the transposition flags
346
347 lapack::gemm( &transa,
348 &transb,
349 &m,
350 &n,
351 &k,
352 &alpha,
353 A.data(),
354 &lda,
355 B.data(),
356 &ldb,
357 &beta,
358 C.data(),
359 &ldc );
360 }
361
362//------------------------------------------------------------------------------
363
364} /* end namespace belfem */
365//------------------------------------------------------------------------------
366#endif //BELFEM_FN_GEMM_HPP
#define BELFEM_ASSERT(aCheck,...)
Definition assert.hpp:244
Dense column-major matrix.
Definition cl_BZ_Matrix.hpp:28
size_t n_rows() const
Definition cl_AR_Matrix.hpp:205
size_t n_cols() const
Definition cl_AR_Matrix.hpp:213
T * data()
Definition cl_AR_Matrix.hpp:135
Shared helpers for the LAPACK wrappers, such as leading_dimension().
Definition fn_gees.hpp:33
void zgemm_(char *transa, char *transb, int_t *m, int_t *n, int_t *k, cplx_double_t *alpha, cplx_double_t *a, int_t *lda, cplx_double_t *b, int_t *ldb, cplx_double_t *beta, cplx_double_t *c, int_t *ldc, fortran_charlen_t lta, fortran_charlen_t ltb)
void sgemm_(char *transa, char *transb, int_t *m, int_t *n, int_t *k, float *alpha, float *a, int_t *lda, float *b, int_t *ldb, float *beta, float *c, int_t *ldc, fortran_charlen_t lta, fortran_charlen_t ltb)
constexpr bool dependent_false
Definition lapacktools.hpp:49
float cplx_float_t
Definition lapacktools.hpp:90
int_t leading_dimension(const Vector< T > &A)
logical length of a vector operand, as passed to LAPACK as LDB; vector storage is contiguous under bo...
Definition lapacktools.hpp:143
void gemm(const char *transa, const char *transb, const int_t *m, const int_t *n, const int_t *k, const T *alpha, const T *a, const int_t *lda, const T *b, const int_t *ldb, const T *beta, T *c, const int_t *ldc)
Definition fn_gemm.hpp:124
size_t fortran_charlen_t
Definition lapacktools.hpp:100
void dgemm_(char *transa, char *transb, int_t *m, int_t *n, int_t *k, double *alpha, double *a, int_t *lda, double *b, int_t *ldb, double *beta, double *c, int_t *ldc, fortran_charlen_t lta, fortran_charlen_t ltb)
double cplx_double_t
Definition lapacktools.hpp:91
void cgemm_(char *transa, char *transb, int_t *m, int_t *n, int_t *k, cplx_float_t *alpha, cplx_float_t *a, int_t *lda, cplx_float_t *b, int_t *ldb, cplx_float_t *beta, cplx_float_t *c, int_t *ldc, fortran_charlen_t lta, fortran_charlen_t ltb)
USER GUIDES:
Definition cl_Capacitor.cpp:16
void gemm(const Matrix< T > &A, const Matrix< T > &B, Matrix< T > &C, const T alpha=1.0, const T beta=0.0, const char transa='N', const char transb='N')
matrix-matrix product C := alpha * op(A) * op(B) + beta * C via BLAS ?gemm, where op is the identity ...
Definition fn_gemm.hpp:298
@ alpha
Definition cl_Material.hpp:161
int32_t int_t
Definition typedefs.hpp:51