mirror of
https://github.com/cp2k/cp2k.git
synced 2026-07-27 21:55:16 -04:00
Use COSMA as fallback, if available
This commit is contained in:
parent
8b101df83d
commit
eeb0f8b27c
2 changed files with 57 additions and 50 deletions
|
|
@ -12,27 +12,27 @@
|
|||
!> \author Fawzi Mohamed
|
||||
! **************************************************************************************************
|
||||
MODULE cp_gemm_interface
|
||||
USE ISO_C_BINDING, ONLY: C_CHAR,&
|
||||
C_DOUBLE,&
|
||||
C_INT,&
|
||||
C_LOC,&
|
||||
C_PTR
|
||||
USE cp_dbcsr_operations, ONLY: copy_dbcsr_to_fm_bc,&
|
||||
copy_fm_to_dbcsr_bc
|
||||
USE cp_fm_basic_linalg, ONLY: cp_fm_gemm
|
||||
USE cp_fm_types, ONLY: cp_fm_get_info,&
|
||||
cp_fm_get_mm_type,&
|
||||
cp_fm_type
|
||||
USE dbcsr_api, ONLY: dbcsr_multiply,&
|
||||
dbcsr_release,&
|
||||
dbcsr_type
|
||||
USE input_constants, ONLY: do_cosma,&
|
||||
do_dbcsr,&
|
||||
do_scalapack
|
||||
USE kinds, ONLY: dp
|
||||
USE message_passing, ONLY: mp_min
|
||||
USE offload_api, ONLY: offload_set_device
|
||||
USE string_utilities, ONLY: uppercase
|
||||
USE ISO_C_BINDING, ONLY: C_CHAR, &
|
||||
C_DOUBLE, &
|
||||
C_INT, &
|
||||
C_LOC, &
|
||||
C_PTR
|
||||
USE cp_dbcsr_operations, ONLY: copy_dbcsr_to_fm_bc, &
|
||||
copy_fm_to_dbcsr_bc
|
||||
USE cp_fm_basic_linalg, ONLY: cp_fm_gemm
|
||||
USE cp_fm_types, ONLY: cp_fm_get_info, &
|
||||
cp_fm_get_mm_type, &
|
||||
cp_fm_type
|
||||
USE dbcsr_api, ONLY: dbcsr_multiply, &
|
||||
dbcsr_release, &
|
||||
dbcsr_type
|
||||
USE input_constants, ONLY: do_cosma, &
|
||||
do_dbcsr, &
|
||||
do_scalapack
|
||||
USE kinds, ONLY: dp
|
||||
USE message_passing, ONLY: mp_min
|
||||
USE offload_api, ONLY: offload_set_device
|
||||
USE string_utilities, ONLY: uppercase
|
||||
#include "./base/base_uses.f90"
|
||||
|
||||
IMPLICIT NONE
|
||||
|
|
@ -79,6 +79,12 @@ CONTAINS
|
|||
|
||||
CHARACTER(LEN=1) :: my_trans
|
||||
INTEGER :: handle, handle1, my_multi
|
||||
INTEGER, PARAMETER :: &
|
||||
#if defined(__COSMA)
|
||||
my_multi_fallback = do_cosma
|
||||
#else
|
||||
my_multi_fallback = do_scalapack
|
||||
#endif
|
||||
INTEGER, DIMENSION(:), POINTER :: a_col_loc, a_row_loc, b_col_loc, &
|
||||
b_row_loc, c_col_loc, c_row_loc
|
||||
TYPE(dbcsr_type) :: a_db, b_db, c_db
|
||||
|
|
@ -88,16 +94,16 @@ CONTAINS
|
|||
my_multi = cp_fm_get_mm_type()
|
||||
|
||||
! catch the special case that matrices have different blocking
|
||||
! SCALAPACK can deal with it but dbcsr doesn't like it
|
||||
! SCALAPACK/COSMA can deal with it but dbcsr doesn't like it
|
||||
CALL cp_fm_get_info(matrix_a, nrow_locals=a_row_loc, ncol_locals=a_col_loc)
|
||||
CALL cp_fm_get_info(matrix_b, nrow_locals=b_row_loc, ncol_locals=b_col_loc)
|
||||
CALL cp_fm_get_info(matrix_c, nrow_locals=c_row_loc, ncol_locals=c_col_loc)
|
||||
IF (PRESENT(a_first_row)) my_multi = do_scalapack
|
||||
IF (PRESENT(a_first_col)) my_multi = do_scalapack
|
||||
IF (PRESENT(b_first_row)) my_multi = do_scalapack
|
||||
IF (PRESENT(b_first_col)) my_multi = do_scalapack
|
||||
IF (PRESENT(c_first_row)) my_multi = do_scalapack
|
||||
IF (PRESENT(c_first_col)) my_multi = do_scalapack
|
||||
IF (PRESENT(a_first_row)) my_multi = my_multi_fallback
|
||||
IF (PRESENT(a_first_col)) my_multi = my_multi_fallback
|
||||
IF (PRESENT(b_first_row)) my_multi = my_multi_fallback
|
||||
IF (PRESENT(b_first_col)) my_multi = my_multi_fallback
|
||||
IF (PRESENT(c_first_row)) my_multi = my_multi_fallback
|
||||
IF (PRESENT(c_first_col)) my_multi = my_multi_fallback
|
||||
|
||||
my_trans = transa; CALL uppercase(my_trans)
|
||||
IF (my_trans == 'T') THEN
|
||||
|
|
@ -109,30 +115,31 @@ CONTAINS
|
|||
CALL cp_fm_get_info(matrix_b, nrow_locals=b_col_loc, ncol_locals=b_row_loc)
|
||||
END IF
|
||||
|
||||
IF (my_multi .NE. do_scalapack) THEN
|
||||
IF (my_multi .NE. do_scalapack .AND. my_multi .NE. do_cosma) THEN
|
||||
IF (SIZE(a_row_loc) == SIZE(c_row_loc)) THEN
|
||||
IF (ANY(a_row_loc - c_row_loc .NE. 0)) my_multi = do_scalapack
|
||||
IF (ANY(a_row_loc - c_row_loc .NE. 0)) my_multi = my_multi_fallback
|
||||
ELSE
|
||||
my_multi = do_scalapack
|
||||
my_multi = my_multi_fallback
|
||||
END IF
|
||||
END IF
|
||||
IF (my_multi .NE. do_scalapack) THEN
|
||||
IF (my_multi .NE. do_scalapack .AND. my_multi .NE. do_cosma) THEN
|
||||
IF (SIZE(b_col_loc) == SIZE(c_col_loc)) THEN
|
||||
IF (ANY(b_col_loc - c_col_loc .NE. 0)) my_multi = do_scalapack
|
||||
IF (ANY(b_col_loc - c_col_loc .NE. 0)) my_multi = my_multi_fallback
|
||||
ELSE
|
||||
my_multi = do_scalapack
|
||||
my_multi = my_multi_fallback
|
||||
END IF
|
||||
END IF
|
||||
IF (my_multi .NE. do_scalapack) THEN
|
||||
IF (my_multi .NE. do_scalapack .AND. my_multi .NE. do_cosma) THEN
|
||||
IF (SIZE(a_col_loc) == SIZE(b_row_loc)) THEN
|
||||
IF (ANY(a_col_loc - b_row_loc .NE. 0)) my_multi = do_scalapack
|
||||
IF (ANY(a_col_loc - b_row_loc .NE. 0)) my_multi = my_multi_fallback
|
||||
ELSE
|
||||
my_multi = do_scalapack
|
||||
my_multi = my_multi_fallback
|
||||
END IF
|
||||
END IF
|
||||
|
||||
! IMPORTANT do_scalapack is lowest value. If one processor has it set make all use it.
|
||||
IF (cp_fm_get_mm_type() .NE. do_scalapack) CALL mp_min(my_multi, matrix_a%matrix_struct%para_env%group)
|
||||
IF (cp_fm_get_mm_type() .NE. do_scalapack .AND. &
|
||||
cp_fm_get_mm_type() .NE. do_cosma) CALL mp_min(my_multi, matrix_a%matrix_struct%para_env%group)
|
||||
|
||||
SELECT CASE (my_multi)
|
||||
CASE (do_scalapack)
|
||||
|
|
@ -145,6 +152,16 @@ CONTAINS
|
|||
c_first_col=c_first_col, &
|
||||
c_first_row=c_first_row)
|
||||
CALL timestop(handle1)
|
||||
CASE (do_cosma)
|
||||
#if defined(__COSMA)
|
||||
CALL timeset("cp_gemm_cosma", handle1)
|
||||
CALL offload_set_device()
|
||||
CALL cosma_pdgemm(transa=transa, transb=transb, m=m, n=n, k=k, alpha=alpha, &
|
||||
matrix_a=matrix_a, matrix_b=matrix_b, beta=beta, matrix_c=matrix_c)
|
||||
CALL timestop(handle1)
|
||||
#else
|
||||
CPABORT("CP2K compiled without the COSMA library.")
|
||||
#endif
|
||||
CASE (do_dbcsr)
|
||||
CALL timeset("cp_gemm_dbcsr_mm", handle1)
|
||||
CALL copy_fm_to_dbcsr_bc(matrix_a, a_db)
|
||||
|
|
@ -158,16 +175,6 @@ CONTAINS
|
|||
CALL dbcsr_release(b_db)
|
||||
CALL dbcsr_release(c_db)
|
||||
CALL timestop(handle1)
|
||||
CASE (do_cosma)
|
||||
#if defined(__COSMA)
|
||||
CALL timeset("cp_gemm_cosma", handle1)
|
||||
CALL offload_set_device()
|
||||
CALL cosma_pdgemm(transa=transa, transb=transb, m=m, n=n, k=k, alpha=alpha, &
|
||||
matrix_a=matrix_a, matrix_b=matrix_b, beta=beta, matrix_c=matrix_c)
|
||||
CALL timestop(handle1)
|
||||
#else
|
||||
CPABORT("CP2K compiled without the COSMA library.")
|
||||
#endif
|
||||
END SELECT
|
||||
CALL timestop(handle)
|
||||
|
||||
|
|
|
|||
|
|
@ -1086,8 +1086,8 @@ MODULE input_constants
|
|||
|
||||
! fm matrix multiplication
|
||||
INTEGER, PARAMETER, PUBLIC :: do_scalapack = 1, &
|
||||
do_dbcsr = 2, &
|
||||
do_cosma = 3
|
||||
do_cosma = 2, &
|
||||
do_dbcsr = 3
|
||||
|
||||
! Dispersion DFTB
|
||||
INTEGER, PARAMETER, PUBLIC :: dispersion_uff = 100, &
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue