Eigen-Contrib  5.0.1
 
Loading...
Searching...
No Matches
CuSolverSupport.h
1// This file is part of Eigen, a lightweight C++ template library
2// for linear algebra.
3//
4// Copyright (C) 2026 Rasmus Munk Larsen <rmlarsen@gmail.com>
5//
6// This Source Code Form is subject to the terms of the Mozilla
7// Public License v. 2.0. If a copy of the MPL was not distributed
8// with this file, You can obtain one at http://mozilla.org/MPL/2.0/.
9// SPDX-License-Identifier: MPL-2.0
10
11// cuSOLVER-specific support types. Generic CUDA runtime utilities (DeviceBuffer,
12// EIGEN_CUDA_RUNTIME_CHECK) live in GpuSupport.h.
13
14#ifndef EIGEN_GPU_CUSOLVER_SUPPORT_H
15#define EIGEN_GPU_CUSOLVER_SUPPORT_H
16
17// IWYU pragma: private
18#include "./InternalHeaderCheck.h"
19
20#include "./GpuSupport.h"
21#include <cusolverDn.h>
22
23namespace Eigen {
24namespace gpu {
25namespace internal {
26
27// cuSOLVER's public API has no cusolverGetErrorString(), so failure reports would
28// otherwise carry a bare numeric code.
29inline const char* cusolver_status_name(cusolverStatus_t s) {
30 switch (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";
49 default:
50 return "CUSOLVER_STATUS_UNKNOWN";
51 }
52}
53
54#define EIGEN_CUSOLVER_CHECK(expr) \
55 do { \
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__); \
59 } while (0)
60
61struct CusolverParams {
62 cusolverDnParams_t p = nullptr;
63
64 CusolverParams() { EIGEN_CUSOLVER_CHECK(cusolverDnCreateParams(&p)); }
65
66 ~CusolverParams() {
67 if (p) (void)cusolverDnDestroyParams(p); // noexcept context: cannot propagate
68 }
69
70 CusolverParams(CusolverParams&& o) noexcept : p(o.p) { o.p = nullptr; }
71 CusolverParams& operator=(CusolverParams&& o) noexcept {
72 if (this != &o) {
73 if (p) (void)cusolverDnDestroyParams(p); // noexcept context: cannot propagate
74 p = o.p;
75 o.p = nullptr;
76 }
77 return *this;
78 }
79
80 CusolverParams(const CusolverParams&) = delete;
81 CusolverParams& operator=(const CusolverParams&) = delete;
82};
83
84// RAII cuSOLVER dense handle; the ownership flag supports handles borrowed from a gpu::Context.
85struct CusolverHandleDeleter {
86 bool owns = true;
87 void operator()(cusolverDnHandle_t h) const noexcept {
88 if (owns && h) (void)cusolverDnDestroy(h);
89 }
90};
91using UniqueCusolverHandle = std::unique_ptr<std::remove_pointer_t<cusolverDnHandle_t>, CusolverHandleDeleter>;
92
93// Alias kept for compatibility; cuda_data_type<> in GpuSupport.h is canonical.
94template <typename Scalar>
95using cusolver_data_type = cuda_data_type<Scalar>;
96
97// cuSOLVER always interprets the matrix as column-major, so callers must pass the
98// triangle that holds the data in column-major layout.
99template <int UpLo>
100struct cusolver_fill_mode;
101
102template <>
103struct cusolver_fill_mode<Lower> {
104 static constexpr cublasFillMode_t value = CUBLAS_FILL_MODE_LOWER;
105};
106template <>
107struct cusolver_fill_mode<Upper> {
108 static constexpr cublasFillMode_t value = CUBLAS_FILL_MODE_UPPER;
109};
110
111// cuSOLVER ships no generic X variant for ormqr/unmqr, so these overloads supply
112// one: real → ormqr (orthogonal Q), complex → unmqr (unitary Q).
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);
117}
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);
122}
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);
130}
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);
138}
139
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);
144}
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);
149}
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),
157 ldc, lwork);
158}
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);
167}
168
169} // namespace internal
170} // namespace gpu
171} // namespace Eigen
172
173#endif // EIGEN_GPU_CUSOLVER_SUPPORT_H
Namespace containing all symbols from the Eigen library.