From eeb0f8b27cc66869d65cbd097caa5ca22ccbdcfb Mon Sep 17 00:00:00 2001 From: Alfio Lazzaro Date: Tue, 20 Jul 2021 09:08:51 -0500 Subject: [PATCH] Use COSMA as fallback, if available --- src/cp_gemm_interface.F | 103 +++++++++++++++++++++------------------- src/input_constants.F | 4 +- 2 files changed, 57 insertions(+), 50 deletions(-) diff --git a/src/cp_gemm_interface.F b/src/cp_gemm_interface.F index 6f46d5205e..f9c9b6fe0e 100644 --- a/src/cp_gemm_interface.F +++ b/src/cp_gemm_interface.F @@ -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) diff --git a/src/input_constants.F b/src/input_constants.F index bfc5c2f0a2..5e4ce2e707 100644 --- a/src/input_constants.F +++ b/src/input_constants.F @@ -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, &