ROOT
master
Reference Guide
Loading...
Searching...
No Matches
ParseBatchNormalization.cxx
Go to the documentation of this file.
1
#include "
TMVA/RModelParser_ONNX.hxx
"
2
#include "
TMVA/ROperator_BatchNormalization.hxx
"
3
#include "
onnx.hxx
"
4
5
namespace
TMVA
{
6
namespace
Experimental {
7
namespace
SOFIE {
8
9
ParserFuncSignature
ParseBatchNormalization
= [](
RModelParser_ONNX
&
parser
,
const
onnx::NodeProto
&
nodeproto
) {
10
ETensorType
input_type
;
11
12
auto
input_name
=
nodeproto
.input(0);
13
if
(
parser
.IsRegisteredTensorType(
input_name
)) {
14
input_type
=
parser
.GetTensorType(
input_name
);
15
}
else
{
16
throw
std::runtime_error(
"TMVA::SOFIE ONNX Parser BatchNorm op has input tensor "
+
input_name
+
17
" but its type is not yet registered"
);
18
}
19
20
std::unique_ptr<ROperator>
op
;
21
std::string
output_name
=
nodeproto
.output(0);
22
float
fepsilon = 1
e
-05;
23
float
fmomentum = 0.9;
24
std::size_t ftraining_mode = 0;
25
for
(
int_t
i = 0; i <
nodeproto
.attribute_size(); i++) {
26
const
std::string &
attribute_name
=
nodeproto
.attribute(i).name();
27
if
(
attribute_name
==
"epsilon"
)
28
fepsilon =
nodeproto
.attribute(i).f();
29
else
if
(
attribute_name
==
"momentum"
)
30
fmomentum =
nodeproto
.attribute(i).f();
31
}
32
33
switch
(
input_type
) {
34
case
ETensorType::FLOAT
:
35
if
(
nodeproto
.input_size() == 5) {
36
op
.reset(
new
ROperator_BatchNormalization<float>
(fepsilon, fmomentum, ftraining_mode,
nodeproto
.input(0),
37
nodeproto
.input(1),
nodeproto
.input(2),
nodeproto
.input(3),
38
nodeproto
.input(4),
output_name
));
39
}
else
{
40
throw
std::runtime_error(
"TMVA::SOFIE ONNX Parser BatchNormalization op requires exactly 5 inputs, got "
+
41
std::to_string(
nodeproto
.input_size()));
42
}
43
break
;
44
default
:
45
throw
std::runtime_error(
"TMVA::SOFIE - Unsupported - Operator BatchNorm does not yet support input type "
+
46
std::to_string(
static_cast<
int
>
(
input_type
)));
47
}
48
49
if
(!
parser
.IsRegisteredTensorType(
output_name
)) {
50
parser
.RegisterTensorType(
output_name
,
input_type
);
51
}
52
53
return
op
;
54
};
55
56
}
// namespace SOFIE
57
}
// namespace Experimental
58
}
// namespace TMVA
RModelParser_ONNX.hxx
ROperator_BatchNormalization.hxx
e
#define e(i)
Definition
RSha256.hxx:103
TRangeDynCast
ROOT::Detail::TRangeCast< T, true > TRangeDynCast
TRangeDynCast is an adapter class that allows the typed iteration through a TCollection.
Definition
TCollection.h:359
ROOT::Detail::TRangeCast
Definition
TCollection.h:312
TMVA::Experimental::SOFIE::RModelParser_ONNX
Definition
RModelParser_ONNX.hxx:30
TMVA::Experimental::SOFIE::onnx::NodeProto
Definition
onnx.hxx:504
TMVA::Experimental::SOFIE::ParseBatchNormalization
ParserFuncSignature ParseBatchNormalization
Definition
ParseBatchNormalization.cxx:9
TMVA::Experimental::SOFIE::ETensorType
ETensorType
Definition
SOFIE_common.hxx:26
TMVA::Experimental::SOFIE::ETensorType::FLOAT
@ FLOAT
TMVA::Experimental::SOFIE::ParserFuncSignature
std::function< std::unique_ptr< ROperator >(RModelParser_ONNX &, const onnx::NodeProto &)> ParserFuncSignature
Definition
RModelParser_ONNX.hxx:25
TMVA::Experimental::SOFIE::int_t
std::int64_t int_t
Definition
SOFIE_common.hxx:53
TMVA
create variable transformations
Definition
GeneticMinimizer.h:22
onnx.hxx
tmva
sofie_parsers
src
ParseBatchNormalization.cxx
ROOTmaster - Reference Guide Generated on Mon Sep 28 2026 15:14:59 (GVA Time) using Doxygen 1.10.0