9 USE iso_c_binding,
ONLY: c_null_ptr, &
12#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
14 USE iso_c_binding,
ONLY: c_associated, &
16 USE spla,
ONLY: spla_op_none, &
18 spla_op_conj_transpose, &
23 spla_ctx_set_op_threshold_gpu, &
30#include "./base/base_uses.f90"
36 CHARACTER(len=*),
PARAMETER,
PRIVATE :: moduleN =
'local_gemm_api'
41 INTEGER,
PARAMETER,
PUBLIC :: &
45 INTEGER,
PRIVATE :: do_dgemm = 1
48 TYPE(c_ptr) :: spla_context = c_null_ptr
49 LOGICAL,
PRIVATE :: timing = .true.
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
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, *)
88#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
91 CHARACTER(LEN=*),
PARAMETER :: routinen =
'local_gemm'
93 IF (m == 0 .OR. n == 0)
RETURN
94 IF (ctx%timing)
CALL timeset(routinen, handle)
95#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
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))
106 CALL dgemm(opa, opb, m, n, k, alpha, a, lda, b, ldb, beta, c, ldc)
107#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
110 IF (ctx%timing)
CALL timestop(handle)
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, *)
140#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
141 INTEGER :: spla_error
143 CHARACTER(LEN=*),
PARAMETER :: routineN =
'local_gemm'
145 IF (m == 0 .OR. n == 0)
RETURN
146 IF (ctx%timing)
CALL timeset(routinen, handle)
147#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
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))
158 CALL zgemm(opa, opb, m, n, k, alpha, a, lda, b, ldb, beta, c, ldc)
159#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
162 IF (ctx%timing)
CALL timestop(handle)
164 END SUBROUTINE local_zgemm
166#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
172 FUNCTION spla_operation(trans)
RESULT(op)
173 CHARACTER,
INTENT(IN) :: trans
180 op = spla_op_transpose
182 op = spla_op_conj_transpose
184 CALL cp_abort(__location__,
"Invalid local GEMM transpose flag.")
186 END FUNCTION spla_operation
195 SUBROUTINE local_gemm_create(ctx, pu, timing)
197 INTEGER,
INTENT(IN) :: pu
198 LOGICAL,
INTENT(IN),
OPTIONAL :: timing
200#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
201 INTEGER :: spla_error
204 IF (
PRESENT(timing)) ctx%timing = timing
205#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
209 spla_error = spla_ctx_create(ctx%spla_context, pu)
210 IF (spla_error /= spla_success) &
211 CALL cp_abort(__location__, &
217 END SUBROUTINE local_gemm_create
223 SUBROUTINE local_gemm_destroy(ctx)
226#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
227 INTEGER :: spla_error
229 IF (c_associated(ctx%spla_context))
THEN
232 spla_error = spla_ctx_destroy(ctx%spla_context)
233 IF (spla_error /= spla_success) &
234 CALL cp_abort(__location__, &
238 ctx%spla_context = c_null_ptr
239 END SUBROUTINE local_gemm_destroy
246 SUBROUTINE local_gemm_set_op_threshold_gpu(ctx, opThresholdGPU)
248 INTEGER,
INTENT(IN) :: opThresholdGPU
250#if defined(__SPLA) && defined(__OFFLOAD_GEMM)
251 INTEGER :: spla_error
253 IF (c_associated(ctx%spla_context))
THEN
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))
263 mark_used(opthresholdgpu)
265 END SUBROUTINE local_gemm_set_op_threshold_gpu
272 INTEGER,
INTENT(IN) :: dgemm_library
274 do_dgemm = dgemm_library
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 ...
Defines the basic variable types.
integer, parameter, public dp
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.
subroutine, public offload_activate_chosen_device()
Activates the device selected via offload_set_chosen_device().