LCOV - code coverage report
Current view: top level - src - local_gemm_api.F (source / functions) Coverage Total Hit
Test: CP2K Regtests (git:71c3ab0) Lines: 94.7 % 19 18
Test Date: 2026-07-25 06:35:44 Functions: 83.3 % 6 5

            Line data    Source code
       1              : !--------------------------------------------------------------------------------------------------!
       2              : !   CP2K: A general program to perform molecular dynamics simulations                              !
       3              : !   Copyright 2000-2026 CP2K developers group <https://cp2k.org>                                   !
       4              : !                                                                                                  !
       5              : !   SPDX-License-Identifier: GPL-2.0-or-later                                                      !
       6              : !--------------------------------------------------------------------------------------------------!
       7              : 
       8              : MODULE local_gemm_api
       9              :    USE ISO_C_BINDING, ONLY: C_NULL_PTR, &
      10              :                             C_PTR
      11              :    USE kinds, ONLY: dp
      12              : #if defined(__SPLA) && defined(__OFFLOAD_GEMM)
      13              :    USE input_constants, ONLY: do_dgemm_spla
      14              :    USE ISO_C_BINDING, ONLY: C_ASSOCIATED, &
      15              :                             C_LOC
      16              :    USE spla, ONLY: SPLA_PU_HOST, &
      17              :                    SPLA_PU_GPU, &
      18              :                    SPLA_OP_NONE, &
      19              :                    SPLA_OP_TRANSPOSE, &
      20              :                    SPLA_OP_CONJ_TRANSPOSE, &
      21              :                    spla_ctx_create, &
      22              :                    spla_ctx_destroy, &
      23              :                    spla_dgemm, &
      24              :                    spla_sgemm, &
      25              :                    spla_cgemm, &
      26              :                    spla_zgemm, &
      27              :                    spla_ctx_set_op_threshold_gpu, &
      28              :                    SPLA_SUCCESS
      29              : #endif
      30              : 
      31              :    USE cp_log_handling, ONLY: cp_to_string
      32              :    USE offload_api, ONLY: offload_activate_chosen_device
      33              : 
      34              : #include "./base/base_uses.f90"
      35              : 
      36              :    IMPLICIT NONE
      37              : 
      38              :    PRIVATE
      39              : 
      40              :    CHARACTER(len=*), PARAMETER, PRIVATE :: moduleN = 'local_gemm_api'
      41              : 
      42              :    PUBLIC :: local_gemm_ctxt_type, &
      43              :              local_gemm_set_library
      44              : 
      45              :    INTEGER, PARAMETER, PUBLIC :: &
      46              :       LOCAL_GEMM_PU_HOST = 0, &
      47              :       LOCAL_GEMM_PU_GPU = 1
      48              : 
      49              :    INTEGER, PRIVATE :: do_dgemm = 1
      50              : 
      51              :    TYPE local_gemm_ctxt_type
      52              :       TYPE(C_PTR) :: spla_context = C_NULL_PTR
      53              :    CONTAINS
      54              :       PROCEDURE, PASS(ctx), NON_OVERRIDABLE :: create => local_gemm_create
      55              :       PROCEDURE, PASS(ctx), NON_OVERRIDABLE :: destroy => local_gemm_destroy
      56              :       PROCEDURE, PASS(ctx), NON_OVERRIDABLE :: set_op_threshold_gpu => local_gemm_set_op_threshold_gpu
      57              :       PROCEDURE, PASS(ctx), NON_OVERRIDABLE :: gemm => local_gemm
      58              :    END TYPE
      59              : 
      60              : CONTAINS
      61              : 
      62              : ! **************************************************************************************************
      63              : !> \brief ...
      64              : !> \param opA ...
      65              : !> \param opB ...
      66              : !> \param m ...
      67              : !> \param n ...
      68              : !> \param k ...
      69              : !> \param alpha ...
      70              : !> \param A ...
      71              : !> \param lda ...
      72              : !> \param B ...
      73              : !> \param ldb ...
      74              : !> \param beta ...
      75              : !> \param C ...
      76              : !> \param ldc ...
      77              : !> \param ctx ...
      78              : ! **************************************************************************************************
      79       106744 :    SUBROUTINE local_gemm(opA, opB, m, n, k, &
      80        53372 :                          alpha, A, lda, B, ldb, &
      81        53372 :                          beta, C, ldc, ctx)
      82              :       CHARACTER, INTENT(in) :: opA
      83              :       CHARACTER, INTENT(in) :: opB
      84              :       INTEGER, INTENT(in) :: m
      85              :       INTEGER, INTENT(in) :: n
      86              :       INTEGER, INTENT(in) :: k
      87              :       REAL(KIND=dp), INTENT(in) :: alpha
      88              : #if defined(__SPLA) && defined(__OFFLOAD_GEMM)
      89              :       REAL(KIND=dp), DIMENSION(*), INTENT(in), TARGET :: A
      90              : #else
      91              :       REAL(KIND=dp), DIMENSION(:, :), INTENT(in), TARGET :: A
      92              : #endif
      93              :       INTEGER, INTENT(in) :: lda
      94              : #if defined(__SPLA) && defined(__OFFLOAD_GEMM)
      95              :       REAL(KIND=dp), DIMENSION(*), INTENT(in), TARGET :: B
      96              : #else
      97              :       REAL(KIND=dp), DIMENSION(:, :), INTENT(in), TARGET :: B
      98              : #endif
      99              : 
     100              :       INTEGER, INTENT(in) :: ldb
     101              :       REAL(KIND=dp), INTENT(in) :: beta
     102              : #if defined(__SPLA) && defined(__OFFLOAD_GEMM)
     103              :       REAL(KIND=dp), DIMENSION(*), INTENT(inout), TARGET ::C
     104              : #else
     105              :       REAL(KIND=dp), DIMENSION(:, :), INTENT(inout), TARGET :: C
     106              : #endif
     107              :       INTEGER, INTENT(in) :: ldc
     108              :       CLASS(local_gemm_ctxt_type), INTENT(inout) :: ctx
     109              : 
     110              :       INTEGER                                            :: handle
     111              : !     no point of using SPLA offloading on CPU ONLY nodes
     112              : #if defined(__SPLA) && defined(__OFFLOAD_GEMM)
     113              :       INTEGER :: spla_op_A, spla_op_B, spla_error
     114              : #endif
     115              :       CHARACTER(LEN=*), PARAMETER :: routineN = 'local_gemm'
     116        53372 :       CALL timeset(routineN, handle)
     117              : 
     118              : !     no point of using SPLA offloading on CPU ONLY nodes
     119              : #if defined(__SPLA) && defined(__OFFLOAD_GEMM)
     120              :       IF (do_dgemm == do_dgemm_spla) THEN
     121              : 
     122              :          IF (opA == 'N') spla_op_A = SPLA_OP_NONE
     123              :          IF (opA == 'T') spla_op_A = SPLA_OP_TRANSPOSE
     124              : 
     125              :          IF (opB == 'N') spla_op_B = SPLA_OP_NONE
     126              :          IF (opB == 'T') spla_op_B = SPLA_OP_TRANSPOSE
     127              : 
     128              : #if __GNUC__ >= 9
     129              :          CPASSERT(IS_CONTIGUOUS(A))
     130              :          CPASSERT(IS_CONTIGUOUS(B))
     131              :          CPASSERT(IS_CONTIGUOUS(C))
     132              : #endif
     133              : 
     134              :          CALL offload_activate_chosen_device()
     135              :          spla_error = spla_dgemm(spla_op_A, spla_op_B, &
     136              :                                  m, n, k, alpha, &
     137              :                                  c_loc(A), lda, &
     138              :                                  c_loc(B), ldb, &
     139              :                                  beta, c_loc(C), ldc, ctx%spla_context)
     140              :          IF (spla_error /= SPLA_SUCCESS) &
     141              :             CPABORT("spla_dgemm failed: "//cp_to_string(spla_error))
     142              :       ELSE
     143              : #endif
     144              :          CALL dgemm(opA, opB, m, n, k, alpha, &
     145              :                     A, lda, &
     146      1523922 :                     B, ldb, beta, C, ldc)
     147              : #if defined(__SPLA) && defined(__OFFLOAD_GEMM)
     148              :       END IF
     149              : #else
     150              :       MARK_USED(ctx)
     151              : #endif
     152        53372 :       CALL timestop(handle)
     153              : 
     154        53372 :    END SUBROUTINE local_gemm
     155              : 
     156              : ! **************************************************************************************************
     157              : !> \brief create a context for handling gemm offloading
     158              : !> \param ctx newly created context
     159              : !> \param pu processing unit to run the (s,d,c,z}dgemm
     160              : ! **************************************************************************************************
     161          412 :    SUBROUTINE local_gemm_create(ctx, pu)
     162              :       CLASS(local_gemm_ctxt_type), INTENT(out) :: ctx
     163              :       INTEGER, INTENT(in) :: pu
     164              : 
     165              : #if defined(__SPLA) && defined(__OFFLOAD_GEMM)
     166              :       INTEGER :: error_
     167              : 
     168              :       IF (.NOT. C_ASSOCIATED(ctx%spla_context)) THEN
     169              :          IF (do_dgemm == do_dgemm_spla) THEN
     170              :             CALL offload_activate_chosen_device()
     171              : 
     172              :             error_ = spla_ctx_create(ctx%spla_context, pu)
     173              :             IF (error_ /= SPLA_SUCCESS) &
     174              :                CPABORT("spla_ctx_create failed: "//cp_to_string(error_))
     175              :          ELSE
     176              :             ctx%spla_context = C_NULL_PTR
     177              :          END IF
     178              :       END IF
     179              : #else
     180              :       MARK_USED(pu)
     181          412 :       ctx%spla_context = C_NULL_PTR
     182              : #endif
     183          412 :    END SUBROUTINE local_gemm_create
     184              : 
     185              : ! **************************************************************************************************
     186              : !> \brief release resources associated to a gemm context
     187              : !> \param ctx handle
     188              : ! **************************************************************************************************
     189          888 :    SUBROUTINE local_gemm_destroy(ctx)
     190              :       CLASS(local_gemm_ctxt_type), INTENT(inout) :: ctx
     191              : 
     192              : #if defined(__SPLA) && defined(__OFFLOAD_GEMM)
     193              :       INTEGER :: error_
     194              : 
     195              :       IF (do_dgemm == do_dgemm_spla) THEN
     196              :          CALL offload_activate_chosen_device()
     197              : 
     198              :          error_ = spla_ctx_destroy(ctx%spla_context)
     199              :          IF (error_ /= SPLA_SUCCESS) &
     200              :             CPABORT("spla_ctx_destroy failed: "//cp_to_string(error_))
     201              :       END IF
     202              : #endif
     203          888 :       ctx%spla_context = C_NULL_PTR
     204          888 :    END SUBROUTINE local_gemm_destroy
     205              : 
     206              : ! **************************************************************************************************
     207              : !> \brief ...
     208              : !> \param ctx ...
     209              : !> \param opThresholdGPU ...
     210              : ! **************************************************************************************************
     211          412 :    SUBROUTINE local_gemm_set_op_threshold_gpu(ctx, opThresholdGPU)
     212              :       CLASS(local_gemm_ctxt_type), INTENT(INOUT)                                        :: ctx
     213              :       INTEGER, INTENT(in)                                :: opThresholdGPU
     214              : 
     215              : #if defined(__SPLA) && defined(__OFFLOAD_GEMM)
     216              :       INTEGER                                            :: error__
     217              : 
     218              :       CALL offload_activate_chosen_device()
     219              :       error__ = spla_ctx_set_op_threshold_gpu(ctx%spla_context, opThresholdGPU)
     220              : #else
     221              :       MARK_USED(ctx)
     222              :       MARK_USED(opThresholdGPU)
     223              : #endif
     224          412 :    END SUBROUTINE local_gemm_set_op_threshold_gpu
     225              : 
     226              : ! **************************************************************************************************
     227              : !> \brief ...
     228              : !> \param dgemm_library ...
     229              : ! **************************************************************************************************
     230        11087 :    SUBROUTINE local_gemm_set_library(dgemm_library)
     231              :       INTEGER, INTENT(IN)                                :: dgemm_library
     232              : 
     233        11087 :       do_dgemm = dgemm_library
     234        11087 :    END SUBROUTINE local_gemm_set_library
     235              : 
     236            0 : END MODULE local_gemm_api
        

Generated by: LCOV version 2.0-1