ROOT
master
Reference Guide
Loading...
Searching...
No Matches
ParseInstanceNormalization.cxx
Go to the documentation of this file.
1
#include "
TMVA/RModelParser_ONNX.hxx
"
2
#include "
TMVA/ROperator_InstanceNormalization.hxx
"
3
#include "
onnx.hxx
"
4
5
namespace
TMVA
{
6
namespace
Experimental {
7
namespace
SOFIE {
8
9
ParserFuncSignature
ParseInstanceNormalization
= [](
RModelParser_ONNX
&
parser
,
10
const
onnx::NodeProto
&
nodeproto
) -> std::unique_ptr<ROperator> {
11
ETensorType
input_type
=
ETensorType::UNDEFINED
;
12
const
std::string
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 InstanceNormalization op has input tensor "
+
input_name
+
17
" but its type is not yet registered"
);
18
}
19
20
float
epsilon = 1
e
-5;
21
for
(int64_t i = 0; i <
nodeproto
.attribute_size(); i++) {
22
if
(
nodeproto
.attribute(i).name() ==
"epsilon"
) {
23
epsilon =
nodeproto
.attribute(i).f();
24
}
25
}
26
27
// Inputs: X (0), scale (1), B (2)
28
const
std::string
name_scale
=
nodeproto
.input(1);
29
const
std::string
name_bias
=
nodeproto
.input(2);
30
const
std::string
output_name
=
nodeproto
.output(0);
31
32
std::unique_ptr<ROperator>
op
;
33
switch
(
input_type
) {
34
case
ETensorType::FLOAT
:
35
op
.reset(
new
ROperator_InstanceNormalization<float>
(epsilon,
input_name
,
name_scale
,
name_bias
,
output_name
));
36
break
;
37
default
:
38
throw
std::runtime_error(
"TMVA::SOFIE ONNX parser Operator with input type "
+
ConvertTypeToString
(
input_type
) +
39
" not supported."
);
40
break
;
41
}
42
43
if
(!
parser
.IsRegisteredTensorType(
output_name
)) {
44
parser
.RegisterTensorType(
output_name
,
input_type
);
45
}
46
47
return
op
;
48
};
49
50
}
// namespace SOFIE
51
}
// namespace Experimental
52
}
// namespace TMVA
RModelParser_ONNX.hxx
ROperator_InstanceNormalization.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::ETensorType
ETensorType
Definition
SOFIE_common.hxx:29
TMVA::Experimental::SOFIE::ETensorType::UNDEFINED
@ UNDEFINED
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::ConvertTypeToString
std::string ConvertTypeToString(ETensorType type)
Definition
SOFIE_common.cxx:66
TMVA::Experimental::SOFIE::ParseInstanceNormalization
ParserFuncSignature ParseInstanceNormalization
Definition
ParseInstanceNormalization.cxx:9
TMVA
create variable transformations
Definition
GeneticMinimizer.h:22
onnx.hxx
tmva
sofie_parsers
src
ParseInstanceNormalization.cxx
ROOTmaster - Reference Guide Generated on Sat Sep 5 2026 04:37:47 (GVA Time) using Doxygen 1.10.0