Block-Structured AMR Software Framework
Loading...
Searching...
No Matches
AMReX_SpMV.H
Go to the documentation of this file.
1#ifndef AMREX_SPMV_H_
2#define AMREX_SPMV_H_
3#include <AMReX_Config.H>
4
5#include <AMReX_AlgVector.H>
6#include <AMReX_AlgVecUtil.H>
7#include <AMReX_GpuComplex.H>
8#include <AMReX_SpMatrix.H>
9
10namespace amrex {
11
18template <typename T, typename I = Long>
30void SpMV (Long nrows, Long ncols, T* AMREX_RESTRICT py, CsrView<T const,I> const& A,
31 T const* AMREX_RESTRICT px)
32{
33 T const* AMREX_RESTRICT mat = A.mat;
34 auto const* AMREX_RESTRICT col = A.col_index;
35 auto const* AMREX_RESTRICT row = A.row_offset;
36
37#if defined(AMREX_USE_GPU)
38
39 Long const nnz = A.nnz;
40
41 // cuSPARSE 12 leaves y untouched when the matrix has no nonzeros.
42 if (nnz == 0) {
43 ParallelFor(nrows, [=] AMREX_GPU_DEVICE (Long i) { py[i] = T(0); });
44 return;
45 }
46
47#if defined(AMREX_USE_CUDA)
48
49 cusparseHandle_t handle;
50 AMREX_CUSPARSE_SAFE_CALL(cusparseCreate(&handle));
51 AMREX_CUSPARSE_SAFE_CALL(cusparseSetStream(handle, Gpu::gpuStream()));
52
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;
62 } else {
63 amrex::Abort("SpMV: unsupported data type");
64 }
65
66 // cuSPARSE needs the same index type for the row offsets and the
67 // column indices (mixed types came with CUDA 13.3).
68 constexpr cusparseIndexType_t index_type = std::is_same_v<I,int>
69 ? CUSPARSE_INDEX_32I : CUSPARSE_INDEX_64I;
70
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,
76 data_type));
77
78 cusparseDnVecDescr_t x_descr;
79 AMREX_CUSPARSE_SAFE_CALL(cusparseCreateDnVec(&x_descr, ncols, (void*)px, data_type));
80
81 cusparseDnVecDescr_t y_descr;
82 AMREX_CUSPARSE_SAFE_CALL(cusparseCreateDnVec(&y_descr, nrows, (void*)py, data_type));
83
84 T alpha = T(1);
85 T beta = T(0);
86
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,
92 &buffer_size));
93
94 auto* pbuffer = (void*)The_Async_Arena()->alloc(buffer_size);
95
97 (cusparseSpMV(handle, CUSPARSE_OPERATION_NON_TRANSPOSE,
98 &alpha, mat_descr, x_descr, &beta, y_descr,
99 data_type, CUSPARSE_SPMV_ALG_DEFAULT, pbuffer));
100
101 AMREX_CUSPARSE_SAFE_CALL(cusparseDestroySpMat(mat_descr));
102 AMREX_CUSPARSE_SAFE_CALL(cusparseDestroyDnVec(x_descr));
103 // No sync needed: the descriptors are host objects, cusparseDestroy defers
104 // the release of GPU resources, and the buffer is freed in stream order.
105 AMREX_CUSPARSE_SAFE_CALL(cusparseDestroyDnVec(y_descr));
106 AMREX_CUSPARSE_SAFE_CALL(cusparseDestroy(handle));
107 The_Async_Arena()->free(pbuffer);
108
109#elif defined(AMREX_USE_HIP)
110
111 rocsparse_handle handle;
112 AMREX_ROCSPARSE_SAFE_CALL(rocsparse_create_handle(&handle));
113 AMREX_ROCSPARSE_SAFE_CALL(rocsparse_set_stream(handle, Gpu::gpuStream()));
114
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;
124 } else {
125 amrex::Abort("SpMV: unsupported data type");
126 }
127
128 constexpr rocsparse_indextype index_type = std::is_same_v<I,int>
129 ? rocsparse_indextype_i32 : rocsparse_indextype_i64;
130
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));
137
138 rocsparse_dnvec_descr x_descr;
139 AMREX_ROCSPARSE_SAFE_CALL(
140 rocsparse_create_dnvec_descr(&x_descr, ncols, (void*)px, data_type));
141
142 rocsparse_dnvec_descr y_descr;
143 AMREX_ROCSPARSE_SAFE_CALL(
144 rocsparse_create_dnvec_descr(&y_descr, nrows, (void*)py, data_type));
145
146 T alpha = T(1.0);
147 T beta = T(0.0);
148
149#if (HIP_VERSION_MAJOR >= 7)
150
151 rocsparse_spmv_descr spmv_descr;
152 AMREX_ROCSPARSE_SAFE_CALL(rocsparse_create_spmv_descr(&spmv_descr));
153
154
155 rocsparse_error p_error[1] = {};
156
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));
161
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));
166
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));
170
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));
174
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, // buffer size for analysis
179 &buffer_size, p_error));
180
181 void* pbuffer = nullptr;
182 if (buffer_size > 0) {
183 pbuffer = (void*)The_Arena()->alloc(buffer_size);
184 }
185
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, // analysis stage
189 buffer_size, pbuffer, p_error));
190
191 if (pbuffer) {
192 The_Arena()->free(pbuffer);
193 }
194
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, // buffer size for compute
198 &buffer_size, p_error));
199
200 if (buffer_size > 0) {
201 pbuffer = (void*)The_Arena()->alloc(buffer_size);
202 } else {
203 pbuffer = nullptr;
204 }
205
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, // compute stage
209 buffer_size, pbuffer, p_error));
210
211#elif (HIP_VERSION_MAJOR == 6)
212
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, // buffer size stage
218 &buffer_size, nullptr));
219
220 void* pbuffer = nullptr;
221 if (buffer_size > 0) {
222 pbuffer = (void*)The_Arena()->alloc(buffer_size);
223 }
224
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, // preprocess stage
229 &buffer_size, pbuffer));
230
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, // compute stage
235 &buffer_size, pbuffer));
236
237#else /* HIP_VERSION_MAJOR < 6 */
238
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));
244
245 void* pbuffer = nullptr;
246 if (buffer_size > 0) {
247 pbuffer = (void*)The_Arena()->alloc(buffer_size);
248 }
249
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));
254
255#endif /* HIP_VERSION_MAJOR */
256
258
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));
262#endif
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));
267 if (pbuffer) {
268 The_Arena()->free(pbuffer);
269 }
270
271#elif defined(AMREX_USE_SYCL)
272
273 mkl::sparse::matrix_handle_t handle{};
274 mkl::sparse::init_matrix_handle(&handle);
275
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);
280#else
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);
283#endif
284 mkl::sparse::gemv(Gpu::Device::streamQueue(), mkl::transpose::nontrans,
285 T(1), handle, px, T(0), py);
286
287 auto ev = mkl::sparse::release_matrix_handle(Gpu::Device::streamQueue(), &handle);
288 ev.wait();
289
290#endif
291
293
294#else
295
297
298#ifdef AMREX_USE_OMP
299#pragma omp parallel for
300#endif
301 for (Long i = 0; i < nrows; ++i) {
302 T r = 0;
303 for (Long j = row[i]; j < row[i+1]; ++j) {
304 r += mat[j] * px[col[j]];
305 }
306 py[i] = r;
307 }
308
309#endif
310}
311
319template <typename T, template<typename> class AllocM, typename AllocV>
321 AlgVector<T,AllocV> const& x)
322{
323 BL_PROFILE("SpMV");
324
325 auto& Am = const_cast<SpMatrix<T,AllocM>&>(A);
326 Am.setColumnPartition(x.partition());
327 Am.startComm_mv(x);
328
329 // Diagonal part
330 SpMV<T,int>(y.numLocalRows(), x.numLocalRows(), y.data(), A.m_csr_local.const_view(),
331 x.data());
332
333 Am.finishComm_mv(y);
334}
335
344template <typename T, template<typename> class AllocM, typename AllocV>
347{
348 SpMV(res, A, x);
349 Xpay(res, T(-1), b);
350}
351
352}
353
354#endif
#define BL_PROFILE(a)
Definition AMReX_BLProfiler.H:562
#define AMREX_RESTRICT
Definition AMReX_Extension.H:37
#define AMREX_CUSPARSE_SAFE_CALL(call)
Definition AMReX_GpuError.H:101
#define AMREX_GPU_ERROR_CHECK()
Definition AMReX_GpuError.H:151
#define AMREX_GPU_DEVICE
Definition AMReX_GpuQualifiers.H:18
GpuArray< Real, 3 > beta
Definition AMReX_MLEBNodeFDLaplacian.cpp:1834
Distributed dense vector that mirrors the layout of an AlgPartition.
Definition AMReX_AlgVector.H:29
virtual void free(void *pt)=0
Free a previously allocated block pointed to by pt.
virtual void * alloc(std::size_t sz)=0
Allocate sz bytes from this arena.
Distributed CSR matrix that manages storage and GPU-friendly partitions.
Definition AMReX_SpMatrix.H:65
void setColumnPartition(AlgPartition const &col_partition)
Set the column partition and split the matrix into local and remote blocks with 32-bit column indices...
Definition AMReX_SpMatrix.H:1019
amrex_long Long
Definition AMReX_INT.H:30
Arena * The_Async_Arena()
Definition AMReX_Arena.cpp:839
Arena * The_Arena()
Definition AMReX_Arena.cpp:829
void streamSynchronize() noexcept
Definition AMReX_GpuDevice.H:310
gpuStream_t gpuStream() noexcept
Definition AMReX_GpuDevice.H:291
Definition AMReX_Amr.cpp:50
__host__ __device__ void ignore_unused(const Ts &...)
No-op helper that marks variables as intentionally unused.
Definition AMReX.H:273
void ParallelFor(TypeList< CTOs... > ctos, std::array< int, sizeof...(CTOs)> const &runtime_options, T N, F &&f)
Definition AMReX_CTOParallelForImpl.H:202
void computeResidual(AlgVector< T, AllocV > &res, SpMatrix< T, AllocM > const &A, AlgVector< T, AllocV > const &x, AlgVector< T, AllocV > const &b)
Compute the residual res = b - A * x.
Definition AMReX_SpMV.H:345
void Xpay(MF &dst, typename MF::value_type a, MF const &src, int scomp, int dcomp, int ncomp, IntVect const &nghost)
dst = src + a * dst
Definition AMReX_FabArrayUtility.H:2206
void Abort(const std::string &msg)
Print a fatal-error message to stderr and abort execution.
Definition AMReX.cpp:244
void SpMV(Long nrows, Long ncols, T *__restrict__ py, CsrView< T const, I > const &A, T const *__restrict__ px)
Perform y = A * x using CSR data (GPU/CPU aware).
Definition AMReX_SpMV.H:30
CsrView< T const, I > const_view() const
Convenience alias for view() const.
Definition AMReX_CSR.H:97
Lightweight non-owning CSR view that can point to host or device buffers.
Definition AMReX_CSR.H:35
CI *__restrict__ col_index
Definition AMReX_CSR.H:39
T *__restrict__ mat
Definition AMReX_CSR.H:38
CI *__restrict__ row_offset
Definition AMReX_CSR.H:40
Long nnz
Definition AMReX_CSR.H:41