mirror of
https://gitlab.com/scemama/qp_plugins_scemama.git
synced 2025-01-03 01:55:52 +01:00
GPU acceleration
This commit is contained in:
parent
7847ccc674
commit
0c6e6a1ca0
@ -105,28 +105,27 @@ void compute_r2_space_chol_gpu(const int nO, const int nV, const int cholesky_mo
|
|||||||
C=d_tmpB1 ; ldc=nV*BLOCK_SIZE;
|
C=d_tmpB1 ; ldc=nV*BLOCK_SIZE;
|
||||||
cublasDgemm(handle, CUBLAS_OP_T, CUBLAS_OP_N, m, n, k, &alpha, A, lda, B, lda, &beta, C, ldc);
|
cublasDgemm(handle, CUBLAS_OP_T, CUBLAS_OP_N, m, n, k, &alpha, A, lda, B, lda, &beta, C, ldc);
|
||||||
|
|
||||||
for (size_t bet=iblock ; bet<(nV < iblock+BLOCK_SIZE ? nV : iblock+BLOCK_SIZE) ; ++bet) {
|
for (size_t bet=iblock ; bet<(nV < iblock+BLOCK_SIZE ? nV : iblock+BLOCK_SIZE) ; ++bet)
|
||||||
|
{
|
||||||
|
|
||||||
alpha = 1.0;
|
alpha = 1.0;
|
||||||
beta = 0.0;
|
beta = 0.0;
|
||||||
A = &(d_tmpB1[nV*(bet-iblock)]); lda = nV*BLOCK_SIZE;
|
A = &(d_tmpB1[nV*(bet-iblock)]); lda = nV*BLOCK_SIZE;
|
||||||
B = d_tmpB1; ldb = nV;
|
B = d_tmpB1; ldb = nV;
|
||||||
C = &(d_B1[nV*nV*(bet-iblock)]) ; ldc = nV;
|
C = &(d_B1[nV*nV*(bet-iblock)]) ; ldc = nV;
|
||||||
cublasDgeam(handle, CUBLAS_OP_N, CUBLAS_OP_N, nV, nV, &alpha, A, lda, &beta, B, ldb, C, ldc);
|
cublasDgeam(handle, CUBLAS_OP_N, CUBLAS_OP_N, nV, nV, &alpha, A, lda, &beta, B, ldb, C, ldc);
|
||||||
|
}
|
||||||
|
|
||||||
alpha=1.0; beta=1.0;
|
alpha=1.0; beta=1.0;
|
||||||
m=nO*nO; n=mbs; k=nV*nV;
|
m=nO*nO; n=mbs; k=nV*nV;
|
||||||
|
|
||||||
A=d_tau; lda=nO*nO;
|
A=d_tau; lda=nO*nO;
|
||||||
B=d_B1 ; ldb=nV*nV;
|
B=d_B1 ; ldb=nV*nV;
|
||||||
C=&(d_r2[nO*nO*(iblock + nV*gam)]); ldc=nO*nO;
|
C=&(d_r2[nO*nO*(iblock + nV*gam)]); ldc=nO*nO;
|
||||||
cublasDgemm(handle, CUBLAS_OP_N, CUBLAS_OP_N, m, n, k, &alpha, A, lda, B, ldb, &beta, C, ldc);
|
cublasDgemm(handle, CUBLAS_OP_N, CUBLAS_OP_N, m, n, k, &alpha, A, lda, B, ldb, &beta, C, ldc);
|
||||||
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
free(B1);
|
|
||||||
free(tmpB1);
|
|
||||||
}
|
}
|
||||||
lda=nO*nO;
|
lda=nO*nO;
|
||||||
cublasGetMatrix(nO*nO, nV*nV, sizeof(double), d_r2, lda, r2, lda);
|
cublasGetMatrix(nO*nO, nV*nV, sizeof(double), d_r2, lda, r2, lda);
|
||||||
|
Loading…
Reference in New Issue
Block a user