VTK  9.7.20261001
vtkONNXInference.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
23#ifndef vtkONNXInference_h
24#define vtkONNXInference_h
25
26#include "vtkFiltersONNXModule.h" // For export macro
28
29#include "vtkDataObject.h" // For AttributeTypes
30
31#include <memory> // For std::unique_ptr
32#include <vector> // For std::vector
33
34VTK_ABI_NAMESPACE_BEGIN
36namespace Ort
37{
38class AllocatorWithDefaultOptions;
39class Value;
40}
41
42class VTKFILTERSONNX_EXPORT vtkONNXInference : public vtkPassInputTypeAlgorithm
43{
44public:
46 void PrintSelf(ostream& os, vtkIndent indent) override;
47
49
51
54 void SetModelFile(const std::string& file);
55 vtkGetMacro(ModelFile, std::string);
57
68
71 void SetTimeStepValues(const std::vector<double>& times);
72
76 void SetTimeStepValue(vtkIdType idx, double timeStepValue);
77
83
89
91
95 vtkSetMacro(TimeStepIndex, int);
96 vtkGetMacro(TimeStepIndex, int);
98
110 void SetInputParameters(const std::vector<float>& params);
111
116 void SetInputParameter(vtkIdType idx, float InputParameter);
117
123
125
130 void SetInputShape(const std::vector<int64_t>& shape);
131 void SetInputShape(vtkIdType idx, int shapeElement);
133 const std::vector<int64_t>& GetInputShape() const;
135
141
147
149
157 void SetInputPermutation(const std::vector<int>& shape);
158 void SetInputPermutationElement(vtkIdType idx, int permutationElement);
159 const std::vector<int>& GetInputPermutation() const;
163
165
173 void SetOutputPermutation(const std::vector<int>& permutation);
174 void SetOutputPermutationElement(vtkIdType idx, int permutationElement);
175 const std::vector<int>& GetOutputPermutation() const;
179
181
185 vtkSetMacro(FieldArrayInput, bool);
186 vtkGetMacro(FieldArrayInput, bool);
187 vtkBooleanMacro(FieldArrayInput, bool);
189
191
194 vtkSetMacro(ProcessedFieldArrayName, const std::string&);
195 vtkGetMacro(ProcessedFieldArrayName, const std::string&);
197
199
202 vtkSetMacro(OutputDimension, int);
203 vtkGetMacro(OutputDimension, int);
205
207
211 vtkSetMacro(ArrayAssociation, int);
212 vtkGetMacro(ArrayAssociation, int);
214
216
223 vtkGetMacro(AutoDetectInputShape, bool);
224 vtkBooleanMacro(AutoDetectInputShape, bool);
226
228
236 vtkSetMacro(AutoDetectPermutation, bool);
237 vtkGetMacro(AutoDetectPermutation, bool);
238 vtkBooleanMacro(AutoDetectPermutation, bool);
240
241protected:
243 ~vtkONNXInference() override = default;
244
249
251
256 int ExecuteData(vtkDataObject* input, vtkDataObject* output, double timevalue);
257
258private:
259 vtkONNXInference(const vtkONNXInference&) = delete;
260 void operator=(const vtkONNXInference&) = delete;
261
267 bool InitializeSession();
268
274 bool ShouldGenerateTimeSteps();
275
280 bool GenerateInputTensorFromParameters(
281 std::vector<float>& parameters, Ort::Value& inputTensor, double timeValue);
282
287 bool GenerateInputTensorFromFieldArray(
288 Ort::Value& inputTensor, vtkDataSetAttributes* inAttributes);
289
294 std::vector<Ort::Value> RunModel(Ort::Value& inputTensor);
295
296 // Input related parameters
297 std::string ModelFile;
298 std::vector<int64_t> InputShape = { 0 };
299 std::vector<float> InputParameters;
300 std::vector<double> TimeStepValues;
301 int TimeStepIndex = -1;
302 bool FieldArrayInput = false;
303 std::string ProcessedFieldArrayName;
304 std::vector<int> InputPermutation;
305
306 // Output related parameters
307 int OutputDimension = 1;
308 std::vector<int> OutputPermutation;
309
310 int ArrayAssociation = vtkDataObject::CELL;
311 std::vector<float> InputDataBuffer;
312
313 bool AutoDetectInputShape = false;
314 bool AutoDetectPermutation = false;
315
316 bool Initialized = false;
317 std::unique_ptr<vtkONNXInferenceInternals> Internals;
318};
319VTK_ABI_NAMESPACE_END
320
321#endif // vtkONNXInference_h
int idx
Definition HexCnBasis.h:79
general representation of visualization data
represent and manipulate attribute data in a dataset
a simple class to control print indentation
Definition vtkIndent.h:108
Store zero or more vtkInformation instances.
Store vtkAlgorithm input/output information.
void ClearInputParameters()
Clear the input parameters vector.
void ClearTimeStepValues()
Clear the time step values vector.
void ClearOutputPermutation()
Set/Get the permutation between the model output and a VTK array.
void SetNumberOfInputPermutationElements(vtkIdType nb)
Set/Get the permutation between the VTK array and the model input.
void SetInputPermutation(const std::vector< int > &shape)
Set/Get the permutation between the VTK array and the model input.
void SetInputShape(vtkIdType idx, int shapeElement)
Set/Get the shape of the input.
int RequestInformation(vtkInformation *, vtkInformationVector **, vtkInformationVector *) override
This is required to inform the pipeline of the time steps.
~vtkONNXInference() override=default
const std::vector< int > & GetOutputPermutation() const
Set/Get the permutation between the model output and a VTK array.
void SetInputShape(const std::vector< int64_t > &shape)
Set/Get the shape of the input.
int RequestData(vtkInformation *, vtkInformationVector **, vtkInformationVector *) override
This is called within ProcessRequest when a request asks the algorithm to do its work.
void SetInputParameters(const std::vector< float > &params)
Input Parameters.
void SetInputParameter(vtkIdType idx, float InputParameter)
Set an input parameter at a given index.
void ClearInputPermutation()
Set/Get the permutation between the VTK array and the model input.
static vtkONNXInference * New()
void SetTimeStepValues(const std::vector< double > &times)
Time Steps.
void SetTimeStepValue(vtkIdType idx, double timeStepValue)
Set a time value at a given index.
void SetInputPermutationElement(vtkIdType idx, int permutationElement)
Set/Get the permutation between the VTK array and the model input.
void SetModelFile(const std::string &file)
Get/Set the path to the ONNX model and load it.
void ClearInputShape()
Clear the input shape vector.
void SetOutputPermutationElement(vtkIdType idx, int permutationElement)
Set/Get the permutation between the model output and a VTK array.
void SetNumberOfOutputPermutationElements(vtkIdType nb)
Set/Get the permutation between the model output and a VTK array.
int ExecuteData(vtkDataObject *input, vtkDataObject *output, double timevalue)
Execute the inference and add the resulting array on the given data object.
void SetNumberOfInputShapeElements(vtkIdType nb)
Set the number of input shape values.
void SetOutputPermutation(const std::vector< int > &permutation)
Set/Get the permutation between the model output and a VTK array.
void PrintSelf(ostream &os, vtkIndent indent) override
Methods invoked by print to print information about the object including superclasses.
const std::vector< int > & GetInputPermutation() const
Set/Get the permutation between the VTK array and the model input.
void SetInputShape(vtkIdType nb)
Set/Get the shape of the input.
void SetNumberOfTimeStepValues(vtkIdType nb)
Set the number of time step values.
const std::vector< int64_t > & GetInputShape() const
Set/Get the shape of the input.
void SetAutoDetectInputShape(bool SetAutoDetectInputShape)
Set/Get whether to automatically detect input shape from the ONNX model.
VTK internal class for hiding ONNX members.
int vtkIdType
Definition vtkType.h:363