16#include <unordered_map>
21namespace Experimental {
109 throw std::runtime_error(
"TMVA::SOFIE - Failed to read float initialized tensor - actual size is " + std::to_string(
tensor->float_data_size()));
111 std::copy(
src.begin(),
src.end(),
static_cast<float *
>(
data));
118 throw std::runtime_error(
"TMVA::SOFIE - Failed to read double initialized tensor - actual size is " + std::to_string(
tensor->double_data_size()));
120 std::copy(
src.begin(),
src.end(),
static_cast<double *
>(
data));
127 throw std::runtime_error(
"TMVA::SOFIE - Failed to read int32 initialized tensor - actual size is " + std::to_string(
tensor->int32_data_size()));
129 std::copy(
src.begin(),
src.end(),
static_cast<int32_t *
>(
data));
136 throw std::runtime_error(
"TMVA::SOFIE - Failed to read int64 initialized tensor - actual size is " + std::to_string(
tensor->int64_data_size()));
138 std::copy(
src.begin(),
src.end(),
static_cast<int64_t *
>(
data));
148template <std::
size_t N>
152 auto dst =
static_cast<unsigned char *
>(
dest);
153 auto src =
static_cast<const unsigned char *
>(
source);
154 for (std::size_t k = 0; k <
nbytes; k +=
N) {
156 std::memcpy(&
v,
src + k,
N);
158 std::memcpy(
dst + k, &
v,
N);
174 throw std::runtime_error(
"Data type " +
ConvertTypeToString(tensor_type) +
" in tensor is not supported!\n");
184 std::shared_ptr<void>
data(
malloc(tensor_size), free);
190 throw std::runtime_error(
"TMVA::SOFIE - Failed to read raw data of initialized tensor - actual raw size is " +
202 switch (tensor_type) {
220 throw std::runtime_error(
"TMVA::SOFIE - ExtractData from TP in BOOL not supported");
224 throw std::runtime_error(
"TMVA::SOFIE - ExtractData from TP in UINT8 not supported");
228 throw std::runtime_error(
"Data type " +
ConvertTypeToString(tensor_type) +
" in weight tensor is not supported!\n");
236 std::string location;
240 if (
kv.key() ==
"location") location =
kv.value();
241 else if (
kv.key() ==
"offset")
offset = std::stoull(
kv.value());
253 throw std::runtime_error(
"TMVA::SOFIE ONNX : tensor " +
tensorproto->name() +
254 " has external data but no data file location is available");
257 std::cout <<
"Initialized data are stored externally in file " <<
dataFileName
258 <<
" at location " << location <<
" offset " <<
offset <<
" and with length " <<
buffer_size << std::endl;
261 throw std::runtime_error(
"TMVA::SOFIE ONNX : invalid stored data size vs tensor size");
269 throw std::runtime_error(
"TMVA::SOFIE ONNX: error reading external weight ONNX data file " +
dataFileName);
377 std::vector<std::string>
ops;
380 ops.emplace_back(it.first);
424std::unique_ptr<ROperator>
427 if (i >= nodes.size())
428 throw std::runtime_error(
"TMVA::SOFIE - Error in parsing ordered operators " + std::to_string(i) +
" is >= " + std::to_string(nodes.size()));
431 const std::string op_type =
nodeproto.op_type();
433 std::cout <<
"Parsing operator " << op_type << std::endl;
439 std::cout <<
"\tFusing operators " <<
graphproto.node(
idx1).name()
456 if (children.size() == 1) {
457 int idx2 = children.front();
458 if (op_type ==
"MatMul") {
476 }
else if (
nodeproto.op_type() ==
"Gemm") {
482 }
else if (
nodeproto.op_type() ==
"BatchNormalization") {
492 std::cout <<
"operator " << op_type <<
" is not supported" << std::endl;
493 throw std::runtime_error(
"TMVA::SOFIE Operator type " + op_type +
" is not yet supported");
496 std::cout <<
"\tCreating operator " << op_type << std::endl;
510 throw std::runtime_error(
"TMVA::SOFIE - Failed to load onnx file " +
filename);
515 std::time_t
ttime = std::time(0);
526 if (
isep != std::string::npos) {
547 throw std::runtime_error(
"TMVA::SOFIE - Failed to parse ONNX model from input stream");
551 std::time_t
ttime = std::time(0);
575 std::fstream
input(
filename, std::ios::in | std::ios::binary);
577 std::cerr <<
"TMVA::SOFIE - Failed to open onnx file " <<
filename << std::endl;
586 auto model = std::make_unique<onnx::ModelProto>();
588 if (!model->ParseFromIstream(&
input)) {
589 std::cerr <<
"TMVA::SOFIE - Failed to parse ONNX model from input stream" << std::endl;
595 std::cout <<
"ONNX Version " << model->ir_version() << std::endl;
602 std::cout <<
"\n" << graph.
name() <<
" Graph operator list\n";
603 for (
int i = 0; i < graph.
node_size(); i++) {
604 const auto & node = graph.
node(i);
605 const std::string
opType = node.op_type();
607 std::cout <<
"\tOperator " << i <<
" : " <<
opType <<
" (" << node.name() <<
"), " << graph.
node(i).input_size()
609 for (
int j = 0;
j < graph.
node(i).input_size();
j++) {
610 std::cout << graph.
node(i).input(
j);
611 if (
j < graph.
node(i).input_size() - 1)
614 std::cout <<
" }" << std::endl;
620 for (
int j = 0;
j < node.attribute_size();
j++) {
621 const auto & attribute = node.attribute(
j);
622 if (attribute.has_g()) {
623 const auto &
subGraph = attribute.g();
635 if (!model)
return false;
640 std::cout <<
"\nModel operator list " << model->producer_name() <<
"\n";
647 std::cout <<
"List of missing operators for model loaded from file " <<
filename << std::endl;
649 std::cout <<
op.first <<
" " <<
op.second << std::endl;
653 std::cout <<
"All operators in the loaded model are supported!\n";
665 std::cout <<
"\nParsing Graph - " <<
graphName << std::endl;
671 std::map<int, std::pair<EFusedOp, int>> &fMap;
672 std::map<int, std::pair<EFusedOp, int>> fSaved;
673 FusedOperatorsGuard(std::map<
int, std::pair<EFusedOp, int>> &map) : fMap(map) { fSaved.swap(fMap); }
683 std::cout <<
"Parsing model inputs...." << std::endl;
685 for (
int i = 0; i < graph.
input_size(); i++) {
690 std::cout <<
"\tgraph input " << i <<
" name " << graph.
input(i).name() <<
" type "
691 << graph.
input(i).type().tensor_type().elem_type() << std::endl;
705 throw std::runtime_error(
"TMVA::SOFIE data node with no shape restrictions is not supported yet");
706 for (
int j = 0;
j <
valueinfoproto.type().tensor_type().shape().dim_size();
j++) {
710 int dim_value =
valueinfoproto.type().tensor_type().shape().dim(
j).dim_value();
718 }
else if (
valueinfoproto.type().tensor_type().shape().dim(
j).value_case() ==
724 throw std::runtime_error(
"TMVA::SOFIE ONNX file error: Valueinfoproto " +
input_name +
725 " has neither dim_value nor dim_param! \n");
729 if (
valueinfoproto.type().tensor_type().shape().dim_size() == 0) {
751 std::cout <<
"\nParsing graph initializer list and fill model initialized tensors" << std::endl;
755 std::vector<std::size_t> shape;
763 std::string tensor_name = graph.
initializer(i).name();
766 std::cout <<
"\t initializer " << i <<
" name " << tensor_name <<
" type " << graph.
initializer(i).data_type()
775 rmodel.AddInitializedTensor(tensor_name, tensor_type, shape,
data);
779 std::cout <<
"add initialized tensor " << tensor_name <<
"with shape " <<
ConvertShapeToString(shape) <<
"and ";
781 std::cout <<
" float data: ";
785 std::cout <<
" int64 data: ";
789 std::cout <<
" uint8 data: ";
793 std::cout <<
" Boolean data: ";
796 std::cout << std::endl;
802 std::cout <<
"\nGraph operator list (ONNX order)\n";
803 for (
int i = 0; i < graph.
node_size(); i++) {
804 std::cout <<
"\tOperator " << i <<
" : " << graph.
node(i).op_type() <<
" , " << graph.
node(i).input_size()
806 for (
int j = 0;
j < graph.
node(i).input_size();
j++) {
807 std::cout << graph.
node(i).input(
j);
808 if (
j < graph.
node(i).input_size() - 1)
811 std::cout <<
" }" << std::endl;
817 std::cout <<
"\n***********************\nRe-Order graph operator list\n*************************\n";
824 for (
int i = 0; i < graph.
input_size(); i++) {
829 for (
int i = 0; i < graph.
node_size(); i++) {
834 int input_size = graph.
node(i).input_size();
837 std::cout <<
"Checking input of Node " << i <<
" : " << graph.
node(i).name() << std::endl;
838 for (
int j = 0;
j < input_size;
j++) {
839 std::string
name = graph.
node(i).input(
j);
845 std::cout <<
"\t\t input " <<
name <<
" "
854 std::cout <<
"skip node " << graph.
node(i).op_type() <<
" " << graph.
node(i).name() <<
" inputs are not existing ";
855 for (
int j = 0;
j < input_size;
j++) {
856 std::cout << graph.
node(i).input(
j) <<
" ";
858 std::cout << std::endl;
865 std::cout <<
"===> New node " << graph.
node(i).op_type() <<
" " << graph.
node(i).name() <<
" order " << i << std::endl;
870 for (
int j = 0;
j < graph.
node(i).output_size();
j++) {
871 if (
fVerbose) std::cout <<
"\toutput : " << graph.
node(i).output(
j) << std::endl;
878 std::cout <<
"cannot find a new node after " << graph.
node(
ilast).op_type() <<
" " << graph.
node(
ilast).name() << std::endl;
879 throw std::runtime_error(
"TMVA::SOFIE - cannot find a new node ");
887 for (
int k = 0; k < graph.
node_size(); k++) {
905 std::cout <<
"\nGraph operator list (re-ordered)\n";
906 for (
int k = 0; k < graph.
node_size(); k++) {
908 std::cout <<
"\tOperator " << i <<
" : " << graph.
node(i).op_type() <<
" , " << graph.
node(i).name() <<
" input tensors : {";
909 for (
int j = 0;
j < graph.
node(i).input_size();
j++) {
910 std::cout << graph.
node(i).input(
j);
911 if (
j < graph.
node(i).input_size() - 1)
915 std::cout <<
" children : {";
919 std::cout <<
"}" << std::endl;
925 std::cout <<
"Fill RModel with operators...\n";
931 for (
int i = 0; i < graph.
node_size(); i++) {
935 std::cout <<
"\t" << i <<
" " <<
nodesOrder[i] <<
" parsing operator " << op_type << std::endl;
941 std::cout <<
"\t\tskipping operator since it is fused with previous one" << std::endl;
951 std::cout <<
"\nParsing Graph output list\n";
954 std::cout <<
"\toutput " << i <<
" name " << graph.
output(i).name() << std::endl;
ROOT::Detail::TRangeCast< T, true > TRangeDynCast
TRangeDynCast is an adapter class that allows the typed iteration through a TCollection.
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void data
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void input
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void char Point_t Rectangle_t dest
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void char Point_t Rectangle_t WindowAttributes_t Float_t Float_t Float_t Int_t Int_t UInt_t UInt_t Rectangle_t Int_t Int_t Window_t TString Int_t GCValues_t GetPrimarySelectionOwner GetDisplay GetScreen GetColormap GetNativeEvent const char const char dpyName wid window const char font_name cursor keysym reg const char only_if_exist regb h Point_t winding char text const char depth char const char Int_t count const char ColorStruct_t color const char filename
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void char Point_t Rectangle_t WindowAttributes_t Float_t Float_t Float_t Int_t Int_t UInt_t UInt_t Rectangle_t Int_t Int_t Window_t TString Int_t GCValues_t GetPrimarySelectionOwner GetDisplay GetScreen GetColormap GetNativeEvent const char const char dpyName wid window const char font_name cursor keysym reg const char only_if_exist regb h Point_t winding char text const char depth char const char Int_t count const char ColorStruct_t color const char Pixmap_t Pixmap_t PictureAttributes_t attr const char char ret_data h unsigned char height h offset
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void char Point_t Rectangle_t WindowAttributes_t Float_t Float_t Float_t Int_t Int_t UInt_t UInt_t Rectangle_t Int_t Int_t Window_t TString Int_t GCValues_t GetPrimarySelectionOwner GetDisplay GetScreen GetColormap GetNativeEvent const char const char dpyName wid window const char font_name cursor keysym reg const char only_if_exist regb h Point_t winding char text const char depth char const char Int_t count const char ColorStruct_t color const char Pixmap_t Pixmap_t PictureAttributes_t attr const char char ret_data h unsigned char height h length
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void char Point_t Rectangle_t src
Option_t Option_t TPoint TPoint const char GetTextMagnitude GetFillStyle GetLineColor GetLineWidth GetMarkerStyle GetTextAlign GetTextColor GetTextSize void char Point_t Rectangle_t WindowAttributes_t Float_t Float_t Float_t Int_t Int_t UInt_t UInt_t Rectangle_t Int_t Int_t Window_t TString Int_t GCValues_t GetPrimarySelectionOwner GetDisplay GetScreen GetColormap GetNativeEvent const char const char dpyName wid window const char font_name cursor keysym reg const char only_if_exist regb h Point_t winding char text const char depth char const char Int_t count const char ColorStruct_t color const char Pixmap_t Pixmap_t PictureAttributes_t attr const char char ret_data h unsigned char height h Atom_t Int_t ULong_t ULong_t unsigned char prop_list Atom_t Atom_t Atom_t Time_t type
const_iterator begin() const
const_iterator end() const
std::string fDefaultDataFileName
std::string fOpenedDataFileName
void RegisterOperator(const std::string &name, ParserFuncSignature func)
std::unique_ptr< ROperator > ParseOperator(const size_t, const onnx::GraphProto &, const std::vector< size_t > &, const std::vector< int > &)
std::string fDataFileName
bool IsRegisteredOperator(const std::string &name)
void CheckGraph(const onnx::GraphProto &g, int &level, std::map< std::string, int > &missingOperators)
void ParseONNXGraph(RModel &model, const onnx::GraphProto &g, std::string name="")
RModelParser_ONNX() noexcept
std::unordered_map< std::string, ETensorType > fTensorTypeMap
RModel Parse(std::string const &filename, bool verbose=false)
std::shared_ptr< void > GetInitializedTensorData(onnx::TensorProto *tensorproto, size_t tensor_length, ETensorType type)
std::map< int, std::pair< EFusedOp, int > > fFusedOperators
bool IsRegisteredTensorType(const std::string &)
void RegisterTensorType(const std::string &, ETensorType)
void ResetExternalDataState()
ETensorType GetTensorType(const std::string &name)
std::string fModelDirectory
std::vector< std::string > GetRegisteredOperators()
std::unique_ptr< onnx::ModelProto > LoadModel(const std::string &filename)
std::unique_ptr< OperatorsMapImpl > fOperatorsMapImpl
bool CheckModel(std::string filename, bool verbose=false)
const ValueInfoProto & input(int i) const
const ValueInfoProto & output(int i) const
int initializer_size() const
const std::string & name() const
const NodeProto & node(int i) const
const TensorProto & initializer(int i) const
std::string Clean_name(std::string input_tensor_name)
ParserFuncSignature ParseBatchNormalization
std::function< std::unique_ptr< ROperator >(RModelParser_ONNX &, const onnx::NodeProto &, const onnx::NodeProto &)> ParserFuseFuncSignature
ParserFuncSignature ParseReshape
ParserFuseFuncSignature ParseFuseConvTransposeAdd
ParserFuseFuncSignature ParseFuseMatMulAdd
ParserFuncSignature ParseGather
ParserFuncSignature ParseWhere
ParserFuncSignature ParseLeakyRelu
std::function< std::unique_ptr< ROperator >(RModelParser_ONNX &, const onnx::NodeProto &)> ParserFuncSignature
ParserFuncSignature ParseEinsum
ParserFuncSignature ParsePool
ParserFuncSignature ParseLayerNormalization
ParserFuncSignature ParseConcat
ParserFuncSignature ParseTopK
void RegisterReduceParsers(RModelParser_ONNX &parser)
ParserFuncSignature ParseIdentity
ParserFuncSignature ParseConvTranspose
ParserFuncSignature ParseNot
ParserFuncSignature ParseSlice
ParserFuncSignature ParseRandom
ParserFuncSignature ParseTranspose
ParserFuncSignature ParseShape
ParserFuncSignature ParseClip
constexpr size_t GetTypeSize(ETensorType type)
ParserFuncSignature ParseScatterND
void RegisterBasicBinaryParsers(RModelParser_ONNX &parser)
ParserFuncSignature ParseGRU
ParserFuncSignature ParseMatMul
ParserFuncSignature ParseErf
ParserFuncSignature ParseNonZero
ParserFuncSignature ParseIf
ParserFuncSignature ParseRange
ParserFuncSignature ParseExpand
ParserFuncSignature ParseRNN
ParserFuncSignature ParseHardSigmoid
ParserFuncSignature ParseLSTM
ParserFuncSignature ParseCast
ParserFuncSignature ParseSwish
ParserFuncSignature ParseSigmoid
ParserFuseFuncSignature ParseFuseConvAdd
ParserFuseFuncSignature ParseFuseBatchnormRelu
ParserFuncSignature ParseSoftmax
void RegisterComparisionParsers(RModelParser_ONNX &parser)
std::string ConvertTypeToString(ETensorType type)
ParserFuncSignature ParseGelu
ParserFuncSignature ParseSplit
ParserFuncSignature ParseConstant
ParserFuncSignature ParseSelu
ParserFuncSignature ParseHardSwish
ParserFuncSignature ParseGatherND
ParserFuncSignature ParseEyeLike
ParserFuncSignature ParsePad
ParserFuncSignature ParseElu
void RegisterBasicIsParsers(RModelParser_ONNX &parser)
std::string ConvertShapeToString(const std::vector< size_t > &shape)
ParserFuncSignature ParseRelu
ParserFuncSignature ParseConv
ParserFuncSignature ParseInstanceNormalization
ParserFuncSignature ParseScatterElements
ParserFuncSignature ParseGemm
ParserFuncSignature ParseTile
ParserFuseFuncSignature ParseFuseGemmRelu
void RegisterBasicUnaryParsers(RModelParser_ONNX &parser)
void RegisterBasicNaryParsers(RModelParser_ONNX &parser)
ParserFuncSignature ParseTanh
create variable transformations
Helper templated class for swapping bytes; specializations for N={2,4,8} are provided below.
std::unordered_map< std::string, ParserFuncSignature > fOperatorsMap