File indexing completed on 2026-10-01 09:10:11
0001 #ifndef TMVA_SOFIE_RMODELPARSER_ONNX
0002 #define TMVA_SOFIE_RMODELPARSER_ONNX
0003
0004 #include "TMVA/RModel.hxx"
0005
0006 #include <memory>
0007 #include <functional>
0008 #include <unordered_map>
0009 #include <fstream>
0010
0011
0012 namespace onnx {
0013 class NodeProto;
0014 class GraphProto;
0015 class ModelProto;
0016 class TensorProto;
0017 }
0018
0019 namespace TMVA {
0020 namespace Experimental {
0021 namespace SOFIE {
0022
0023 class RModelParser_ONNX;
0024
0025 using ParserFuncSignature =
0026 std::function<std::unique_ptr<ROperator>(RModelParser_ONNX & , const onnx::NodeProto & )>;
0027 using ParserFuseFuncSignature =
0028 std::function<std::unique_ptr<ROperator> (RModelParser_ONNX& , const onnx::NodeProto& , const onnx::NodeProto& )>;
0029
0030 class RModelParser_ONNX {
0031 public:
0032 struct OperatorsMapImpl;
0033
0034 enum EFusedOp { kMatMulAdd, kConvAdd, kConvTransAdd, kGemmRelu, kBatchnormRelu};
0035
0036 private:
0037
0038 bool fVerbose = false;
0039
0040 std::unique_ptr<OperatorsMapImpl> fOperatorsMapImpl;
0041
0042 std::unordered_map<std::string, ETensorType> fTensorTypeMap;
0043
0044
0045 std::map<int, std::pair<EFusedOp, int>> fFusedOperators;
0046
0047
0048 std::ifstream fDataFile;
0049
0050 std::string fDataFileName;
0051
0052
0053 public:
0054
0055 void RegisterOperator(const std::string &name, ParserFuncSignature func);
0056
0057
0058 bool IsRegisteredOperator(const std::string &name);
0059
0060
0061 std::vector<std::string> GetRegisteredOperators();
0062
0063
0064 void RegisterTensorType(const std::string & , ETensorType );
0065
0066
0067 bool IsRegisteredTensorType(const std::string & );
0068
0069
0070 bool Verbose() const {
0071 return fVerbose;
0072 }
0073
0074
0075 ETensorType GetTensorType(const std::string &name);
0076
0077
0078 std::unique_ptr<ROperator> ParseOperator(const size_t , const onnx::GraphProto & ,
0079 const std::vector<size_t> & , const std::vector<int> & );
0080
0081
0082 void CheckGraph(const onnx::GraphProto & g, int & level, std::map<std::string, int> & missingOperators);
0083
0084
0085 void ParseONNXGraph(RModel & model, const onnx::GraphProto & g, std::string name = "");
0086
0087 std::unique_ptr<onnx::ModelProto> LoadModel(const std::string &filename);
0088 std::unique_ptr<onnx::ModelProto> LoadModel(std::istream &input);
0089
0090 std::shared_ptr<void> GetInitializedTensorData(onnx::TensorProto *tensorproto, size_t tensor_length, ETensorType type );
0091
0092 public:
0093
0094 RModelParser_ONNX() noexcept;
0095
0096 RModel Parse(std::string const &filename, bool verbose = false);
0097 RModel Parse(std::istream &input, std::string const &name, bool verbose = false);
0098
0099
0100 bool CheckModel(std::string filename, bool verbose = false);
0101
0102
0103
0104 void SetExternalDataFile(const std::string & dataFileName) {
0105 fDataFileName = dataFileName;
0106 }
0107
0108 ~RModelParser_ONNX();
0109 };
0110
0111 }
0112 }
0113 }
0114
0115 #endif