VTK  9.7.20261003
vtkONNXInternalUtils.h
Go to the documentation of this file.
1// SPDX-FileCopyrightText: Copyright (c) Ken Martin, Will Schroeder, Bill Lorensen
2// SPDX-License-Identifier: BSD-3-Clause
12
13#ifndef vtkONNXInternalUtils_h
14#define vtkONNXInternalUtils_h
15
16#include "vtkSMPTools.h"
17
18#include <algorithm>
19#include <cstdint>
20#include <iostream>
21#include <numeric>
22#include <vector>
23
24VTK_ABI_NAMESPACE_BEGIN
26{
27
31inline int64_t TensorNumberOfElements(const std::vector<int64_t>& shape)
32{
33 return std::accumulate(shape.begin(), shape.end(), 1LL, std::multiplies<>());
34}
35
39inline bool IsPermutation(const std::vector<int>& permutation)
40{
41 std::cout << std::endl;
42 std::vector<int> identity(permutation.size());
43 std::iota(identity.begin(), identity.end(), 0);
44 return std::is_permutation(identity.begin(), identity.end(), permutation.begin());
45}
46
51inline std::vector<int> InversePermutation(const std::vector<int>& permutation)
52{
53 std::vector<int> inversePermutation(permutation.size());
54 for (size_t i = 0; i < permutation.size(); ++i)
55 {
56 inversePermutation[permutation[i]] = i;
57 }
58 return inversePermutation;
59}
60
65inline void Permute(
66 float* data, const std::vector<int64_t>& outputShape, const std::vector<int>& permutation)
67{
68 const size_t nDim = outputShape.size();
69 int64_t numElements = TensorNumberOfElements(outputShape);
70
71 // Compute input shape
72 std::vector<int> inversePermutation = InversePermutation(permutation);
73 std::vector<int64_t> inputShape(outputShape.size());
74
75 for (size_t i = 0; i < nDim; ++i)
76 {
77 inputShape[i] = outputShape[inversePermutation[i]];
78 }
79
80 // Compute input/output memory strides
81 auto computeStrides = [nDim](const std::vector<int64_t>& shape)
82 {
83 std::vector<int64_t> strides(nDim);
84 strides[nDim - 1] = 1;
85 for (int i = static_cast<int>(nDim) - 2; i >= 0; --i)
86 {
87 strides[i] = strides[i + 1] * shape[i + 1];
88 }
89 return strides;
90 };
91
92 std::vector<int64_t> inputStrides = computeStrides(inputShape);
93 std::vector<int64_t> outputStrides = computeStrides(outputShape);
94
95 // Permutation loop
96 std::vector<float> buffer(numElements);
97 vtkSMPTools::For(0, numElements,
98 [&](int64_t begin, int64_t end)
99 {
100 std::vector<int64_t> inputCoords(nDim);
101 std::vector<int64_t> outputCoords(nDim);
102
103 for (int64_t inputIndex = begin; inputIndex < end; ++inputIndex)
104 {
105 int tmpInputIndex = inputIndex;
106 // Input shape coords
107 for (size_t i = 0; i < nDim; ++i)
108 {
109 inputCoords[i] = tmpInputIndex / inputStrides[i];
110 tmpInputIndex %= inputStrides[i];
111 }
112
113 // Apply permutation
114 for (size_t i = 0; i < nDim; ++i)
115 {
116 outputCoords[i] = inputCoords[permutation[i]];
117 }
118
119 // Output shape coords
120 int outputIndex = 0;
121 for (size_t i = 0; i < nDim; ++i)
122 {
123 outputIndex += outputCoords[i] * outputStrides[i];
124 }
125
126 buffer[outputIndex] = data[inputIndex];
127 }
128 });
129
130 std::copy(buffer.begin(), buffer.end(), data);
131}
132
137inline std::vector<int> FindMatchingPermutation(
138 int64_t numTuples, int64_t numComponents, const std::vector<int64_t> modelShape)
139{
140 if (modelShape.empty())
141 {
142 return {};
143 }
144
145 const int64_t vtkShape[2] = { numTuples, numComponents };
146 constexpr int vtkRank = 2;
147 const int modelRank = static_cast<int>(modelShape.size());
148
149 const int64_t modelTotalElements = TensorNumberOfElements(modelShape);
150 const int64_t vtkTotalElements = numTuples * numComponents;
151
152 if (modelTotalElements != vtkTotalElements)
153 {
154 return {};
155 }
156
157 if (modelRank <= 1)
158 {
159 return {};
160 }
161
162 if (modelRank == 2)
163 {
164 if (modelShape[0] == numTuples && modelShape[1] == numComponents)
165 {
166 return {};
167 }
168
169 if (modelShape[0] == numComponents && modelShape[1] == numTuples)
170 {
171 return { 1, 0 };
172 }
173 }
174
175 // For high rank cases, first look for direct matches
176 std::array<std::vector<int>, vtkRank> vtkToModelMapping;
177
178 std::vector<bool> usedModel(modelRank, false);
179 std::vector<bool> usedVTK(vtkRank, false);
180
181 for (int i = 0; i < modelRank; ++i)
182 {
183 for (int j = 0; j < vtkRank; ++j)
184 {
185 if (modelShape[i] == vtkShape[j] && !usedVTK[j] && !usedModel[i])
186 {
187 vtkToModelMapping[j].push_back(i);
188 usedModel[i] = true;
189 usedVTK[j] = true;
190 }
191 }
192 }
193
194 // This iterates through each bit of the `productMask` and apply `func` if true
195 auto forEachSelectedDimension = [&](int productMask, auto&& func)
196 {
197 int bitShift = 0;
198
199 for (size_t i = 0; i < usedModel.size(); ++i)
200 {
201 if (!usedModel[i])
202 {
203 if (productMask & (1 << bitShift))
204 {
205 func(i);
206 }
207 ++bitShift;
208 }
209 }
210 };
211
212 // Then, brute-force remaining dimensions
213 for (int j = 0; j < vtkRank; ++j)
214 {
215 if (usedVTK[j])
216 {
217 continue;
218 }
219 const int64_t dimension = vtkShape[j];
220 int unusedShapeElements =
221 std::count_if(usedModel.begin(), usedModel.end(), [](bool used) { return !used; });
222
223 for (int productMask = (1 << unusedShapeElements); productMask > 0; --productMask)
224 {
225 int64_t extractedDimension = 1;
226
227 forEachSelectedDimension(productMask, [&](size_t i) { extractedDimension *= modelShape[i]; });
228
229 if (extractedDimension == dimension)
230 {
231 usedVTK[j] = true;
232
233 forEachSelectedDimension(productMask,
234 [&](size_t i)
235 {
236 vtkToModelMapping[j].push_back(i);
237 usedModel[i] = true;
238 });
239
240 break;
241 }
242 }
243 }
244
245 if (vtkToModelMapping[0].empty() || vtkToModelMapping[1].empty())
246 {
247 return {};
248 }
249
250 std::vector<int> permutation;
251 permutation.reserve(modelRank);
252 for (int i = 0; i < vtkRank; ++i)
253 {
254 for (int j = 0; j < static_cast<int>(vtkToModelMapping[i].size()); ++j)
255 {
256 permutation.push_back(vtkToModelMapping[i][j]);
257 }
258 }
259
260 if (!IsPermutation(permutation))
261 {
262 return {};
263 }
264
265 return permutation;
266}
267
268} // namespace vtkONNXInternalUtils
269VTK_ABI_NAMESPACE_END
270#endif
static void For(vtkIdType first, vtkIdType last, vtkIdType grain, Functor &f)
Execute a for operation in parallel.
std::vector< int > FindMatchingPermutation(int64_t numTuples, int64_t numComponents, const std::vector< int64_t > modelShape)
This function tries to find the permutation required to match the dimensions of a VTK array to any N ...
void Permute(float *data, const std::vector< int64_t > &outputShape, const std::vector< int > &permutation)
This reorders the memory pointed by data so that it matches the layout defined by outputShape and per...
int64_t TensorNumberOfElements(const std::vector< int64_t > &shape)
Helper to find the total number of elements given the list of dimensions of a tensor.
bool IsPermutation(const std::vector< int > &permutation)
This checks if a sequence actually represents a permutation.
std::vector< int > InversePermutation(const std::vector< int > &permutation)
Computes the inverse of the input permutation, in other words the permutation you need to apply after...