Tpetra parallel linear algebra Version of the Day
Loading...
Searching...
No Matches
TpetraExt_PointToBlockDiagPermute_def.hpp
1// @HEADER
2// *****************************************************************************
3// TpetraExt: Tpetra Extended - Linear Algebra Services Package
4//
5// Copyright 2025 NTESS
6// SPDX-License-Identifier: BSD-3-Clause
7// *****************************************************************************
8// @HEADER
9
10#ifndef TPETRAEXT_POINTTOBLOCKDIAGPERMUTE_DEF_HPP
11#define TPETRAEXT_POINTTOBLOCKDIAGPERMUTE_DEF_HPP
12
13#include "TpetraExt_PointToBlockDiagPermute_decl.hpp"
14
15#include "Teuchos_OrdinalTraits.hpp"
16#include "Tpetra_Import.hpp"
17#include "Tpetra_Export.hpp"
18
19#include <vector>
20#include <algorithm>
21#include <cmath>
22#include <stdexcept>
23
24namespace Tpetra::Ext {
25
26template <class Scalar, class LocalOrdinal, class GlobalOrdinal, class Node>
27PointToBlockDiagPermute<Scalar, LocalOrdinal, GlobalOrdinal, Node>::PointToBlockDiagPermute(
28 const crs_type& A)
29 : matrix_(&A)
30 , purelyLocalMode_(true)
31 , contiguousBlockMode_(false)
32 , contiguousBlockSize_(0)
33 , numBlocks_(0) {}
34
35template <class Scalar, class LocalOrdinal, class GlobalOrdinal, class Node>
36void PointToBlockDiagPermute<Scalar, LocalOrdinal, GlobalOrdinal, Node>::cleanupBlockInfo() {
37 blockStarts_.clear();
38 blockGids_.clear();
39 numBlocks_ = 0;
40}
41
42template <class Scalar, class LocalOrdinal, class GlobalOrdinal, class Node>
43int PointToBlockDiagPermute<Scalar, LocalOrdinal, GlobalOrdinal, Node>::setParameters(
44 Teuchos::ParameterList& list) {
45 cleanupBlockInfo();
46 list_ = list;
47
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;
53
54 if (contiguousBlockMode_) {
55 return setupContiguousMode();
56 }
57
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");
61
62 if (numBlocks_ == 0) {
63 blockStarts_.assign(1, 0);
64 blockGids_.clear();
65 return 0;
66 }
67
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");
74
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");
79
80 return 0;
81}
82
83template <class Scalar, class LocalOrdinal, class GlobalOrdinal, class Node>
84int PointToBlockDiagPermute<Scalar, LocalOrdinal, GlobalOrdinal, Node>::setupContiguousMode() {
85 if (!contiguousBlockMode_) return 0;
86
87 const auto rowMap = matrix_->getRowMap();
88 if (rowMap->getLocalNumElements() == 0) {
89 numBlocks_ = 0;
90 blockStarts_.assign(1, 0);
91 blockGids_.clear();
92 return 0;
93 }
94
95 const GlobalOrdinal minMyGID = rowMap->getMinGlobalIndex();
96 const GlobalOrdinal maxMyGID = rowMap->getMaxGlobalIndex();
97 const GlobalOrdinal base = rowMap->getIndexBase();
98
99 const GlobalOrdinal myFirstBlockGID =
100 static_cast<GlobalOrdinal>(contiguousBlockSize_ *
101 std::floor(static_cast<double>(minMyGID - base) /
102 static_cast<double>(contiguousBlockSize_)) +
103 base);
104
105 numBlocks_ = static_cast<int>(
106 std::ceil(static_cast<double>(maxMyGID - myFirstBlockGID + 1.0) /
107 static_cast<double>(contiguousBlockSize_)));
108
109 blockStarts_.resize(numBlocks_ + 1);
110
111 const size_t numBlockEntries =
112 static_cast<size_t>(numBlocks_) * static_cast<size_t>(contiguousBlockSize_);
113 blockGids_.resize(numBlockEntries);
114
115 blockStarts_[numBlocks_] = numBlocks_ * contiguousBlockSize_;
116
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;
121 }
122 }
123
124 return 0;
125}
126
127template <class Scalar, class LocalOrdinal, class GlobalOrdinal, class Node>
128int PointToBlockDiagPermute<Scalar, LocalOrdinal, GlobalOrdinal, Node>::compute() {
129 return extractBlockDiagonal();
130}
131
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");
136
137 const auto rowMap = matrix_->getRowMap();
138 const auto colMap = matrix_->getColMap();
139
140 const LocalOrdinal numMyRows = static_cast<LocalOrdinal>(rowMap->getLocalNumElements());
141
142 std::vector<int> localToBlock(numMyRows, -1);
143
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;
150 }
151 }
152 }
153
154 auto blockDiag =
155 Teuchos::rcp(new crs_type(rowMap, contiguousBlockSize_ > 0 ? contiguousBlockSize_ : 8));
156
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());
161
162 std::vector<GlobalOrdinal> outCols;
163 std::vector<Scalar> outVals;
164 outCols.reserve(contiguousBlockSize_ > 0 ? contiguousBlockSize_ : 8);
165 outVals.reserve(contiguousBlockSize_ > 0 ? contiguousBlockSize_ : 8);
166
167 for (LocalOrdinal lrow = 0; lrow < numMyRows; ++lrow) {
168 int blockNum = localToBlock[lrow];
169 if (blockNum < 0) continue;
170
171 GlobalOrdinal rowGid = rowMap->getGlobalElement(lrow);
172
173 size_t numEntries = Teuchos::OrdinalTraits<size_t>::invalid();
174 matrix_->getLocalRowCopy(lrow, inds, vals, numEntries);
175
176 outCols.clear();
177 outVals.clear();
178
179 for (size_t k = 0; k < numEntries; ++k) {
180 LocalOrdinal lcol = inds(k);
181 if (lcol == Teuchos::OrdinalTraits<LocalOrdinal>::invalid()) continue;
182
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;
188
189 outCols.push_back(colGid);
190 outVals.push_back(vals(k));
191 }
192
193 if (outCols.empty()) {
194 outCols.push_back(rowGid);
195 outVals.push_back(Teuchos::ScalarTraits<Scalar>::one());
196 }
197
198 blockDiag->insertGlobalValues(rowGid, Teuchos::ArrayView<const GlobalOrdinal>(outCols),
199 Teuchos::ArrayView<const Scalar>(outVals));
200 }
201
202 blockDiag->fillComplete(matrix_->getDomainMap(), matrix_->getRangeMap());
203 compatibleMap_ = rowMap;
204 blockDiagMatrix_ = blockDiag;
205
206 return 0;
207}
208
209} // namespace Tpetra::Ext
210
211#endif
Namespace for external Tpetra functionality.