Skip to content

Commit 8faedb0

Browse files
committed
add ispin to transferT0MatrixToGPUHIP
1 parent cb447fb commit 8faedb0

3 files changed

Lines changed: 7 additions & 5 deletions

File tree

src/MultipleScattering/calculateTauMatrix.cpp

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -464,7 +464,7 @@ void calculateTauMatrix(LSMSSystemParameters &lsms, LocalTypeInfo &local,
464464
devM = deviceStorage->getDevM();
465465
transferMatrixToGPUHip(devM, m);
466466
devT0 = deviceStorage->getDevT0();
467-
transferT0MatrixToGPUHip(devT0, lsms, local, atom, iie);
467+
transferT0MatrixToGPUHip(devT0, lsms, local, atom, iie, ispin);
468468
break;
469469
#endif
470470
default:
@@ -520,7 +520,7 @@ void calculateTauMatrix(LSMSSystemParameters &lsms, LocalTypeInfo &local,
520520
break;
521521
case MST_LINEAR_SOLVER_ZGETRF_ROCSOLVER:
522522
devT0 = deviceStorage->getDevT0();
523-
transferT0MatrixToGPUHip(devT0, lsms, local, atom, iie);
523+
transferT0MatrixToGPUHip(devT0, lsms, local, atom, iie, ispin);
524524
break;
525525
#ifdef ACCELERATOR_CUDA_C
526526
case MST_LINEAR_SOLVER_ZGETRF_CUBLAS:

src/MultipleScattering/linearSolvers.hpp

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -145,7 +145,7 @@ void unitMatrixCuda(T *devM, int lDim, int nCol) {
145145
#ifdef ACCELERATOR_HIP
146146
void transferMatrixToGPUHip(Complex *devM, Matrix<Complex> &m);
147147
void transferMatrixFromGPUHip(Matrix<Complex> &m, hipDoubleComplex *devM);
148-
void transferT0MatrixToGPUHip(Complex *devT0, LSMSSystemParameters &lsms, LocalTypeInfo &local, AtomData &atom, int iie);
148+
void transferT0MatrixToGPUHip(Complex *devT0, LSMSSystemParameters &lsms, LocalTypeInfo &local, AtomData &atom, int iie, int ispin);
149149
void transferFullTMatrixToGPUHip(Complex *devT, LSMSSystemParameters &lsms, LocalTypeInfo &local,
150150
AtomData &atom, int ispin);
151151

src/MultipleScattering/linearSolvers_HIP.cpp

Lines changed: 4 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -71,10 +71,12 @@ __global__ void zeroDiagonalBlocksKernelHip(T *devM, int lDim, int nCol,
7171
}
7272

7373
void transferT0MatrixToGPUHip(Complex *devT0, LSMSSystemParameters &lsms,
74-
LocalTypeInfo &local, AtomData &atom, int iie) {
74+
LocalTypeInfo &local, AtomData &atom, int iie, int ispin) {
7575
int kkrsz_ns = lsms.n_spin_cant * atom.kkrsz;
76+
int jsm = kkrsz_ns * kkrsz_ns * ispin;
77+
7678
hipError_t ret = hipMemcpy(devT0,
77-
&local.tmatStore(iie * local.blkSizeTmatStore, atom.LIZStoreIdx[0]),
79+
&local.tmatStore(iie * local.blkSizeTmatStore + jsm, atom.LIZStoreIdx[0]),
7880
kkrsz_ns * kkrsz_ns * sizeof(hipDoubleComplex),
7981
hipMemcpyHostToDevice);
8082
}

0 commit comments

Comments
 (0)