(git:f2099e5)
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_op_none, &
17 spla_op_transpose, &
18 spla_op_conj_transpose, &
19 spla_ctx_create, &
20 spla_ctx_destroy, &
21 spla_dgemm, &
22 spla_zgemm, &
23 spla_ctx_set_op_threshold_gpu, &
24 spla_success
25#endif
26
29
30#include "./base/base_uses.f90"
31
32 IMPLICIT NONE
33
34 PRIVATE
35
36 CHARACTER(len=*), PARAMETER, PRIVATE :: moduleN = 'local_gemm_api'
37
38 PUBLIC :: local_gemm_ctxt_type, &
40
41 INTEGER, PARAMETER, PUBLIC :: &
44
45 INTEGER, PRIVATE :: do_dgemm = 1
46
48 TYPE(c_ptr) :: spla_context = c_null_ptr
49 LOGICAL, PRIVATE :: timing = .true.
50 CONTAINS
51 PROCEDURE, pass(ctx), non_overridable :: create => local_gemm_create
52 PROCEDURE, pass(ctx), non_overridable :: destroy => local_gemm_destroy
53 PROCEDURE, pass(ctx), non_overridable :: set_op_threshold_gpu => local_gemm_set_op_threshold_gpu
54 PROCEDURE, pass(ctx), non_overridable, PRIVATE :: gemm_d => local_dgemm
55 PROCEDURE, pass(ctx), non_overridable, PRIVATE :: gemm_z => local_zgemm
56 generic :: gemm => gemm_d, gemm_z
57 END TYPE
58
59CONTAINS
60
61! **************************************************************************************************
62!> \brief Local GEMM on contiguous arrays, using BLAS or the configured SPLA backend.
63!> Each concurrent caller must own its context. No distributed matrix metadata is used.
64!> \param opA operation on A (N/T/C, case insensitive)
65!> \param opB operation on B (N/T/C, case insensitive)
66!> \param m output rows
67!> \param n output columns
68!> \param k contraction dimension
69!> \param alpha product scale
70!> \param A left operand
71!> \param lda leading dimension of A
72!> \param B right operand
73!> \param ldb leading dimension of B
74!> \param beta output scale
75!> \param C output, must not overlap A or B
76!> \param ldc leading dimension of C
77!> \param ctx caller-owned context
78! **************************************************************************************************
79 SUBROUTINE local_dgemm(opA, opB, m, n, k, alpha, A, lda, B, ldb, beta, C, ldc, ctx)
80 CHARACTER, INTENT(IN) :: opA, opB
81 INTEGER, INTENT(IN) :: m, n, k, lda, ldb, ldc
82 REAL (KIND=dp), INTENT(IN) :: alpha, beta
83 REAL (KIND=dp), INTENT(IN), TARGET :: a(lda, *), b(ldb, *)
84 REAL (KIND=dp), INTENT(INOUT), TARGET :: c(ldc, *)
85 CLASS(local_gemm_ctxt_type), INTENT(INOUT) :: ctx
86
87 INTEGER :: handle
88#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
89 INTEGER :: spla_error
90#endif
91 CHARACTER(LEN=*), PARAMETER :: routinen = 'local_gemm'
92
93 IF (m == 0 .OR. n == 0) RETURN
94 IF (ctx%timing) CALL timeset(routinen, handle)
95#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
96 IF (do_dgemm == do_dgemm_spla) THEN
97 cpassert(c_associated(ctx%spla_context))
99 spla_error = spla_dgemm(spla_operation(opa), spla_operation(opb), m, n, k, alpha, &
100 c_loc(a(1, 1)), lda, c_loc(b(1, 1)), ldb, beta, &
101 c_loc(c(1, 1)), ldc, ctx%spla_context)
102 IF (spla_error /= spla_success) &
103 CALL cp_abort(__location__, "spla_dgemm failed: "//cp_to_string(spla_error))
104 ELSE
105#endif
106 CALL dgemm(opa, opb, m, n, k, alpha, a, lda, b, ldb, beta, c, ldc)
107#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
108 END IF
109#endif
110 IF (ctx%timing) CALL timestop(handle)
111
112 END SUBROUTINE local_dgemm
113! **************************************************************************************************
114!> \brief Local GEMM on contiguous arrays, using BLAS or the configured SPLA backend.
115!> Each concurrent caller must own its context. No distributed matrix metadata is used.
116!> \param opA operation on A (N/T/C, case insensitive)
117!> \param opB operation on B (N/T/C, case insensitive)
118!> \param m output rows
119!> \param n output columns
120!> \param k contraction dimension
121!> \param alpha product scale
122!> \param A left operand
123!> \param lda leading dimension of A
124!> \param B right operand
125!> \param ldb leading dimension of B
126!> \param beta output scale
127!> \param C output, must not overlap A or B
128!> \param ldc leading dimension of C
129!> \param ctx caller-owned context
130! **************************************************************************************************
131 SUBROUTINE local_zgemm(opA, opB, m, n, k, alpha, A, lda, B, ldb, beta, C, ldc, ctx)
132 CHARACTER, INTENT(IN) :: opa, opb
133 INTEGER, INTENT(IN) :: m, n, k, lda, ldb, ldc
134 COMPLEX (KIND=dp), INTENT(IN) :: alpha, beta
135 COMPLEX (KIND=dp), INTENT(IN), TARGET :: a(lda, *), B(ldb, *)
136 COMPLEX (KIND=dp), INTENT(INOUT), TARGET :: C(ldc, *)
137 CLASS(local_gemm_ctxt_type), INTENT(INOUT) :: ctx
138
139 INTEGER :: handle
140#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
141 INTEGER :: spla_error
142#endif
143 CHARACTER(LEN=*), PARAMETER :: routineN = 'local_gemm'
144
145 IF (m == 0 .OR. n == 0) RETURN
146 IF (ctx%timing) CALL timeset(routinen, handle)
147#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
148 IF (do_dgemm == do_dgemm_spla) THEN
149 cpassert(c_associated(ctx%spla_context))
151 spla_error = spla_zgemm(spla_operation(opa), spla_operation(opb), m, n, k, alpha, &
152 c_loc(a(1, 1)), lda, c_loc(b(1, 1)), ldb, beta, &
153 c_loc(c(1, 1)), ldc, ctx%spla_context)
154 IF (spla_error /= spla_success) &
155 CALL cp_abort(__location__, "spla_zgemm failed: "//cp_to_string(spla_error))
156 ELSE
157#endif
158 CALL zgemm(opa, opb, m, n, k, alpha, a, lda, b, ldb, beta, c, ldc)
159#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
160 END IF
161#endif
162 IF (ctx%timing) CALL timestop(handle)
163
164 END SUBROUTINE local_zgemm
165
166#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
167! **************************************************************************************************
168!> \brief Translate a BLAS transpose flag for SPLA, including complex conjugation.
169!> \param trans BLAS operation
170!> \return SPLA operation
171! **************************************************************************************************
172 FUNCTION spla_operation(trans) RESULT(op)
173 CHARACTER, INTENT(IN) :: trans
174 INTEGER :: op
175
176 SELECT CASE (trans)
177 CASE ('N', 'n')
178 op = spla_op_none
179 CASE ('T', 't')
180 op = spla_op_transpose
181 CASE ('C', 'c')
182 op = spla_op_conj_transpose
183 CASE DEFAULT
184 CALL cp_abort(__location__, "Invalid local GEMM transpose flag.")
185 END SELECT
186 END FUNCTION spla_operation
187#endif
188
189! **************************************************************************************************
190!> \brief Create a local GEMM context; destroy an existing context before recreating it.
191!> \param ctx newly created context, with timing enabled by default
192!> \param pu processing unit for local GEMM
193!> \param timing collect per-GEMM timings (default true); disable for timed batches
194! **************************************************************************************************
195 SUBROUTINE local_gemm_create(ctx, pu, timing)
196 CLASS(local_gemm_ctxt_type), INTENT(OUT) :: ctx
197 INTEGER, INTENT(IN) :: pu
198 LOGICAL, INTENT(IN), OPTIONAL :: timing
199
200#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
201 INTEGER :: spla_error
202#endif
203
204 IF (PRESENT(timing)) ctx%timing = timing
205#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
206 IF (do_dgemm == do_dgemm_spla) THEN
208
209 spla_error = spla_ctx_create(ctx%spla_context, pu)
210 IF (spla_error /= spla_success) &
211 CALL cp_abort(__location__, &
212 "spla_ctx_create failed: "//cp_to_string(spla_error))
213 END IF
214#else
215 mark_used(pu)
216#endif
217 END SUBROUTINE local_gemm_create
218
219! **************************************************************************************************
220!> \brief Release an owned SPLA context, independently of the current backend preference.
221!> \param ctx context to release; an empty context is allowed
222! **************************************************************************************************
223 SUBROUTINE local_gemm_destroy(ctx)
224 CLASS(local_gemm_ctxt_type), INTENT(INOUT) :: ctx
225
226#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
227 INTEGER :: spla_error
228
229 IF (c_associated(ctx%spla_context)) THEN
231
232 spla_error = spla_ctx_destroy(ctx%spla_context)
233 IF (spla_error /= spla_success) &
234 CALL cp_abort(__location__, &
235 "spla_ctx_destroy failed: "//cp_to_string(spla_error))
236 END IF
237#endif
238 ctx%spla_context = c_null_ptr
239 END SUBROUTINE local_gemm_destroy
240
241! **************************************************************************************************
242!> \brief Set the SPLA GPU operation threshold; no-op when no SPLA context is allocated.
243!> \param ctx local GEMM context
244!> \param opThresholdGPU operation-count threshold for GPU offloading
245! **************************************************************************************************
246 SUBROUTINE local_gemm_set_op_threshold_gpu(ctx, opThresholdGPU)
247 CLASS(local_gemm_ctxt_type), INTENT(INOUT) :: ctx
248 INTEGER, INTENT(IN) :: opThresholdGPU
249
250#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
251 INTEGER :: spla_error
252
253 IF (c_associated(ctx%spla_context)) THEN
255
256 spla_error = spla_ctx_set_op_threshold_gpu(ctx%spla_context, opthresholdgpu)
257 IF (spla_error /= spla_success) &
258 CALL cp_abort(__location__, &
259 "spla_ctx_set_op_threshold_gpu failed: "//cp_to_string(spla_error))
260 END IF
261#else
262 mark_used(ctx)
263 mark_used(opthresholdgpu)
264#endif
265 END SUBROUTINE local_gemm_set_op_threshold_gpu
266
267! **************************************************************************************************
268!> \brief Select the backend for subsequent local GEMM calls and context creation.
269!> \param dgemm_library backend selector from input_constants (SPLA or BLAS)
270! **************************************************************************************************
271 SUBROUTINE local_gemm_set_library(dgemm_library)
272 INTEGER, INTENT(IN) :: dgemm_library
273
274 do_dgemm = dgemm_library
275 END SUBROUTINE local_gemm_set_library
276
277END 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)
Select the backend for subsequent local GEMM calls and context creation.
integer, parameter, public local_gemm_pu_gpu
subroutine local_dgemm(opa, opb, m, n, k, alpha, a, lda, b, ldb, beta, c, ldc, ctx)
Local GEMM on contiguous arrays, using BLAS or the configured SPLA backend. Each concurrent caller mu...
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().