Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
ParseBasicBinary.cxx
Go to the documentation of this file.
3#include "onnx.hxx"
4
5namespace TMVA {
6namespace Experimental {
7namespace SOFIE {
8
9template <EBasicBinaryOperator Op>
11{
13
14 for (int i = 0; i < 2; ++i) {
15 auto input_name = nodeproto.input(i);
16 if (parser.IsRegisteredTensorType(input_name)) {
17 // according to ONNX both inputs have same type
18 if (i == 0)
19 input_type = parser.GetTensorType(input_name);
20 else {
22 if (input_type2 != input_type) {
23 throw
24 std::runtime_error("TMVA::SOFIE ONNX parser Binary op has input tensors of different types: " +
26 " and " + nodeproto.input(0) + " : " + ConvertTypeToString(input_type));
27 }
28 }
29 } else {
30 throw std::runtime_error("TMVA::SOFIE ONNX Parser Binary op has input tensor " + input_name +
31 " but its type is not yet registered");
32 }
33 }
34
35 std::unique_ptr<ROperator> op;
36 std::string output_name = nodeproto.output(0);
37
38 switch (input_type) {
41 break;
44 break;
47 break;
50 break;
51 default:
52 throw std::runtime_error("TMVA::SOFIE - Unsupported - Binary Operator does not yet support input type " +
53 std::to_string(static_cast<int>(input_type)));
54 }
55
56 // Infer the output type
57 if (!parser.IsRegisteredTensorType(output_name)) {
58 parser.RegisterTensorType(output_name, input_type);
59 }
60
61 return op;
62};
63
64
65// Mod (and fmod) is a special case di BasicBinary
66
68
70 for (int i = 0; i < 2; ++i) {
71 auto input_name = nodeproto.input(i);
72 if (parser.IsRegisteredTensorType(input_name)) {
73 // according to ONNX both inputs have same type
74 if (i == 0)
75 input_type = parser.GetTensorType(input_name);
76 else {
78 if (input_type2 != input_type) {
79 throw
80 std::runtime_error("TMVA::SOFIE ONNX parser Binary op has input tensors of different types: " +
82 " and " + nodeproto.input(0) + " : " + ConvertTypeToString(input_type));
83 }
84 }
85 } else {
86 throw std::runtime_error("TMVA::SOFIE ONNX Parser Binary op has input tensor " + input_name +
87 " but its type is not yet registered");
88 }
89 }
90 // in case of Mod there can be an attribute
91 int fmod = 0;
92 if (nodeproto.attribute_size() > 0) {
93 fmod = nodeproto.attribute(0).i();
94 }
95 std::unique_ptr<ROperator> op;
96 std::string output_name = nodeproto.output(0);
97
98 switch (input_type) {
101 break;
104 break;
106 if (fmod == 1)
108 else
110 break;
112 if (fmod == 1)
114 else
116 break;
117 default:
118 throw std::runtime_error("TMVA::SOFIE - Unsupported - Binary Operator does not yet support input type " +
119 std::to_string(static_cast<int>(input_type)));
120 }
121
122 // Infer the output type
123 if (!parser.IsRegisteredTensorType(output_name)) {
124 parser.RegisterTensorType(output_name, input_type);
125 }
126
127 return op;
128};
129
131{
137 parser.RegisterOperator("Mod", ParseMod);
138}
139
140} // namespace SOFIE
141} // namespace Experimental
142} // namespace TMVA
ROOT::Detail::TRangeCast< T, true > TRangeDynCast
TRangeDynCast is an adapter class that allows the typed iteration through a TCollection.
std::function< std::unique_ptr< ROperator >(RModelParser_ONNX &, const onnx::NodeProto &)> ParserFuncSignature
void RegisterBasicBinaryParsers(RModelParser_ONNX &parser)
std::unique_ptr< ROperator > ParseBasicBinary(RModelParser_ONNX &parser, const onnx::NodeProto &nodeproto)
ParserFuncSignature ParseMod
std::string ConvertTypeToString(ETensorType type)
create variable transformations