14#ifndef EIGEN_GPU_CUBLAS_SUPPORT_H
15#define EIGEN_GPU_CUBLAS_SUPPORT_H
18#include "./InternalHeaderCheck.h"
20#include "./GpuSupport.h"
32inline void cublas_check_failed(cublasStatus_t status,
const char* expression,
const char* file,
int line) {
33#if defined(CUBLAS_VERSION) && CUBLAS_VERSION >= 110601
35 EIGEN_UNUSED_VARIABLE(status);
36 EIGEN_UNUSED_VARIABLE(expression);
37 EIGEN_UNUSED_VARIABLE(file);
38 EIGEN_UNUSED_VARIABLE(line);
39 EIGEN_GPU_CHECK_FAILED(cublasGetStatusName(status), expression, file, line);
41 gpu_check_failed_code(
"cuBLAS",
static_cast<int>(status), expression, file, line);
45#define EIGEN_CUBLAS_CHECK(expr) \
47 const cublasStatus_t _s = (expr); \
48 if (_s != CUBLAS_STATUS_SUCCESS) ::Eigen::gpu::internal::cublas_check_failed(_s, #expr, __FILE__, __LINE__); \
51constexpr cublasOperation_t to_cublas_op(GpuOp op) {
55 case GpuOp::ConjTrans:
70#if defined(CUBLAS_VERSION) && CUBLAS_VERSION >= 120000
71#define EIGEN_CUBLAS_FN(name) name##_64
72inline int64_t to_blas_dim(int64_t v) {
return v; }
74#define EIGEN_CUBLAS_FN(name) name
75inline int to_blas_dim(int64_t v) {
return to_blas_int(v); }
79struct CublasHandleDeleter {
81 void operator()(cublasHandle_t h)
const noexcept {
82 if (owns && h) (void)cublasDestroy(h);
85using UniqueCublasHandle = std::unique_ptr<std::remove_pointer_t<cublasHandle_t>, CublasHandleDeleter>;
87struct CublasLtHandleDeleter {
88 void operator()(cublasLtHandle_t h)
const noexcept {
89 if (h) (void)cublasLtDestroy(h);
92using UniqueCublasLtHandle = std::unique_ptr<std::remove_pointer_t<cublasLtHandle_t>, CublasLtHandleDeleter>;
103namespace cuda_compute_type_detail {
104#if defined(EIGEN_NO_CUDA_TENSOR_OPS)
105constexpr cublasComputeType_t kFloat = CUBLAS_COMPUTE_32F_PEDANTIC;
106constexpr cublasComputeType_t kDouble = CUBLAS_COMPUTE_64F_PEDANTIC;
107#elif defined(EIGEN_CUDA_TF32)
108constexpr cublasComputeType_t kFloat = CUBLAS_COMPUTE_32F_FAST_TF32;
109constexpr cublasComputeType_t kDouble = CUBLAS_COMPUTE_64F;
111constexpr cublasComputeType_t kFloat = CUBLAS_COMPUTE_32F;
112constexpr cublasComputeType_t kDouble = CUBLAS_COMPUTE_64F;
116template <
typename Scalar>
117struct cuda_compute_type;
120struct cuda_compute_type<float> {
121 static constexpr cublasComputeType_t value = cuda_compute_type_detail::kFloat;
124struct cuda_compute_type<double> {
125 static constexpr cublasComputeType_t value = cuda_compute_type_detail::kDouble;
128struct cuda_compute_type<std::complex<float>> {
129 static constexpr cublasComputeType_t value = cuda_compute_type_detail::kFloat;
132struct cuda_compute_type<std::complex<double>> {
133 static constexpr cublasComputeType_t value = cuda_compute_type_detail::kDouble;
136#define EIGEN_CUBLASLT_CHECK(expr) EIGEN_CUBLAS_CHECK(expr)
141#ifndef EIGEN_CUDA_CUBLASLT_MAX_WORKSPACE_BYTES
142#define EIGEN_CUDA_CUBLASLT_MAX_WORKSPACE_BYTES (32 * 1024 * 1024)
144static constexpr size_t kCublasLtMaxWorkspaceBytes = EIGEN_CUDA_CUBLASLT_MAX_WORKSPACE_BYTES;
150#ifndef EIGEN_CUDA_CUBLAS_WORKSPACE_BYTES
151#define EIGEN_CUDA_CUBLAS_WORKSPACE_BYTES (4 * 1024 * 1024)
153static constexpr size_t kCublasWorkspaceBytes = EIGEN_CUDA_CUBLAS_WORKSPACE_BYTES;
156constexpr cublasGemmAlgo_t cuda_gemm_algo() {
157#ifdef EIGEN_NO_CUDA_TENSOR_OPS
158 return CUBLAS_GEMM_DEFAULT;
160 return CUBLAS_GEMM_DEFAULT_TENSOR_OP;
168static constexpr std::size_t kCublasLtPlanCacheCapacity = 8;
170struct CublasLtPlanKey {
172 int64_t lda, ldb, ldc;
173 cudaDataType_t dtype;
174 cublasOperation_t transA, transB;
176 bool operator==(
const CublasLtPlanKey& o)
const {
177 return m == o.m && n == o.n && k == o.k && lda == o.lda && ldb == o.ldb && ldc == o.ldc && dtype == o.dtype &&
178 transA == o.transA && transB == o.transB;
182struct CublasLtPlanKeyHash {
183 std::size_t operator()(
const CublasLtPlanKey& k)
const noexcept {
185 auto mix = [](std::size_t a, std::size_t b) {
return a ^ (b + 0x9e3779b97f4a7c15ULL + (a << 6) + (a >> 2)); };
186 std::size_t r = std::hash<int64_t>{}(k.m);
187 r = mix(r, std::hash<int64_t>{}(k.n));
188 r = mix(r, std::hash<int64_t>{}(k.k));
189 r = mix(r, std::hash<int64_t>{}(k.lda));
190 r = mix(r, std::hash<int64_t>{}(k.ldb));
191 r = mix(r, std::hash<int64_t>{}(k.ldc));
192 r = mix(r, std::hash<int>{}(
static_cast<int>(k.dtype)));
193 r = mix(r, std::hash<int>{}(
static_cast<int>(k.transA)));
194 r = mix(r, std::hash<int>{}(
static_cast<int>(k.transB)));
203class CublasLtPlanEntry {
208 CublasLtPlanEntry(cublasLtHandle_t lt_handle,
const CublasLtPlanKey& key, cublasComputeType_t compute,
209 cudaDataType_t alpha_type, std::size_t max_workspace_bytes) {
210 EIGEN_CUBLASLT_CHECK(cublasLtMatmulDescCreate(&matmul_desc, compute, alpha_type));
211 EIGEN_CUBLASLT_CHECK(
212 cublasLtMatmulDescSetAttribute(matmul_desc, CUBLASLT_MATMUL_DESC_TRANSA, &key.transA,
sizeof(key.transA)));
213 EIGEN_CUBLASLT_CHECK(
214 cublasLtMatmulDescSetAttribute(matmul_desc, CUBLASLT_MATMUL_DESC_TRANSB, &key.transB,
sizeof(key.transB)));
219 const int64_t a_rows = (key.transA == CUBLAS_OP_N) ? key.m : key.k;
220 const int64_t a_cols = (key.transA == CUBLAS_OP_N) ? key.k : key.m;
221 const int64_t b_rows = (key.transB == CUBLAS_OP_N) ? key.k : key.n;
222 const int64_t b_cols = (key.transB == CUBLAS_OP_N) ? key.n : key.k;
223 EIGEN_CUBLASLT_CHECK(cublasLtMatrixLayoutCreate(&layout_A, key.dtype, a_rows, a_cols, key.lda));
224 EIGEN_CUBLASLT_CHECK(cublasLtMatrixLayoutCreate(&layout_B, key.dtype, b_rows, b_cols, key.ldb));
225 EIGEN_CUBLASLT_CHECK(cublasLtMatrixLayoutCreate(&layout_C, key.dtype, key.m, key.n, key.ldc));
227 cublasLtMatmulPreference_t preference =
nullptr;
228 EIGEN_CUBLASLT_CHECK(cublasLtMatmulPreferenceCreate(&preference));
229 EIGEN_CUBLASLT_CHECK(cublasLtMatmulPreferenceSetAttribute(preference, CUBLASLT_MATMUL_PREF_MAX_WORKSPACE_BYTES,
230 &max_workspace_bytes,
sizeof(max_workspace_bytes)));
232 cublasLtMatmulHeuristicResult_t result;
233 int returned_results = 0;
234 cublasStatus_t heuristic_status = cublasLtMatmulAlgoGetHeuristic(
235 lt_handle, matmul_desc, layout_A, layout_B, layout_C, layout_C, preference, 1, &result, &returned_results);
237 EIGEN_CUBLASLT_CHECK(cublasLtMatmulPreferenceDestroy(preference));
241 if (heuristic_status == CUBLAS_STATUS_SUCCESS && returned_results > 0 && result.state == CUBLAS_STATUS_SUCCESS) {
243 workspace_size = result.workspaceSize;
248 ~CublasLtPlanEntry() { destroy(); }
250 CublasLtPlanEntry(
const CublasLtPlanEntry&) =
delete;
251 CublasLtPlanEntry& operator=(
const CublasLtPlanEntry&) =
delete;
253 CublasLtPlanEntry(CublasLtPlanEntry&& o) noexcept
254 : matmul_desc(o.matmul_desc),
255 layout_A(o.layout_A),
256 layout_B(o.layout_B),
257 layout_C(o.layout_C),
259 workspace_size(o.workspace_size),
260 use_cublaslt(o.use_cublaslt) {
261 o.matmul_desc =
nullptr;
262 o.layout_A = o.layout_B = o.layout_C =
nullptr;
263 o.use_cublaslt =
false;
266 CublasLtPlanEntry& operator=(CublasLtPlanEntry&& o)
noexcept {
269 matmul_desc = o.matmul_desc;
270 layout_A = o.layout_A;
271 layout_B = o.layout_B;
272 layout_C = o.layout_C;
274 workspace_size = o.workspace_size;
275 use_cublaslt = o.use_cublaslt;
276 o.matmul_desc =
nullptr;
277 o.layout_A = o.layout_B = o.layout_C =
nullptr;
278 o.use_cublaslt =
false;
283 cublasLtMatmulDesc_t matmul_desc =
nullptr;
284 cublasLtMatrixLayout_t layout_A =
nullptr;
285 cublasLtMatrixLayout_t layout_B =
nullptr;
286 cublasLtMatrixLayout_t layout_C =
nullptr;
287 cublasLtMatmulAlgo_t algo{};
288 std::size_t workspace_size = 0;
289 bool use_cublaslt =
false;
292 void destroy() noexcept {
293 if (layout_C) cublasLtMatrixLayoutDestroy(layout_C);
294 if (layout_B) cublasLtMatrixLayoutDestroy(layout_B);
295 if (layout_A) cublasLtMatrixLayoutDestroy(layout_A);
296 if (matmul_desc) cublasLtMatmulDescDestroy(matmul_desc);
300using CublasLtPlanCache = Eigen::internal::LruCache<CublasLtPlanKey, CublasLtPlanEntry, CublasLtPlanKeyHash>;
307template <
typename Scalar>
308void cublaslt_gemm(cublasLtHandle_t lt_handle, cublasHandle_t cublas_handle, cublasOperation_t transA,
309 cublasOperation_t transB, int64_t m, int64_t n, int64_t k,
const Scalar* alpha,
const Scalar* A,
310 int64_t lda,
const Scalar* B, int64_t ldb,
const Scalar* beta, Scalar* C, int64_t ldc,
311 DeviceBuffer& workspace, CublasLtPlanCache& plan_cache, std::size_t max_workspace_bytes,
312 cudaStream_t stream) {
313 constexpr cudaDataType_t dtype = cuda_data_type<Scalar>::value;
314 constexpr cublasComputeType_t compute = cuda_compute_type<Scalar>::value;
315 constexpr cudaDataType_t alpha_type = cuda_data_type<Scalar>::value;
319 const CublasLtPlanKey key{m, n, k, lda, ldb, ldc, dtype, transA, transB};
320 CublasLtPlanEntry* entry = plan_cache.find(key);
322 entry = plan_cache.insert(key, CublasLtPlanEntry(lt_handle, key, compute, alpha_type, max_workspace_bytes));
330 alignas(cuDoubleComplex)
const Scalar alpha_val = *alpha;
331 alignas(cuDoubleComplex)
const Scalar beta_val = *beta;
333 if (entry->use_cublaslt) {
334 const size_t needed = entry->workspace_size;
335 if (needed > workspace.size()) {
337 if (workspace.get()) EIGEN_CUDA_RUNTIME_CHECK(cudaStreamSynchronize(stream));
341 EIGEN_CUBLASLT_CHECK(cublasLtMatmul(lt_handle, entry->matmul_desc, &alpha_val, A, entry->layout_A, B,
342 entry->layout_B, &beta_val, C, entry->layout_C, C, entry->layout_C,
343 &entry->algo, workspace.get(), needed, stream));
346 EIGEN_CUBLAS_CHECK(EIGEN_CUBLAS_FN(cublasGemmEx)(cublas_handle, transA, transB, to_blas_dim(m), to_blas_dim(n),
347 to_blas_dim(k), &alpha_val, A, dtype, to_blas_dim(lda), B, dtype,
348 to_blas_dim(ldb), &beta_val, C, dtype, to_blas_dim(ldc), compute,
355inline cublasStatus_t cublasXgemm(cublasHandle_t h, cublasOperation_t transA, cublasOperation_t transB, int64_t m,
356 int64_t n, int64_t k,
const float* alpha,
const float* A, int64_t lda,
const float* B,
357 int64_t ldb,
const float* beta,
float* C, int64_t ldc) {
358 return EIGEN_CUBLAS_FN(cublasSgemm)(h, transA, transB, to_blas_dim(m), to_blas_dim(n), to_blas_dim(k), alpha, A,
359 to_blas_dim(lda), B, to_blas_dim(ldb), beta, C, to_blas_dim(ldc));
361inline cublasStatus_t cublasXgemm(cublasHandle_t h, cublasOperation_t transA, cublasOperation_t transB, int64_t m,
362 int64_t n, int64_t k,
const double* alpha,
const double* A, int64_t lda,
363 const double* B, int64_t ldb,
const double* beta,
double* C, int64_t ldc) {
364 return EIGEN_CUBLAS_FN(cublasDgemm)(h, transA, transB, to_blas_dim(m), to_blas_dim(n), to_blas_dim(k), alpha, A,
365 to_blas_dim(lda), B, to_blas_dim(ldb), beta, C, to_blas_dim(ldc));
367static_assert(
sizeof(cuComplex) ==
sizeof(std::complex<float>),
"cuComplex and std::complex<float> layout mismatch");
368static_assert(
sizeof(cuDoubleComplex) ==
sizeof(std::complex<double>),
369 "cuDoubleComplex and std::complex<double> layout mismatch");
376inline cublasStatus_t cublasXgemm(cublasHandle_t h, cublasOperation_t transA, cublasOperation_t transB, int64_t m,
377 int64_t n, int64_t k,
const std::complex<float>* alpha,
const std::complex<float>* A,
378 int64_t lda,
const std::complex<float>* B, int64_t ldb,
379 const std::complex<float>* beta, std::complex<float>* C, int64_t ldc) {
381 std::memcpy(&a, alpha,
sizeof(a));
382 std::memcpy(&b, beta,
sizeof(b));
383 return EIGEN_CUBLAS_FN(cublasCgemm)(h, transA, transB, to_blas_dim(m), to_blas_dim(n), to_blas_dim(k), &a,
384 reinterpret_cast<const cuComplex*
>(A), to_blas_dim(lda),
385 reinterpret_cast<const cuComplex*
>(B), to_blas_dim(ldb), &b,
386 reinterpret_cast<cuComplex*
>(C), to_blas_dim(ldc));
388inline cublasStatus_t cublasXgemm(cublasHandle_t h, cublasOperation_t transA, cublasOperation_t transB, int64_t m,
389 int64_t n, int64_t k,
const std::complex<double>* alpha,
390 const std::complex<double>* A, int64_t lda,
const std::complex<double>* B,
391 int64_t ldb,
const std::complex<double>* beta, std::complex<double>* C, int64_t ldc) {
392 cuDoubleComplex a, b;
393 std::memcpy(&a, alpha,
sizeof(a));
394 std::memcpy(&b, beta,
sizeof(b));
395 return EIGEN_CUBLAS_FN(cublasZgemm)(h, transA, transB, to_blas_dim(m), to_blas_dim(n), to_blas_dim(k), &a,
396 reinterpret_cast<const cuDoubleComplex*
>(A), to_blas_dim(lda),
397 reinterpret_cast<const cuDoubleComplex*
>(B), to_blas_dim(ldb), &b,
398 reinterpret_cast<cuDoubleComplex*
>(C), to_blas_dim(ldc));
401inline cublasStatus_t cublasXtrsm(cublasHandle_t h, cublasSideMode_t side, cublasFillMode_t uplo,
402 cublasOperation_t trans, cublasDiagType_t diag, int64_t m, int64_t n,
403 const float* alpha,
const float* A, int64_t lda,
float* B, int64_t ldb) {
404 return EIGEN_CUBLAS_FN(cublasStrsm)(h, side, uplo, trans, diag, to_blas_dim(m), to_blas_dim(n), alpha, A,
405 to_blas_dim(lda), B, to_blas_dim(ldb));
407inline cublasStatus_t cublasXtrsm(cublasHandle_t h, cublasSideMode_t side, cublasFillMode_t uplo,
408 cublasOperation_t trans, cublasDiagType_t diag, int64_t m, int64_t n,
409 const double* alpha,
const double* A, int64_t lda,
double* B, int64_t ldb) {
410 return EIGEN_CUBLAS_FN(cublasDtrsm)(h, side, uplo, trans, diag, to_blas_dim(m), to_blas_dim(n), alpha, A,
411 to_blas_dim(lda), B, to_blas_dim(ldb));
413inline cublasStatus_t cublasXtrsm(cublasHandle_t h, cublasSideMode_t side, cublasFillMode_t uplo,
414 cublasOperation_t trans, cublasDiagType_t diag, int64_t m, int64_t n,
415 const std::complex<float>* alpha,
const std::complex<float>* A, int64_t lda,
416 std::complex<float>* B, int64_t ldb) {
418 std::memcpy(&a, alpha,
sizeof(a));
419 return EIGEN_CUBLAS_FN(cublasCtrsm)(h, side, uplo, trans, diag, to_blas_dim(m), to_blas_dim(n), &a,
420 reinterpret_cast<const cuComplex*
>(A), to_blas_dim(lda),
421 reinterpret_cast<cuComplex*
>(B), to_blas_dim(ldb));
423inline cublasStatus_t cublasXtrsm(cublasHandle_t h, cublasSideMode_t side, cublasFillMode_t uplo,
424 cublasOperation_t trans, cublasDiagType_t diag, int64_t m, int64_t n,
425 const std::complex<double>* alpha,
const std::complex<double>* A, int64_t lda,
426 std::complex<double>* B, int64_t ldb) {
428 std::memcpy(&a, alpha,
sizeof(a));
429 return EIGEN_CUBLAS_FN(cublasZtrsm)(h, side, uplo, trans, diag, to_blas_dim(m), to_blas_dim(n), &a,
430 reinterpret_cast<const cuDoubleComplex*
>(A), to_blas_dim(lda),
431 reinterpret_cast<cuDoubleComplex*
>(B), to_blas_dim(ldb));
435inline cublasStatus_t cublasXsymm(cublasHandle_t h, cublasSideMode_t side, cublasFillMode_t uplo, int64_t m, int64_t n,
436 const float* alpha,
const float* A, int64_t lda,
const float* B, int64_t ldb,
437 const float* beta,
float* C, int64_t ldc) {
438 return EIGEN_CUBLAS_FN(cublasSsymm)(h, side, uplo, to_blas_dim(m), to_blas_dim(n), alpha, A, to_blas_dim(lda), B,
439 to_blas_dim(ldb), beta, C, to_blas_dim(ldc));
441inline cublasStatus_t cublasXsymm(cublasHandle_t h, cublasSideMode_t side, cublasFillMode_t uplo, int64_t m, int64_t n,
442 const double* alpha,
const double* A, int64_t lda,
const double* B, int64_t ldb,
443 const double* beta,
double* C, int64_t ldc) {
444 return EIGEN_CUBLAS_FN(cublasDsymm)(h, side, uplo, to_blas_dim(m), to_blas_dim(n), alpha, A, to_blas_dim(lda), B,
445 to_blas_dim(ldb), beta, C, to_blas_dim(ldc));
447inline cublasStatus_t cublasXsymm(cublasHandle_t h, cublasSideMode_t side, cublasFillMode_t uplo, int64_t m, int64_t n,
448 const std::complex<float>* alpha,
const std::complex<float>* A, int64_t lda,
449 const std::complex<float>* B, int64_t ldb,
const std::complex<float>* beta,
450 std::complex<float>* C, int64_t ldc) {
452 std::memcpy(&a, alpha,
sizeof(a));
453 std::memcpy(&b, beta,
sizeof(b));
454 return EIGEN_CUBLAS_FN(cublasChemm)(
455 h, side, uplo, to_blas_dim(m), to_blas_dim(n), &a,
reinterpret_cast<const cuComplex*
>(A), to_blas_dim(lda),
456 reinterpret_cast<const cuComplex*
>(B), to_blas_dim(ldb), &b,
reinterpret_cast<cuComplex*
>(C), to_blas_dim(ldc));
458inline cublasStatus_t cublasXsymm(cublasHandle_t h, cublasSideMode_t side, cublasFillMode_t uplo, int64_t m, int64_t n,
459 const std::complex<double>* alpha,
const std::complex<double>* A, int64_t lda,
460 const std::complex<double>* B, int64_t ldb,
const std::complex<double>* beta,
461 std::complex<double>* C, int64_t ldc) {
462 cuDoubleComplex a, b;
463 std::memcpy(&a, alpha,
sizeof(a));
464 std::memcpy(&b, beta,
sizeof(b));
465 return EIGEN_CUBLAS_FN(cublasZhemm)(h, side, uplo, to_blas_dim(m), to_blas_dim(n), &a,
466 reinterpret_cast<const cuDoubleComplex*
>(A), to_blas_dim(lda),
467 reinterpret_cast<const cuDoubleComplex*
>(B), to_blas_dim(ldb), &b,
468 reinterpret_cast<cuDoubleComplex*
>(C), to_blas_dim(ldc));
472inline cublasStatus_t cublasXgeam(cublasHandle_t h, cublasOperation_t transA, cublasOperation_t transB, int64_t m,
473 int64_t n,
const float* alpha,
const float* A, int64_t lda,
const float* beta,
474 const float* B, int64_t ldb,
float* C, int64_t ldc) {
475 return EIGEN_CUBLAS_FN(cublasSgeam)(h, transA, transB, to_blas_dim(m), to_blas_dim(n), alpha, A, to_blas_dim(lda),
476 beta, B, to_blas_dim(ldb), C, to_blas_dim(ldc));
478inline cublasStatus_t cublasXgeam(cublasHandle_t h, cublasOperation_t transA, cublasOperation_t transB, int64_t m,
479 int64_t n,
const double* alpha,
const double* A, int64_t lda,
const double* beta,
480 const double* B, int64_t ldb,
double* C, int64_t ldc) {
481 return EIGEN_CUBLAS_FN(cublasDgeam)(h, transA, transB, to_blas_dim(m), to_blas_dim(n), alpha, A, to_blas_dim(lda),
482 beta, B, to_blas_dim(ldb), C, to_blas_dim(ldc));
484inline cublasStatus_t cublasXgeam(cublasHandle_t h, cublasOperation_t transA, cublasOperation_t transB, int64_t m,
485 int64_t n,
const std::complex<float>* alpha,
const std::complex<float>* A,
486 int64_t lda,
const std::complex<float>* beta,
const std::complex<float>* B,
487 int64_t ldb, std::complex<float>* C, int64_t ldc) {
489 std::memcpy(&a, alpha,
sizeof(a));
490 std::memcpy(&b, beta,
sizeof(b));
491 return EIGEN_CUBLAS_FN(cublasCgeam)(
492 h, transA, transB, to_blas_dim(m), to_blas_dim(n), &a,
reinterpret_cast<const cuComplex*
>(A), to_blas_dim(lda),
493 &b,
reinterpret_cast<const cuComplex*
>(B), to_blas_dim(ldb),
reinterpret_cast<cuComplex*
>(C), to_blas_dim(ldc));
495inline cublasStatus_t cublasXgeam(cublasHandle_t h, cublasOperation_t transA, cublasOperation_t transB, int64_t m,
496 int64_t n,
const std::complex<double>* alpha,
const std::complex<double>* A,
497 int64_t lda,
const std::complex<double>* beta,
const std::complex<double>* B,
498 int64_t ldb, std::complex<double>* C, int64_t ldc) {
499 cuDoubleComplex a, b;
500 std::memcpy(&a, alpha,
sizeof(a));
501 std::memcpy(&b, beta,
sizeof(b));
502 return EIGEN_CUBLAS_FN(cublasZgeam)(h, transA, transB, to_blas_dim(m), to_blas_dim(n), &a,
503 reinterpret_cast<const cuDoubleComplex*
>(A), to_blas_dim(lda), &b,
504 reinterpret_cast<const cuDoubleComplex*
>(B), to_blas_dim(ldb),
505 reinterpret_cast<cuDoubleComplex*
>(C), to_blas_dim(ldc));
509inline cublasStatus_t cublasXsyrk(cublasHandle_t h, cublasFillMode_t uplo, cublasOperation_t trans, int64_t n,
510 int64_t k,
const float* alpha,
const float* A, int64_t lda,
const float* beta,
511 float* C, int64_t ldc) {
512 return EIGEN_CUBLAS_FN(cublasSsyrk)(h, uplo, trans, to_blas_dim(n), to_blas_dim(k), alpha, A, to_blas_dim(lda), beta,
513 C, to_blas_dim(ldc));
515inline cublasStatus_t cublasXsyrk(cublasHandle_t h, cublasFillMode_t uplo, cublasOperation_t trans, int64_t n,
516 int64_t k,
const double* alpha,
const double* A, int64_t lda,
const double* beta,
517 double* C, int64_t ldc) {
518 return EIGEN_CUBLAS_FN(cublasDsyrk)(h, uplo, trans, to_blas_dim(n), to_blas_dim(k), alpha, A, to_blas_dim(lda), beta,
519 C, to_blas_dim(ldc));
521inline cublasStatus_t cublasXsyrk(cublasHandle_t h, cublasFillMode_t uplo, cublasOperation_t trans, int64_t n,
522 int64_t k,
const float* alpha,
const std::complex<float>* A, int64_t lda,
523 const float* beta, std::complex<float>* C, int64_t ldc) {
524 return EIGEN_CUBLAS_FN(cublasCherk)(h, uplo, trans, to_blas_dim(n), to_blas_dim(k), alpha,
525 reinterpret_cast<const cuComplex*
>(A), to_blas_dim(lda), beta,
526 reinterpret_cast<cuComplex*
>(C), to_blas_dim(ldc));
528inline cublasStatus_t cublasXsyrk(cublasHandle_t h, cublasFillMode_t uplo, cublasOperation_t trans, int64_t n,
529 int64_t k,
const double* alpha,
const std::complex<double>* A, int64_t lda,
530 const double* beta, std::complex<double>* C, int64_t ldc) {
531 return EIGEN_CUBLAS_FN(cublasZherk)(h, uplo, trans, to_blas_dim(n), to_blas_dim(k), alpha,
532 reinterpret_cast<const cuDoubleComplex*
>(A), to_blas_dim(lda), beta,
533 reinterpret_cast<cuDoubleComplex*
>(C), to_blas_dim(ldc));
538inline cublasStatus_t cublasXscal(cublasHandle_t h, int64_t n,
const float* alpha,
float* x, int64_t incx) {
539 return EIGEN_CUBLAS_FN(cublasSscal)(h, to_blas_dim(n), alpha, x, to_blas_dim(incx));
541inline cublasStatus_t cublasXscal(cublasHandle_t h, int64_t n,
const double* alpha,
double* x, int64_t incx) {
542 return EIGEN_CUBLAS_FN(cublasDscal)(h, to_blas_dim(n), alpha, x, to_blas_dim(incx));
544inline cublasStatus_t cublasXscal(cublasHandle_t h, int64_t n,
const float* alpha, std::complex<float>* x,
546 return EIGEN_CUBLAS_FN(cublasCsscal)(h, to_blas_dim(n), alpha,
reinterpret_cast<cuComplex*
>(x), to_blas_dim(incx));
548inline cublasStatus_t cublasXscal(cublasHandle_t h, int64_t n,
const double* alpha, std::complex<double>* x,
550 return EIGEN_CUBLAS_FN(cublasZdscal)(h, to_blas_dim(n), alpha,
reinterpret_cast<cuDoubleComplex*
>(x),
556inline cublasStatus_t cublasXscal(cublasHandle_t h, int64_t n,
float alpha,
float* x, int64_t incx) {
557 return EIGEN_CUBLAS_FN(cublasSscal)(h, to_blas_dim(n), &alpha, x, to_blas_dim(incx));
559inline cublasStatus_t cublasXscal(cublasHandle_t h, int64_t n,
double alpha,
double* x, int64_t incx) {
560 return EIGEN_CUBLAS_FN(cublasDscal)(h, to_blas_dim(n), &alpha, x, to_blas_dim(incx));
562inline cublasStatus_t cublasXscal(cublasHandle_t h, int64_t n,
float alpha, std::complex<float>* x, int64_t incx) {
563 return EIGEN_CUBLAS_FN(cublasCsscal)(h, to_blas_dim(n), &alpha,
reinterpret_cast<cuComplex*
>(x), to_blas_dim(incx));
565inline cublasStatus_t cublasXscal(cublasHandle_t h, int64_t n,
double alpha, std::complex<double>* x, int64_t incx) {
566 return EIGEN_CUBLAS_FN(cublasZdscal)(h, to_blas_dim(n), &alpha,
reinterpret_cast<cuDoubleComplex*
>(x),
573inline cublasStatus_t cublasXdgmm(cublasHandle_t h, cublasSideMode_t side, int64_t m, int64_t n,
const float* A,
574 int64_t lda,
const float* x, int64_t incx,
float* C, int64_t ldc) {
575 return EIGEN_CUBLAS_FN(cublasSdgmm)(h, side, to_blas_dim(m), to_blas_dim(n), A, to_blas_dim(lda), x,
576 to_blas_dim(incx), C, to_blas_dim(ldc));
578inline cublasStatus_t cublasXdgmm(cublasHandle_t h, cublasSideMode_t side, int64_t m, int64_t n,
const double* A,
579 int64_t lda,
const double* x, int64_t incx,
double* C, int64_t ldc) {
580 return EIGEN_CUBLAS_FN(cublasDdgmm)(h, side, to_blas_dim(m), to_blas_dim(n), A, to_blas_dim(lda), x,
581 to_blas_dim(incx), C, to_blas_dim(ldc));
583inline cublasStatus_t cublasXdgmm(cublasHandle_t h, cublasSideMode_t side, int64_t m, int64_t n,
584 const std::complex<float>* A, int64_t lda,
const std::complex<float>* x, int64_t incx,
585 std::complex<float>* C, int64_t ldc) {
586 return EIGEN_CUBLAS_FN(cublasCdgmm)(h, side, to_blas_dim(m), to_blas_dim(n),
reinterpret_cast<const cuComplex*
>(A),
587 to_blas_dim(lda),
reinterpret_cast<const cuComplex*
>(x), to_blas_dim(incx),
588 reinterpret_cast<cuComplex*
>(C), to_blas_dim(ldc));
590inline cublasStatus_t cublasXdgmm(cublasHandle_t h, cublasSideMode_t side, int64_t m, int64_t n,
591 const std::complex<double>* A, int64_t lda,
const std::complex<double>* x,
592 int64_t incx, std::complex<double>* C, int64_t ldc) {
593 return EIGEN_CUBLAS_FN(cublasZdgmm)(h, side, to_blas_dim(m), to_blas_dim(n),
594 reinterpret_cast<const cuDoubleComplex*
>(A), to_blas_dim(lda),
595 reinterpret_cast<const cuDoubleComplex*
>(x), to_blas_dim(incx),
596 reinterpret_cast<cuDoubleComplex*
>(C), to_blas_dim(ldc));
607template <
typename CuComplexType>
608struct Blas1ComplexAlpha {
609 template <
typename Complex>
610 Blas1ComplexAlpha(cublasHandle_t h,
const Complex* alpha) {
611 static_assert(
sizeof(Complex) ==
sizeof(CuComplexType),
"complex alpha layout mismatch");
612 cublasPointerMode_t mode = CUBLAS_POINTER_MODE_HOST;
613 status = cublasGetPointerMode(h, &mode);
614 if (status != CUBLAS_STATUS_SUCCESS || mode == CUBLAS_POINTER_MODE_DEVICE) {
615 ptr =
reinterpret_cast<const CuComplexType*
>(alpha);
617 std::memcpy(&value, alpha,
sizeof(value));
622 Blas1ComplexAlpha(
const Blas1ComplexAlpha&) =
delete;
623 Blas1ComplexAlpha& operator=(
const Blas1ComplexAlpha&) =
delete;
626 const CuComplexType* ptr;
627 cublasStatus_t status;
631inline cublasStatus_t cublasXdot(cublasHandle_t h, int64_t n,
const float* x, int64_t incx,
const float* y,
632 int64_t incy,
float* result) {
633 return EIGEN_CUBLAS_FN(cublasSdot)(h, to_blas_dim(n), x, to_blas_dim(incx), y, to_blas_dim(incy), result);
635inline cublasStatus_t cublasXdot(cublasHandle_t h, int64_t n,
const double* x, int64_t incx,
const double* y,
636 int64_t incy,
double* result) {
637 return EIGEN_CUBLAS_FN(cublasDdot)(h, to_blas_dim(n), x, to_blas_dim(incx), y, to_blas_dim(incy), result);
639inline cublasStatus_t cublasXdot(cublasHandle_t h, int64_t n,
const std::complex<float>* x, int64_t incx,
640 const std::complex<float>* y, int64_t incy, std::complex<float>* result) {
641 return EIGEN_CUBLAS_FN(cublasCdotc)(h, to_blas_dim(n),
reinterpret_cast<const cuComplex*
>(x), to_blas_dim(incx),
642 reinterpret_cast<const cuComplex*
>(y), to_blas_dim(incy),
643 reinterpret_cast<cuComplex*
>(result));
645inline cublasStatus_t cublasXdot(cublasHandle_t h, int64_t n,
const std::complex<double>* x, int64_t incx,
646 const std::complex<double>* y, int64_t incy, std::complex<double>* result) {
647 return EIGEN_CUBLAS_FN(cublasZdotc)(h, to_blas_dim(n),
reinterpret_cast<const cuDoubleComplex*
>(x), to_blas_dim(incx),
648 reinterpret_cast<const cuDoubleComplex*
>(y), to_blas_dim(incy),
649 reinterpret_cast<cuDoubleComplex*
>(result));
653inline cublasStatus_t cublasXnrm2(cublasHandle_t h, int64_t n,
const float* x, int64_t incx,
float* result) {
654 return EIGEN_CUBLAS_FN(cublasSnrm2)(h, to_blas_dim(n), x, to_blas_dim(incx), result);
656inline cublasStatus_t cublasXnrm2(cublasHandle_t h, int64_t n,
const double* x, int64_t incx,
double* result) {
657 return EIGEN_CUBLAS_FN(cublasDnrm2)(h, to_blas_dim(n), x, to_blas_dim(incx), result);
659inline cublasStatus_t cublasXnrm2(cublasHandle_t h, int64_t n,
const std::complex<float>* x, int64_t incx,
661 return EIGEN_CUBLAS_FN(cublasScnrm2)(h, to_blas_dim(n),
reinterpret_cast<const cuComplex*
>(x), to_blas_dim(incx),
664inline cublasStatus_t cublasXnrm2(cublasHandle_t h, int64_t n,
const std::complex<double>* x, int64_t incx,
666 return EIGEN_CUBLAS_FN(cublasDznrm2)(h, to_blas_dim(n),
reinterpret_cast<const cuDoubleComplex*
>(x),
667 to_blas_dim(incx), result);
671inline cublasStatus_t cublasXaxpy(cublasHandle_t h, int64_t n,
const float* alpha,
const float* x, int64_t incx,
672 float* y, int64_t incy) {
673 return EIGEN_CUBLAS_FN(cublasSaxpy)(h, to_blas_dim(n), alpha, x, to_blas_dim(incx), y, to_blas_dim(incy));
675inline cublasStatus_t cublasXaxpy(cublasHandle_t h, int64_t n,
const double* alpha,
const double* x, int64_t incx,
676 double* y, int64_t incy) {
677 return EIGEN_CUBLAS_FN(cublasDaxpy)(h, to_blas_dim(n), alpha, x, to_blas_dim(incx), y, to_blas_dim(incy));
679inline cublasStatus_t cublasXaxpy(cublasHandle_t h, int64_t n,
const std::complex<float>* alpha,
680 const std::complex<float>* x, int64_t incx, std::complex<float>* y, int64_t incy) {
681 const Blas1ComplexAlpha<cuComplex> a(h, alpha);
682 if (a.status != CUBLAS_STATUS_SUCCESS)
return a.status;
683 return EIGEN_CUBLAS_FN(cublasCaxpy)(h, to_blas_dim(n), a.ptr,
reinterpret_cast<const cuComplex*
>(x),
684 to_blas_dim(incx),
reinterpret_cast<cuComplex*
>(y), to_blas_dim(incy));
686inline cublasStatus_t cublasXaxpy(cublasHandle_t h, int64_t n,
const std::complex<double>* alpha,
687 const std::complex<double>* x, int64_t incx, std::complex<double>* y, int64_t incy) {
688 const Blas1ComplexAlpha<cuDoubleComplex> a(h, alpha);
689 if (a.status != CUBLAS_STATUS_SUCCESS)
return a.status;
690 return EIGEN_CUBLAS_FN(cublasZaxpy)(h, to_blas_dim(n), a.ptr,
reinterpret_cast<const cuDoubleComplex*
>(x),
691 to_blas_dim(incx),
reinterpret_cast<cuDoubleComplex*
>(y), to_blas_dim(incy));
695inline cublasStatus_t cublasXscal(cublasHandle_t h, int64_t n,
const std::complex<float>* alpha, std::complex<float>* x,
697 const Blas1ComplexAlpha<cuComplex> a(h, alpha);
698 if (a.status != CUBLAS_STATUS_SUCCESS)
return a.status;
699 return EIGEN_CUBLAS_FN(cublasCscal)(h, to_blas_dim(n), a.ptr,
reinterpret_cast<cuComplex*
>(x), to_blas_dim(incx));
701inline cublasStatus_t cublasXscal(cublasHandle_t h, int64_t n,
const std::complex<double>* alpha,
702 std::complex<double>* x, int64_t incx) {
703 const Blas1ComplexAlpha<cuDoubleComplex> a(h, alpha);
704 if (a.status != CUBLAS_STATUS_SUCCESS)
return a.status;
705 return EIGEN_CUBLAS_FN(cublasZscal)(h, to_blas_dim(n), a.ptr,
reinterpret_cast<cuDoubleComplex*
>(x),
710inline cublasStatus_t cublasXcopy(cublasHandle_t h, int64_t n,
const float* x, int64_t incx,
float* y, int64_t incy) {
711 return EIGEN_CUBLAS_FN(cublasScopy)(h, to_blas_dim(n), x, to_blas_dim(incx), y, to_blas_dim(incy));
713inline cublasStatus_t cublasXcopy(cublasHandle_t h, int64_t n,
const double* x, int64_t incx,
double* y, int64_t incy) {
714 return EIGEN_CUBLAS_FN(cublasDcopy)(h, to_blas_dim(n), x, to_blas_dim(incx), y, to_blas_dim(incy));
716inline cublasStatus_t cublasXcopy(cublasHandle_t h, int64_t n,
const std::complex<float>* x, int64_t incx,
717 std::complex<float>* y, int64_t incy) {
718 return EIGEN_CUBLAS_FN(cublasCcopy)(h, to_blas_dim(n),
reinterpret_cast<const cuComplex*
>(x), to_blas_dim(incx),
719 reinterpret_cast<cuComplex*
>(y), to_blas_dim(incy));
721inline cublasStatus_t cublasXcopy(cublasHandle_t h, int64_t n,
const std::complex<double>* x, int64_t incx,
722 std::complex<double>* y, int64_t incy) {
723 return EIGEN_CUBLAS_FN(cublasZcopy)(h, to_blas_dim(n),
reinterpret_cast<const cuDoubleComplex*
>(x), to_blas_dim(incx),
724 reinterpret_cast<cuDoubleComplex*
>(y), to_blas_dim(incy));
Internal RAII owner for an untyped GPU device allocation.
Definition GpuSupport.h:293
Namespace containing all symbols from the Eigen library.