Use COSMA as fallback, if available

This commit is contained in:
Alfio Lazzaro 2021-07-20 09:08:51 -05:00
parent 8b101df83d
commit eeb0f8b27c
2 changed files with 57 additions and 50 deletions

View file

@ -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)

View file

@ -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, &