37#if defined(AMREX_USE_GPU)
47#if defined(AMREX_USE_CUDA)
49 cusparseHandle_t handle;
53 cudaDataType data_type;
54 if constexpr (std::is_same_v<T,float>) {
55 data_type = CUDA_R_32F;
56 }
else if constexpr (std::is_same_v<T,double>) {
57 data_type = CUDA_R_64F;
58 }
else if constexpr (std::is_same_v<T,GpuComplex<float>>) {
59 data_type = CUDA_C_32F;
60 }
else if constexpr (std::is_same_v<T,GpuComplex<double>>) {
61 data_type = CUDA_C_64F;
68 constexpr cusparseIndexType_t index_type = std::is_same_v<I,int>
69 ? CUSPARSE_INDEX_32I : CUSPARSE_INDEX_64I;
71 cusparseSpMatDescr_t mat_descr;
73 (cusparseCreateCsr(&mat_descr, nrows, ncols, nnz,
74 (
void*)row, (
void*)col, (
void*)mat,
75 index_type, index_type, CUSPARSE_INDEX_BASE_ZERO,
78 cusparseDnVecDescr_t x_descr;
81 cusparseDnVecDescr_t y_descr;
87 std::size_t buffer_size;
89 (cusparseSpMV_bufferSize(handle, CUSPARSE_OPERATION_NON_TRANSPOSE,
90 &alpha, mat_descr, x_descr, &
beta, y_descr,
91 data_type, CUSPARSE_SPMV_ALG_DEFAULT,
97 (cusparseSpMV(handle, CUSPARSE_OPERATION_NON_TRANSPOSE,
98 &alpha, mat_descr, x_descr, &
beta, y_descr,
99 data_type, CUSPARSE_SPMV_ALG_DEFAULT, pbuffer));
109#elif defined(AMREX_USE_HIP)
111 rocsparse_handle handle;
112 AMREX_ROCSPARSE_SAFE_CALL(rocsparse_create_handle(&handle));
113 AMREX_ROCSPARSE_SAFE_CALL(rocsparse_set_stream(handle,
Gpu::gpuStream()));
115 rocsparse_datatype data_type;
116 if constexpr (std::is_same_v<T,float>) {
117 data_type = rocsparse_datatype_f32_r;
118 }
else if constexpr (std::is_same_v<T,double>) {
119 data_type = rocsparse_datatype_f64_r;
120 }
else if constexpr (std::is_same_v<T,GpuComplex<float>>) {
121 data_type = rocsparse_datatype_f32_c;
122 }
else if constexpr (std::is_same_v<T,GpuComplex<double>>) {
123 data_type = rocsparse_datatype_f64_c;
128 constexpr rocsparse_indextype index_type = std::is_same_v<I,int>
129 ? rocsparse_indextype_i32 : rocsparse_indextype_i64;
131 rocsparse_spmat_descr mat_descr;
132 AMREX_ROCSPARSE_SAFE_CALL(
133 rocsparse_create_csr_descr(&mat_descr, nrows, ncols, nnz,
134 (
void*)row, (
void*)col, (
void*)mat,
135 index_type, index_type,
136 rocsparse_index_base_zero, data_type));
138 rocsparse_dnvec_descr x_descr;
139 AMREX_ROCSPARSE_SAFE_CALL(
140 rocsparse_create_dnvec_descr(&x_descr, ncols, (
void*)px, data_type));
142 rocsparse_dnvec_descr y_descr;
143 AMREX_ROCSPARSE_SAFE_CALL(
144 rocsparse_create_dnvec_descr(&y_descr, nrows, (
void*)py, data_type));
149#if (HIP_VERSION_MAJOR >= 7)
151 rocsparse_spmv_descr spmv_descr;
152 AMREX_ROCSPARSE_SAFE_CALL(rocsparse_create_spmv_descr(&spmv_descr));
155 rocsparse_error p_error[1] = {};
157 const rocsparse_spmv_alg spmv_alg = rocsparse_spmv_alg_csr_adaptive;
158 AMREX_ROCSPARSE_SAFE_CALL(
159 rocsparse_spmv_set_input(handle, spmv_descr, rocsparse_spmv_input_alg,
160 &spmv_alg,
sizeof(spmv_alg), p_error));
162 const rocsparse_operation spmv_operation = rocsparse_operation_none;
163 AMREX_ROCSPARSE_SAFE_CALL(
164 rocsparse_spmv_set_input(handle, spmv_descr, rocsparse_spmv_input_operation,
165 &spmv_operation,
sizeof(spmv_operation), p_error));
167 AMREX_ROCSPARSE_SAFE_CALL(
168 rocsparse_spmv_set_input(handle, spmv_descr, rocsparse_spmv_input_scalar_datatype,
169 &data_type,
sizeof(data_type), p_error));
171 AMREX_ROCSPARSE_SAFE_CALL(
172 rocsparse_spmv_set_input(handle, spmv_descr, rocsparse_spmv_input_compute_datatype,
173 &data_type,
sizeof(data_type), p_error));
175 std::size_t buffer_size = 0;
176 AMREX_ROCSPARSE_SAFE_CALL(
177 rocsparse_v2_spmv_buffer_size(handle, spmv_descr, mat_descr, x_descr, y_descr,
178 rocsparse_v2_spmv_stage_analysis,
179 &buffer_size, p_error));
181 void* pbuffer =
nullptr;
182 if (buffer_size > 0) {
186 AMREX_ROCSPARSE_SAFE_CALL(
187 rocsparse_v2_spmv(handle, spmv_descr, &alpha, mat_descr, x_descr, &
beta, y_descr,
188 rocsparse_v2_spmv_stage_analysis,
189 buffer_size, pbuffer, p_error));
195 AMREX_ROCSPARSE_SAFE_CALL(
196 rocsparse_v2_spmv_buffer_size(handle, spmv_descr, mat_descr, x_descr, y_descr,
197 rocsparse_v2_spmv_stage_compute,
198 &buffer_size, p_error));
200 if (buffer_size > 0) {
206 AMREX_ROCSPARSE_SAFE_CALL(
207 rocsparse_v2_spmv(handle, spmv_descr, &alpha, mat_descr, x_descr, &
beta, y_descr,
208 rocsparse_v2_spmv_stage_compute,
209 buffer_size, pbuffer, p_error));
211#elif (HIP_VERSION_MAJOR == 6)
213 std::size_t buffer_size = 0;
214 AMREX_ROCSPARSE_SAFE_CALL(
215 rocsparse_spmv(handle, rocsparse_operation_none, &alpha, mat_descr, x_descr,
216 &
beta, y_descr, data_type, rocsparse_spmv_alg_default,
217 rocsparse_spmv_stage_buffer_size,
218 &buffer_size,
nullptr));
220 void* pbuffer =
nullptr;
221 if (buffer_size > 0) {
225 AMREX_ROCSPARSE_SAFE_CALL(
226 rocsparse_spmv(handle, rocsparse_operation_none, &alpha, mat_descr, x_descr,
227 &
beta, y_descr, data_type, rocsparse_spmv_alg_default,
228 rocsparse_spmv_stage_preprocess,
229 &buffer_size, pbuffer));
231 AMREX_ROCSPARSE_SAFE_CALL(
232 rocsparse_spmv(handle, rocsparse_operation_none, &alpha, mat_descr, x_descr,
233 &
beta, y_descr, data_type, rocsparse_spmv_alg_default,
234 rocsparse_spmv_stage_compute,
235 &buffer_size, pbuffer));
239 std::size_t buffer_size = 0;
240 AMREX_ROCSPARSE_SAFE_CALL(
241 rocsparse_spmv(handle, rocsparse_operation_none, &alpha, mat_descr, x_descr,
242 &
beta, y_descr, data_type, rocsparse_spmv_alg_default,
243 &buffer_size,
nullptr));
245 void* pbuffer =
nullptr;
246 if (buffer_size > 0) {
250 AMREX_ROCSPARSE_SAFE_CALL(
251 rocsparse_spmv(handle, rocsparse_operation_none, &alpha, mat_descr, x_descr,
252 &
beta, y_descr, data_type, rocsparse_spmv_alg_default,
253 &buffer_size, pbuffer));
259#if (HIP_VERSION_MAJOR >= 7)
260 AMREX_ROCSPARSE_SAFE_CALL(rocsparse_destroy_error(p_error[0]));
261 AMREX_ROCSPARSE_SAFE_CALL(rocsparse_destroy_spmv_descr(spmv_descr));
263 AMREX_ROCSPARSE_SAFE_CALL(rocsparse_destroy_spmat_descr(mat_descr));
264 AMREX_ROCSPARSE_SAFE_CALL(rocsparse_destroy_dnvec_descr(x_descr));
265 AMREX_ROCSPARSE_SAFE_CALL(rocsparse_destroy_dnvec_descr(y_descr));
266 AMREX_ROCSPARSE_SAFE_CALL(rocsparse_destroy_handle(handle));
271#elif defined(AMREX_USE_SYCL)
273 mkl::sparse::matrix_handle_t handle{};
274 mkl::sparse::init_matrix_handle(&handle);
276#if defined(INTEL_MKL_VERSION) && (INTEL_MKL_VERSION < 20250300)
278 mkl::sparse::set_csr_data(Gpu::Device::streamQueue(), handle, I(nrows), I(ncols),
279 mkl::index_base::zero, (I*)row, (I*)col, (T*)mat);
281 mkl::sparse::set_csr_data(Gpu::Device::streamQueue(), handle, I(nrows), I(ncols),
282 I(nnz), mkl::index_base::zero, (I*)row, (I*)col, (T*)mat);
284 mkl::sparse::gemv(Gpu::Device::streamQueue(), mkl::transpose::nontrans,
285 T(1), handle, px, T(0), py);
287 auto ev = mkl::sparse::release_matrix_handle(Gpu::Device::streamQueue(), &handle);
299#pragma omp parallel for
301 for (
Long i = 0; i < nrows; ++i) {
303 for (
Long j = row[i]; j < row[i+1]; ++j) {
304 r += mat[j] * px[col[j]];