10#ifndef TPETRAEXT_POINTTOBLOCKDIAGPERMUTE_DEF_HPP
11#define TPETRAEXT_POINTTOBLOCKDIAGPERMUTE_DEF_HPP
13#include "TpetraExt_PointToBlockDiagPermute_decl.hpp"
15#include "Teuchos_OrdinalTraits.hpp"
16#include "Tpetra_Import.hpp"
17#include "Tpetra_Export.hpp"
26template <
class Scalar,
class LocalOrdinal,
class GlobalOrdinal,
class Node>
27PointToBlockDiagPermute<Scalar, LocalOrdinal, GlobalOrdinal, Node>::PointToBlockDiagPermute(
30 , purelyLocalMode_(true)
31 , contiguousBlockMode_(false)
32 , contiguousBlockSize_(0)
35template <
class Scalar,
class LocalOrdinal,
class GlobalOrdinal,
class Node>
36void PointToBlockDiagPermute<Scalar, LocalOrdinal, GlobalOrdinal, Node>::cleanupBlockInfo() {
42template <
class Scalar,
class LocalOrdinal,
class GlobalOrdinal,
class Node>
43int PointToBlockDiagPermute<Scalar, LocalOrdinal, GlobalOrdinal, Node>::setParameters(
44 Teuchos::ParameterList& list) {
48 contiguousBlockSize_ = list_.get(
"contiguous block size", 0);
49 TEUCHOS_TEST_FOR_EXCEPTION(contiguousBlockSize_ < 0, std::runtime_error,
50 "PointToBlockDiagPermute: contiguous block size must be non-negative");
51 contiguousBlockMode_ = (contiguousBlockSize_ != 0);
52 purelyLocalMode_ =
true;
54 if (contiguousBlockMode_) {
55 return setupContiguousMode();
58 numBlocks_ = list_.get(
"number of local blocks", 0);
59 TEUCHOS_TEST_FOR_EXCEPTION(numBlocks_ < 0, std::runtime_error,
60 "PointToBlockDiagPermute: invalid number of local blocks");
62 if (numBlocks_ == 0) {
63 blockStarts_.assign(1, 0);
68 TEUCHOS_TEST_FOR_EXCEPTION(!list_.isParameter(
"block start index"), std::runtime_error,
69 "PointToBlockDiagPermute: missing block start index");
70 TEUCHOS_TEST_FOR_EXCEPTION(!list_.isParameter(
"block entry gids"), std::runtime_error,
71 "PointToBlockDiagPermute: missing block entry gids");
72 blockStarts_ = list_.get<Teuchos::Array<int>>(
"block start index");
73 blockGids_ = list_.get<Teuchos::Array<GlobalOrdinal>>(
"block entry gids");
75 TEUCHOS_TEST_FOR_EXCEPTION(blockStarts_.size() < numBlocks_ + 1, std::runtime_error,
76 "PointToBlockDiagPermute: block start index array size is too small");
77 TEUCHOS_TEST_FOR_EXCEPTION(blockGids_.size() < blockStarts_[numBlocks_], std::runtime_error,
78 "PointToBlockDiagPermute: block entry gids array size is too small");
83template <
class Scalar,
class LocalOrdinal,
class GlobalOrdinal,
class Node>
84int PointToBlockDiagPermute<Scalar, LocalOrdinal, GlobalOrdinal, Node>::setupContiguousMode() {
85 if (!contiguousBlockMode_)
return 0;
87 const auto rowMap = matrix_->getRowMap();
88 if (rowMap->getLocalNumElements() == 0) {
90 blockStarts_.assign(1, 0);
95 const GlobalOrdinal minMyGID = rowMap->getMinGlobalIndex();
96 const GlobalOrdinal maxMyGID = rowMap->getMaxGlobalIndex();
97 const GlobalOrdinal base = rowMap->getIndexBase();
99 const GlobalOrdinal myFirstBlockGID =
100 static_cast<GlobalOrdinal
>(contiguousBlockSize_ *
101 std::floor(
static_cast<double>(minMyGID - base) /
102 static_cast<double>(contiguousBlockSize_)) +
105 numBlocks_ =
static_cast<int>(
106 std::ceil(
static_cast<double>(maxMyGID - myFirstBlockGID + 1.0) /
107 static_cast<double>(contiguousBlockSize_)));
109 blockStarts_.resize(numBlocks_ + 1);
111 const size_t numBlockEntries =
112 static_cast<size_t>(numBlocks_) *
static_cast<size_t>(contiguousBlockSize_);
113 blockGids_.resize(numBlockEntries);
115 blockStarts_[numBlocks_] = numBlocks_ * contiguousBlockSize_;
117 for (
int i = 0, ct = 0; i < numBlocks_; i++) {
118 blockStarts_[i] = ct;
119 for (
int j = 0; j < contiguousBlockSize_; j++, ct++) {
120 blockGids_[ct] = myFirstBlockGID + ct;
127template <
class Scalar,
class LocalOrdinal,
class GlobalOrdinal,
class Node>
128int PointToBlockDiagPermute<Scalar, LocalOrdinal, GlobalOrdinal, Node>::compute() {
129 return extractBlockDiagonal();
132template <
class Scalar,
class LocalOrdinal,
class GlobalOrdinal,
class Node>
133int PointToBlockDiagPermute<Scalar, LocalOrdinal, GlobalOrdinal, Node>::extractBlockDiagonal() {
134 TEUCHOS_TEST_FOR_EXCEPTION(matrix_ ==
nullptr, std::runtime_error,
135 "PointToBlockDiagPermute: null matrix");
137 const auto rowMap = matrix_->getRowMap();
138 const auto colMap = matrix_->getColMap();
140 const LocalOrdinal numMyRows =
static_cast<LocalOrdinal
>(rowMap->getLocalNumElements());
142 std::vector<int> localToBlock(numMyRows, -1);
144 for (
int b = 0; b < numBlocks_; ++b) {
145 for (
int j = blockStarts_[b]; j < blockStarts_[b + 1]; ++j) {
146 GlobalOrdinal gid = blockGids_[j];
147 LocalOrdinal lid = rowMap->getLocalElement(gid);
148 if (lid != Teuchos::OrdinalTraits<LocalOrdinal>::invalid()) {
149 localToBlock[lid] = b;
155 Teuchos::rcp(
new crs_type(rowMap, contiguousBlockSize_ > 0 ? contiguousBlockSize_ : 8));
157 typename crs_type::nonconst_local_inds_host_view_type inds(
158 "inds", matrix_->getLocalMaxNumRowEntries());
159 typename crs_type::nonconst_values_host_view_type vals(
160 "vals", matrix_->getLocalMaxNumRowEntries());
162 std::vector<GlobalOrdinal> outCols;
163 std::vector<Scalar> outVals;
164 outCols.reserve(contiguousBlockSize_ > 0 ? contiguousBlockSize_ : 8);
165 outVals.reserve(contiguousBlockSize_ > 0 ? contiguousBlockSize_ : 8);
167 for (LocalOrdinal lrow = 0; lrow < numMyRows; ++lrow) {
168 int blockNum = localToBlock[lrow];
169 if (blockNum < 0)
continue;
171 GlobalOrdinal rowGid = rowMap->getGlobalElement(lrow);
173 size_t numEntries = Teuchos::OrdinalTraits<size_t>::invalid();
174 matrix_->getLocalRowCopy(lrow, inds, vals, numEntries);
179 for (
size_t k = 0; k < numEntries; ++k) {
180 LocalOrdinal lcol = inds(k);
181 if (lcol == Teuchos::OrdinalTraits<LocalOrdinal>::invalid())
continue;
183 GlobalOrdinal colGid = colMap->getGlobalElement(lcol);
184 if (colGid == Teuchos::OrdinalTraits<GlobalOrdinal>::invalid())
continue;
185 LocalOrdinal rowLid = rowMap->getLocalElement(colGid);
186 if (rowLid == Teuchos::OrdinalTraits<LocalOrdinal>::invalid())
continue;
187 if (localToBlock[rowLid] != blockNum)
continue;
189 outCols.push_back(colGid);
190 outVals.push_back(vals(k));
193 if (outCols.empty()) {
194 outCols.push_back(rowGid);
195 outVals.push_back(Teuchos::ScalarTraits<Scalar>::one());
198 blockDiag->insertGlobalValues(rowGid, Teuchos::ArrayView<const GlobalOrdinal>(outCols),
199 Teuchos::ArrayView<const Scalar>(outVals));
202 blockDiag->fillComplete(matrix_->getDomainMap(), matrix_->getRangeMap());
203 compatibleMap_ = rowMap;
204 blockDiagMatrix_ = blockDiag;
Namespace for external Tpetra functionality.