(git:98357aa)
Loading...
Searching...
No Matches
local_gemm_api.F
Go to the documentation of this file.
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
9 USE iso_c_binding, ONLY: c_null_ptr, &
10 c_ptr
11 USE kinds, ONLY: dp
12#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
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
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, &
44
45 INTEGER, PARAMETER, PUBLIC :: &
48
49 INTEGER, PRIVATE :: do_dgemm = 1
50
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
60CONTAINS
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 SUBROUTINE local_gemm(opA, opB, m, n, k, &
80 alpha, A, lda, B, ldb, &
81 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 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
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 b, ldb, beta, c, ldc)
147#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
148 END IF
149#else
150 mark_used(ctx)
151#endif
152 CALL timestop(handle)
153
154 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 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
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 ctx%spla_context = c_null_ptr
182#endif
183 END SUBROUTINE local_gemm_create
184
185! **************************************************************************************************
186!> \brief release resources associated to a gemm context
187!> \param ctx handle
188! **************************************************************************************************
189 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
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 ctx%spla_context = c_null_ptr
204 END SUBROUTINE local_gemm_destroy
205
206! **************************************************************************************************
207!> \brief ...
208!> \param ctx ...
209!> \param opThresholdGPU ...
210! **************************************************************************************************
211 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
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 END SUBROUTINE local_gemm_set_op_threshold_gpu
225
226! **************************************************************************************************
227!> \brief ...
228!> \param dgemm_library ...
229! **************************************************************************************************
230 SUBROUTINE local_gemm_set_library(dgemm_library)
231 INTEGER, INTENT(IN) :: dgemm_library
232
233 do_dgemm = dgemm_library
234 END SUBROUTINE local_gemm_set_library
235
236END MODULE local_gemm_api
static void dgemm(const char transa, const char transb, const int m, const int n, const int k, const double alpha, const double *a, const int lda, const double *b, const int ldb, const double beta, double *c, const int ldc)
Convenient wrapper to hide Fortran nature of dgemm_, swapping a and b.
various routines to log and control the output. The idea is that decisions about where to log should ...
collects all constants needed in input so that they can be used without circular dependencies
integer, parameter, public do_dgemm_spla
Defines the basic variable types.
Definition kinds.F:23
integer, parameter, public dp
Definition kinds.F:34
subroutine, public local_gemm_set_library(dgemm_library)
...
integer, parameter, public local_gemm_pu_gpu
subroutine local_gemm(opa, opb, m, n, k, alpha, a, lda, b, ldb, beta, c, ldc, ctx)
...
integer, parameter, public local_gemm_pu_host
Fortran API for the offload package, which is written in C.
Definition offload_api.F:12
subroutine, public offload_activate_chosen_device()
Activates the device selected via offload_set_chosen_device()