@@ -21,93 +21,6 @@ namespace cuSolverConnector {
2121
2222// Legacy trtri using cuSOLVER legacy API (CUDA < 11.4)
2323// The legacy functions (cusolverDnStrtri, cusolverDnDtrtri, etc.) were removed in CUDA 11.4
24- #if CUDA_VERSION < 11040
25- static inline
26- void trtri (cusolverDnHandle_t& cusolver_handle, const char & uplo, const char & diag, const int & n, float * A, const int & lda)
27- {
28- int lwork = 0 ;
29- int h_info = 0 ;
30- int * d_info = nullptr ;
31- float * d_work = nullptr ;
32- cudaErrcheck (cudaMalloc ((void **)&d_info, sizeof (int )));
33-
34- cusolverErrcheck (cusolverDnStrtri_bufferSize (cusolver_handle, cublas_fill_mode (uplo), cublas_diag_type (diag), n, A, lda, &lwork));
35- cudaErrcheck (cudaMalloc ((void **)&d_work, sizeof (float ) * lwork));
36- cusolverErrcheck (cusolverDnStrtri (cusolver_handle, cublas_fill_mode (uplo), cublas_diag_type (diag), n, A, lda, d_work, lwork, d_info));
37-
38- cudaErrcheck (cudaMemcpy (&h_info, d_info, sizeof (int ), cudaMemcpyDeviceToHost));
39- if (h_info != 0 ) {
40- throw std::runtime_error (" trtri: failed to invert matrix" );
41- }
42- cudaErrcheck (cudaFree (d_work));
43- cudaErrcheck (cudaFree (d_info));
44- }
45-
46- static inline
47- void trtri (cusolverDnHandle_t& cusolver_handle, const char & uplo, const char & diag, const int & n, double * A, const int & lda)
48- {
49- int lwork = 0 ;
50- int h_info = 0 ;
51- int * d_info = nullptr ;
52- double * d_work = nullptr ;
53- cudaErrcheck (cudaMalloc ((void **)&d_info, sizeof (int )));
54-
55- cusolverErrcheck (cusolverDnDtrtri_bufferSize (cusolver_handle, cublas_fill_mode (uplo), cublas_diag_type (diag), n, A, lda, &lwork));
56- cudaErrcheck (cudaMalloc ((void **)&d_work, sizeof (double ) * lwork));
57- cusolverErrcheck (cusolverDnDtrtri (cusolver_handle, cublas_fill_mode (uplo), cublas_diag_type (diag), n, A, lda, d_work, lwork, d_info));
58-
59- cudaErrcheck (cudaMemcpy (&h_info, d_info, sizeof (int ), cudaMemcpyDeviceToHost));
60- if (h_info != 0 ) {
61- throw std::runtime_error (" trtri: failed to invert matrix" );
62- }
63- cudaErrcheck (cudaFree (d_work));
64- cudaErrcheck (cudaFree (d_info));
65- }
66-
67- static inline
68- void trtri (cusolverDnHandle_t& cusolver_handle, const char & uplo, const char & diag, const int & n, std::complex <float >* A, const int & lda)
69- {
70- int lwork = 0 ;
71- int h_info = 0 ;
72- int * d_info = nullptr ;
73- cuComplex* d_work = nullptr ;
74- cudaErrcheck (cudaMalloc ((void **)&d_info, sizeof (int )));
75-
76- cusolverErrcheck (cusolverDnCtrtri_bufferSize (cusolver_handle, cublas_fill_mode (uplo), cublas_diag_type (diag), n, reinterpret_cast <cuComplex*>(A), lda, &lwork));
77- cudaErrcheck (cudaMalloc ((void **)&d_work, sizeof (cuComplex) * lwork));
78- cusolverErrcheck (cusolverDnCtrtri (cusolver_handle, cublas_fill_mode (uplo), cublas_diag_type (diag), n, reinterpret_cast <cuComplex*>(A), lda, d_work, lwork, d_info));
79-
80- cudaErrcheck (cudaMemcpy (&h_info, d_info, sizeof (int ), cudaMemcpyDeviceToHost));
81- if (h_info != 0 ) {
82- throw std::runtime_error (" trtri: failed to invert matrix" );
83- }
84- cudaErrcheck (cudaFree (d_work));
85- cudaErrcheck (cudaFree (d_info));
86- }
87-
88- static inline
89- void trtri (cusolverDnHandle_t& cusolver_handle, const char & uplo, const char & diag, const int & n, std::complex <double >* A, const int & lda)
90- {
91- int lwork = 0 ;
92- int h_info = 0 ;
93- int * d_info = nullptr ;
94- cuDoubleComplex* d_work = nullptr ;
95- cudaErrcheck (cudaMalloc ((void **)&d_info, sizeof (int )));
96-
97- cusolverErrcheck (cusolverDnZtrtri_bufferSize (cusolver_handle, cublas_fill_mode (uplo), cublas_diag_type (diag), n, reinterpret_cast <cuDoubleComplex*>(A), lda, &lwork));
98- cudaErrcheck (cudaMalloc ((void **)&d_work, sizeof (cuDoubleComplex) * lwork));
99- cusolverErrcheck (cusolverDnZtrtri (cusolver_handle, cublas_fill_mode (uplo), cublas_diag_type (diag), n, reinterpret_cast <cuDoubleComplex*>(A), lda, d_work, lwork, d_info));
100-
101- cudaErrcheck (cudaMemcpy (&h_info, d_info, sizeof (int ), cudaMemcpyDeviceToHost));
102- if (h_info != 0 ) {
103- throw std::runtime_error (" trtri: failed to invert matrix" );
104- }
105- cudaErrcheck (cudaFree (d_work));
106- cudaErrcheck (cudaFree (d_info));
107- }
108- #else
109- // Generic trtri using cuSOLVER generic API (CUDA >= 11.4)
110- // The generic API (cusolverDnXtrtri) was introduced in CUDA 11.4 to replace the legacy functions
11124template <typename T>
11225static inline
11326void trtri (cusolverDnHandle_t& cusolver_handle, const char & uplo, const char & diag, const int & n, T* A, const int & lda)
@@ -136,7 +49,6 @@ void trtri (cusolverDnHandle_t& cusolver_handle, const char& uplo, const char& d
13649 cudaErrcheck (cudaFree (d_work));
13750 cudaErrcheck (cudaFree (d_info));
13851}
139- #endif
14052
14153static inline
14254void potri (cusolverDnHandle_t& cusolver_handle, const char & uplo, const char & diag, const int & n, float * A, const int & lda)
0 commit comments