35#ifndef MADNESS_TENSOR_MXM_H__INCLUDED
36#define MADNESS_TENSOR_MXM_H__INCLUDED
63 template <
typename T,
typename Q,
typename S>
74 for (
long i=0; i<dimi; ++i) {
75 for (
long k=0;
k<dimk; ++
k) {
76 for (
long j=0; j<dimj; ++j) {
77 c[i*dimj+j] +=
a[i*dimk+
k]*
b[
k*dimj+j];
85 template <
typename T,
typename Q,
typename S>
99 for (
long k=0;
k<dimk; ++
k) {
100 for (
long j=0; j<dimj; ++j) {
101 for (
long i=0; i<dimi; ++i) {
102 c[i*dimj+j] +=
a[
k*dimi+i]*
b[
k*dimj+j];
109 template <
typename T,
typename Q,
typename S>
122 for (
long i=0; i<dimi; ++i) {
123 for (
long j=0; j<dimj; ++j) {
125 for (
long k=0;
k<dimk; ++
k) {
126 sum +=
a[i*dimk+
k]*
b[j*dimk+
k];
134 template <
typename T,
typename Q,
typename S>
145 for (
long i=0; i<dimi; ++i) {
146 for (
long j=0; j<dimj; ++j) {
147 for (
long k=0;
k<dimk; ++
k) {
148 c[i*dimj+j] +=
a[
k*dimi+i]*
b[j*dimk+
k];
155#if defined(HAVE_FAST_BLAS) && !defined(HAVE_INTEL_MKL)
164 template <
typename T>
165 void mxm(
long dimi,
long dimj,
long dimk,
168 cblas::gemm(
cblas::NoTrans,
cblas::NoTrans,dimj,dimi,dimk,
one,
b,dimj,
a,dimk,
one,
c,dimj);
177 template <
typename T>
178 void mTxm(
long dimi,
long dimj,
long dimk,
181 cblas::gemm(
cblas::NoTrans,
cblas::Trans,dimj,dimi,dimk,
one,
b,dimj,
a,dimi,
one,
c,dimj);
190 template <
typename T>
191 void mxmT(
long dimi,
long dimj,
long dimk,
194 cblas::gemm(
cblas::Trans,
cblas::NoTrans,dimj,dimi,dimk,
one,
b,dimk,
a,dimk,
one,
c,dimj);
203 template <
typename T>
204 void mTxmT(
long dimi,
long dimj,
long dimk,
207 cblas::gemm(
cblas::Trans,
cblas::Trans,dimj,dimi,dimk,
one,
b,dimk,
a,dimi,
one,
c,dimj);
219 template <
typename aT,
typename bT,
typename cT>
220 void mxm(
long dimi,
long dimj,
long dimk,
223 cblas::gemm(
cblas::NoTrans,
cblas::NoTrans,dimj,dimi,dimk,
one,
b,dimj,
a,dimk,
one,
c,dimj);
232 template <
typename aT,
typename bT,
typename cT>
233 void mTxm(
long dimi,
long dimj,
long dimk,
236 cblas::gemm(
cblas::NoTrans,
cblas::Trans,dimj,dimi,dimk,
one,
b,dimj,
a,dimi,
one,
c,dimj);
245 template <
typename aT,
typename bT,
typename cT>
246 void mxmT(
long dimi,
long dimj,
long dimk,
249 cblas::gemm(
cblas::Trans,
cblas::NoTrans,dimj,dimi,dimk,
one,
b,dimk,
a,dimk,
one,
c,dimj);
258 template <
typename aT,
typename bT,
typename cT>
259 void mTxmT(
long dimi,
long dimj,
long dimk,
262 cblas::gemm(
cblas::Trans,
cblas::Trans,dimj,dimi,dimk,
one,
b,dimk,
a,dimi,
one,
c,dimj);
269 template <
typename T,
typename Q,
typename S>
270 static inline void mxm(
long dimi,
long dimj,
long dimk,
276 template <
typename T,
typename Q,
typename S>
278 void mTxm(
long dimi,
long dimj,
long dimk,
284 template <
typename T,
typename Q,
typename S>
285 static inline void mxmT(
long dimi,
long dimj,
long dimk,
291 template <
typename T,
typename Q,
typename S>
292 static inline void mTxmT(
long dimi,
long dimj,
long dimk,
303 inline void mTxm(
long dimi,
long dimj,
long dimk,
318 long dimk4 = (dimk/4)*4;
319 for (
long i=0; i<dimi; ++i,
c+=dimj) {
320 const double* ai =
a+i;
322 for (
long k=0;
k<dimk4;
k+=4,ai+=4*dimi,
p+=4*dimj) {
323 double ak0i = ai[0 ];
324 double ak1i = ai[dimi];
325 double ak2i = ai[dimi+dimi];
326 double ak3i = ai[dimi+dimi+dimi];
327 const double* bk0 =
p;
328 const double* bk1 =
p+dimj;
329 const double* bk2 =
p+dimj+dimj;
330 const double* bk3 =
p+dimj+dimj+dimj;
331 for (
long j=0; j<dimj; ++j) {
332 c[j] += ak0i*bk0[j] + ak1i*bk1[j] + ak2i*bk2[j] + ak3i*bk3[j];
335 for (
long k=dimk4;
k<dimk; ++
k) {
336 double aki =
a[
k*dimi+i];
337 const double* bk =
b+
k*dimj;
338 for (
long j=0; j<dimj; ++j) {
350 inline void mxmT(
long dimi,
long dimj,
long dimk,
365 long dimi2 = (dimi/2)*2;
366 for (
long i=0; i<dimi2; i+=2) {
367 const double* ai0 =
a+i*dimk;
368 const double* ai1 =
a+i*dimk+dimk;
371 for (
long j=0; j<dimj; ++j) {
374 const double* bj =
b + j*dimk;
375 for (
long k=0;
k<dimk; ++
k) {
376 sum0 += ai0[
k]*bj[
k];
377 sum1 += ai1[
k]*bj[
k];
383 for (
long i=dimi2; i<dimi; ++i) {
384 const double* ai =
a+i*dimk;
386 for (
long j=0; j<dimj; ++j) {
388 const double* bj =
b+j*dimk;
389 for (
long k=0;
k<dimk; ++
k) {
399 inline void mxm(
long dimi,
long dimj,
long dimk,
411 long dimk4 = (dimk/4)*4;
412 for (
long i=0; i<dimi; ++i,
c+=dimj,
a+=dimk) {
414 for (
long k=0;
k<dimk4;
k+=4,
p+=4*dimj) {
416 double aik1 =
a[
k+1];
417 double aik2 =
a[
k+2];
418 double aik3 =
a[
k+3];
419 const double* bk0 =
p;
420 const double* bk1 = bk0+dimj;
421 const double* bk2 = bk1+dimj;
422 const double* bk3 = bk2+dimj;
423 for (
long j=0; j<dimj; ++j) {
424 c[j] += aik0*bk0[j] + aik1*bk1[j] + aik2*bk2[j] + aik3*bk3[j];
427 for (
long k=dimk4;
k<dimk; ++
k) {
429 for (
long j=0; j<dimj; ++j) {
430 c[j] += aik*
b[
k*dimj+j];
438 inline void mTxmT(
long dimi,
long dimj,
long dimk,
451 long dimj2 = (dimj/2)*2;
453 for (
long klo=0; klo<dimk; klo+=ktile, asave+=ktile*dimi,
b+=ktile) {
454 long khi = klo+ktile;
455 if (khi > dimk) khi = dimk;
460 for (
long i=0; i<dimi; ++i,
c+=dimj,++
a) {
462 for (
long k=0;
k<nk; ++
k,
q+=dimi) ai[
k] = *
q;
464 const double* bj0 =
b;
465 for (
long j=0; j<dimj2; j+=2,bj0+=2*dimk) {
466 const double* bj1 = bj0+dimk;
469 for (
long k=0;
k<nk; ++
k) {
470 sum0 += ai[
k]*bj0[
k];
471 sum1 += ai[
k]*bj1[
k];
477 for (
long j=dimj2; j<dimj; ++j,bj0+=dimk) {
479 for (
long k=0;
k<nk; ++
k) {
double q(double t)
Definition DKops.h:18
Define BLAS like functions.
char * p(char *buf, const char *name, int k, int initial_level, double thresh, int order)
Definition derivatives.cc:72
#define MADNESS_RESTRICT
Definition mTxmq.h:37
Macros and tools pertaining to the configuration of MADNESS.
void gemm(const CBLAS_TRANSPOSE OpA, const CBLAS_TRANSPOSE OpB, const integer m, const integer n, const integer k, const float alpha, const float *a, const integer lda, const float *b, const integer ldb, const float beta, float *c, const integer ldc)
Multiplies a matrix by a vector.
Definition cblas.h:352
@ NoTrans
Definition cblas_types.h:65
@ Trans
Definition cblas_types.h:66
Namespace for all elements and tools of MADNESS.
Definition DFConvergence.h:9
static void mTxm_reference(long dimi, long dimj, long dimk, T *MADNESS_RESTRICT c, const Q *MADNESS_RESTRICT a, const S *MADNESS_RESTRICT b)
Matrix += Matrix transpose * matrix ... reference implementation (slow but correct)
Definition mxm.h:87
void mTxmT(long dimi, long dimj, long dimk, T *MADNESS_RESTRICT c, const T *a, const T *b)
Matrix += Matrix transpose * matrix transpose ... MKL interface version.
Definition mxm.h:204
void mxm(long dimi, long dimj, long dimk, T *MADNESS_RESTRICT c, const T *a, const T *b)
Matrix += Matrix * matrix ... BLAS/MKL interface version.
Definition mxm.h:165
static void mxm_reference(long dimi, long dimj, long dimk, T *MADNESS_RESTRICT c, const Q *MADNESS_RESTRICT a, const S *MADNESS_RESTRICT b)
Matrix += Matrix * matrix reference implementation (slow but correct)
Definition mxm.h:64
void mTxm(long dimi, long dimj, long dimk, T *MADNESS_RESTRICT c, const T *a, const T *b)
Matrix += Matrix transpose * matrix ... MKL interface version.
Definition mxm.h:178
static void mxmT_reference(long dimi, long dimj, long dimk, T *MADNESS_RESTRICT c, const Q *MADNESS_RESTRICT a, const S *MADNESS_RESTRICT b)
Matrix += Matrix * matrix transpose ... reference implementation (slow but correct)
Definition mxm.h:110
static void mTxmT_reference(long dimi, long dimj, long dimk, T *MADNESS_RESTRICT c, const Q *MADNESS_RESTRICT a, const S *MADNESS_RESTRICT b)
Matrix += Matrix transpose * matrix transpose reference implementation (slow but correct)
Definition mxm.h:135
void mxmT(long dimi, long dimj, long dimk, T *MADNESS_RESTRICT c, const T *a, const T *b)
Matrix += Matrix * matrix transpose ... MKL interface version.
Definition mxm.h:191
static const double b
Definition nonlinschro.cc:119
static const double a
Definition nonlinschro.cc:118
double Q(double a)
Definition relops.cc:20
static const double c
Definition relops.cc:10
static const long k
Definition rk.cc:44
AtomicInt sum
Definition test_atomicint.cc:46
constexpr coord_t one(1.0)