Logo ROOT  
Reference Guide
 
Loading...
Searching...
No Matches
ParseFuseBatchnormRelu.cxx
Go to the documentation of this file.
3#include "onnx.hxx"
4
5namespace TMVA {
6namespace Experimental {
7namespace SOFIE {
8
12
13 auto input_name = batchnormnode.input(0);
14 if (parser.IsRegisteredTensorType(input_name)) {
15 input_type = parser.GetTensorType(input_name);
16 } else {
17 throw std::runtime_error("TMVA::SOFIE ONNX Parser BatchNorm op has input tensor " + input_name +
18 " but its type is not yet registered");
19 }
20
21 std::unique_ptr<ROperator> op;
22 std::string output_name = relunode.output(0);
23 float fepsilon = 1e-05;
24 float fmomentum = 0.9;
25 std::size_t ftraining_mode = 0;
26 for (int_t i = 0; i < batchnormnode.attribute_size(); i++) {
27 const std::string &attribute_name = batchnormnode.attribute(i).name();
28 if (attribute_name == "epsilon")
29 fepsilon = batchnormnode.attribute(i).f();
30 else if (attribute_name == "momentum")
31 fmomentum = batchnormnode.attribute(i).f();
32 }
33
34 switch (input_type) {
36 if (batchnormnode.input_size() == 5) {
37 op.reset(new ROperator_BatchNormalization<float>(fepsilon, fmomentum, ftraining_mode, batchnormnode.input(0),
38 batchnormnode.input(1), batchnormnode.input(2), batchnormnode.input(3),
40 }
41 break;
42 default:
43 throw std::runtime_error("TMVA::SOFIE - Unsupported - Operator BatchNorm does not yet support input type " +
44 std::to_string(static_cast<int>(input_type)));
45 }
46
47 if (!parser.IsRegisteredTensorType(output_name)) {
48 parser.RegisterTensorType(output_name, input_type);
49 }
50
51 return op;
52};
53
54} // namespace SOFIE
55} // namespace Experimental
56} // namespace TMVA
#define e(i)
Definition RSha256.hxx:103
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 &, const onnx::NodeProto &)> ParserFuseFuncSignature
ParserFuseFuncSignature ParseFuseBatchnormRelu
create variable transformations