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>
10std::unique_ptr<ROperator> ParseBasicBinary(RModelParser_ONNX &parser, const onnx::NodeProto &nodeproto)
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 {
21 ETensorType input_type2 = parser.GetTensorType(input_name);
22 if (input_type2 != input_type) {
23 throw
24 std::runtime_error("TMVA::SOFIE ONNX parser Binary op has input tensors of different types: " +
25 input_name + " : " + ConvertTypeToString(input_type2) +
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) {
40 op.reset(new ROperator_BasicBinary<float, Op>(nodeproto.input(0), nodeproto.input(1), output_name));
41 break;
43 op.reset(new ROperator_BasicBinary<double, Op>(nodeproto.input(0), nodeproto.input(1), output_name));
44 break;
46 op.reset(new ROperator_BasicBinary<int32_t, Op>(nodeproto.input(0), nodeproto.input(1), output_name));
47 break;
49 op.reset(new ROperator_BasicBinary<int64_t, Op>(nodeproto.input(0), nodeproto.input(1), output_name));
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 {
77 ETensorType input_type2 = parser.GetTensorType(input_name);
78 if (input_type2 != input_type) {
79 throw
80 std::runtime_error("TMVA::SOFIE ONNX parser Binary op has input tensors of different types: " +
81 input_name + " : " + ConvertTypeToString(input_type2) +
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) {
100 op.reset(new ROperator_BasicBinary<float,EBasicBinaryOperator::FMod >(nodeproto.input(0), nodeproto.input(1), output_name));
101 break;
103 op.reset(new ROperator_BasicBinary<double, EBasicBinaryOperator::FMod>(nodeproto.input(0), nodeproto.input(1), output_name));
104 break;
106 if (fmod == 1)
107 op.reset(new ROperator_BasicBinary<int32_t, EBasicBinaryOperator::FMod>(nodeproto.input(0), nodeproto.input(1), output_name));
108 else
109 op.reset(new ROperator_BasicBinary<int32_t, EBasicBinaryOperator::Mod>(nodeproto.input(0), nodeproto.input(1), output_name));
110 break;
112 if (fmod == 1)
113 op.reset(new ROperator_BasicBinary<int64_t, EBasicBinaryOperator::FMod>(nodeproto.input(0), nodeproto.input(1), output_name));
114 else
115 op.reset(new ROperator_BasicBinary<int64_t, EBasicBinaryOperator::Mod>(nodeproto.input(0), nodeproto.input(1), output_name));
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{
132 parser.RegisterOperator("Add", ParseBasicBinary<EBasicBinaryOperator::Add>);
133 parser.RegisterOperator("Sub", ParseBasicBinary<EBasicBinaryOperator::Sub>);
134 parser.RegisterOperator("Mul", ParseBasicBinary<EBasicBinaryOperator::Mul>);
135 parser.RegisterOperator("Div", ParseBasicBinary<EBasicBinaryOperator::Div>);
136 parser.RegisterOperator("Pow", ParseBasicBinary<EBasicBinaryOperator::Pow>);
137 parser.RegisterOperator("Mod", ParseMod);
138}
139
140} // namespace SOFIE
141} // namespace Experimental
142} // namespace TMVA
void RegisterOperator(const std::string &name, ParserFuncSignature func)
void RegisterTensorType(const std::string &, ETensorType)
ETensorType GetTensorType(const std::string &name)
const std::string & input(int i) const
Definition onnx.hxx:509
const std::string & output(int i) const
Definition onnx.hxx:512
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