1#ifndef AMREX_SP_MATRIX_H_
2#define AMREX_SP_MATRIX_H_
3#include <AMReX_Config.H>
12#if defined(AMREX_USE_CUDA)
14#elif defined(AMREX_USE_HIP)
15# include <rocsparse/rocsparse.h>
16#elif defined(AMREX_USE_SYCL)
17# include <mkl_version.h>
18# include <oneapi/mkl/spblas.hpp>
51 explicit operator bool()
const {
return b; }
57 explicit operator bool()
const {
return b; }
208 return m_csr.
mat.data();
264 template <
typename F>
302 template <
typename U,
template<
typename>
class M,
typename N>
friend
305 template <
typename U,
template<
typename>
class M>
friend
308 template <
typename U,
template<
typename>
class M>
friend
312 template <
typename U,
template<
typename>
class M,
typename F>
friend
316 template <
typename U,
template<
typename>
class M>
friend
332 template <
typename I>
334 Long nentries,
Long const* row_offset);
354 template <
typename U>
357 template <
typename U>
377 Long const* remote_cols,
Long nremote);
381 void set_num_neighbors ();
385 Long m_row_begin = 0;
387 Long m_col_begin = 0;
395 bool m_split =
false;
442 int m_num_neighbors = -1;
490 template <
typename VL,
typename VI>
519 col_index(std::exchange(rhs.col_index,
nullptr)),
520 mat(std::exchange(rhs.mat,
nullptr)),
526 col_index = std::exchange(rhs.col_index,
nullptr);
527 mat = std::exchange(rhs.mat,
nullptr);
567inline bool spmat_comm_is_local (AlgPartition
const& row_partition,
568 AlgPartition
const& col_partition)
570 int const rp = row_partition.singleActiveProc();
571 int const cp = col_partition.singleActiveProc();
572 return row_partition.numActiveProcs() <= 1
573 && col_partition.numActiveProcs() <= 1
574 && (rp < 0 || cp < 0 || rp == cp);
578template <
typename T,
typename I>
579void extract_diagonal (T* p, CsrView<T const,I>
const& csr,
Long offset)
587 for (
Long j = row[i]; j < row[i+1]; ++j) {
597template <
typename T,
typename IO,
typename II>
598void transpose (CsrView<T,IO>
const& csrt, CsrView<T const,II>
const& csr)
600 Long nrows = csr.nrows;
601 Long ncols = csrt.nrows;
604 if (nrows <= 0 || ncols <= 0 || nnz <= 0) {
605 auto* p = csrt.row_offset;
612 if constexpr (!std::is_same_v<IO,Long>) {
614 nnz <
Long(std::numeric_limits<IO>::max()));
619#if defined(AMREX_USE_CUDA)
621 cusparseHandle_t handle;
625 cudaDataType data_type;
626 if constexpr (std::is_same_v<T,float>) {
627 data_type = CUDA_R_32F;
628 }
else if constexpr (std::is_same_v<T,double>) {
629 data_type = CUDA_R_64F;
630 }
else if constexpr (std::is_same_v<T,GpuComplex<float>>) {
631 data_type = CUDA_C_32F;
632 }
else if constexpr (std::is_same_v<T,GpuComplex<double>>) {
633 data_type = CUDA_C_64F;
635 amrex::Abort(
"SpMatrix transpose: unsupported data type");
639 ncols <
Long(std::numeric_limits<int>::max()) &&
640 nnz <
Long(std::numeric_limits<int>::max()));
643 constexpr bool in32 = std::is_same_v<II,int>;
644 constexpr bool out32 = std::is_same_v<IO,int>;
645 CsrIndex<int,Gpu::AsyncVector> ci, cit;
646 int const* csr_col_index;
647 int const* csr_row_offset;
649 int* csrt_row_offset;
650 if constexpr (in32) {
651 csr_col_index = csr.col_index;
652 csr_row_offset = csr.row_offset;
655 csr_col_index = ci.col_index.data();
656 csr_row_offset = ci.row_offset.data();
658 if constexpr (out32) {
659 csrt_col_index = csrt.col_index;
660 csrt_row_offset = csrt.row_offset;
662 cit.col_index.resize(csrt.nnz);
663 cit.row_offset.resize(csrt.nrows+1);
664 csrt_col_index = cit.col_index.data();
665 csrt_row_offset = cit.row_offset.data();
668 std::size_t buffer_size;
670 cusparseCsr2cscEx2_bufferSize(handle,
int(nrows),
int(ncols),
int(nnz),
671 csr.mat, csr_row_offset, csr_col_index,
672 csrt.mat, csrt_row_offset, csrt_col_index,
673 data_type, CUSPARSE_ACTION_NUMERIC,
674 CUSPARSE_INDEX_BASE_ZERO,
675 CUSPARSE_CSR2CSC_ALG1,
681 cusparseCsr2cscEx2(handle,
int(nrows),
int(ncols),
int(nnz),
682 csr.mat, csr_row_offset, csr_col_index,
683 csrt.mat, csrt_row_offset, csrt_col_index,
684 data_type, CUSPARSE_ACTION_NUMERIC,
685 CUSPARSE_INDEX_BASE_ZERO,
686 CUSPARSE_CSR2CSC_ALG1,
689 if constexpr (!out32) {
698#elif defined(AMREX_USE_HIP)
700 rocsparse_handle handle;
701 AMREX_ROCSPARSE_SAFE_CALL(rocsparse_create_handle(&handle));
702 AMREX_ROCSPARSE_SAFE_CALL(rocsparse_set_stream(handle,
Gpu::gpuStream()));
704 constexpr bool same_int = (
sizeof(rocsparse_int) ==
sizeof(II)) &&
705 (
sizeof(rocsparse_int) ==
sizeof(IO));
707 rocsparse_int
const* csr_col_index;
708 rocsparse_int
const* csr_row_offset;
709 rocsparse_int* csrt_col_index;
710 rocsparse_int* csrt_row_offset;
711 CsrIndex<rocsparse_int,Gpu::AsyncVector> ci, cit;
712 if constexpr (same_int) {
713 csr_col_index =
reinterpret_cast<rocsparse_int const*
>(csr.col_index);
714 csr_row_offset =
reinterpret_cast<rocsparse_int const*
>(csr.row_offset);
715 csrt_col_index =
reinterpret_cast<rocsparse_int*
>(csrt.col_index);
716 csrt_row_offset =
reinterpret_cast<rocsparse_int*
>(csrt.row_offset);
720 cit.col_index.resize(csrt.nnz);
721 cit.row_offset.resize(csrt.nrows+1);
722 csr_col_index = ci.col_index.data();
723 csr_row_offset = ci.row_offset.data();
724 csrt_col_index = cit.col_index.data();
725 csrt_row_offset = cit.row_offset.data();
728 std::size_t buffer_size;
729 AMREX_ROCSPARSE_SAFE_CALL(
730 rocsparse_csr2csc_buffer_size(handle, rocsparse_int(nrows),
731 rocsparse_int(ncols), rocsparse_int(nnz),
732 csr_row_offset, csr_col_index,
733 rocsparse_action_numeric,
738 if constexpr (std::is_same_v<T,float>) {
739 AMREX_ROCSPARSE_SAFE_CALL(
740 rocsparse_scsr2csc(handle, rocsparse_int(nrows),
741 rocsparse_int(ncols), rocsparse_int(nnz),
742 csr.mat, csr_row_offset, csr_col_index,
743 csrt.mat, csrt_col_index, csrt_row_offset,
744 rocsparse_action_numeric,
745 rocsparse_index_base_zero,
747 }
else if constexpr (std::is_same_v<T,double>) {
748 AMREX_ROCSPARSE_SAFE_CALL(
749 rocsparse_dcsr2csc(handle, rocsparse_int(nrows),
750 rocsparse_int(ncols), rocsparse_int(nnz),
751 csr.mat, csr_row_offset, csr_col_index,
752 csrt.mat, csrt_col_index, csrt_row_offset,
753 rocsparse_action_numeric,
754 rocsparse_index_base_zero,
756 }
else if constexpr (std::is_same_v<T,GpuComplex<float>>) {
757 AMREX_ROCSPARSE_SAFE_CALL(
758 rocsparse_ccsr2csc(handle, rocsparse_int(nrows),
759 rocsparse_int(ncols), rocsparse_int(nnz),
760 (rocsparse_float_complex*)csr.mat, csr_row_offset, csr_col_index,
761 (rocsparse_float_complex*)csrt.mat, csrt_col_index, csrt_row_offset,
762 rocsparse_action_numeric,
763 rocsparse_index_base_zero,
765 }
else if constexpr (std::is_same_v<T,GpuComplex<double>>) {
766 AMREX_ROCSPARSE_SAFE_CALL(
767 rocsparse_zcsr2csc(handle, rocsparse_int(nrows),
768 rocsparse_int(ncols), rocsparse_int(nnz),
769 (rocsparse_double_complex*)csr.mat, csr_row_offset, csr_col_index,
770 (rocsparse_double_complex*)csrt.mat, csrt_col_index, csrt_row_offset,
771 rocsparse_action_numeric,
772 rocsparse_index_base_zero,
775 amrex::Abort(
"SpMatrix transpose: unsupported data type");
778 if constexpr (!same_int) {
783 AMREX_ROCSPARSE_SAFE_CALL(rocsparse_destroy_handle(handle));
786#elif defined(AMREX_USE_SYCL)
788 mkl::sparse::matrix_handle_t handle_in{};
789 mkl::sparse::matrix_handle_t handle_out{};
790 mkl::sparse::init_matrix_handle(&handle_in);
791 mkl::sparse::init_matrix_handle(&handle_out);
794 Gpu::DeviceVector<Long> lrow_in, lcol_in, lrow_out, lcol_out;
799 if constexpr (std::is_same_v<II,Long>) {
800 prow_in = csr.row_offset;
801 pcol_in = csr.col_index;
803 lrow_in.resize(nrows+1);
805 auto* pr = lrow_in.data();
auto const* sr = csr.row_offset;
806 auto* pc = lcol_in.data();
auto const* sc = csr.col_index;
808 if (i < nrows+1) { pr[i] =
Long(sr[i]); }
809 if (i < nnz) { pc[i] =
Long(sc[i]); }
811 prow_in = lrow_in.data();
812 pcol_in = lcol_in.data();
814 if constexpr (std::is_same_v<IO,Long>) {
815 prow_out = csrt.row_offset;
816 pcol_out = csrt.col_index;
818 lrow_out.resize(ncols+1);
819 lcol_out.resize(nnz);
820 prow_out = lrow_out.data();
821 pcol_out = lcol_out.data();
823#if defined(INTEL_MKL_VERSION) && (INTEL_MKL_VERSION < 20250300)
824 mkl::sparse::set_csr_data(Gpu::Device::streamQueue(), handle_in, nrows, ncols,
825 mkl::index_base::zero, (
Long*)prow_in, (
Long*)pcol_in,
827 mkl::sparse::set_csr_data(Gpu::Device::streamQueue(), handle_out, ncols, nrows,
828 mkl::index_base::zero, prow_out, pcol_out, (T*)csrt.mat);
830 mkl::sparse::set_csr_data(Gpu::Device::streamQueue(), handle_in, nrows, ncols, nnz,
831 mkl::index_base::zero, (
Long*)prow_in, (
Long*)pcol_in,
833 mkl::sparse::set_csr_data(Gpu::Device::streamQueue(), handle_out, ncols, nrows, nnz,
834 mkl::index_base::zero, prow_out, pcol_out, (T*)csrt.mat);
837 mkl::sparse::omatcopy(Gpu::Device::streamQueue(), mkl::transpose::trans,
838 handle_in, handle_out);
840 mkl::sparse::release_matrix_handle(Gpu::Device::streamQueue(), &handle_in);
841 auto ev = mkl::sparse::release_matrix_handle(Gpu::Device::streamQueue(), &handle_out);
843 if constexpr (!std::is_same_v<IO,Long>) {
844 auto* po = csrt.col_index;
auto const* so = lcol_out.data();
845 auto* pr = csrt.row_offset;
auto const* sr = lrow_out.data();
847 if (i < nnz) { po[i] = IO(so[i]); }
848 if (i < ncols+1) { pr[i] = IO(sr[i]); }
859 auto* p = csrt.row_offset;
865#pragma omp parallel for
867 for (
Long i = 0; i < nnz; ++i) {
868 auto col = csr.col_index[i];
870#pragma omp atomic update
876 Vector<Long> current_pos(ncols+1);
878 for (
Long i = 0; i < ncols; ++i) {
880 current_pos[i+1] = p[i+1];
885 for (
Long i = 0; i < nrows; ++i) {
886 for (
Long idx = csr.row_offset[i]; idx < csr.row_offset[i+1]; ++idx) {
887 auto col = csr.col_index[idx];
888 Long dest = current_pos[col]++;
889 csrt.mat[dest] = csr.mat[idx];
890 csrt.col_index[dest] = IO(i);
898template <
typename T,
template<
typename>
class Allocator>
900 : m_partition(std::move(partition)),
901 m_row_begin(m_partition[ParallelContext::MyProcSub()]),
902 m_row_end(m_partition[ParallelContext::MyProcSub()+1])
907template <
typename T,
template<
typename>
class Allocator>
909 : m_partition(std::move(partition)),
910 m_row_begin(m_partition[ParallelContext::MyProcSub()]),
911 m_row_end(m_partition[ParallelContext::MyProcSub()+1]),
913 m_csr(std::move(csr))
916template <
typename T,
template<
typename>
class Allocator>
919 m_partition = std::move(partition);
923 define_doit(nnz_per_row);
926template <
typename T,
template<
typename>
class Allocator>
932 m_partition = std::move(partition);
936 m_csr = std::move(csr);
938 if (! is_sorted) { m_csr.
sort(); }
941template <
typename T,
template<
typename>
class Allocator>
947 nnz_per_row = std::max(nnz_per_row, 0);
948 Long nlocalrows = this->numLocalRows();
949 m_nnz = nlocalrows*nnz_per_row;
950 m_csr.
mat.resize(m_nnz);
958 poffset[lrow] = lrow*nnz_per_row;
962template <
typename T,
template<
typename>
class Allocator>
965 Long const* col_index,
Long nentries,
971 m_partition = std::move(partition);
980 Long nlocalrows = this->numLocalRows();
981 m_csr.
mat.resize(nentries);
984 m_csr.
nnz = nentries;
991 if (nentries <
Long(std::numeric_limits<int>::max())) {
992 define_and_filter_doit<int>(mat, col_index, nentries, row_offset);
994 define_and_filter_doit<Long>(mat, col_index, nentries, row_offset);
1009template <
typename T,
template<
typename>
class Allocator>
1017template <
typename T,
template<
typename>
class Allocator>
1021 split_csr(col_partition);
1024template <
typename T,
template<
typename>
class Allocator>
1025template <
typename I>
1028 Long nentries,
Long const* row_offset)
1031 auto* ps = psum.
data();
1032 m_nnz = Scan::PrefixSum<I>(I(nentries),
1034 return col_index[i] >= 0 && mat[i] != 0; },
1038 Long nlocalrows = this->numLocalRows();
1039 m_csr.
mat.resize(m_nnz);
1043 auto* pmat = m_csr.
mat.data();
1046 auto actual_nnz = m_nnz;
1050 if (col_index[i] >= 0 && mat[i] != 0) {
1051 pmat[ps[i]] = mat[i];
1052 pcol[ps[i]] = col_index[i];
1055 if (i <= nlocalrows) {
1056 prow[i] = (i < nlocalrows && row_offset[i] < nentries)
1057 ?
Long(ps[row_offset[i]]) : actual_nnz;
1063template <
typename T,
template<
typename>
class Allocator>
1068 ofs << m_row_begin <<
" " << m_row_end <<
" " << m_nnz <<
"\n";
1069 Long const nrows = numLocalRows();
1077 auto const& csr = m_csr;
1079 for (
Long i = 0; i < nrows; ++i) {
1081 ofs << i+m_row_begin <<
" " << csr.
col_index[j] <<
" " << csr.
mat[j] <<
"\n";
1090# ifdef AMREX_USE_MPI
1098 auto const& csr = m_csr_local;
1099# ifdef AMREX_USE_MPI
1100 auto const& csr_r = m_csr_remote;
1101 auto const& ri_ltor = m_ri_ltor;
1105 for (
Long i = 0; i < nrows; ++i) {
1107 ofs << i+m_row_begin <<
" " <<
Long(csr.
col_index[j])+m_col_begin
1108 <<
" " << csr.
mat[j] <<
"\n";
1111 if (i <
Long(ri_ltor.
size()) && ri_ltor[i] >= 0) {
1112 Long ii = ri_ltor[i];
1114 ofs << i+m_row_begin <<
" " << m_remote_cols_v[csr_r.
col_index[j]]
1115 <<
" " << csr_r.
mat[j] <<
"\n";
1122template <
typename T,
template<
typename>
class Allocator>
1123template <
typename F>
1131 Long nlocalrows = this->numLocalRows();
1132 Long rowbegin = this->globalRowBegin();
1133 auto* pmat = m_csr.
mat.data();
1134 auto* pcolindex = m_csr.
col_index.data();
1138 f(rowbegin+lrow, pcolindex+prowoffset[lrow], pmat+prowoffset[lrow]);
1141 if (! is_sorted) { m_csr.
sort(); }
1144template <
typename T,
template<
typename>
class Allocator>
1147 if (m_diagonal.
empty()) {
1151 "SpMatrix::diagonalVector: column partition differs from row partition");
1153 m_diagonal.
define(this->partition());
1157 detail::extract_diagonal(m_diagonal.
data(), m_csr.
const_view(), m_row_begin);
1163template <
typename T,
template<
typename>
class Allocator>
1173 for (
auto idx = c.row_offset[i]; idx < c.row_offset[i+1]; ++idx) {
1180 auto const& a = this->const_parcsr();
1184 for (
auto idx = a.csr0.row_offset[i];
1185 idx < a.csr0.row_offset[i+1]; ++idx) {
1186 s += a.csr0.mat[idx];
1188 if (detail::has_remote_row(a, i)) {
1189 auto ii = a.row_map[i];
1190 for (
auto idx = a.csr1.row_offset[ii];
1191 idx < a.csr1.row_offset[ii+1]; ++idx) {
1192 s += a.csr1.mat[idx];
1200template <
typename T,
template<
typename>
class Allocator>
1206 m_csr_remote.
view(),
1214# ifdef AMREX_USE_GPU
1215 m_remote_cols_dv.
data()
1217 m_remote_cols_v.data()
1225template <
typename T,
template<
typename>
class Allocator>
1240# ifdef AMREX_USE_GPU
1241 m_remote_cols_dv.
data()
1243 m_remote_cols_v.data()
1251template <
typename T,
template<
typename>
class Allocator>
1254 return this->const_parcsr();
1257template <
typename T,
template<
typename>
class Allocator>
1261#ifndef AMREX_USE_MPI
1264 if (detail::spmat_comm_is_local(this->partition(),
x.partition())) {
return; }
1266 this->prepare_comm_mv(
x.partition());
1272 auto const nrecvs =
int(m_comm_mv.recv_from.size());
1276 auto* p_recv = m_comm_mv.recv_buffer;
1277 for (
int irecv = 0; irecv < nrecvs; ++irecv) {
1278 BL_MPI_REQUIRE(MPI_Irecv(p_recv,
1279 m_comm_mv.recv_counts[irecv], mpi_t_type,
1280 m_comm_mv.recv_from[irecv], mpi_tag, mpi_comm,
1281 &(m_comm_mv.recv_reqs[irecv])));
1282 p_recv += m_comm_mv.recv_counts[irecv];
1284 AMREX_ASSERT(p_recv == m_comm_mv.recv_buffer + m_comm_mv.total_counts_recv);
1287 auto const nsends =
int(m_comm_mv.send_to.size());
1295 auto* p_send = m_comm_mv.send_buffer;
1296 for (
int isend = 0; isend < nsends; ++isend) {
1297 auto count = m_comm_mv.send_counts[isend];
1298 BL_MPI_REQUIRE(MPI_Isend(p_send, count, mpi_t_type, m_comm_mv.send_to[isend],
1299 mpi_tag, mpi_comm, &(m_comm_mv.send_reqs[isend])));
1302 AMREX_ASSERT(p_send == m_comm_mv.send_buffer + m_comm_mv.total_counts_send);
1307template <
typename T,
template<
typename>
class Allocator>
1311#ifndef AMREX_USE_MPI
1314 if (detail::spmat_comm_is_local(this->partition(), m_col_partition)) {
return; }
1316 if ( ! m_comm_mv.recv_reqs.empty()) {
1318 BL_MPI_REQUIRE(MPI_Waitall(
int(m_comm_mv.recv_reqs.size()),
1319 m_comm_mv.recv_reqs.data(),
1320 mpi_statuses.data()));
1323 unpack_buffer_mv(
y);
1325 if ( ! m_comm_mv.send_reqs.empty()) {
1327 BL_MPI_REQUIRE(MPI_Waitall(
int(m_comm_mv.send_reqs.size()),
1328 m_comm_mv.send_reqs.data(),
1329 mpi_statuses.data()));
1335 m_comm_mv.send_reqs.clear();
1336 m_comm_mv.recv_reqs.clear();
1340template <
typename T,
template<
typename>
class Allocator>
1341template <
typename U>
1349template <
typename T,
template<
typename>
class Allocator>
1350template <
typename U>
1354#ifndef AMREX_USE_MPI
1359 if (detail::spmat_comm_is_local(this->partition(), m_col_partition)) { r.
clear();
return; }
1361 this->prepare_comm_mv(m_col_partition);
1367 auto const nrecvs =
int(m_comm_mv.recv_from.size());
1368 U* recv_buffer =
nullptr;
1372 auto* p_recv = recv_buffer;
1373 for (
int irecv = 0; irecv < nrecvs; ++irecv) {
1374 BL_MPI_REQUIRE(MPI_Irecv(p_recv, m_comm_mv.recv_counts[irecv], mpi_type,
1375 m_comm_mv.recv_from[irecv], mpi_tag, mpi_comm,
1376 &recv_reqs[irecv]));
1377 p_recv += m_comm_mv.recv_counts[irecv];
1381 auto const nsends =
int(m_comm_mv.send_to.size());
1382 U* send_buffer =
nullptr;
1388 auto const col_begin = m_col_begin;
1391 pdst[i] =
x[pidx[i]-col_begin];
1395 auto* p_send = send_buffer;
1396 for (
int isend = 0; isend < nsends; ++isend) {
1397 auto count = m_comm_mv.send_counts[isend];
1398 BL_MPI_REQUIRE(MPI_Isend(p_send, count, mpi_type, m_comm_mv.send_to[isend],
1399 mpi_tag, mpi_comm, &send_reqs[isend]));
1404 r.
resize(m_comm_mv.total_counts_recv);
1407 BL_MPI_REQUIRE(MPI_Waitall(nrecvs, recv_reqs.data(), mpi_statuses.data()));
1418 BL_MPI_REQUIRE(MPI_Waitall(nsends, send_reqs.data(), mpi_statuses.data()));
1427template <
typename T,
template<
typename>
class Allocator>
1432 if (detail::spmat_comm_is_local(this->partition(), col_partition)) {
return; }
1434 this->split_csr(col_partition);
1443 if (m_csr_remote.
nnz > 0) {
1444 m_comm_tr.csrt.nnz = m_csr_remote.
nnz;
1445 m_comm_tr.csrt.nrows = m_remote_cols_v.
size();
1447 (
sizeof(T)*m_comm_tr.csrt.nnz);
1449 (
sizeof(
Long)*m_comm_tr.csrt.nnz);
1451 (
sizeof(
Long)*(m_comm_tr.csrt.nrows+1));
1454 csr_comm.
resize(m_comm_tr.csrt.nrows, m_comm_tr.csrt.nnz);
1455 auto const& csrv_comm = csr_comm.
view();
1457 auto const& csrv_comm = m_comm_tr.csrt;
1459 detail::transpose(csrv_comm, m_csr_remote.
const_view());
1460 auto row_begin = m_row_begin;
1461 auto ri_rtol = m_ri_rtol.
data();
1462 auto* col_index = csrv_comm.col_index;
1465 auto gjt =ri_rtol[col_index[idx]] + row_begin;
1466 col_index[idx] = gjt;
1471 csrv_comm. mat + csrv_comm.nnz,
1472 m_comm_tr.csrt.mat);
1474 csrv_comm. col_index,
1475 csrv_comm. col_index + csrv_comm.nnz,
1476 m_comm_tr.csrt.col_index);
1478 csrv_comm. row_offset,
1479 csrv_comm. row_offset + csrv_comm.nrows+1,
1480 m_comm_tr.csrt.row_offset);
1485 if (m_num_neighbors < 0) { set_num_neighbors(); }
1491 mpi_requests.reserve(nprocs);
1492 if (m_csr_remote.
nnz > 0) {
1494 for (
int iproc = 0; iproc < nprocs; ++iproc) {
1496 for (
Long i = 0; i <
Long(m_remote_cols_vv[iproc].size()); ++i) {
1497 n += m_comm_tr.csrt.row_offset[it+1] - m_comm_tr.csrt.row_offset[it];
1503 std::array<int,2> nn{
int(n),
int(m_remote_cols_vv[iproc].size())};
1504 BL_MPI_REQUIRE(MPI_Isend(nn.data(), 2, MPI_INT, iproc, mpi_tag,
1505 mpi_comm, &(mpi_requests.back())));
1506 m_comm_tr.send_to.push_back(iproc);
1507 m_comm_tr.send_counts.push_back(nn);
1515 for (
int irecv = 0; irecv < m_num_neighbors; ++irecv) {
1517 BL_MPI_REQUIRE(MPI_Probe(MPI_ANY_SOURCE, mpi_tag, mpi_comm, &mpi_status));
1518 int sender = mpi_status.MPI_SOURCE;
1519 std::array<int,2> nn;
1520 BL_MPI_REQUIRE(MPI_Recv(nn.data(), 2, MPI_INT, sender, mpi_tag,
1521 mpi_comm, &mpi_status));
1522 m_comm_tr.recv_from.push_back(sender);
1523 m_comm_tr.recv_counts.push_back(nn);
1524 m_comm_tr.total_counts_recv[0] += nn[0];
1525 m_comm_tr.total_counts_recv[1] += nn[1];
1528 if (! mpi_requests.empty()) {
1530 BL_MPI_REQUIRE(MPI_Waitall(
int(mpi_requests.
size()), mpi_requests.data(),
1531 mpi_statuses.data()));
1543 auto const nrecvs =
int(m_comm_tr.recv_from.size());
1546 (
sizeof(T) * m_comm_tr.total_counts_recv[0]);
1548 (
sizeof(
Long) * m_comm_tr.total_counts_recv[0]);
1550 (
sizeof(
Long) * (m_comm_tr.total_counts_recv[1]+nrecvs));
1552 (
sizeof(
Long) * m_comm_tr.total_counts_recv[1]);
1553 m_comm_tr.recv_buffer_offset.push_back({0,0,0,0});
1555 for (
int irecv = 0; irecv < nrecvs; ++irecv) {
1556 auto [os0, os1, os2, os3] = m_comm_tr.recv_buffer_offset.back();
1557 auto [n0, n1] = m_comm_tr.recv_counts[irecv];
1558 auto recv_from_rank = m_comm_tr.recv_from[irecv];
1559 BL_MPI_REQUIRE(MPI_Irecv(m_comm_tr.recv_buffer_mat + os0,
1565 &(m_comm_tr.recv_reqs[irecv*4])));
1566 BL_MPI_REQUIRE(MPI_Irecv(m_comm_tr.recv_buffer_col_index + os1,
1572 &(m_comm_tr.recv_reqs[irecv*4+1])));
1573 BL_MPI_REQUIRE(MPI_Irecv(m_comm_tr.recv_buffer_row_offset + os2,
1579 &(m_comm_tr.recv_reqs[irecv*4+2])));
1580 BL_MPI_REQUIRE(MPI_Irecv(m_comm_tr.recv_buffer_idx_map + os3,
1586 &(m_comm_tr.recv_reqs[irecv*4+3])));
1587 m_comm_tr.recv_buffer_offset.push_back({os0 + n0,
1594 auto const nsends =
int(m_comm_tr.send_to.size());
1597 Long os0 = 0, os1 = 0;
1598 for (
int isend = 0; isend < nsends; ++isend) {
1599 auto [n0, n1] = m_comm_tr.send_counts[isend];
1600 auto send_to_rank = m_comm_tr.send_to[isend];
1601 BL_MPI_REQUIRE(MPI_Isend(m_comm_tr.csrt.mat + os0,
1607 &(m_comm_tr.send_reqs[isend*4])));
1608 BL_MPI_REQUIRE(MPI_Isend(m_comm_tr.csrt.col_index + os0,
1614 &(m_comm_tr.send_reqs[isend*4+1])));
1615 BL_MPI_REQUIRE(MPI_Isend(m_comm_tr.csrt.row_offset + os1,
1621 &(m_comm_tr.send_reqs[isend*4+2])));
1622 BL_MPI_REQUIRE(MPI_Isend(m_remote_cols_vv[send_to_rank].data(),
1628 &(m_comm_tr.send_reqs[isend*4+3])));
1638template <
typename T,
template<
typename>
class Allocator>
1643 if (detail::spmat_comm_is_local(this->partition(), AT.
partition())) {
return; }
1645 this->comm_tr_recv_wait();
1649 this->comm_tr_clear();
1657template <
typename T,
template<
typename>
class Allocator>
1660 if (! m_comm_tr.recv_reqs.empty()) {
1662 BL_MPI_REQUIRE(MPI_Waitall(
int(m_comm_tr.recv_reqs.size()),
1663 m_comm_tr.recv_reqs.data(),
1664 mpi_statuses.data()));
1668template <
typename T,
template<
typename>
class Allocator>
1671 if (! m_comm_tr.send_reqs.empty()) {
1673 BL_MPI_REQUIRE(MPI_Waitall(
int(m_comm_tr.send_reqs.size()),
1674 m_comm_tr.send_reqs.data(),
1675 mpi_statuses.data()));
1678 if (m_comm_tr.csrt.nnz > 0) {
1683 if (m_comm_tr.recv_buffer_mat) {
1694template <
typename T,
template<
typename>
class Allocator>
1698 Long const* remote_cols,
Long nremote)
1704 m_partition = std::move(partition);
1708 AMREX_ASSERT(nremote >= 0 && (nremote == 0 || remote_cols !=
nullptr));
1710 m_col_partition = col_partition;
1714 m_nnz = compact.nnz;
1717 Long const nlocalrows = this->numLocalRows();
1718 Long const nnz = compact.nnz;
1720#if defined(AMREX_USE_MPI) && !defined(AMREX_USE_GPU)
1725 int const nl =
int(nlocal);
1729 Long remote_nnz = 0;
1730 for (
Long i = 0; i < nnz; ++i) { remote_nnz += (pcol[i] >= nl); }
1735 m_csr_remote.
mat.resize(remote_nnz);
1737 Long nloc = 0, nrem = 0;
1738 for (
Long i = 0; i < nlocalrows; ++i) {
1739 Long const b = pro[i];
1740 Long const e = pro[i+1];
1742 pro_r[i] =
int(nrem);
1743 for (
Long idx = b; idx < e; ++idx) {
1744 int const c = pcol[idx];
1747 pmat[nloc] = pmat[idx];
1750 pcol_r[nrem] = remote_cols[c - nl];
1751 pmat_r[nrem] = pmat[idx];
1756 pro[nlocalrows] =
int(nloc);
1757 pro_r[nlocalrows] =
int(nrem);
1758 compact.col_index.resize(nloc);
1759 compact.mat.resize(nloc);
1761 m_csr_local = std::move(compact);
1762 if (remote_nnz > 0) {
1763 m_csr_remote.
col_index.resize(remote_nnz);
1764 m_csr_remote.
row_offset = std::move(remote_row_offset);
1765 m_csr_remote.
nnz = remote_nnz;
1770 update_remote_col_index(remote_gcols, m_csr_remote.
col_index,
true);
1779 auto const* pcol = compact.col_index.
data();
1781 int const nl =
int(nlocal);
1782 Long local_nnz = nnz;
1785 auto* p_pfsum = pfsum.
data();
1786 local_nnz = Scan::PrefixSum<Long>(nnz,
1791 auto const* p_pfsum = pfsum.
data();
1792 Long const remote_nnz = nnz - local_nnz;
1794#ifndef AMREX_USE_MPI
1796 "SpMatrix::define_split: remote columns without MPI");
1799 if (remote_nnz == 0) {
1800 m_csr_local = std::move(compact);
1802 m_csr_local.
resize(nlocalrows, local_nnz);
1806 auto const* pmat = compact.mat.data();
1807 auto* pmat_l = m_csr_local.
mat.data();
1808 auto* pcol_l = m_csr_local.
col_index.data();
1809 auto* pmat_r = remote_mat.
data();
1810 auto* pcol_r = remote_gcols.
data();
1813 auto ps = p_pfsum[i];
1815 pmat_l[ps] = pmat[i];
1816 pcol_l[ps] = pcol[i];
1818 pmat_r[i-ps] = pmat[i];
1819 pcol_r[i-ps] = remote_cols[pcol[i] - nl];
1822 auto const noffset = nlocalrows+1;
1823 auto const* pro = compact.row_offset.data();
1825 auto* pro_r = remote_row_offset.
data();
1828 Long ro_l = (i < noffset-1)
1829 ? ((pro[i] < nnz) ? p_pfsum[pro[i]] : local_nnz) : local_nnz;
1830 pro_l[i] =
int(ro_l);
1831 pro_r[i] =
int(pro[i] - ro_l);
1836 m_csr_remote.
mat = std::move(remote_mat);
1837 m_csr_remote.
col_index.resize(remote_nnz);
1838 m_csr_remote.
row_offset = std::move(remote_row_offset);
1839 m_csr_remote.
nnz = remote_nnz;
1841 update_remote_col_index(remote_gcols, m_csr_remote.
col_index,
true);
1845 if (remote_nnz == 0) {
1854template <
typename T,
template<
typename>
class Allocator>
1867 m_col_partition = col_partition;
1871 (m_col_end - m_col_begin <
Long(std::numeric_limits<int>::max()),
1872 "SpMatrix: the local column block is too large for 32-bit indices");
1880 Long const nlocalrows = this->numLocalRows();
1882 if (m_col_begin == 0 && m_col_end == col_partition.
numGlobalRows()) {
1885 "SpMatrix: too many local nonzeros for 32-bit row offsets");
1886 Long const nnz = m_nnz;
1890 auto const* pcol = m_csr.
col_index.data();
1892 auto* pcol_l = m_csr_local.
col_index.data();
1896 if (i < nnz) { pcol_l[i] =
int(pcol[i]); }
1897 if (i <= nlocalrows) { pro_l[i] = (i < noffset) ?
int(pro[i]) : 0; }
1900 m_csr_local.
mat = std::move(m_csr.
mat);
1901 m_csr_local.
nnz = m_nnz;
1910#if defined(AMREX_USE_MPI) && !defined(AMREX_USE_GPU)
1913 Long const nnz = m_nnz;
1918 Long remote_nnz = 0;
1919 for (
Long i = 0; i < nnz; ++i) {
1920 remote_nnz += (pcol[i] < m_col_begin || pcol[i] >= m_col_end);
1923 (nnz - remote_nnz <
Long(std::numeric_limits<int>::max()) &&
1924 remote_nnz <
Long(std::numeric_limits<int>::max()),
1925 "SpMatrix: too many local nonzeros for 32-bit row offsets");
1926 m_csr_local.
col_index.resize(nnz - remote_nnz);
1936 Long nloc = 0, nrem = 0;
1937 for (
Long i = 0; i < nlocalrows; ++i) {
1938 pro_l[i] =
int(nloc);
1939 pro_r[i] =
int(nrem);
1940 if (i+1 >= noffset) {
continue; }
1941 for (
Long idx = pro[i]; idx < pro[i+1]; ++idx) {
1942 Long const c = pcol[idx];
1943 if (c >= m_col_begin && c < m_col_end) {
1944 pcol_l[nloc] =
int(c - m_col_begin);
1945 pmat[nloc] = pmat[idx];
1949 pmat_r[nrem] = pmat[idx];
1954 pro_l[nlocalrows] =
int(nloc);
1955 pro_r[nlocalrows] =
int(nrem);
1956 m_csr.
mat.resize(nloc);
1957 m_csr_local.
mat = std::move(m_csr.
mat);
1958 m_csr_local.
nnz = nloc;
1960 if (remote_nnz > 0) {
1961 m_csr_remote.
mat = std::move(remote_mat);
1962 m_csr_remote.
col_index.resize(remote_nnz);
1963 m_csr_remote.
row_offset = std::move(remote_row_offset);
1964 m_csr_remote.
nnz = remote_nnz;
1967 update_remote_col_index(remote_gcols, m_csr_remote.
col_index,
true);
1975 auto* p_pfsum = pfsum.
data();
1976 auto col_begin = m_col_begin;
1977 auto col_end = m_col_end;
1978 auto const* pcol = m_csr.
col_index.data();
1979 if (m_csr.
nnz <
Long(std::numeric_limits<int>::max())) {
1980 local_nnz = Scan::PrefixSum<int>(
int(m_nnz),
1982 return (pcol[i] >= col_begin &&
1983 pcol[i] < col_end); },
1988 local_nnz = Scan::PrefixSum<Long>(m_nnz,
1990 return (pcol[i] >= col_begin &&
1991 pcol[i] < col_end); },
1996 Long const remote_nnz = m_nnz - local_nnz;
1998 (local_nnz <
Long(std::numeric_limits<int>::max()) &&
1999 remote_nnz <
Long(std::numeric_limits<int>::max()),
2000 "SpMatrix: too many local nonzeros for 32-bit row offsets");
2002 m_csr_local.
resize(nlocalrows, local_nnz);
2007 remote_gcols.
resize(remote_nnz);
2008 remote_mat.
resize(remote_nnz);
2009 remote_row_offset.
resize(nlocalrows+1);
2012 "SpMatrix: column index outside the column partition");
2015 auto const* pmat = m_csr.
mat.data();
2016 auto* pmat_l = m_csr_local.
mat.data();
2017 auto* pcol_l = m_csr_local.
col_index.data();
2018 auto* pmat_r = remote_mat.
data();
2019 auto* pcol_r = remote_gcols.
data();
2022 auto ps = p_pfsum[i];
2023 auto local = (pcol[i] >= col_begin && pcol[i] < col_end);
2025 pmat_l[ps] = pmat[i];
2026 pcol_l[ps] =
int(pcol[i] - col_begin);
2027 }
else if (pmat_r) {
2028 pmat_r[i-ps] = pmat[i];
2029 pcol_r[i-ps] = pcol[i];
2035 auto* pro_r = remote_row_offset.
data();
2036 auto total_nnz = m_nnz;
2040 if (i < noffset-1) {
2041 ro_l = (pro[i] < total_nnz) ? p_pfsum[pro[i]] : local_nnz;
2045 pro_l[i] =
int(ro_l);
2046 if (pro_r) { pro_r[i] =
int(pro[i] - ro_l); }
2054 if (remote_nnz > 0) {
2055 m_csr_remote.
mat = std::move(remote_mat);
2056 m_csr_remote.
col_index.resize(remote_nnz);
2057 m_csr_remote.
row_offset = std::move(remote_row_offset);
2058 m_csr_remote.
nnz = remote_nnz;
2061 update_remote_col_index(remote_gcols, m_csr_remote.
col_index,
true);
2069template <
typename T,
template<
typename>
class Allocator>
2077 m_ri_ltor.
resize(old_size-1);
2078 m_ri_rtol.
resize(old_size-1);
2079 auto* p_ltor = m_ri_ltor.
data();
2080 auto* p_rtol = m_ri_rtol.
data();
2082 auto const* p_ro = m_csr_remote.
row_offset.data();
2083 auto* p_tro = trimmed_row_offset.
data();
2085 Long new_size = Scan::PrefixSum<Long>(old_size,
2087 if (i+1 < old_size) {
2088 return (p_ro[i+1] > p_ro[i]);
2096 }
else if (p_ro[i] > p_ro[i-1]) {
2099 if (i+1 < old_size) {
2100 if (p_ro[i+1] > p_ro[i]) {
2109 m_ri_rtol.
resize(new_size-1);
2110 trimmed_row_offset.
resize(new_size);
2115 m_remote_row_offset = std::move(trimmed_row_offset);
2116 std::swap(m_csr_remote.
row_offset, m_remote_row_offset);
2119template <
typename T,
template<
typename>
class Allocator>
2120template <
typename VL,
typename VI>
2122 bool in_device_memory)
2128 m_remote_cols_v.clear();
2129 m_remote_cols_vv.clear();
2130 m_remote_cols_vv.resize(nprocs);
2132 m_remote_cols_dv.
clear();
2137 if (n == 0) {
return; }
2143 if (in_device_memory) {
2150 h_gcols.assign(gcols.begin(), gcols.end());
2153 m_remote_cols_v = h_gcols;
2156 (
Long(m_remote_cols_v.
size()) <
Long(std::numeric_limits<int>::max()),
2157 "SpMatrix: too many remote columns for 32-bit indices");
2160 m_remote_cols_dv.
resize(m_remote_cols_v.
size());
2162 m_remote_cols_v.begin(),
2163 m_remote_cols_v.end(),
2164 m_remote_cols_dv.
data());
2168 auto const& cp = this->m_col_partition.
dataVector();
2170 m_remote_cols_v.back() < cp.back());
2171 auto it = cp.cbegin();
2172 for (
auto c : m_remote_cols_v) {
2173 it = std::find_if(it, cp.cend(), [&] (
auto x) { return x > c; });
2174 if (it != cp.cend()) {
2175 int iproc =
int(std::distance(cp.cbegin(),it)) - 1;
2176 m_remote_cols_vv[iproc].push_back(c);
2178 amrex::Abort(
"SpMatrix::update_remote_col_index: how did this happen?");
2184 for (
Long i = 0; i < n; ++i) {
2185 h_lcols[i] =
int(std::lower_bound(m_remote_cols_v.begin(), m_remote_cols_v.end(),
2186 h_gcols[i]) - m_remote_cols_v.begin());
2190 if (in_device_memory) {
2196 std::copy(h_lcols.
begin(), h_lcols.
end(), lcols.begin());
2200template <
typename T,
template<
typename>
class Allocator>
2203 if (m_num_neighbors >= 0) {
return; }
2210 for (
int iproc = 0; iproc < nprocs; ++iproc) {
2211 connection[iproc] = m_remote_cols_vv[iproc].empty() ? 0 : 1;
2214 m_num_neighbors = 0;
2215 BL_MPI_REQUIRE(MPI_Reduce_scatter
2216 (connection.data(), &m_num_neighbors, reduce_scatter_counts.data(),
2217 mpi_int, MPI_SUM, mpi_comm));
2220template <
typename T,
template<
typename>
class Allocator>
2224 if (m_comm_mv.prepared) {
return; }
2228 this->split_csr(col_partition);
2235 if (m_num_neighbors < 0) { set_num_neighbors(); }
2238 mpi_requests.reserve(nprocs);
2239 for (
int iproc = 0; iproc < nprocs; ++iproc) {
2240 if ( ! m_remote_cols_vv[iproc].empty()) {
2242 auto const sz = m_remote_cols_vv[iproc].
size();
2243 if (sz >
static_cast<Long>(std::numeric_limits<int>::max())) {
2244 amrex::Abort(
"SpMatrix::prepare_comm_mv: remote column payload exceeds MPI int count range.");
2246 auto const msg_count =
static_cast<int>(sz);
2248 BL_MPI_REQUIRE(MPI_Isend(m_remote_cols_vv[iproc].data(),
2250 mpi_long, iproc, mpi_tag, mpi_comm,
2251 &(mpi_requests.back())));
2252 m_comm_mv.recv_from.push_back(iproc);
2253 m_comm_mv.recv_counts.push_back(msg_count);
2257 m_comm_mv.total_counts_recv =
Long(m_remote_cols_v.
size());
2260 m_comm_mv.total_counts_send = 0;
2261 for (
int isend = 0; isend < m_num_neighbors; ++isend) {
2263 BL_MPI_REQUIRE(MPI_Probe(MPI_ANY_SOURCE, mpi_tag, mpi_comm, &mpi_status));
2264 int receiver = mpi_status.MPI_SOURCE;
2266 BL_MPI_REQUIRE(MPI_Get_count(&mpi_status, mpi_long, &count));
2267 m_comm_mv.send_to.push_back(receiver);
2268 m_comm_mv.send_counts.push_back(count);
2269 send_indices[isend].resize(count);
2270 BL_MPI_REQUIRE(MPI_Recv(send_indices[isend].data(), count, mpi_long,
2271 receiver, mpi_tag, mpi_comm, &mpi_status));
2272 m_comm_mv.total_counts_send += count;
2275 m_comm_mv.send_indices.resize(m_comm_mv.total_counts_send);
2277 send_indices_all.
reserve(m_comm_mv.total_counts_send);
2278 for (
auto const& vl : send_indices) {
2284 m_comm_mv.send_indices.begin());
2287 if (! mpi_requests.empty()) {
2289 BL_MPI_REQUIRE(MPI_Waitall(
int(mpi_requests.
size()), mpi_requests.data(),
2290 mpi_statuses.data()));
2293 m_comm_mv.prepared =
true;
2296template <
typename T,
template<
typename>
class Allocator>
2299 auto*
pdst = m_comm_mv.send_buffer;
2300 auto* pidx = m_comm_mv.send_indices.data();
2301 auto const& vv = v.
view();
2302 auto const nsends =
Long(m_comm_mv.send_indices.size());
2305 pdst[i] = vv(pidx[i]);
2309template <
typename T,
template<
typename>
class Allocator>
2312 auto const& csr = m_csr_remote;
2318 auto const* rtol = m_ri_rtol.
data();
2323 auto const nrr =
Long(csr.row_offset.size())-1;
2327 for (
Long j = row[i]; j < row[i+1]; ++j) {
2328 r += mat[j] * px[col[j]];
2335template <
typename T,
template<
typename>
class Allocator>
2340 m_col_partition = col_partition;
2344 m_ri_ltor.
resize(numLocalRows(), -1);
2348 if (nnz == 0) {
return; }
2350 "SpMatrix: too many received nonzeros for 32-bit row offsets");
2367 auto& csrr = m_csr_remote;
2370 csrr.
mat.resize(nnz);
2378 std::iota(order.begin(), order.end(), 0);
2379 std::sort(order.begin(), order.end(), [&] (
int a,
int b) {
2380 return ctr.recv_from[a] < ctr.recv_from[b]; });
2384 for (
int i : order) {
2393 for (
int lr = 0; lr < nrow_i; ++lr) {
2394 Long const gr = idx_map[lr];
2395 while (p < nrows && ri_map[p] < gr) { ++p; }
2398 row_nnz[p] +=
int(row_offset[lr+1] - row_offset[lr]);
2402 std::exclusive_scan(row_nnz.begin(), row_nnz.end(), csrr.
row_offset.begin(), 0);
2407 for (
int i : order) {
2419 for (
int lr = 0; lr < nrow_i; ++lr) {
2420 Long const gr = idx_map[lr];
2421 while (p < nrows && ri_map[p] < gr) { ++p; }
2424 auto os_src = row_offset[lr] - row_offset[0];
2425 auto nvals = row_offset[lr+1] - row_offset[lr];
2426 auto os_dst = rowpos[p];
2427 std::memcpy(csrr.
mat.data()+os_dst, mat+os_src,
sizeof(T)*nvals);
2428 std::memcpy(gcols.data()+os_dst, col_index+os_src,
sizeof(
Long)*nvals);
2429 rowpos[p] +=
int(nvals);
2436 auto row_begin = m_row_begin;
2440 rtol[i] -= row_begin;
2445 update_remote_col_index(gcols, csrr.
col_index,
false);
2453template <
typename T,
template<
typename>
class Allocator>
2459 this->prepare_comm_mv(B.partition());
2461 auto const& cm = m_comm_mv;
2462 auto const nrecvs =
int(cm.recv_from.size());
2463 auto const nsends =
int(cm.send_to.size());
2469 ext.
nrows = cm.total_counts_recv;
2472 auto const b0 = B.m_csr_local.const_view();
2473 auto const b1 = B.remote_full_const_view();
2474 Long const b_row_begin = B.m_row_begin;
2475 Long const nsend_rows = cm.total_counts_send;
2482 auto* pcnt = d_cnt.
data();
2485 Long const lr = send_idx[i] - b_row_begin;
2486 pcnt[i] = (b0.row_offset[lr+1] - b0.row_offset[lr])
2487 + (b1.row_offset[lr+1] - b1.row_offset[lr]);
2497 reqs.reserve(nrecvs+nsends);
2499 for (
int i = 0; i < nrecvs; ++i) {
2501 BL_MPI_REQUIRE(MPI_Irecv(h_recv.
data()+os, cm.recv_counts[i], mpi_long,
2502 cm.recv_from[i], tag, mpi_comm, &reqs.back()));
2503 os += cm.recv_counts[i];
2506 for (
int i = 0; i < nsends; ++i) {
2508 BL_MPI_REQUIRE(MPI_Isend(h_send.
data()+os, cm.send_counts[i], mpi_long,
2509 cm.send_to[i], tag, mpi_comm, &reqs.back()));
2510 os += cm.send_counts[i];
2512 if (! reqs.empty()) {
2514 BL_MPI_REQUIRE(MPI_Waitall(
int(reqs.
size()), reqs.data(), stats.data()));
2521 for (
Long i = 0; i < n; ++i) {
2522 Long const c = v[i];
2528 to_offsets(h_recv, ext.
nrows);
2529 to_offsets(h_send, nsend_rows);
2531 Long const send_nnz = h_send[nsend_rows];
2534 for (
int i = 0, r = 0; i < nrecvs; ++i) {
2535 recv_nnz[i] = h_recv[r+cm.recv_counts[i]] - h_recv[r];
2536 r += cm.recv_counts[i];
2537 if (recv_nnz[i] >=
Long(std::numeric_limits<int>::max())) {
2538 amrex::Abort(
"SpMatrix::fetch_remote_rows_mm: message exceeds MPI int count range.");
2541 for (
int i = 0, r = 0; i < nsends; ++i) {
2542 send_nnz_v[i] = h_send[r+cm.send_counts[i]] - h_send[r];
2543 r += cm.send_counts[i];
2544 if (send_nnz_v[i] >=
Long(std::numeric_limits<int>::max())) {
2545 amrex::Abort(
"SpMatrix::fetch_remote_rows_mm: message exceeds MPI int count range.");
2561 rreqs.reserve(2*nrecvs);
2563 for (
int i = 0; i < nrecvs; ++i) {
2564 if (recv_nnz[i] > 0) {
2565 auto n =
int(recv_nnz[i]);
2567 BL_MPI_REQUIRE(MPI_Irecv(ext.
col_index+os, n, mpi_long, cm.recv_from[i],
2568 tag_c, mpi_comm, &rreqs.back()));
2570 BL_MPI_REQUIRE(MPI_Irecv(ext.
mat+os, n, mpi_t, cm.recv_from[i],
2571 tag_m, mpi_comm, &rreqs.back()));
2577 Long* send_col =
nullptr;
2578 T* send_mat =
nullptr;
2582 auto const* poff = d_send_off.
data();
2583 Long const b_col_begin = B.m_col_begin;
2585 auto const* b_rcols = B.m_remote_cols_dv.data();
2587 auto const* b_rcols = B.m_remote_cols_v.data();
2593 constexpr Long gmax = std::numeric_limits<Long>::max();
2594 Long const lr = send_idx[i] - b_row_begin;
2596 Long q0 = b0.row_offset[lr];
2597 Long q1 = b1.row_offset[lr];
2598 Long const e0 = b0.row_offset[lr+1];
2599 Long const e1 = b1.row_offset[lr+1];
2600 while (q0 < e0 || q1 < e1) {
2601 Long const g0 = (q0 < e0) ? b0.col_index[q0] + b_col_begin : gmax;
2602 Long const g1 = (q1 < e1) ? b_rcols[b1.col_index[q1]] : gmax;
2605 send_mat[p] = b0.mat[q0];
2609 send_mat[p] = b1.mat[q1];
2617 sreqs.reserve(2*nsends);
2619 for (
int i = 0; i < nsends; ++i) {
2620 if (send_nnz_v[i] > 0) {
2621 auto n =
int(send_nnz_v[i]);
2623 BL_MPI_REQUIRE(MPI_Isend(send_col+os, n, mpi_long, cm.send_to[i],
2624 tag_c, mpi_comm, &sreqs.back()));
2626 BL_MPI_REQUIRE(MPI_Isend(send_mat+os, n, mpi_t, cm.send_to[i],
2627 tag_m, mpi_comm, &sreqs.back()));
2628 os += send_nnz_v[i];
2633 if (! rreqs.empty()) {
2635 BL_MPI_REQUIRE(MPI_Waitall(
int(rreqs.
size()), rreqs.data(), stats.data()));
2637 if (! sreqs.empty()) {
2639 BL_MPI_REQUIRE(MPI_Waitall(
int(sreqs.
size()), sreqs.data(), stats.data()));
2648template <
typename T,
template<
typename>
class Allocator>
2651 if (! m_remote_row_offset.
empty()) {
return; }
2655 auto nrows = numLocalRows();
2656 m_remote_row_offset.
resize(nrows+1);
2657 auto* pro_full = m_remote_row_offset.
data();
2658 auto const* pro_comp = m_csr_remote.
row_offset.data();
2659 auto const* ri_ltor = m_ri_ltor.
data();
2660 auto nnz_r = m_csr_remote.
nnz;
2661 if (nnz_r == 0 || m_ri_ltor.
empty()) {
2666 Scan::PrefixSum<int>(nrows,
2668 Long rrow = ri_ltor[i];
2672 return int(pro_comp[rrow+1]-pro_comp[rrow]);
2684template <
typename T,
template<
typename>
class Allocator>
2687 if (m_remote_row_offset.
empty()) {
2688 expand_remote_row_offset();
2691 csr_view.row_offset = m_remote_row_offset.
data();
2692 csr_view.nrows = numLocalRows();
#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
#define AMREX_FORCE_INLINE
Definition AMReX_Extension.H:124
#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
#define AMREX_GPU_HOST_DEVICE
Definition AMReX_GpuQualifiers.H:20
Convenience header for the core AMReX GPU facilities.
Array4< int const > offset
Definition AMReX_HypreMLABecLap.cpp:1139
Real * pdst
Definition AMReX_HypreMLABecLap.cpp:1140
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
bool empty() const
True if the partition contains no rows.
Definition AMReX_AlgPartition.H:45
Vector< Long > const & dataVector() const
Underlying array describing row offsets (size nproc+1).
Definition AMReX_AlgPartition.H:72
Distributed dense vector that mirrors the layout of an AlgPartition.
Definition AMReX_AlgVector.H:29
Long numLocalRows() const
Number of entries stored on this rank.
Definition AMReX_AlgVector.H:74
bool empty() const
True if the local storage holds zero entries.
Definition AMReX_AlgVector.H:68
T const * data() const
Definition AMReX_AlgVector.H:85
void define(Long global_size)
Resize/repartition the vector to span global_size rows.
Definition AMReX_AlgVector.H:255
Table1D< T const, Long > view() const
Definition AMReX_AlgVector.H:94
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.
Dynamically allocated vector for trivially copyable data.
Definition AMReX_PODVector.H:308
void reserve(size_type a_capacity, GrowthStrategy strategy=GrowthStrategy::Poisson)
Definition AMReX_PODVector.H:819
size_type size() const noexcept
Definition AMReX_PODVector.H:654
void shrink_to_fit()
Definition AMReX_PODVector.H:826
iterator begin() noexcept
Definition AMReX_PODVector.H:680
void resize(size_type a_new_size, GrowthStrategy strategy=GrowthStrategy::Poisson)
Definition AMReX_PODVector.H:734
iterator end() noexcept
Definition AMReX_PODVector.H:684
void clear() noexcept
Definition AMReX_PODVector.H:652
T * data() noexcept
Definition AMReX_PODVector.H:672
bool empty() const noexcept
Definition AMReX_PODVector.H:658
void push_back(const T &a_value)
Definition AMReX_PODVector.H:633
Distributed CSR matrix that manages storage and GPU-friendly partitions.
Definition AMReX_SpMatrix.H:65
void finishComm_tr(SpMatrix< T, Allocator > &AT)
Complete transpose communication, writing the assembled matrix into AT.
Definition AMReX_SpMatrix.H:1639
void split_csr(AlgPartition const &col_partition)
Split into local and remote blocks with 32-bit column indices.
Definition AMReX_SpMatrix.H:1855
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 globalRowBegin() const
Inclusive global index begin on this process.
Definition AMReX_SpMatrix.H:201
void define_and_filter_doit(T const *mat, Long const *col_index, Long nentries, Long const *row_offset)
Private helper (exposed for CUDA) that copies/filters CSR arrays into device storage.
Definition AMReX_SpMatrix.H:1027
Long * rowOffset()
Don't use this beyond initial setup.
Definition AMReX_SpMatrix.H:218
CsrView< T const, int > remote_full_const_view() const
Off-diagonal part with the full (untrimmed) row offsets.
Definition AMReX_SpMatrix.H:2685
void sortCSR()
Definition AMReX_SpMatrix.H:1011
void pack_buffer_mv(AlgVector< T, AllocT > const &v)
Definition AMReX_SpMatrix.H:2297
void unpack_buffer_mv(AlgVector< T, AllocT > &v)
Definition AMReX_SpMatrix.H:2310
void startComm_tr(AlgPartition const &col_partition)
Initiate communication required to build the transpose with column partition col_partition.
Definition AMReX_SpMatrix.H:1428
void comm_tr_clear()
Definition AMReX_SpMatrix.H:1669
void define_doit(int nnz_per_row)
Private helper (exposed for CUDA) that allocates fixed-connectivity matrices with nnz_per_row entries...
Definition AMReX_SpMatrix.H:943
T * data()
Don't use this beyond initial setup.
Definition AMReX_SpMatrix.H:206
Long numGlobalRows() const
Global row count.
Definition AMReX_SpMatrix.H:196
Gpu::DeviceVector< U > gatherRemote(U const *x)
Gather values at the columns of the off-diagonal block.
Definition AMReX_SpMatrix.H:1342
SpMatrix & operator=(SpMatrix const &)=delete
T value_type
Definition AMReX_SpMatrix.H:67
RemoteRowsMM fetch_remote_rows_mm(SpMatrix< T, Allocator > const &B)
Definition AMReX_SpMatrix.H:2454
friend SpMatrix< U, M > SpGEMM(SpMatrix< U, M > const &A, SpMatrix< U, M > const &B, AlgPartition const &col_partition, F const &row_post)
Allocator< U > allocator_type
Definition AMReX_SpMatrix.H:68
Long globalRowEnd() const
Exclusive global index end on this process.
Definition AMReX_SpMatrix.H:203
void trim_remote_rows()
Definition AMReX_SpMatrix.H:2070
Long numLocalNonZeros() const
Number of nonzeros stored locally.
Definition AMReX_SpMatrix.H:198
struct amrex::SpMatrix::CommMV m_comm_mv
AlgVector< T, AllocT > rowSum() const
Sum the values in each local row and return the result as an AlgVector.
Definition AMReX_SpMatrix.H:1164
void define_split(AlgPartition partition, AlgPartition const &col_partition, local_csr_type &&compact, Long nlocal, Long const *remote_cols, Long nremote)
Define directly in split form from a CSR with compact 32-bit columns: c < nlocal is the local column ...
Definition AMReX_SpMatrix.H:1695
void update_remote_col_index(VL const &gcols, VI &lcols, bool in_device_memory)
Definition AMReX_SpMatrix.H:2121
AlgPartition const & columnPartition() const
Return the column partition used for matrix-vector and matrix-matrix multiplications.
Definition AMReX_SpMatrix.H:191
void printToFile(std::string const &file) const
Definition AMReX_SpMatrix.H:1065
ParCsr< T const > const_parcsr() const
Const-qualified alias of parcsr() for convenience.
Definition AMReX_SpMatrix.H:1226
SpMatrix(SpMatrix const &)=delete
ParCsr< T > parcsr()
Build GPU-friendly CSR views split into diagonal/off-diagonal blocks.
Definition AMReX_SpMatrix.H:1201
SpMatrix(SpMatrix &&)=default
Long numLocalRows() const
Number of rows owned by this rank.
Definition AMReX_SpMatrix.H:194
void define(AlgPartition partition, csr_type csr, CsrSorted is_sorted)
Define a default-constructed matrix from a given CSR.
Definition AMReX_SpMatrix.H:927
AlgPartition const & partition() const
Row partition describing how matrix rows are distributed across ranks.
Definition AMReX_SpMatrix.H:182
SpMatrix(AlgPartition partition, csr_type csr)
Construct a sparse matrix from a given Partition and CSR.
Definition AMReX_SpMatrix.H:908
friend SpMatrix< U, M > RAP(SpMatrix< U, M > const &R, SpMatrix< U, M > const &A, SpMatrix< U, M > const &P, AlgPartition const &col_partition)
AlgVector< T, AllocT > const & diagonalVector() const
Return diagonal elements in a square matrix.
Definition AMReX_SpMatrix.H:1145
Long * columnIndex()
Don't use this beyond initial setup.
Definition AMReX_SpMatrix.H:212
void comm_tr_recv_wait()
Definition AMReX_SpMatrix.H:1658
void define(AlgPartition partition, T const *mat, Long const *col_index, Long nentries, Long const *row_offset, CsrSorted is_sorted, CsrValid is_valid)
Define a default-constructed matrix from given CSR arrays.
Definition AMReX_SpMatrix.H:964
void finishComm_mv(AlgVector< T, AllocT > &y)
Finish halo exchanges and accumulate contributions into y.
Definition AMReX_SpMatrix.H:1308
Allocator< T > AllocT
Definition AMReX_SpMatrix.H:72
friend SpMatrix< U, M > SpGEMM(SpMatrix< U, M > const &A, SpMatrix< U, M > const &B, AlgPartition const &col_partition)
struct amrex::SpMatrix::CommTR m_comm_tr
void expand_remote_row_offset() const
Definition AMReX_SpMatrix.H:2649
void gatherRemote(U const *x, Gpu::DeviceVector< U > &r)
As above, into r (reused across calls).
Definition AMReX_SpMatrix.H:1351
void prepare_comm_mv(AlgPartition const &col_partition)
Definition AMReX_SpMatrix.H:2221
ParCsr< T const > parcsr() const
Const variant of parcsr().
Definition AMReX_SpMatrix.H:1252
void startComm_mv(AlgVector< T, AllocT > const &x)
Prepare halo exchanges for a subsequent SpMV using x as the source vector.
Definition AMReX_SpMatrix.H:1258
friend SpMatrix< U, M > transpose(SpMatrix< U, M > const &A, AlgPartition const &col_partition)
void unpack_buffer_tr(CommTR const &ctr, AlgPartition const &col_partition)
Definition AMReX_SpMatrix.H:2336
friend void SpMV(AlgVector< U, N > &y, SpMatrix< U, M > const &A, AlgVector< U, N > const &x)
SpMatrix(AlgPartition partition, int nnz_per_row)
Construct a sparse matrix with a fixed number of nonzeros per row.
Definition AMReX_SpMatrix.H:899
void setVal(F const &f, CsrSorted is_sorted)
Initialize matrix entries using a row-wise functor.
Definition AMReX_SpMatrix.H:1124
void define(AlgPartition partition, int nnz_per_row)
Allocate storage for a default-constructed matrix with a fixed number of nonzeros per row.
Definition AMReX_SpMatrix.H:917
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_Comms_Arena()
Definition AMReX_Arena.cpp:889
Arena * The_Pinned_Arena()
Definition AMReX_Arena.cpp:869
Arena * The_Async_Arena()
Definition AMReX_Arena.cpp:839
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 DeviceToHost deviceToHost
Definition AMReX_GpuContainers.H:106
static constexpr HostToDevice hostToDevice
Definition AMReX_GpuContainers.H:105
void streamSynchronize() noexcept
Definition AMReX_GpuDevice.H:310
gpuStream_t gpuStream() noexcept
Definition AMReX_GpuDevice.H:291
MPI_Comm CommunicatorSub() noexcept
sub-communicator for current frame
Definition AMReX_ParallelContext.H:70
int MyProcSub() noexcept
my sub-rank in current frame
Definition AMReX_ParallelContext.H:76
int NProcsSub() noexcept
number of ranks in current frame
Definition AMReX_ParallelContext.H:74
int SeqNum() noexcept
Returns sequential message sequence numbers, usually used as tags for send/recv.
Definition AMReX_ParallelDescriptor.H:678
static constexpr struct amrex::Scan::Type::Exclusive exclusive
static constexpr struct amrex::Scan::Type::Inclusive inclusive
static constexpr RetSum noRetSum
Definition AMReX_Scan.H:35
static constexpr RetSum retSum
Definition AMReX_Scan.H:34
static constexpr int MPI_REQUEST_NULL
Definition AMReX_ccse-mpi.H:57
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
amrex::ArenaAllocator< T > DefaultAllocator
Definition AMReX_GpuAllocators.H:205
void duplicateCSR(C c, CSR< T, AD, I > &dst, CSR< T, AS, I > const &src)
Definition AMReX_CSR.H:125
void ParallelFor(TypeList< CTOs... > ctos, std::array< int, sizeof...(CTOs)> const &runtime_options, T N, F &&f)
Definition AMReX_CTOParallelForImpl.H:202
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
CsrView< T const, I > const_view() const
Convenience alias for view() const.
Definition AMReX_CSR.H:97
V< I > row_offset
Definition AMReX_CSR.H:57
CsrView< T, I > view()
Mutable view of the underlying buffers.
Definition AMReX_CSR.H:82
Long nnz
Definition AMReX_CSR.H:58
V< I > col_index
Definition AMReX_CSR.H:56
void resize(Long num_rows, Long num_non_zeros)
Resize the storage to accommodate num_rows and num_non_zeros entries.
Definition AMReX_CSR.H:74
void sort()
Sort each row by column index. Uses GPU acceleration when possible.
Definition AMReX_CSR.H:146
V< T > mat
Definition AMReX_CSR.H:55
Sorted CSR means for each row the column indices are sorted.
Definition AMReX_SpMatrix.H:49
bool b
Definition AMReX_SpMatrix.H:50
Valid CSR means all entries are valid. It may be sorted ro unsorted.
Definition AMReX_SpMatrix.H:55
bool b
Definition AMReX_SpMatrix.H:56
Lightweight non-owning CSR view that can point to host or device buffers.
Definition AMReX_CSR.H:35
Long nnz
Definition AMReX_CSR.H:41
Definition AMReX_SpMatrix.H:39
Long const *__restrict__ col_map
Definition AMReX_SpMatrix.H:45
CsrView< T, int > csr1
Definition AMReX_SpMatrix.H:41
Long const *__restrict__ row_map
Definition AMReX_SpMatrix.H:44
Long col_begin
Definition AMReX_SpMatrix.H:43
Long row_begin
Definition AMReX_SpMatrix.H:42
CsrView< T, int > csr0
Definition AMReX_SpMatrix.H:40
static MPI_Datatype type()
Definition AMReX_SpMatrix.H:446
T * send_buffer
Definition AMReX_SpMatrix.H:455
bool prepared
Definition AMReX_SpMatrix.H:462
Vector< int > recv_counts
Definition AMReX_SpMatrix.H:452
Long total_counts_recv
Definition AMReX_SpMatrix.H:460
Vector< int > recv_from
Definition AMReX_SpMatrix.H:451
T * recv_buffer
Definition AMReX_SpMatrix.H:459
Vector< int > send_counts
Definition AMReX_SpMatrix.H:448
Long total_counts_send
Definition AMReX_SpMatrix.H:456
Gpu::DeviceVector< Long > send_indices
Definition AMReX_SpMatrix.H:449
Vector< MPI_Request > recv_reqs
Definition AMReX_SpMatrix.H:458
Vector< int > send_to
Definition AMReX_SpMatrix.H:447
Vector< MPI_Request > send_reqs
Definition AMReX_SpMatrix.H:454
Definition AMReX_SpMatrix.H:465
Vector< std::array< int, 2 > > send_counts
Definition AMReX_SpMatrix.H:469
Long * recv_buffer_col_index
Definition AMReX_SpMatrix.H:482
Vector< MPI_Request > send_reqs
Definition AMReX_SpMatrix.H:470
Vector< MPI_Request > recv_reqs
Definition AMReX_SpMatrix.H:474
Vector< int > send_to
Definition AMReX_SpMatrix.H:468
std::array< Long, 2 > total_counts_recv
Definition AMReX_SpMatrix.H:476
Long * recv_buffer_row_offset
Definition AMReX_SpMatrix.H:483
Vector< std::array< int, 2 > > recv_counts
Definition AMReX_SpMatrix.H:473
CsrView< T > csrt
Definition AMReX_SpMatrix.H:466
T * recv_buffer_mat
Definition AMReX_SpMatrix.H:481
Vector< std::array< Long, 4 > > recv_buffer_offset
Definition AMReX_SpMatrix.H:477
Vector< int > recv_from
Definition AMReX_SpMatrix.H:472
Long * recv_buffer_idx_map
Definition AMReX_SpMatrix.H:484
Rows of another matrix fetched for SpGEMM. Column indices are global.
Definition AMReX_SpMatrix.H:506
container_type< Long > row_offset
Definition AMReX_SpMatrix.H:507
RemoteRowsMM(RemoteRowsMM &&rhs) noexcept
Definition AMReX_SpMatrix.H:517
RemoteRowsMM(RemoteRowsMM const &)=delete
Long nnz
Definition AMReX_SpMatrix.H:511
~RemoteRowsMM()
Definition AMReX_SpMatrix.H:514
Long * col_index
Definition AMReX_SpMatrix.H:508
RemoteRowsMM & operator=(RemoteRowsMM const &)=delete
Long nrows
Definition AMReX_SpMatrix.H:510
void clear()
Definition AMReX_SpMatrix.H:533
T * mat
Definition AMReX_SpMatrix.H:509
Definition AMReX_ccse-mpi.H:55