14#ifndef EIGEN_GPU_CUSOLVER_SUPPORT_H
15#define EIGEN_GPU_CUSOLVER_SUPPORT_H
18#include "./InternalHeaderCheck.h"
20#include "./GpuSupport.h"
21#include <cusolverDn.h>
29inline const char* cusolver_status_name(cusolverStatus_t s) {
31 case CUSOLVER_STATUS_SUCCESS:
32 return "CUSOLVER_STATUS_SUCCESS";
33 case CUSOLVER_STATUS_NOT_INITIALIZED:
34 return "CUSOLVER_STATUS_NOT_INITIALIZED";
35 case CUSOLVER_STATUS_ALLOC_FAILED:
36 return "CUSOLVER_STATUS_ALLOC_FAILED";
37 case CUSOLVER_STATUS_INVALID_VALUE:
38 return "CUSOLVER_STATUS_INVALID_VALUE";
39 case CUSOLVER_STATUS_ARCH_MISMATCH:
40 return "CUSOLVER_STATUS_ARCH_MISMATCH";
41 case CUSOLVER_STATUS_EXECUTION_FAILED:
42 return "CUSOLVER_STATUS_EXECUTION_FAILED";
43 case CUSOLVER_STATUS_INTERNAL_ERROR:
44 return "CUSOLVER_STATUS_INTERNAL_ERROR";
45 case CUSOLVER_STATUS_MATRIX_TYPE_NOT_SUPPORTED:
46 return "CUSOLVER_STATUS_MATRIX_TYPE_NOT_SUPPORTED";
47 case CUSOLVER_STATUS_NOT_SUPPORTED:
48 return "CUSOLVER_STATUS_NOT_SUPPORTED";
50 return "CUSOLVER_STATUS_UNKNOWN";
54#define EIGEN_CUSOLVER_CHECK(expr) \
56 const cusolverStatus_t _s = (expr); \
57 if (_s != CUSOLVER_STATUS_SUCCESS) \
58 EIGEN_GPU_CHECK_FAILED(::Eigen::gpu::internal::cusolver_status_name(_s), #expr, __FILE__, __LINE__); \
61struct CusolverParams {
62 cusolverDnParams_t p =
nullptr;
64 CusolverParams() { EIGEN_CUSOLVER_CHECK(cusolverDnCreateParams(&p)); }
67 if (p) (void)cusolverDnDestroyParams(p);
70 CusolverParams(CusolverParams&& o) noexcept : p(o.p) { o.p =
nullptr; }
71 CusolverParams& operator=(CusolverParams&& o)
noexcept {
73 if (p) (void)cusolverDnDestroyParams(p);
80 CusolverParams(
const CusolverParams&) =
delete;
81 CusolverParams& operator=(
const CusolverParams&) =
delete;
85struct CusolverHandleDeleter {
87 void operator()(cusolverDnHandle_t h)
const noexcept {
88 if (owns && h) (void)cusolverDnDestroy(h);
91using UniqueCusolverHandle = std::unique_ptr<std::remove_pointer_t<cusolverDnHandle_t>, CusolverHandleDeleter>;
94template <
typename Scalar>
95using cusolver_data_type = cuda_data_type<Scalar>;
100struct cusolver_fill_mode;
103struct cusolver_fill_mode<
Lower> {
104 static constexpr cublasFillMode_t value = CUBLAS_FILL_MODE_LOWER;
107struct cusolver_fill_mode<
Upper> {
108 static constexpr cublasFillMode_t value = CUBLAS_FILL_MODE_UPPER;
113inline cusolverStatus_t cusolverDnXormqr(cusolverDnHandle_t h, cublasSideMode_t side, cublasOperation_t trans,
int m,
114 int n,
int k,
const float* A,
int lda,
const float* tau,
float* C,
int ldc,
115 float* work,
int lwork,
int* info) {
116 return cusolverDnSormqr(h, side, trans, m, n, k, A, lda, tau, C, ldc, work, lwork, info);
118inline cusolverStatus_t cusolverDnXormqr(cusolverDnHandle_t h, cublasSideMode_t side, cublasOperation_t trans,
int m,
119 int n,
int k,
const double* A,
int lda,
const double* tau,
double* C,
int ldc,
120 double* work,
int lwork,
int* info) {
121 return cusolverDnDormqr(h, side, trans, m, n, k, A, lda, tau, C, ldc, work, lwork, info);
123inline cusolverStatus_t cusolverDnXormqr(cusolverDnHandle_t h, cublasSideMode_t side, cublasOperation_t trans,
int m,
124 int n,
int k,
const std::complex<float>* A,
int lda,
125 const std::complex<float>* tau, std::complex<float>* C,
int ldc,
126 std::complex<float>* work,
int lwork,
int* info) {
127 return cusolverDnCunmqr(h, side, trans, m, n, k,
reinterpret_cast<const cuComplex*
>(A), lda,
128 reinterpret_cast<const cuComplex*
>(tau),
reinterpret_cast<cuComplex*
>(C), ldc,
129 reinterpret_cast<cuComplex*
>(work), lwork, info);
131inline cusolverStatus_t cusolverDnXormqr(cusolverDnHandle_t h, cublasSideMode_t side, cublasOperation_t trans,
int m,
132 int n,
int k,
const std::complex<double>* A,
int lda,
133 const std::complex<double>* tau, std::complex<double>* C,
int ldc,
134 std::complex<double>* work,
int lwork,
int* info) {
135 return cusolverDnZunmqr(h, side, trans, m, n, k,
reinterpret_cast<const cuDoubleComplex*
>(A), lda,
136 reinterpret_cast<const cuDoubleComplex*
>(tau),
reinterpret_cast<cuDoubleComplex*
>(C), ldc,
137 reinterpret_cast<cuDoubleComplex*
>(work), lwork, info);
140inline cusolverStatus_t cusolverDnXormqr_bufferSize(cusolverDnHandle_t h, cublasSideMode_t side,
141 cublasOperation_t trans,
int m,
int n,
int k,
const float* A,
142 int lda,
const float* tau,
const float* C,
int ldc,
int* lwork) {
143 return cusolverDnSormqr_bufferSize(h, side, trans, m, n, k, A, lda, tau, C, ldc, lwork);
145inline cusolverStatus_t cusolverDnXormqr_bufferSize(cusolverDnHandle_t h, cublasSideMode_t side,
146 cublasOperation_t trans,
int m,
int n,
int k,
const double* A,
147 int lda,
const double* tau,
const double* C,
int ldc,
int* lwork) {
148 return cusolverDnDormqr_bufferSize(h, side, trans, m, n, k, A, lda, tau, C, ldc, lwork);
150inline cusolverStatus_t cusolverDnXormqr_bufferSize(cusolverDnHandle_t h, cublasSideMode_t side,
151 cublasOperation_t trans,
int m,
int n,
int k,
152 const std::complex<float>* A,
int lda,
153 const std::complex<float>* tau,
const std::complex<float>* C,
154 int ldc,
int* lwork) {
155 return cusolverDnCunmqr_bufferSize(h, side, trans, m, n, k,
reinterpret_cast<const cuComplex*
>(A), lda,
156 reinterpret_cast<const cuComplex*
>(tau),
reinterpret_cast<const cuComplex*
>(C),
159inline cusolverStatus_t cusolverDnXormqr_bufferSize(cusolverDnHandle_t h, cublasSideMode_t side,
160 cublasOperation_t trans,
int m,
int n,
int k,
161 const std::complex<double>* A,
int lda,
162 const std::complex<double>* tau,
const std::complex<double>* C,
163 int ldc,
int* lwork) {
164 return cusolverDnZunmqr_bufferSize(h, side, trans, m, n, k,
reinterpret_cast<const cuDoubleComplex*
>(A), lda,
165 reinterpret_cast<const cuDoubleComplex*
>(tau),
166 reinterpret_cast<const cuDoubleComplex*
>(C), ldc, lwork);
Namespace containing all symbols from the Eigen library.