Block-Structured AMR Software Framework
Loading...
Searching...
No Matches
AMReX_SpGEMM.H
Go to the documentation of this file.
1#ifndef AMREX_SPGEMM_H_
2#define AMREX_SPGEMM_H_
3#include <AMReX_Config.H>
4
5#include <AMReX_Algorithm.H>
6#include <AMReX_BLProfiler.H>
7#include <AMReX_Exception.H>
8#include <AMReX_GpuComplex.H>
9#include <AMReX_OpenMP.H>
10#include <AMReX_SpMatrix.H>
11
12#include <algorithm>
13#include <limits>
14#include <type_traits>
15#include <utility>
16
17namespace amrex::detail {
18
20struct NoRowPost {
21 template <typename I, typename T>
22 Long operator() (Long, I*, T*, Long n) const { return n; }
23};
24
26template <typename T, template<typename> class V, typename I = Long>
27CSR<T,V,I> spgemm_empty (Long nrows)
28{
29 CSR<T,V,I> C;
30 C.row_offset.resize(nrows+1);
31 auto* p = C.row_offset.data();
32 ParallelForOMP(nrows+1, [=] AMREX_GPU_DEVICE (Long i) { p[i] = 0; });
34 return C;
35}
36
37#if !defined(AMREX_USE_GPU)
38
39// Sorts the n entries of a row by column; `tmp` is scratch for long rows.
40template <typename T, typename I>
41void sort_row_cpu (I* AMREX_RESTRICT col, T* AMREX_RESTRICT val, Long n,
42 Vector<std::pair<I,T>>& tmp)
43{
44 if (n <= 32) {
45 for (Long k = 1; k < n; ++k) {
46 I const c = col[k];
47 T const v = val[k];
48 Long m = k;
49 for (; m > 0 && col[m-1] > c; --m) {
50 col[m] = col[m-1];
51 val[m] = val[m-1];
52 }
53 col[m] = c;
54 val[m] = v;
55 }
56 } else {
57 tmp.resize(n);
58 for (Long k = 0; k < n; ++k) { tmp[k] = {col[k], val[k]}; }
59 std::sort(tmp.begin(), tmp.end(),
60 [] (auto const& x, auto const& y) { return x.first < y.first; });
61 for (Long k = 0; k < n; ++k) { col[k] = tmp[k].first; val[k] = tmp[k].second; }
62 }
63}
64
75template <typename T, template<typename> class V, typename I, typename W, typename F>
76CSR<T,V,I> csr_from_rows_cpu (Long nrows, W const& work, F const& make_row)
77{
78 CSR<T,V,I> C;
79 C.row_offset.resize(nrows+1);
80 I* AMREX_RESTRICT crow = C.row_offset.data();
81 AMREX_ASSUME(crow != nullptr); // gcc -Wnull-dereference false positive
82 crow[0] = 0;
83
84 int const nblocks = int(std::max(Long(1), std::min(Long(OpenMP::get_max_threads()), nrows)));
85 // Built with push_back: gcc -Wnull-dereference flags writes through data().
86 Vector<Long> rbegin{0};
87 rbegin.reserve(nblocks+1);
88 if (nblocks > 1) {
89 Vector<Long> wsum{0};
90 wsum.reserve(nrows+1);
91 for (Long i = 0; i < nrows; ++i) { wsum.push_back(wsum.back() + work(i)); }
92 for (int t = 1; t < nblocks; ++t) {
93 rbegin.push_back(std::lower_bound(wsum.begin(), wsum.end(), wsum.back()/nblocks*t)
94 - wsum.begin());
95 }
96 }
97 rbegin.push_back(nrows);
98
99 if (nblocks == 1) {
100 // Rows go straight into the result. The capacity is a guess from
101 // `work`; pages that are never written are never touched.
102 Long cap = 0;
103 for (Long i = 0; i < nrows; ++i) { cap += work(i); }
104 C.col_index.resize(cap);
105 C.mat.resize(cap);
106 auto row = make_row(0);
107 Long p = 0;
108 for (Long i = 0; i < nrows; ++i) {
109 I const* rc = nullptr;
110 T const* rv = nullptr;
111 Long const n = row(i, rc, rv);
112 if (p + n > Long(C.col_index.size())) {
113 Long const newcap = std::max(p + n, Long(C.col_index.size())*3/2);
114 C.col_index.resize(p); // only the filled part is copied
115 C.mat.resize(p);
116 C.col_index.resize(newcap);
117 C.mat.resize(newcap);
118 }
119 std::copy(rc, rc+n, C.col_index.data()+p);
120 std::copy(rv, rv+n, C.mat.data()+p);
121 p += n;
122 AMREX_ALWAYS_ASSERT_WITH_MESSAGE(p <= Long(std::numeric_limits<I>::max()),
123 "SpGEMM: too many nonzeros for the index type");
124 crow[i+1] = I(p);
125 }
126 C.col_index.resize(p);
127 C.mat.resize(p);
128 C.nnz = p;
129 return C;
130 }
131
132 struct Buffer {
133 V<I> col;
134 V<T> val;
135 Long n = 0;
136 };
137 Vector<Vector<Buffer>> buffers(nblocks);
138 constexpr Long buffer_size = Long(1) << 20;
139
140#ifdef AMREX_USE_OMP
141#pragma omp parallel for schedule(static,1)
142#endif
143 for (int t = 0; t < nblocks; ++t) {
144 auto row = make_row(t);
145 auto& bufs = buffers[t];
146 for (Long i = rbegin[t]; i < rbegin[t+1]; ++i) {
147 I const* rc = nullptr;
148 T const* rv = nullptr;
149 Long const n = row(i, rc, rv);
150 if (bufs.empty() || bufs.back().n + n > Long(bufs.back().col.size())) {
151 bufs.emplace_back();
152 bufs.back().col.resize(std::max(buffer_size, n));
153 bufs.back().val.resize(std::max(buffer_size, n));
154 }
155 auto& buf = bufs.back();
156 std::copy(rc, rc+n, buf.col.data()+buf.n);
157 std::copy(rv, rv+n, buf.val.data()+buf.n);
158 buf.n += n;
159 crow[i+1] = I(n);
160 }
161 }
162
163 // Exclusive scan in Long: the result may not fit the index type.
164 Long total = 0;
165 for (Long i = 0; i < nrows; ++i) {
166 Long const cnt = crow[i+1];
167 crow[i] = I(total);
168 total += cnt;
169 }
170 AMREX_ALWAYS_ASSERT_WITH_MESSAGE(total <= Long(std::numeric_limits<I>::max()),
171 "SpGEMM: too many nonzeros for the index type");
172 crow[nrows] = I(total);
173 C.nnz = total;
174 C.col_index.resize(C.nnz);
175 C.mat.resize(C.nnz);
176 I* AMREX_RESTRICT ccol = C.col_index.data();
177 T* AMREX_RESTRICT cmat = C.mat.data();
178
179#ifdef AMREX_USE_OMP
180#pragma omp parallel for schedule(static,1)
181#endif
182 for (int t = 0; t < nblocks; ++t) {
183 Long p = crow[rbegin[t]];
184 for (auto& buf : buffers[t]) {
185 std::copy(buf.col.data(), buf.col.data()+buf.n, ccol+p);
186 std::copy(buf.val.data(), buf.val.data()+buf.n, cmat+p);
187 p += buf.n;
188 buf = Buffer{};
189 }
190 }
191
192 return C;
193}
194
195// Gustavson's algorithm with a dense marker holding the position of each
196// column in the current row. Output rows are sorted, then passed to
197// `row_post(i, col, val, n)`, which may change them in place and returns
198// the new length.
199template <typename T, template<typename> class V, typename I, typename F = NoRowPost>
200CSR<T,V,I> spgemm_local_cpu (Long nrows, Long ncols,
201 CsrView<T const,I> const& A, CsrView<T const,I> const& B,
202 F const& row_post = {})
203{
204 auto nprod = [&] (Long i) {
205 Long n = 0;
206 for (Long ap = A.row_offset[i]; ap < A.row_offset[i+1]; ++ap) {
207 Long const k = A.col_index[ap];
208 n += B.row_offset[k+1] - B.row_offset[k];
209 }
210 return n;
211 };
212 return csr_from_rows_cpu<T,V,I>(nrows, nprod, [&] (int) {
213 return [&, marker = Vector<I>(ncols, I(-1)), col = Vector<I>(), val = Vector<T>(),
214 tmp = Vector<std::pair<I,T>>()]
215 (Long i, I const*& rc, T const*& rv) mutable -> Long
216 {
217 Long const maxlen = std::min(nprod(i), ncols);
218 if (Long(col.size()) < maxlen) {
219 col.resize(maxlen);
220 val.resize(maxlen);
221 }
222 Long n = 0;
223 for (Long ap = A.row_offset[i]; ap < A.row_offset[i+1]; ++ap) {
224 Long const k = A.col_index[ap];
225 T const a = A.mat[ap];
226 for (Long bp = B.row_offset[k]; bp < B.row_offset[k+1]; ++bp) {
227 I const j = B.col_index[bp];
228 Long const m = marker[j]; // may be left over from an earlier row
229 if (m >= 0 && m < n && col[m] == j) {
230 val[m] += a * B.mat[bp];
231 } else {
232 marker[j] = I(n);
233 col[n] = j;
234 val[n] = a * B.mat[bp];
235 ++n;
236 }
237 }
238 }
239 sort_row_cpu(col.data(), val.data(), n, tmp);
240 n = row_post(i, col.data(), val.data(), n);
241 rc = col.data();
242 rv = val.data();
243 return n;
244 };
245 });
246}
247
248#ifdef AMREX_USE_MPI
259template <typename T, template<typename> class V, typename I, typename F>
260CSR<T,V,I> spgemm_split_cpu (Long nrows, Long ncols, Long nb,
261 CsrView<T const,I> const& A0, CsrView<T const,I> const& A1,
262 CsrView<T const,I> const& B0, CsrView<T const,I> const& B1,
263 I const* b1map, CsrView<T const,I> const& E, F const& row_post)
264{
265 auto blen = [&] (Long k) -> Long {
266 if (k < nb) {
267 return (B0.row_offset[k+1] - B0.row_offset[k])
268 + ((B1.nnz > 0) ? (B1.row_offset[k+1] - B1.row_offset[k]) : 0);
269 } else {
270 return E.row_offset[k-nb+1] - E.row_offset[k-nb];
271 }
272 };
273 auto nprod = [&] (Long i) {
274 Long n = 0;
275 for (Long ap = A0.row_offset[i]; ap < A0.row_offset[i+1]; ++ap) { n += blen(A0.col_index[ap]); }
276 if (A1.nnz > 0) {
277 for (Long ap = A1.row_offset[i]; ap < A1.row_offset[i+1]; ++ap) { n += blen(nb + A1.col_index[ap]); }
278 }
279 return n;
280 };
281 return csr_from_rows_cpu<T,V,I>(nrows, nprod, [&] (int) {
282 return [&, marker = Vector<I>(ncols, I(-1)), col = Vector<I>(), val = Vector<T>(),
283 tmp = Vector<std::pair<I,T>>()]
284 (Long i, I const*& rc, T const*& rv) mutable -> Long
285 {
286 Long const maxlen = std::min(nprod(i), ncols);
287 if (Long(col.size()) < maxlen) {
288 col.resize(maxlen);
289 val.resize(maxlen);
290 }
291 Long n = 0;
292 auto add = [&] (I j, T x) {
293 Long const m = marker[j]; // may be left over from an earlier row
294 if (m >= 0 && m < n && col[m] == j) {
295 val[m] += x;
296 } else {
297 marker[j] = I(n);
298 col[n] = j;
299 val[n] = x;
300 ++n;
301 }
302 };
303 auto brow = [&] (Long k, T a) {
304 if (k < nb) {
305 for (Long bp = B0.row_offset[k]; bp < B0.row_offset[k+1]; ++bp) {
306 add(B0.col_index[bp], a * B0.mat[bp]);
307 }
308 if (B1.nnz > 0) {
309 for (Long bp = B1.row_offset[k]; bp < B1.row_offset[k+1]; ++bp) {
310 add(b1map[B1.col_index[bp]], a * B1.mat[bp]);
311 }
312 }
313 } else {
314 for (Long bp = E.row_offset[k-nb]; bp < E.row_offset[k-nb+1]; ++bp) {
315 add(E.col_index[bp], a * E.mat[bp]);
316 }
317 }
318 };
319 for (Long ap = A0.row_offset[i]; ap < A0.row_offset[i+1]; ++ap) {
320 brow(A0.col_index[ap], A0.mat[ap]);
321 }
322 if (A1.nnz > 0) {
323 for (Long ap = A1.row_offset[i]; ap < A1.row_offset[i+1]; ++ap) {
324 brow(nb + A1.col_index[ap], A1.mat[ap]);
325 }
326 }
327 sort_row_cpu(col.data(), val.data(), n, tmp);
328 n = row_post(i, col.data(), val.data(), n);
329 rc = col.data();
330 rv = val.data();
331 return n;
332 };
333 });
334}
335#endif
336
346template <typename T, typename I>
347struct APRowsCpu
348{
349 CsrView<T const,I> A0, A1, P0, P1, PE;
350 I const* p1map = nullptr;
351
352 [[nodiscard]] Long plen (Long k) const {
353 return (P0.row_offset[k+1] - P0.row_offset[k])
354 + ((P1.nnz > 0) ? (P1.row_offset[k+1] - P1.row_offset[k]) : 0);
355 }
356
358 [[nodiscard]] Long maxlen (Long i) const {
359 Long n = 0;
360 for (Long ap = A0.row_offset[i]; ap < A0.row_offset[i+1]; ++ap) { n += plen(A0.col_index[ap]); }
361 if (A1.nnz > 0) {
362 for (Long ap = A1.row_offset[i]; ap < A1.row_offset[i+1]; ++ap) {
363 Long const k = A1.col_index[ap];
364 n += PE.row_offset[k+1] - PE.row_offset[k];
365 }
366 }
367 return n;
368 }
369
372 Long row (Long i, I* marker, I* col, T* val, Long p) const {
373 Long const rs = p;
374 auto add = [&] (I j, T x) {
375 Long const m = marker[j]; // may be left over from elsewhere
376 if (m >= rs && m < p && col[m] == j) {
377 val[m] += x;
378 } else {
379 marker[j] = I(p);
380 col[p] = j;
381 val[p] = x;
382 ++p;
383 }
384 };
385 for (Long ap = A0.row_offset[i]; ap < A0.row_offset[i+1]; ++ap) {
386 Long const k = A0.col_index[ap];
387 T const a = A0.mat[ap];
388 for (Long pp = P0.row_offset[k]; pp < P0.row_offset[k+1]; ++pp) {
389 add(P0.col_index[pp], a * P0.mat[pp]);
390 }
391 if (P1.nnz > 0) {
392 for (Long pp = P1.row_offset[k]; pp < P1.row_offset[k+1]; ++pp) {
393 add(p1map[P1.col_index[pp]], a * P1.mat[pp]);
394 }
395 }
396 }
397 if (A1.nnz > 0) {
398 for (Long ap = A1.row_offset[i]; ap < A1.row_offset[i+1]; ++ap) {
399 Long const k = A1.col_index[ap];
400 T const a = A1.mat[ap];
401 for (Long pp = PE.row_offset[k]; pp < PE.row_offset[k+1]; ++pp) {
402 add(PE.col_index[pp], a * PE.mat[pp]);
403 }
404 }
405 }
406 return p;
407 }
408};
409
420template <typename T, template<typename> class V, typename I>
421CSR<T,V,I> rap_local_cpu (Long nf, Long nc, Long ncols,
422 CsrView<T const,I> const& R0, CsrView<T const,I> const& R1,
423 CsrView<T const,I> const& APE, APRowsCpu<T,I> const& ap)
424{
425 constexpr Long chunk_rows = 1024;
426 Long const nchunks = (nf + chunk_rows - 1) / chunk_rows;
427
428 // About the length of row c: also the capacity guess of csr_from_rows_cpu.
429 auto work = [&] (Long c) {
430 Long n = 0;
431 for (Long rp = R0.row_offset[c]; rp < R0.row_offset[c+1]; ++rp) {
432 Long const i = R0.col_index[rp];
433 n += ap.A0.row_offset[i+1] - ap.A0.row_offset[i];
434 if (ap.A1.nnz > 0) { n += ap.A1.row_offset[i+1] - ap.A1.row_offset[i]; }
435 }
436 if (R1.nnz > 0) {
437 for (Long rp = R1.row_offset[c]; rp < R1.row_offset[c+1]; ++rp) {
438 Long const j = R1.col_index[rp];
439 n += APE.row_offset[j+1] - APE.row_offset[j];
440 }
441 }
442 return n;
443 };
444
445 // Last row of R that needs a chunk: the largest local column of P in it.
446 Vector<Long> chunk_last(nchunks, -1);
447 for (Long q = 0; q < nchunks; ++q) {
448 for (Long i = q*chunk_rows; i < std::min(nf, (q+1)*chunk_rows); ++i) {
449 if (ap.P0.row_offset[i+1] > ap.P0.row_offset[i]) {
450 chunk_last[q] = std::max(chunk_last[q],
451 Long(ap.P0.col_index[ap.P0.row_offset[i+1]-1]));
452 }
453 }
454 }
455
456 struct Chunk {
457 Vector<Long> off;
458 Vector<I> col;
459 Vector<T> val;
460 };
461
462 return csr_from_rows_cpu<T,V,I>(nc, work, [&] (int) {
463 return [&, cmarker = Vector<I>(ncols, I(-1)), marker = Vector<I>(ncols, I(-1)),
464 slot = Vector<int>(nchunks, -1), chunks = Vector<Chunk>(),
465 live = Vector<Long>(), free_slots = Vector<int>(),
466 next_release = std::numeric_limits<Long>::max(),
467 col = Vector<I>(), val = Vector<T>(), tmp = Vector<std::pair<I,T>>()]
468 (Long c, I const*& rc, T const*& rv) mutable -> Long
469 {
470 auto compute_chunk = [&] (Long q) {
471 int s;
472 if (free_slots.empty()) {
473 s = int(chunks.size());
474 chunks.emplace_back();
475 } else {
476 s = free_slots.back();
477 free_slots.pop_back();
478 }
479 slot[q] = s;
480 live.push_back(q);
481 next_release = std::min(next_release, chunk_last[q]);
482 auto& ch = chunks[s];
483 Long const i0 = q*chunk_rows;
484 Long const i1 = std::min(nf, i0+chunk_rows);
485 ch.off.resize(i1-i0+1);
486 ch.off[0] = 0;
487 Long p = 0;
488 for (Long i = i0; i < i1; ++i) {
489 Long const maxlen = ap.maxlen(i);
490 if (Long(ch.col.size()) < p + maxlen) {
491 ch.col.resize(std::max(p + maxlen, Long(ch.col.size())*2));
492 ch.val.resize(ch.col.size());
493 }
494 p = ap.row(i, cmarker.data(), ch.col.data(), ch.val.data(), p);
495 ch.off[i-i0+1] = p;
496 }
497 };
498
499 // Chunks no later row needs are recycled.
500 if (c > next_release) {
501 next_release = std::numeric_limits<Long>::max();
502 for (Long pos = 0; pos < Long(live.size()); ) {
503 Long const q = live[pos];
504 if (chunk_last[q] < c) {
505 free_slots.push_back(slot[q]);
506 slot[q] = -1;
507 live[pos] = live.back();
508 live.pop_back();
509 } else {
510 next_release = std::min(next_release, chunk_last[q]);
511 ++pos;
512 }
513 }
514 }
515
516 Long maxlen = 0;
517 for (Long rp = R0.row_offset[c]; rp < R0.row_offset[c+1]; ++rp) {
518 Long const i = R0.col_index[rp];
519 Long const q = i / chunk_rows;
520 if (slot[q] < 0) { compute_chunk(q); }
521 auto const& ch = chunks[slot[q]];
522 maxlen += ch.off[i-q*chunk_rows+1] - ch.off[i-q*chunk_rows];
523 }
524 if (R1.nnz > 0) {
525 for (Long rp = R1.row_offset[c]; rp < R1.row_offset[c+1]; ++rp) {
526 Long const j = R1.col_index[rp];
527 maxlen += APE.row_offset[j+1] - APE.row_offset[j];
528 }
529 }
530 maxlen = std::min(maxlen, ncols);
531 if (Long(col.size()) < maxlen) {
532 col.resize(maxlen);
533 val.resize(maxlen);
534 }
535
536 Long n = 0;
537 auto add = [&] (I j, T x) {
538 Long const m = marker[j]; // may be left over from an earlier row
539 if (m >= 0 && m < n && col[m] == j) {
540 val[m] += x;
541 } else {
542 marker[j] = I(n);
543 col[n] = j;
544 val[n] = x;
545 ++n;
546 }
547 };
548 for (Long rp = R0.row_offset[c]; rp < R0.row_offset[c+1]; ++rp) {
549 Long const i = R0.col_index[rp];
550 T const r = R0.mat[rp];
551 Long const q = i / chunk_rows;
552 auto const& ch = chunks[slot[q]];
553 for (Long x = ch.off[i-q*chunk_rows]; x < ch.off[i-q*chunk_rows+1]; ++x) {
554 add(ch.col[x], r * ch.val[x]);
555 }
556 }
557 if (R1.nnz > 0) {
558 for (Long rp = R1.row_offset[c]; rp < R1.row_offset[c+1]; ++rp) {
559 Long const j = R1.col_index[rp];
560 T const r = R1.mat[rp];
561 for (Long x = APE.row_offset[j]; x < APE.row_offset[j+1]; ++x) {
562 add(APE.col_index[x], r * APE.mat[x]);
563 }
564 }
565 }
566 sort_row_cpu(col.data(), val.data(), n, tmp);
567 rc = col.data();
568 rv = val.data();
569 return n;
570 };
571 });
572}
573
574#elif defined(AMREX_USE_CUDA)
575
576inline void spgemm_cusparse_check (cusparseStatus_t status)
577{
578 if (status == CUSPARSE_STATUS_ALLOC_FAILED ||
579 status == CUSPARSE_STATUS_INSUFFICIENT_RESOURCES) {
580 (void)cudaGetLastError(); // clear a failed internal cudaMalloc
581 throw OutOfMemoryError("SpGEMM: cuSPARSE ran out of memory");
582 }
584}
585
586template <typename T, template<typename> class V, typename I>
587CSR<T,V,I> spgemm_local_cusparse (Long nrows, Long ncols,
588 CsrView<T const,I> const& A, CsrView<T const,I> const& B)
589{
590 static_assert(std::is_same_v<I,int>, "spgemm_local_cusparse: 32-bit indices only");
591
592 cudaDataType data_type;
593 if constexpr (std::is_same_v<T,float>) {
594 data_type = CUDA_R_32F;
595 } else if constexpr (std::is_same_v<T,double>) {
596 data_type = CUDA_R_64F;
597 } else if constexpr (std::is_same_v<T,GpuComplex<float>>) {
598 data_type = CUDA_C_32F;
599 } else if constexpr (std::is_same_v<T,GpuComplex<double>>) {
600 data_type = CUDA_C_64F;
601 } else {
602 amrex::Abort("SpGEMM: unsupported data type");
603 }
604
605 AMREX_ALWAYS_ASSERT(nrows < Long(std::numeric_limits<int>::max()) &&
606 ncols < Long(std::numeric_limits<int>::max()));
607
608 CSR<T,V,I> C;
609 C.row_offset.resize(nrows+1);
610
611 // Releases everything, also when an OutOfMemoryError is thrown.
612 struct Resources {
613 cusparseHandle_t handle = nullptr;
614 cusparseSpMatDescr_t mat_A = nullptr, mat_B = nullptr, mat_C = nullptr;
615 cusparseSpGEMMDescr_t descr = nullptr;
616 void* buffer1 = nullptr;
617 void* buffer2 = nullptr;
618 Resources () = default;
619 Resources (Resources const&) = delete;
620 Resources& operator= (Resources const&) = delete;
621 ~Resources () {
622 Gpu::streamSynchronize();
623 if (descr) { cusparseSpGEMM_destroyDescr(descr); }
624 if (mat_A) { cusparseDestroySpMat(mat_A); }
625 if (mat_B) { cusparseDestroySpMat(mat_B); }
626 if (mat_C) { cusparseDestroySpMat(mat_C); }
627 if (handle) { cusparseDestroy(handle); }
628 if (buffer1) { The_Arena()->free(buffer1); }
629 if (buffer2) { The_Arena()->free(buffer2); }
630 }
631 } r;
632
633 AMREX_CUSPARSE_SAFE_CALL(cusparseCreate(&r.handle));
634 AMREX_CUSPARSE_SAFE_CALL(cusparseSetStream(r.handle, Gpu::gpuStream()));
635
636 constexpr cusparseIndexType_t index_type = CUSPARSE_INDEX_32I;
637 void* rowA = (void*)A.row_offset;
638 void* colA = (void*)A.col_index;
639 void* rowB = (void*)B.row_offset;
640 void* colB = (void*)B.col_index;
641 void* rowC = (void*)C.row_offset.data();
642
644 (cusparseCreateCsr(&r.mat_A, nrows, B.nrows, A.nnz, rowA, colA, (void*)A.mat,
645 index_type, index_type, CUSPARSE_INDEX_BASE_ZERO, data_type));
647 (cusparseCreateCsr(&r.mat_B, B.nrows, ncols, B.nnz, rowB, colB, (void*)B.mat,
648 index_type, index_type, CUSPARSE_INDEX_BASE_ZERO, data_type));
650 (cusparseCreateCsr(&r.mat_C, nrows, ncols, 0, rowC, nullptr, nullptr,
651 index_type, index_type, CUSPARSE_INDEX_BASE_ZERO, data_type));
652
653 AMREX_CUSPARSE_SAFE_CALL(cusparseSpGEMM_createDescr(&r.descr));
654
655 T alpha = T(1);
656 T beta = T(0);
657 cusparseOperation_t op = CUSPARSE_OPERATION_NON_TRANSPOSE;
658 auto const alg = CUSPARSE_SPGEMM_DEFAULT;
659
660 std::size_t buffer_size1 = 0;
661 spgemm_cusparse_check
662 (cusparseSpGEMM_workEstimation(r.handle, op, op, &alpha, r.mat_A, r.mat_B, &beta,
663 r.mat_C, data_type, alg, r.descr,
664 &buffer_size1, nullptr));
665 r.buffer1 = (void*)The_Arena()->alloc(buffer_size1);
666 spgemm_cusparse_check
667 (cusparseSpGEMM_workEstimation(r.handle, op, op, &alpha, r.mat_A, r.mat_B, &beta,
668 r.mat_C, data_type, alg, r.descr,
669 &buffer_size1, r.buffer1));
670
671 std::size_t buffer_size2 = 0;
672 spgemm_cusparse_check
673 (cusparseSpGEMM_compute(r.handle, op, op, &alpha, r.mat_A, r.mat_B, &beta,
674 r.mat_C, data_type, alg, r.descr,
675 &buffer_size2, nullptr));
676 r.buffer2 = (void*)The_Arena()->alloc(buffer_size2);
677 spgemm_cusparse_check
678 (cusparseSpGEMM_compute(r.handle, op, op, &alpha, r.mat_A, r.mat_B, &beta,
679 r.mat_C, data_type, alg, r.descr,
680 &buffer_size2, r.buffer2));
681
682 std::int64_t c_nrows, c_ncols, c_nnz;
683 AMREX_CUSPARSE_SAFE_CALL(cusparseSpMatGetSize(r.mat_C, &c_nrows, &c_ncols, &c_nnz));
684 AMREX_ALWAYS_ASSERT(c_nrows == nrows && c_ncols == ncols);
685 AMREX_ALWAYS_ASSERT(c_nnz < std::int64_t(std::numeric_limits<int>::max()));
686
687 C.mat.resize(c_nnz);
688 C.col_index.resize(c_nnz);
689 C.nnz = c_nnz;
690 void* colC = (void*)C.col_index.data();
691 AMREX_CUSPARSE_SAFE_CALL(cusparseCsrSetPointers(r.mat_C, rowC, colC, (void*)C.mat.data()));
692
693 // cuSPARSE guarantees sorted column indices in C.
695 (cusparseSpGEMM_copy(r.handle, op, op, &alpha, r.mat_A, r.mat_B, &beta, r.mat_C,
696 data_type, alg, r.descr));
697
698 return C;
699}
700
701#elif defined(AMREX_USE_HIP)
702
703inline void spgemm_rocsparse_check (rocsparse_status status)
704{
705 if (status == rocsparse_status_memory_error) {
706 (void)hipGetLastError(); // clear a failed internal hipMalloc
707 throw OutOfMemoryError("SpGEMM: rocSPARSE ran out of memory");
708 }
709 AMREX_ROCSPARSE_SAFE_CALL(status);
710}
711
712template <typename T, template<typename> class V, typename I>
713CSR<T,V,I> spgemm_local_rocsparse (Long nrows, Long ncols,
714 CsrView<T const,I> const& A, CsrView<T const,I> const& B)
715{
716 static_assert(std::is_same_v<I,int>, "spgemm_local_rocsparse: 32-bit indices only");
717
718 AMREX_ALWAYS_ASSERT(nrows < Long(std::numeric_limits<int>::max()) &&
719 ncols < Long(std::numeric_limits<int>::max()));
720
721 rocsparse_datatype data_type;
722 if constexpr (std::is_same_v<T,float>) {
723 data_type = rocsparse_datatype_f32_r;
724 } else if constexpr (std::is_same_v<T,double>) {
725 data_type = rocsparse_datatype_f64_r;
726 } else if constexpr (std::is_same_v<T,GpuComplex<float>>) {
727 data_type = rocsparse_datatype_f32_c;
728 } else if constexpr (std::is_same_v<T,GpuComplex<double>>) {
729 data_type = rocsparse_datatype_f64_c;
730 } else {
731 amrex::Abort("SpGEMM: unsupported data type");
732 }
733
734 constexpr rocsparse_indextype index_type = rocsparse_indextype_i32;
735 constexpr rocsparse_index_base index_base = rocsparse_index_base_zero;
736
737 CSR<T,V,I> C;
738 C.row_offset.resize(nrows+1);
739
740 // Releases everything, also when an OutOfMemoryError is thrown.
741 struct Resources {
742 rocsparse_handle handle = nullptr;
743 rocsparse_spmat_descr mat_A = nullptr, mat_B = nullptr, mat_C = nullptr,
744 mat_D = nullptr;
745 void* buffer = nullptr;
746 Resources () = default;
747 Resources (Resources const&) = delete;
748 Resources& operator= (Resources const&) = delete;
749 ~Resources () {
750 Gpu::streamSynchronize();
751 for (auto m : {mat_A, mat_B, mat_C, mat_D}) {
752 if (m) { rocsparse_destroy_spmat_descr(m); }
753 }
754 if (handle) { rocsparse_destroy_handle(handle); }
755 if (buffer) { The_Arena()->free(buffer); }
756 }
757 } r;
758
759 AMREX_ROCSPARSE_SAFE_CALL(rocsparse_create_handle(&r.handle));
760 AMREX_ROCSPARSE_SAFE_CALL(rocsparse_set_stream(r.handle, Gpu::gpuStream()));
761
762 AMREX_ROCSPARSE_SAFE_CALL
763 (rocsparse_create_csr_descr(&r.mat_A, nrows, B.nrows, A.nnz,
764 (void*)A.row_offset, (void*)A.col_index, (void*)A.mat,
765 index_type, index_type, index_base, data_type));
766 AMREX_ROCSPARSE_SAFE_CALL
767 (rocsparse_create_csr_descr(&r.mat_B, B.nrows, ncols, B.nnz,
768 (void*)B.row_offset, (void*)B.col_index, (void*)B.mat,
769 index_type, index_type, index_base, data_type));
770 AMREX_ROCSPARSE_SAFE_CALL
771 (rocsparse_create_csr_descr(&r.mat_C, nrows, ncols, 0,
772 (void*)C.row_offset.data(), nullptr, nullptr,
773 index_type, index_type, index_base, data_type));
774 // D is unused because beta is zero, but a valid descriptor is required.
775 AMREX_ROCSPARSE_SAFE_CALL
776 (rocsparse_create_csr_descr(&r.mat_D, 0, 0, 0, nullptr, nullptr, nullptr,
777 index_type, index_type, index_base, data_type));
778
779 T alpha = T(1);
780 T beta = T(0);
781 auto const op = rocsparse_operation_none;
782 auto const alg = rocsparse_spgemm_alg_default;
783
784 std::size_t buffer_size = 0;
785 spgemm_rocsparse_check
786 (rocsparse_spgemm(r.handle, op, op, &alpha, r.mat_A, r.mat_B, &beta, r.mat_D, r.mat_C,
787 data_type, alg, rocsparse_spgemm_stage_buffer_size,
788 &buffer_size, nullptr));
789 r.buffer = (void*)The_Arena()->alloc(buffer_size);
790
791 // This stage fills C's row offsets.
792 spgemm_rocsparse_check
793 (rocsparse_spgemm(r.handle, op, op, &alpha, r.mat_A, r.mat_B, &beta, r.mat_D, r.mat_C,
794 data_type, alg, rocsparse_spgemm_stage_nnz,
795 &buffer_size, r.buffer));
796
797 std::int64_t c_nrows, c_ncols, c_nnz;
798 AMREX_ROCSPARSE_SAFE_CALL(rocsparse_spmat_get_size(r.mat_C, &c_nrows, &c_ncols, &c_nnz));
799 AMREX_ALWAYS_ASSERT(c_nrows == nrows && c_ncols == ncols);
800 AMREX_ALWAYS_ASSERT(c_nnz < std::int64_t(std::numeric_limits<int>::max()));
801
802 C.mat.resize(c_nnz);
803 C.col_index.resize(c_nnz);
804 C.nnz = c_nnz;
805 AMREX_ROCSPARSE_SAFE_CALL
806 (rocsparse_csr_set_pointers(r.mat_C, (void*)C.row_offset.data(),
807 (void*)C.col_index.data(), (void*)C.mat.data()));
808
809 spgemm_rocsparse_check
810 (rocsparse_spgemm(r.handle, op, op, &alpha, r.mat_A, r.mat_B, &beta, r.mat_D, r.mat_C,
811 data_type, alg, rocsparse_spgemm_stage_compute,
812 &buffer_size, r.buffer));
813
815 The_Arena()->free(r.buffer);
816 r.buffer = nullptr;
817
818 C.sort(); // rocSPARSE does not promise sorted rows for the generic API.
819
820 return C;
821}
822
823#elif defined(AMREX_USE_SYCL)
824
825template <typename T, template<typename> class V, typename I>
826CSR<T,V,I> spgemm_local_onemkl (Long nrows, Long ncols,
827 CsrView<T const,I> const& A, CsrView<T const,I> const& B)
828{
829 auto& q = Gpu::Device::streamQueue();
830
831 CSR<T,V,I> C;
832 C.row_offset.resize(nrows+1);
833 // oneMKL wants valid pointers for C before the nnz is known.
834 V<I> dummy_col(1);
835 V<T> dummy_mat(1);
836
837 // Releases everything, also when an OutOfMemoryError is thrown.
838 struct Resources {
839 sycl::queue& q;
840 mkl::sparse::matrix_handle_t hA{}, hB{}, hC{};
841 mkl::sparse::matmat_descr_t descr = nullptr;
842 std::int64_t* size_buf = nullptr;
843 void* buffer1 = nullptr;
844 void* buffer2 = nullptr;
845 explicit Resources (sycl::queue& a_q) : q(a_q) {}
846 Resources (Resources const&) = delete;
847 Resources& operator= (Resources const&) = delete;
848 ~Resources () {
849 q.wait();
850 if (descr) { mkl::sparse::release_matmat_descr(&descr); }
851 mkl::sparse::release_matrix_handle(q, &hA);
852 mkl::sparse::release_matrix_handle(q, &hB);
853 mkl::sparse::release_matrix_handle(q, &hC).wait();
854 Gpu::streamSynchronize();
855 if (buffer1) { The_Arena()->free(buffer1); }
856 if (buffer2) { The_Arena()->free(buffer2); }
857 if (size_buf) { The_Pinned_Arena()->free(size_buf); }
858 }
859 } r(q);
860
861 mkl::sparse::init_matrix_handle(&r.hA);
862 mkl::sparse::init_matrix_handle(&r.hB);
863 mkl::sparse::init_matrix_handle(&r.hC);
864
865#if defined(INTEL_MKL_VERSION) && (INTEL_MKL_VERSION < 20250300)
866 mkl::sparse::set_csr_data(q, r.hA, I(nrows), I(B.nrows), mkl::index_base::zero,
867 (I*)A.row_offset, (I*)A.col_index, (T*)A.mat);
868 mkl::sparse::set_csr_data(q, r.hB, I(B.nrows), I(ncols), mkl::index_base::zero,
869 (I*)B.row_offset, (I*)B.col_index, (T*)B.mat);
870 mkl::sparse::set_csr_data(q, r.hC, I(nrows), I(ncols), mkl::index_base::zero,
871 C.row_offset.data(), dummy_col.data(), dummy_mat.data());
872#else
873 mkl::sparse::set_csr_data(q, r.hA, I(nrows), I(B.nrows), I(A.nnz), mkl::index_base::zero,
874 (I*)A.row_offset, (I*)A.col_index, (T*)A.mat);
875 mkl::sparse::set_csr_data(q, r.hB, I(B.nrows), I(ncols), I(B.nnz), mkl::index_base::zero,
876 (I*)B.row_offset, (I*)B.col_index, (T*)B.mat);
877 mkl::sparse::set_csr_data(q, r.hC, I(nrows), I(ncols), I(0), mkl::index_base::zero,
878 C.row_offset.data(), dummy_col.data(), dummy_mat.data());
879#endif
880
881 mkl::sparse::init_matmat_descr(&r.descr);
882 mkl::sparse::set_matmat_data(r.descr,
883 mkl::sparse::matrix_view_descr::general,
884 mkl::transpose::nontrans,
885 mkl::sparse::matrix_view_descr::general,
886 mkl::transpose::nontrans,
887 mkl::sparse::matrix_view_descr::general);
888
889 using req = mkl::sparse::matmat_request;
890 r.size_buf = (std::int64_t*)The_Pinned_Arena()->alloc(sizeof(std::int64_t));
891 auto* size_buf = r.size_buf;
892
893 mkl::sparse::matmat(q, r.hA, r.hB, r.hC, req::get_work_estimation_buf_size, r.descr,
894 size_buf, nullptr, {}).wait();
895 r.buffer1 = (void*)The_Arena()->alloc(std::size_t(*size_buf));
896 mkl::sparse::matmat(q, r.hA, r.hB, r.hC, req::work_estimation, r.descr,
897 size_buf, r.buffer1, {}).wait();
898
899 mkl::sparse::matmat(q, r.hA, r.hB, r.hC, req::get_compute_buf_size, r.descr,
900 size_buf, nullptr, {}).wait();
901 r.buffer2 = (void*)The_Arena()->alloc(std::size_t(*size_buf));
902 mkl::sparse::matmat(q, r.hA, r.hB, r.hC, req::compute, r.descr,
903 size_buf, r.buffer2, {}).wait();
904
905 mkl::sparse::matmat(q, r.hA, r.hB, r.hC, req::get_nnz, r.descr,
906 size_buf, nullptr, {}).wait();
907 Long const c_nnz = *size_buf;
908 AMREX_ALWAYS_ASSERT(c_nnz < Long(std::numeric_limits<I>::max()));
909
910 C.mat.resize(c_nnz);
911 C.col_index.resize(c_nnz);
912 C.nnz = c_nnz;
913#if defined(INTEL_MKL_VERSION) && (INTEL_MKL_VERSION < 20250300)
914 mkl::sparse::set_csr_data(q, r.hC, I(nrows), I(ncols), mkl::index_base::zero,
915 C.row_offset.data(), C.col_index.data(), C.mat.data());
916#else
917 mkl::sparse::set_csr_data(q, r.hC, I(nrows), I(ncols), I(c_nnz), mkl::index_base::zero,
918 C.row_offset.data(), C.col_index.data(), C.mat.data());
919#endif
920
921 mkl::sparse::matmat(q, r.hA, r.hB, r.hC, req::finalize, r.descr,
922 size_buf, nullptr, {}).wait();
923
925 The_Arena()->free(r.buffer1);
926 The_Arena()->free(r.buffer2);
927 r.buffer1 = nullptr;
928 r.buffer2 = nullptr;
929
930 C.sort(); // oneMKL does not sort the output.
931
932 return C;
933}
934
935#endif
936
937#ifdef AMREX_USE_GPU
938
939template <typename T, template<typename> class V, typename I>
940using SpGEMMLocalFn = CSR<T,V,I> (*) (Long, Long, CsrView<T const,I> const&,
941 CsrView<T const,I> const&);
942
943// Copies block Cb into rows [r0,r1) of C, starting at nonzero p.
944template <typename T, template<typename> class V, typename I>
945void spgemm_copy_block (CSR<T,V,I>& C, Long r0, Long r1, CSR<T,V,I> const& Cb, Long p)
946{
947 auto* crow = C.row_offset.data() + r0;
948 auto const* src = Cb.row_offset.data();
949 auto const offset = I(p);
950 ParallelFor(r1 - r0 + 1, [=] AMREX_GPU_DEVICE (Long i) {
951 crow[i] = src[i] + offset;
952 });
953 Gpu::copyAsync(Gpu::deviceToDevice, Cb.col_index.begin(), Cb.col_index.end(),
954 C.col_index.begin() + p);
955 Gpu::copyAsync(Gpu::deviceToDevice, Cb.mat.begin(), Cb.mat.end(),
956 C.mat.begin() + p);
957}
958
959// Multiplies row blocks A(r0:r1,:) * B by `f`, covering `blocks` in order.
960// A block that runs out of memory is split in half. Each block is copied
961// into C at nonzero p if C is given; p is advanced by its nonzeros.
962// Returns the blocks used.
963template <typename T, template<typename> class V, typename I>
964Vector<std::pair<Long,Long>>
965spgemm_blocks (Vector<std::pair<Long,Long>> const& blocks, Long ncols,
966 CsrView<T const,I> const& A, CsrView<T const,I> const& B,
967 SpGEMMLocalFn<T,V,I> f, CSR<T,V,I>* C, Long& p)
968{
969 Vector<std::pair<Long,Long>> done;
970 Vector<std::pair<Long,Long>> todo(blocks.rbegin(), blocks.rend()); // processed from the back
971 while (!todo.empty()) {
972 auto const [r0, r1] = todo.back();
973 todo.pop_back();
974 Long const n = r1 - r0;
975 if (n == 0) { continue; }
976
977 I a[2];
978 Gpu::dtoh_memcpy_async(a , A.row_offset+r0, sizeof(I));
979 Gpu::dtoh_memcpy_async(a+1, A.row_offset+r1, sizeof(I));
981 I const a0 = a[0];
982 Long const annz = Long(a[1]) - Long(a0);
983
984 try {
985 CSR<T,V,I> Cb;
986 if (annz == 0) {
987 Cb = spgemm_empty<T,V,I>(n);
988 } else {
989 V<I> row(n+1);
990 auto* prow = row.data();
991 auto const* arow = A.row_offset + r0;
992 ParallelFor(n+1, [=] AMREX_GPU_DEVICE (Long i) { prow[i] = arow[i] - a0; });
993 CsrView<T const,I> Ab{A.mat + a0, A.col_index + a0, row.data(), annz, n};
994 Cb = f(n, ncols, Ab, B);
995 }
996 if (C) { spgemm_copy_block(*C, r0, r1, Cb, p); }
997 p += Cb.nnz;
998 } catch (OutOfMemoryError const&) {
999 if (n == 1) { throw; }
1000 todo.emplace_back(r0 + n/2, r1);
1001 todo.emplace_back(r0, r0 + n/2);
1002 continue;
1003 }
1004 done.emplace_back(r0, r1);
1005 }
1006 return done;
1007}
1008
1014template <typename T, template<typename> class V, typename I>
1015CSR<T,V,I> spgemm_local_chunked (Long nrows, Long ncols,
1016 CsrView<T const,I> const& A, CsrView<T const,I> const& B,
1017 SpGEMMLocalFn<T,V,I> f)
1018{
1019 Long total = 0;
1020 auto const blocks = spgemm_blocks<T,V,I>({{0, nrows/2}, {nrows/2, nrows}}, ncols, A, B, f,
1021 nullptr, total);
1022 AMREX_ALWAYS_ASSERT(total <= Long(std::numeric_limits<I>::max()));
1023
1024 CSR<T,V,I> C;
1025 C.resize(nrows, total);
1026 Long p = 0;
1027 spgemm_blocks<T,V,I>(blocks, ncols, A, B, f, &C, p);
1029
1030 return C;
1031}
1032
1033#endif
1034
1043template <typename T, template<typename> class V, typename I, typename F = NoRowPost>
1044CSR<T,V,I> spgemm_local (Long nrows, Long ncols,
1045 CsrView<T const,I> const& A, CsrView<T const,I> const& B,
1046 F const& row_post = {})
1047{
1048 AMREX_ASSERT(A.nrows == nrows);
1049
1050 if (nrows <= 0 || ncols <= 0 || A.nnz <= 0 || B.nnz <= 0 || B.nrows <= 0) {
1051 return spgemm_empty<T,V,I>(nrows);
1052 }
1053
1054#if !defined(AMREX_USE_GPU)
1055 return spgemm_local_cpu<T,V,I>(nrows, ncols, A, B, row_post);
1056#else
1057 amrex::ignore_unused(row_post);
1058#endif
1059#if defined(AMREX_USE_CUDA)
1060 try {
1061 return spgemm_local_cusparse<T,V,I>(nrows, ncols, A, B);
1062 } catch (OutOfMemoryError const&) {
1063 return spgemm_local_chunked<T,V,I>(nrows, ncols, A, B,
1064 spgemm_local_cusparse<T,V,I>);
1065 }
1066#elif defined(AMREX_USE_HIP)
1067 try {
1068 return spgemm_local_rocsparse<T,V,I>(nrows, ncols, A, B);
1069 } catch (OutOfMemoryError const&) {
1070 return spgemm_local_chunked<T,V,I>(nrows, ncols, A, B,
1071 spgemm_local_rocsparse<T,V,I>);
1072 }
1073#elif defined(AMREX_USE_SYCL)
1074 try {
1075 return spgemm_local_onemkl<T,V,I>(nrows, ncols, A, B);
1076 } catch (OutOfMemoryError const&) {
1077 return spgemm_local_chunked<T,V,I>(nrows, ncols, A, B,
1078 spgemm_local_onemkl<T,V,I>);
1079 }
1080#endif
1081}
1082
1083
1084#ifdef AMREX_USE_MPI
1085
1086// Row i of the result is row i of X0 followed by row i of X1, with column
1087// indices transformed by map0 and map1. Both views must have full row offsets.
1088template <typename T, template<typename> class V, typename I, typename F0, typename F1>
1089CSR<T,V,int> concat_csr_cols (Long nrows, CsrView<T const,I> const& X0, CsrView<T const,I> const& X1,
1090 F0 const& map0, F1 const& map1)
1091{
1092 AMREX_ALWAYS_ASSERT(X0.nnz + X1.nnz < Long(std::numeric_limits<int>::max()));
1093 CSR<T,V,int> C;
1094 C.resize(nrows, X0.nnz + X1.nnz);
1095 int* AMREX_RESTRICT crow = C.row_offset.data();
1096 int* AMREX_RESTRICT ccol = C.col_index.data();
1097 T* AMREX_RESTRICT cmat = C.mat.data();
1098 ParallelForOMP(nrows+1, [=] AMREX_GPU_DEVICE (Long i)
1099 {
1100 crow[i] = X0.row_offset[i] + X1.row_offset[i];
1101 if (i < nrows) {
1102 Long p = crow[i];
1103 for (Long q = X0.row_offset[i]; q < X0.row_offset[i+1]; ++q) {
1104 ccol[p] = map0(X0.col_index[q]);
1105 cmat[p] = X0.mat[q];
1106 ++p;
1107 }
1108 for (Long q = X1.row_offset[i]; q < X1.row_offset[i+1]; ++q) {
1109 ccol[p] = map1(X1.col_index[q]);
1110 cmat[p] = X1.mat[q];
1111 ++p;
1112 }
1113 }
1114 });
1115 return C;
1116}
1117
1118// Append rows with global column indices, sorted within each row, to Bh
1119// whose first rows are already in the compact column space. Each row is
1120// rotated so that the local block [c0,c1) comes first, which keeps it
1121// sorted in the compact space [local | remote_union].
1122template <typename T, template<typename> class V>
1123void append_ext_rows (CSR<T,V,int>& Bh, Long const* ext_row_offset, Long const* ext_col,
1124 T const* ext_mat, Long n_ext, Long nnz_ext,
1125 Long c0, Long c1, Long const* ru, Long nru)
1126{
1127 Long const nb = Bh.nrows();
1128 Long const nnz_b = Bh.nnz;
1129 AMREX_ALWAYS_ASSERT(nnz_b + nnz_ext < Long(std::numeric_limits<int>::max()));
1130 Bh.mat.resize(nnz_b + nnz_ext);
1131 Bh.col_index.resize(nnz_b + nnz_ext);
1132 Bh.row_offset.resize(nb + n_ext + 1);
1133 Bh.nnz = nnz_b + nnz_ext;
1134 int* AMREX_RESTRICT brow = Bh.row_offset.data();
1135 int* AMREX_RESTRICT bcol = Bh.col_index.data();
1136 T* AMREX_RESTRICT bmat = Bh.mat.data();
1137 Long const nlocal = c1 - c0;
1138 ParallelForOMP(n_ext, [=] AMREX_GPU_DEVICE (Long r)
1139 {
1140 Long const b = ext_row_offset[r];
1141 Long const e = ext_row_offset[r+1];
1142 brow[nb+r+1] = int(nnz_b + e);
1143 Long const p0 = amrex::lower_bound(ext_col+b, ext_col+e, c0) - ext_col;
1144 Long const p1 = amrex::lower_bound(ext_col+b, ext_col+e, c1) - ext_col;
1145 Long p = nnz_b + b;
1146 for (Long q = p0; q < p1; ++q) {
1147 bcol[p] = int(ext_col[q] - c0);
1148 bmat[p] = ext_mat[q];
1149 ++p;
1150 }
1151 for (Long q = b; q < e; ++q) {
1152 if (q < p0 || q >= p1) {
1153 bcol[p] = int(nlocal + (amrex::lower_bound(ru, ru+nru, ext_col[q]) - ru));
1154 bmat[p] = ext_mat[q];
1155 ++p;
1156 }
1157 }
1158 });
1159}
1160
1161#endif
1162
1163}
1164
1165namespace amrex {
1166
1167template <typename T, template <typename> class Allocator, typename F>
1168SpMatrix<T,Allocator>
1169SpGEMM (SpMatrix<T,Allocator> const& A, SpMatrix<T,Allocator> const& B,
1170 AlgPartition const& col_partition, F const& row_post);
1171
1187template <typename T, template <typename> class Allocator>
1188SpMatrix<T,Allocator>
1190 AlgPartition const& col_partition)
1191{
1192 return SpGEMM(A, B, col_partition, detail::NoRowPost{});
1193}
1194
1210template <typename T, template <typename> class Allocator, typename F>
1211SpMatrix<T,Allocator>
1213 AlgPartition const& col_partition, F const& row_post)
1214{
1215 BL_PROFILE("SpGEMM");
1216#ifdef AMREX_USE_GPU
1217 static_assert(std::is_same_v<F,detail::NoRowPost>, "SpGEMM: row_post is CPU only");
1218 amrex::ignore_unused(row_post);
1219#endif
1220
1221 using SpMat = SpMatrix<T,Allocator>;
1222
1223 auto& Am = const_cast<SpMat&>(A);
1224 auto& Bm = const_cast<SpMat&>(B);
1225 Long const nrows = Am.numLocalRows();
1226
1227#ifdef AMREX_USE_MPI
1228 using LongVec = typename SpMat::template container_type<Long>;
1229
1230 Am.setColumnPartition(Bm.partition());
1231 Bm.setColumnPartition(col_partition);
1232
1233 Long const nb = Bm.numLocalRows();
1234 Long const c0 = col_partition.globalRowBegin();
1235 Long const c1 = col_partition.globalRowEnd();
1236 Long const nlocal = c1 - c0;
1237
1238 // Rows of B needed by A's off-diagonal columns, global column indices.
1239 typename SpMat::RemoteRowsMM ext;
1240 if (! detail::spmat_comm_is_local(Am.partition(), Bm.partition())) {
1241 ext = Am.fetch_remote_rows_mm(Bm);
1242 }
1243
1244 if (ext.nrows == 0 && Am.m_remote_cols_v.empty() && Bm.m_remote_cols_v.empty()) {
1245 auto Ch = detail::spgemm_local<T,SpMat::template container_type,int>
1246 (nrows, nlocal, Am.m_csr_local.const_view(), Bm.m_csr_local.const_view(), row_post);
1247 SpMat C;
1248 C.define_split(Am.partition(), col_partition, std::move(Ch), nlocal, nullptr, 0);
1249 return C;
1250 }
1251
1252 // Compact column space: [local columns | sorted remote columns].
1253 // TODO: this serial host loop over ext.nnz limits scaling on GPUs.
1254 Vector<Long> ru_h = Bm.m_remote_cols_v;
1255 for (Long i = 0; i < ext.nnz; ++i) {
1256 auto g = ext.col_index[i];
1257 if (g < c0 || g >= c1) { ru_h.push_back(g); }
1258 }
1259 RemoveDuplicates(ru_h);
1260 Long const nru = Long(ru_h.size());
1261 LongVec ru_d(nru);
1262 Gpu::copyAsync(Gpu::hostToDevice, ru_h.begin(), ru_h.end(), ru_d.begin());
1263 Long const* ru = ru_d.data();
1264 Long const ncols_hat = nlocal + nru;
1265 AMREX_ALWAYS_ASSERT(ncols_hat < Long(std::numeric_limits<int>::max()) &&
1266 nb + Long(Am.m_remote_cols_v.size()) < Long(std::numeric_limits<int>::max()));
1267
1268 using local_csr_type = typename SpMat::local_csr_type;
1269#ifndef AMREX_USE_GPU
1270 {
1271 // The rows of B fetched from other processes, in the compact
1272 // column space; B's own remote columns are mapped once each.
1273 local_csr_type E;
1274 E.resize(0, 0);
1275 E.row_offset[0] = 0;
1276 if (ext.nrows > 0) {
1277 detail::append_ext_rows(E, ext.row_offset.data(), ext.col_index, ext.mat,
1278 ext.nrows, ext.nnz, c0, c1, ru, nru);
1279 }
1280 ext.clear();
1281 Vector<int> b1map(Bm.m_remote_cols_v.size());
1282 for (Long j = 0; j < Long(b1map.size()); ++j) {
1283 b1map[j] = int(nlocal + (std::lower_bound(ru, ru+nru, Bm.m_remote_cols_v[j]) - ru));
1284 }
1285 auto const a1 = Am.m_remote_cols_v.empty() ? CsrView<T const,int>{}
1286 : Am.remote_full_const_view();
1287 auto const b1 = Bm.m_remote_cols_v.empty() ? CsrView<T const,int>{}
1288 : Bm.remote_full_const_view();
1289 auto Ch = detail::spgemm_split_cpu<T,SpMat::template container_type,int>
1290 (nrows, ncols_hat, nb, Am.m_csr_local.const_view(), a1,
1291 Bm.m_csr_local.const_view(), b1, b1map.data(), E.const_view(), row_post);
1292 SpMat C;
1293 C.define_split(Am.partition(), col_partition, std::move(Ch), nlocal, ru, nru);
1294 return C;
1295 }
1296#endif
1297 local_csr_type Ah = detail::concat_csr_cols<T,SpMat::template container_type>
1298 (nrows, Am.m_csr_local.const_view(), Am.remote_full_const_view(),
1299 [=] AMREX_GPU_DEVICE (int c) { return c; },
1300 [=] AMREX_GPU_DEVICE (int j) { return int(nb + j); });
1301
1302#ifdef AMREX_USE_GPU
1303 Long const* b_rcols = Bm.m_remote_cols_dv.data();
1304#else
1305 Long const* b_rcols = Bm.m_remote_cols_v.data();
1306#endif
1307 local_csr_type Bh = detail::concat_csr_cols<T,SpMat::template container_type>
1308 (nb, Bm.m_csr_local.const_view(), Bm.remote_full_const_view(),
1309 [=] AMREX_GPU_DEVICE (int c) { return c; },
1310 [=] AMREX_GPU_DEVICE (int j) {
1311 return int(nlocal + (amrex::lower_bound(ru, ru+nru, b_rcols[j]) - ru)); });
1312
1313 if (ext.nrows > 0) {
1314 LongVec ext_col_d(ext.nnz);
1315 Gpu::copyAsync(Gpu::hostToDevice, ext.col_index, ext.col_index+ext.nnz,
1316 ext_col_d.begin());
1317 detail::append_ext_rows(Bh, ext.row_offset.data(), ext_col_d.data(), ext.mat,
1318 ext.nrows, ext.nnz, c0, c1, ru, nru);
1320 }
1321 ext.clear();
1322
1323 local_csr_type Ch = detail::spgemm_local<T,SpMat::template container_type,int>
1324 (nrows, ncols_hat, Ah.const_view(), Bh.const_view());
1325 Ah = local_csr_type{};
1326 Bh = local_csr_type{};
1327
1328 // The compact column space is [local | sorted remote], so the result
1329 // can be defined directly in split form.
1330 SpMat C;
1331 C.define_split(Am.partition(), col_partition, std::move(Ch), nlocal, ru, nru);
1332 return C;
1333
1334#else
1335
1336 Am.setColumnPartition(Bm.partition());
1337 Bm.setColumnPartition(col_partition);
1338
1339 // One rank: local column indices are the global ones.
1340 Long const ncols = col_partition.numGlobalRows();
1341 AMREX_ALWAYS_ASSERT(ncols < Long(std::numeric_limits<int>::max()));
1342 auto Ch = detail::spgemm_local<T,SpMat::template container_type,int>
1343 (nrows, ncols, Am.m_csr_local.const_view(), Bm.m_csr_local.const_view(), row_post);
1344 SpMat C;
1345 C.define_split(Am.partition(), col_partition, std::move(Ch), ncols, nullptr, 0);
1346 return C;
1347
1348#endif
1349}
1350
1362template <typename T, template <typename> class Allocator>
1363SpMatrix<T,Allocator>
1365 SpMatrix<T,Allocator> const& P, AlgPartition const& col_partition)
1366{
1367 BL_PROFILE("RAP");
1368#ifdef AMREX_USE_GPU
1369 auto AP = SpGEMM(A, P, col_partition);
1370 return SpGEMM(R, AP, col_partition);
1371#else
1372 using SpMat = SpMatrix<T,Allocator>;
1373 using local_csr_type = typename SpMat::local_csr_type;
1374 auto& Rm = const_cast<SpMat&>(R);
1375 auto& Am = const_cast<SpMat&>(A);
1376 auto& Pm = const_cast<SpMat&>(P);
1377 Am.setColumnPartition(Pm.partition());
1378 Pm.setColumnPartition(col_partition);
1379 Rm.setColumnPartition(Am.partition());
1380
1381 Long const nf = Am.numLocalRows();
1382 Long const nc = col_partition.numLocalRows();
1383 AMREX_ALWAYS_ASSERT(Rm.numLocalRows() == nc);
1384
1385 detail::APRowsCpu<T,int> ap;
1386 ap.A0 = Am.m_csr_local.const_view();
1387 ap.P0 = Pm.m_csr_local.const_view();
1389 local_csr_type pe, ape;
1390 pe.resize(0, 0);
1391 pe.row_offset[0] = 0;
1392 ape.resize(0, 0);
1393 ape.row_offset[0] = 0;
1394 ap.PE = pe.const_view();
1395 Vector<int> p1map;
1396 Vector<Long> ru;
1397
1398#ifdef AMREX_USE_MPI
1399 Long const c0 = col_partition.globalRowBegin();
1400 Long const c1 = col_partition.globalRowEnd();
1401 if (! Am.m_remote_cols_v.empty()) { ap.A1 = Am.remote_full_const_view(); }
1402 if (! Pm.m_remote_cols_v.empty()) { ap.P1 = Pm.remote_full_const_view(); }
1403
1404 // Remote columns [ru_first, ru_last) as sorted unique global indices,
1405 // P's remote columns mapped into the compact space of `cols`, and the
1406 // fetched rows `ext` in that space.
1407 auto add_remote = [&] (Vector<Long>& cols, typename SpMat::RemoteRowsMM const& ext) {
1408 for (Long i = 0; i < ext.nnz; ++i) {
1409 auto g = ext.col_index[i];
1410 if (g < c0 || g >= c1) { cols.push_back(g); }
1411 }
1412 RemoveDuplicates(cols);
1413 };
1414 auto compact_ext = [&] (local_csr_type& E, typename SpMat::RemoteRowsMM const& ext,
1415 Vector<Long> const& cols) {
1416 E.resize(0, 0);
1417 E.row_offset[0] = 0;
1418 if (ext.nrows > 0) {
1419 detail::append_ext_rows(E, ext.row_offset.data(), ext.col_index, ext.mat,
1420 ext.nrows, ext.nnz, c0, c1, cols.data(), Long(cols.size()));
1421 }
1422 };
1423 auto map_p1 = [&] (Vector<Long> const& cols) {
1424 // Built with push_back: gcc < 15 -Wnull-dereference false positive in resize.
1425 p1map.clear();
1426 p1map.reserve(Pm.m_remote_cols_v.size());
1427 for (auto const g : Pm.m_remote_cols_v) {
1428 p1map.push_back(int(nc + (std::lower_bound(cols.begin(), cols.end(), g)
1429 - cols.begin())));
1430 }
1431 ap.p1map = p1map.data();
1432 };
1433
1434 // Rows of P needed by A's remote columns.
1435 typename SpMat::RemoteRowsMM pext;
1436 if (! detail::spmat_comm_is_local(Am.partition(), Pm.partition())) {
1437 pext = Am.fetch_remote_rows_mm(Pm);
1438 }
1439 Vector<Long> ru_ap = Pm.m_remote_cols_v;
1440 add_remote(ru_ap, pext);
1441
1442 // The rows of A P that other processes need: those where P has remote
1443 // columns, since the transpose of P then has a remote column there.
1444 SpMat APb;
1445 {
1446 compact_ext(pe, pext, ru_ap);
1447 ap.PE = pe.const_view();
1448 map_p1(ru_ap);
1449 Long const ncols = nc + Long(ru_ap.size());
1450 AMREX_ALWAYS_ASSERT(ncols < Long(std::numeric_limits<int>::max()));
1451 auto has_remote = [&] (Long i) {
1452 return ap.P1.nnz > 0 && ap.P1.row_offset[i+1] > ap.P1.row_offset[i];
1453 };
1454 auto csr = detail::csr_from_rows_cpu<T,SpMat::template container_type,int>
1455 (nf, [&] (Long i) { return has_remote(i) ? ap.maxlen(i) : Long(0); },
1456 [&] (int) {
1457 return [&, marker = Vector<int>(ncols, -1), col = Vector<int>(), val = Vector<T>(),
1458 tmp = Vector<std::pair<int,T>>()]
1459 (Long i, int const*& rc, T const*& rv) mutable -> Long
1460 {
1461 Long n = 0;
1462 if (has_remote(i)) {
1463 Long const maxlen = ap.maxlen(i);
1464 if (Long(col.size()) < maxlen) {
1465 col.resize(maxlen);
1466 val.resize(maxlen);
1467 }
1468 n = ap.row(i, marker.data(), col.data(), val.data(), 0);
1469 // fetch_remote_rows_mm sends rows sorted by column.
1470 detail::sort_row_cpu(col.data(), val.data(), n, tmp);
1471 }
1472 rc = col.data();
1473 rv = val.data();
1474 return n;
1475 };
1476 });
1477 APb.define_split(Am.partition(), col_partition, std::move(csr), nc,
1478 ru_ap.data(), Long(ru_ap.size()));
1479 }
1480
1481 // Rows of A P needed by R's remote columns.
1482 typename SpMat::RemoteRowsMM apext;
1483 if (! detail::spmat_comm_is_local(Rm.partition(), APb.partition())) {
1484 apext = Rm.fetch_remote_rows_mm(APb);
1485 }
1486 ru = ru_ap;
1487 add_remote(ru, apext);
1488 AMREX_ALWAYS_ASSERT(nc + Long(ru.size()) < Long(std::numeric_limits<int>::max()));
1489 compact_ext(pe, pext, ru);
1490 compact_ext(ape, apext, ru);
1491 ap.PE = pe.const_view();
1492 map_p1(ru);
1493 pext.clear();
1494 apext.clear();
1495 if (! Rm.m_remote_cols_v.empty()) { r1 = Rm.remote_full_const_view(); }
1496#endif
1497
1498 Long const nru = Long(ru.size());
1499 auto csr = detail::rap_local_cpu<T,SpMat::template container_type,int>
1500 (nf, nc, nc+nru, Rm.m_csr_local.const_view(), r1, ape.const_view(), ap);
1501 SpMat C;
1502 C.define_split(Rm.partition(), col_partition, std::move(csr), nc, ru.data(), nru);
1503 return C;
1504#endif
1505}
1506
1507}
1508
1509#endif
General-purpose algorithm utilities available on both host and device.
#define BL_PROFILE(a)
Definition AMReX_BLProfiler.H:562
#define AMREX_ALWAYS_ASSERT_WITH_MESSAGE(EX, MSG)
Definition AMReX_BLassert.H:49
#define AMREX_ASSERT(EX)
Definition AMReX_BLassert.H:38
#define AMREX_ALWAYS_ASSERT(EX)
Definition AMReX_BLassert.H:50
AMReX exception types.
#define AMREX_ASSUME(ASSUMPTION)
Definition AMReX_Extension.H:287
#define AMREX_RESTRICT
Definition AMReX_Extension.H:37
#define AMREX_CUSPARSE_SAFE_CALL(call)
Definition AMReX_GpuError.H:101
#define AMREX_GPU_DEVICE
Definition AMReX_GpuQualifiers.H:18
amrex::ParmParse pp
Input file parser instance for the given namespace.
Definition AMReX_HypreIJIface.cpp:18
Array4< int const > offset
Definition AMReX_HypreMLABecLap.cpp:1139
GpuArray< Real, 3 > beta
Definition AMReX_MLEBNodeFDLaplacian.cpp:1834
GpuArray< MultiArray4< Real const >, 3 > s
Definition AMReX_MLEBNodeFDLaplacian.cpp:214
Definition AMReX_AlgPartition.H:26
Long numGlobalRows() const
Total number of rows covered by the partition.
Definition AMReX_AlgPartition.H:55
Long globalRowEnd() const
Exclusive global index end on this process.
Definition AMReX_AlgPartition.H:67
Long globalRowBegin() const
Inclusive global index begin on this process.
Definition AMReX_AlgPartition.H:62
Long numLocalRows() const
Number of local rows.
Definition AMReX_AlgPartition.H:50
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
Long numLocalRows() const
Number of rows owned by this rank.
Definition AMReX_SpMatrix.H:194
This class is a thin wrapper around std::vector. Unlike vector, Vector::operator[] provides bound che...
Definition AMReX_Vector.H:29
Long size() const noexcept
Definition AMReX_Vector.H:54
amrex_long Long
Definition AMReX_INT.H:30
void ParallelForOMP(T n, L const &f) noexcept
Performance-portable kernel launch function with optional OpenMP threading.
Definition AMReX_GpuLaunch.H:328
Arena * The_Pinned_Arena()
Definition AMReX_Arena.cpp:869
Arena * The_Arena()
Definition AMReX_Arena.cpp:829
__host__ __device__ ItType lower_bound(ItType first, ItType last, const ValType &val)
Return an iterator to the first element not less than a given value.
Definition AMReX_Algorithm.H:298
void copyAsync(HostToDevice, InIter begin, InIter end, OutIter result) noexcept
A host-to-device copy routine. Note this is just a wrapper around memcpy, so it assumes contiguous st...
Definition AMReX_GpuContainers.H:228
static constexpr DeviceToDevice deviceToDevice
Definition AMReX_GpuContainers.H:107
static constexpr HostToDevice hostToDevice
Definition AMReX_GpuContainers.H:105
void streamSynchronize() noexcept
Definition AMReX_GpuDevice.H:310
void dtoh_memcpy_async(void *p_h, const void *p_d, const std::size_t sz) noexcept
Definition AMReX_GpuDevice.H:435
gpuStream_t gpuStream() noexcept
Definition AMReX_GpuDevice.H:291
constexpr int get_max_threads()
Definition AMReX_OpenMP.H:36
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
SpMatrix< T, Allocator > RAP(SpMatrix< T, Allocator > const &R, SpMatrix< T, Allocator > const &A, SpMatrix< T, Allocator > const &P, AlgPartition const &col_partition)
Galerkin product R (A P), with the result of SpGEMM(R, SpGEMM(A, P, col_partition),...
Definition AMReX_SpGEMM.H:1364
SpMatrix< T, Allocator > SpGEMM(SpMatrix< T, Allocator > const &A, SpMatrix< T, Allocator > const &B, AlgPartition const &col_partition, F const &row_post)
SpGEMM with a callback on each row of the product.
Definition AMReX_SpGEMM.H:1212
void Abort(const std::string &msg)
Print a fatal-error message to stderr and abort execution.
Definition AMReX.cpp:244
const int[]
Definition AMReX_BLProfiler.cpp:1665
void RemoveDuplicates(Vector< T > &vec)
Definition AMReX_Vector.H:210
Lightweight non-owning CSR view that can point to host or device buffers.
Definition AMReX_CSR.H:35
CI *__restrict__ row_offset
Definition AMReX_CSR.H:40