Skip to content

Commit 4faa441

Browse files
committed
Fix: cusolver.h
1 parent 6a0c45a commit 4faa441

1 file changed

Lines changed: 0 additions & 88 deletions

File tree

  • source/source_base/module_container/base/third_party

source/source_base/module_container/base/third_party/cusolver.h

Lines changed: 0 additions & 88 deletions
Original file line numberDiff line numberDiff line change
@@ -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
11124
template <typename T>
11225
static inline
11326
void 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

14153
static inline
14254
void potri (cusolverDnHandle_t& cusolver_handle, const char& uplo, const char& diag, const int& n, float * A, const int& lda)

0 commit comments

Comments
 (0)