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
|