ROOT
master
Reference Guide
Loading...
Searching...
No Matches
ParseFuseBatchnormRelu.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
ParserFuseFuncSignature
ParseFuseBatchnormRelu
= [](
RModelParser_ONNX
&parser,
const
onnx::NodeProto
&
batchnormnode
,
10
const
onnx::NodeProto
&
relunode
) {
11
ETensorType
input_type
;
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 = 1
e
-05;
24
float
fmomentum = 0.9;
25
std::size_t ftraining_mode = 0;
26
27
switch
(
input_type
) {
28
case
ETensorType::FLOAT
:
29
if
(
batchnormnode
.input_size() == 5) {
30
op
.reset(
new
ROperator_BatchNormalization<float>
(fepsilon, fmomentum, ftraining_mode,
batchnormnode
.input(0),
31
batchnormnode
.input(1),
batchnormnode
.input(2),
batchnormnode
.input(3),
32
batchnormnode
.input(4),
output_name
,
EActivationType::RELU
));
33
}
34
break
;
35
default
:
36
throw
std::runtime_error(
"TMVA::SOFIE - Unsupported - Operator BatchNorm does not yet support input type "
+
37
std::to_string(
static_cast<
int
>
(
input_type
)));
38
}
39
40
if
(!parser.
IsRegisteredTensorType
(
output_name
)) {
41
parser.
RegisterTensorType
(
output_name
,
input_type
);
42
}
43
44
return
op
;
45
};
46
47
}
// namespace SOFIE
48
}
// namespace Experimental
49
}
// 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::RModelParser_ONNX::IsRegisteredTensorType
bool IsRegisteredTensorType(const std::string &)
Definition
RModelParser_ONNX.cxx:426
TMVA::Experimental::SOFIE::RModelParser_ONNX::RegisterTensorType
void RegisterTensorType(const std::string &, ETensorType)
Definition
RModelParser_ONNX.cxx:421
TMVA::Experimental::SOFIE::RModelParser_ONNX::GetTensorType
ETensorType GetTensorType(const std::string &name)
Definition
RModelParser_ONNX.cxx:431
onnx::NodeProto
Definition
onnx.hxx:492
TMVA::Experimental::SOFIE::ParserFuseFuncSignature
std::function< std::unique_ptr< ROperator >(RModelParser_ONNX &, const onnx::NodeProto &, const onnx::NodeProto &)> ParserFuseFuncSignature
Definition
RModelParser_ONNX.hxx:27
TMVA::Experimental::SOFIE::ETensorType
ETensorType
Definition
SOFIE_common.hxx:29
TMVA::Experimental::SOFIE::ETensorType::FLOAT
@ FLOAT
TMVA::Experimental::SOFIE::ParseFuseBatchnormRelu
ParserFuseFuncSignature ParseFuseBatchnormRelu
Definition
ParseFuseBatchnormRelu.cxx:9
TMVA::Experimental::SOFIE::EActivationType::RELU
@ RELU
TMVA
create variable transformations
Definition
GeneticMinimizer.h:22
onnx.hxx
tmva
sofie_parsers
src
ParseFuseBatchnormRelu.cxx
ROOTmaster - Reference Guide Generated on Sun Aug 16 2026 04:54:07 (GVA Time) using Doxygen 1.10.0