12#include <unordered_map>
17namespace Experimental {
136 throw std::runtime_error(
"TMVA::SOFIE - Failed to read float initialized tensor - actual size is " + std::to_string(
tensor->float_data_size()));
138 std::copy(
src.begin(),
src.end(),
static_cast<float *
>(
data));
145 throw std::runtime_error(
"TMVA::SOFIE - Failed to read double initialized tensor - actual size is " + std::to_string(
tensor->double_data_size()));
147 std::copy(
src.begin(),
src.end(),
static_cast<double *
>(
data));
154 throw std::runtime_error(
"TMVA::SOFIE - Failed to read int32 initialized tensor - actual size is " + std::to_string(
tensor->int32_data_size()));
156 std::copy(
src.begin(),
src.end(),
static_cast<int32_t *
>(
data));
163 throw std::runtime_error(
"TMVA::SOFIE - Failed to read int64 initialized tensor - actual size is " + std::to_string(
tensor->int64_data_size()));
165 std::copy(
src.begin(),
src.end(),
static_cast<int64_t *
>(
data));
175template <std::
size_t N>
179 auto dst =
static_cast<unsigned char *
>(
dest);
180 auto src =
static_cast<const unsigned char *
>(
source);
181 for (std::size_t k = 0; k <
nbytes; k +=
N) {
183 std::memcpy(&
v,
src + k,
N);
185 std::memcpy(
dst + k, &
v,
N);
201 throw std::runtime_error(
"Data type " +
ConvertTypeToString(tensor_type) +
" in tensor is not supported!\n");
211 std::shared_ptr<void>
data(
malloc(tensor_size), free);
217 throw std::runtime_error(
"TMVA::SOFIE - Failed to read raw data of initialized tensor - actual raw size is " +
229 switch (tensor_type) {
247 throw std::runtime_error(
"TMVA::SOFIE - ExtractData from TP in BOOL not supported");
251 throw std::runtime_error(
"TMVA::SOFIE - ExtractData from TP in UINT8 not supported");
255 throw std::runtime_error(
"Data type " +
ConvertTypeToString(tensor_type) +
" in weight tensor is not supported!\n");
263 std::string location;
267 if (
kv.key() ==
"location") location =
kv.value();
268 else if (
kv.key() ==
"offset")
offset = std::stoull(
kv.value());
280 throw std::runtime_error(
"TMVA::SOFIE ONNX : tensor " +
tensorproto->name() +
281 " has external data but no data file location is available");
284 std::cout <<
"Initialized data are stored externally in file " <<
dataFileName
285 <<
" at location " << location <<
" offset " <<
offset <<
" and with length " <<
buffer_size << std::endl;
288 throw std::runtime_error(
"TMVA::SOFIE ONNX : invalid stored data size vs tensor size");
296 throw std::runtime_error(
"TMVA::SOFIE ONNX: error reading external weight ONNX data file " +
dataFileName);
435 std::vector<std::string>
ops;
438 ops.emplace_back(it.first);
461std::unique_ptr<ROperator>
464 if (i >= nodes.size())
465 throw std::runtime_error(
"TMVA::SOFIE - Error in parsing ordered operators " + std::to_string(i) +
" is >= " + std::to_string(nodes.size()));
468 const std::string op_type =
nodeproto.op_type();
470 std::cout <<
"Parsing operator " << op_type << std::endl;
476 std::cout <<
"\tFusing operators " <<
graphproto.node(
idx1).name()
493 if (children.size() == 1) {
494 int idx2 = children.front();
495 if (op_type ==
"MatMul") {
512 }
else if (
nodeproto.op_type() ==
"Gemm") {
518 }
else if (
nodeproto.op_type() ==
"BatchNormalization") {
528 std::cout <<
"operator " << op_type <<
" is not supported" << std::endl;
529 throw std::runtime_error(
"TMVA::SOFIE Operator type " + op_type +
" is not yet supported");
532 std::cout <<
"\tCreating operator " << op_type << std::endl;
546 throw std::runtime_error(
"TMVA::SOFIE - Failed to load onnx file " +
filename);
551 std::time_t
ttime = std::time(0);
562 if (
isep != std::string::npos) {
583 throw std::runtime_error(
"TMVA::SOFIE - Failed to parse ONNX model from input stream");
587 std::time_t
ttime = std::time(0);
611 std::fstream
input(
filename, std::ios::in | std::ios::binary);
613 std::cerr <<
"TMVA::SOFIE - Failed to open onnx file " <<
filename << std::endl;
622 auto model = std::make_unique<onnx::ModelProto>();
624 if (!model->ParseFromIstream(&
input)) {
625 std::cerr <<
"TMVA::SOFIE - Failed to parse ONNX model from input stream" << std::endl;
631 std::cout <<
"ONNX Version " << model->ir_version() << std::endl;
638 std::cout <<
"\n" << graph.
name() <<
" Graph operator list\n";
639 for (
int i = 0; i < graph.
node_size(); i++) {
640 const auto & node = graph.
node(i);
641 const std::string
opType = node.op_type();
643 std::cout <<
"\tOperator " << i <<
" : " <<
opType <<
" (" << node.name() <<
"), " << graph.
node(i).input_size()
645 for (
int j = 0;
j < graph.
node(i).input_size();
j++) {
646 std::cout << graph.
node(i).input(
j);
647 if (
j < graph.
node(i).input_size() - 1)
650 std::cout <<
" }" << std::endl;
656 for (
int j = 0;
j < node.attribute_size();
j++) {
657 const auto & attribute = node.attribute(
j);
658 if (attribute.has_g()) {
659 const auto &
subGraph = attribute.g();
671 if (!model)
return false;
676 std::cout <<
"\nModel operator list " << model->producer_name() <<
"\n";
683 std::cout <<
"List of missing operators for model loaded from file " <<
filename << std::endl;
685 std::cout <<
op.first <<
" " <<
op.second << std::endl;
689 std::cout <<
"All operators in the loaded model are supported!\n";
701 std::cout <<
"\nParsing Graph - " <<
graphName << std::endl;
707 std::map<int, std::pair<EFusedOp, int>> &fMap;
708 std::map<int, std::pair<EFusedOp, int>> fSaved;
709 FusedOperatorsGuard(std::map<
int, std::pair<EFusedOp, int>> &map) : fMap(map) { fSaved.swap(fMap); }
719 std::cout <<
"Parsing model inputs...." << std::endl;
721 for (
int i = 0; i < graph.
input_size(); i++) {
726 std::cout <<
"\tgraph input " << i <<
" name " << graph.
input(i).name() <<
" type "
727 << graph.
input(i).type().tensor_type().elem_type() << std::endl;
741 throw std::runtime_error(
"TMVA::SOFIE data node with no shape restrictions is not supported yet");
742 for (
int j = 0;
j <
valueinfoproto.type().tensor_type().shape().dim_size();
j++) {
746 int dim_value =
valueinfoproto.type().tensor_type().shape().dim(
j).dim_value();
754 }
else if (
valueinfoproto.type().tensor_type().shape().dim(
j).value_case() ==
760 throw std::runtime_error(
"TMVA::SOFIE ONNX file error: Valueinfoproto " +
input_name +
761 " has neither dim_value nor dim_param! \n");
765 if (
valueinfoproto.type().tensor_type().shape().dim_size() == 0) {
787 std::cout <<
"\nParsing graph initializer list and fill model initialized tensors" << std::endl;
791 std::vector<std::size_t> shape;
799 std::string tensor_name = graph.
initializer(i).name();
802 std::cout <<
"\t initializer " << i <<
" name " << tensor_name <<
" type " << graph.
initializer(i).data_type()
811 rmodel.AddInitializedTensor(tensor_name, tensor_type, shape,
data);
815 std::cout <<
"add initialized tensor " << tensor_name <<
"with shape " <<
ConvertShapeToString(shape) <<
"and ";
817 std::cout <<
" float data: ";
821 std::cout <<
" int64 data: ";
825 std::cout <<
" uint8 data: ";
829 std::cout <<
" Boolean data: ";
832 std::cout << std::endl;
838 std::cout <<
"\nGraph operator list (ONNX order)\n";
839 for (
int i = 0; i < graph.
node_size(); i++) {
840 std::cout <<
"\tOperator " << i <<
" : " << graph.
node(i).op_type() <<
" , " << graph.
node(i).input_size()
842 for (
int j = 0;
j < graph.
node(i).input_size();
j++) {
843 std::cout << graph.
node(i).input(
j);
844 if (
j < graph.
node(i).input_size() - 1)
847 std::cout <<
" }" << std::endl;
853 std::cout <<
"\n***********************\nRe-Order graph operator list\n*************************\n";
860 for (
int i = 0; i < graph.
input_size(); i++) {
865 for (
int i = 0; i < graph.
node_size(); i++) {
870 int input_size = graph.
node(i).input_size();
873 std::cout <<
"Checking input of Node " << i <<
" : " << graph.
node(i).name() << std::endl;
874 for (
int j = 0;
j < input_size;
j++) {
875 std::string
name = graph.
node(i).input(
j);
881 std::cout <<
"\t\t input " <<
name <<
" "
890 std::cout <<
"skip node " << graph.
node(i).op_type() <<
" " << graph.
node(i).name() <<
" inputs are not existing ";
891 for (
int j = 0;
j < input_size;
j++) {
892 std::cout << graph.
node(i).input(
j) <<
" ";
894 std::cout << std::endl;
901 std::cout <<
"===> New node " << graph.
node(i).op_type() <<
" " << graph.
node(i).name() <<
" order " << i << std::endl;
906 for (
int j = 0;
j < graph.
node(i).output_size();
j++) {
907 if (
fVerbose) std::cout <<
"\toutput : " << graph.
node(i).output(
j) << std::endl;
914 std::cout <<
"cannot find a new node after " << graph.
node(
ilast).op_type() <<
" " << graph.
node(
ilast).name() << std::endl;
915 throw std::runtime_error(
"TMVA::SOFIE - cannot find a new node ");
923 for (
int k = 0; k < graph.
node_size(); k++) {
941 std::cout <<
"\nGraph operator list (re-ordered)\n";
942 for (
int k = 0; k < graph.
node_size(); k++) {
944 std::cout <<
"\tOperator " << i <<
" : " << graph.
node(i).op_type() <<
" , " << graph.
node(i).name() <<
" input tensors : {";
945 for (
int j = 0;
j < graph.
node(i).input_size();
j++) {
946 std::cout << graph.
node(i).input(
j);
947 if (
j < graph.
node(i).input_size() - 1)
951 std::cout <<
" children : {";
955 std::cout <<
"}" << std::endl;
961 std::cout <<
"Fill RModel with operators...\n";
967 for (
int i = 0; i < graph.
node_size(); i++) {
971 std::cout <<
"\t" << i <<
" " <<
nodesOrder[i] <<
" parsing operator " << op_type << std::endl;
977 std::cout <<
"\t\tskipping operator since it is fused with previous one" << std::endl;
987 std::cout <<
"\nParsing Graph output list\n";
990 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 ParseIsNaN
ParserFuncSignature ParseSqrt
ParserFuncSignature ParseBatchNormalization
ParserFuncSignature ParseGreater
std::function< std::unique_ptr< ROperator >(RModelParser_ONNX &, const onnx::NodeProto &, const onnx::NodeProto &)> ParserFuseFuncSignature
ParserFuncSignature ParseReshape
ParserFuseFuncSignature ParseFuseConvTransposeAdd
ParserFuncSignature ParseReduceMean
ParserFuseFuncSignature ParseFuseMatMulAdd
ParserFuncSignature ParseGather
ParserFuncSignature ParseNeg
ParserFuncSignature ParseWhere
ParserFuncSignature ParseCos
ParserFuncSignature ParseLog
ParserFuncSignature ParseLeakyRelu
ParserFuncSignature ParseExp
std::function< std::unique_ptr< ROperator >(RModelParser_ONNX &, const onnx::NodeProto &)> ParserFuncSignature
ParserFuncSignature ParseEinsum
ParserFuncSignature ParsePool
ParserFuncSignature ParseDiv
ParserFuncSignature ParseLayerNormalization
ParserFuncSignature ParseConcat
ParserFuncSignature ParseTopK
ParserFuncSignature ParseMax
ParserFuncSignature ParseEq
ParserFuncSignature ParseIdentity
ParserFuncSignature ParseConvTranspose
ParserFuncSignature ParseReduceProd
ParserFuncSignature ParseNot
ParserFuncSignature ParseSlice
ParserFuncSignature ParseRandom
ParserFuncSignature ParseTranspose
ParserFuncSignature ParseLess
ParserFuncSignature ParseShape
ParserFuncSignature ParseClip
constexpr size_t GetTypeSize(ETensorType type)
ParserFuncSignature ParseScatterND
ParserFuncSignature ParseGRU
ParserFuncSignature ParseMatMul
ParserFuncSignature ParseErf
ParserFuncSignature ParseSub
ParserFuncSignature ParseAdd
ParserFuncSignature ParseNonZero
ParserFuncSignature ParseIf
ParserFuncSignature ParseRange
ParserFuncSignature ParseSoftplus
ParserFuncSignature ParseExpand
ParserFuncSignature ParseRNN
ParserFuncSignature ParseHardSigmoid
ParserFuncSignature ParseLSTM
ParserFuncSignature ParseCast
ParserFuncSignature ParseReciprocal
ParserFuncSignature ParseSwish
ParserFuncSignature ParseSigmoid
ParserFuseFuncSignature ParseFuseConvAdd
ParserFuncSignature ParseAtan
ParserFuncSignature ParseReduceMax
ParserFuncSignature ParseFloor
ParserFuseFuncSignature ParseFuseBatchnormRelu
ParserFuncSignature ParseIsInf
ParserFuncSignature ParseSoftmax
ParserFuncSignature ParseGreaterEq
ParserFuncSignature ParseMod
std::string ConvertTypeToString(ETensorType type)
ParserFuncSignature ParseGelu
ParserFuncSignature ParseMean
ParserFuncSignature ParseSplit
ParserFuncSignature ParseConstant
ParserFuncSignature ParseSelu
ParserFuncSignature ParseAsinh
ParserFuncSignature ParseLessEq
ParserFuncSignature ParseAcosh
ParserFuncSignature ParseHardSwish
ParserFuncSignature ParseGatherND
ParserFuncSignature ParseSum
ParserFuncSignature ParseEyeLike
ParserFuncSignature ParsePad
ParserFuncSignature ParseElu
std::string ConvertShapeToString(const std::vector< size_t > &shape)
ParserFuncSignature ParseMin
ParserFuncSignature ParseRelu
ParserFuncSignature ParseReduceSum
ParserFuncSignature ParseConv
ParserFuncSignature ParseInstanceNormalization
ParserFuncSignature ParseScatterElements
ParserFuncSignature ParseGemm
ParserFuncSignature ParseTile
ParserFuncSignature ParseMul
ParserFuseFuncSignature ParseFuseGemmRelu
ParserFuncSignature ParsePow
ParserFuncSignature ParseAbs
ParserFuncSignature ParseSin
ParserFuncSignature ParseAtanh
ParserFuncSignature ParseReduceSumSquare
ParserFuncSignature ParseTanh
ParserFuncSignature ParseReduceMin
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