Back to home page

EIC code displayed by LXR

 
 

    


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 // forward declaration
0012 namespace onnx {
0013 class NodeProto;
0014 class GraphProto;
0015 class ModelProto;
0016 class TensorProto;
0017 } // namespace onnx
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 & /*parser*/, const onnx::NodeProto & /*nodeproto*/)>;
0027 using ParserFuseFuncSignature =
0028    std::function<std::unique_ptr<ROperator> (RModelParser_ONNX& /*parser*/, const onnx::NodeProto& /*firstnode*/, const onnx::NodeProto& /*secondnode*/)>;
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    // Registered operators
0040    std::unique_ptr<OperatorsMapImpl> fOperatorsMapImpl;
0041    // Type of the tensors
0042    std::unordered_map<std::string, ETensorType> fTensorTypeMap;
0043 
0044    // List of fused operators storing as key the second operator and a value a pair of fusion type and parent operator
0045    std::map<int, std::pair<EFusedOp, int>> fFusedOperators;
0046 
0047    //  weight data file
0048    std::ifstream fDataFile;
0049    // filename of model
0050    std::string fDataFileName;
0051 
0052 
0053 public:
0054    // Register an ONNX operator
0055    void RegisterOperator(const std::string &name, ParserFuncSignature func);
0056 
0057    // Check if the operator is registered
0058    bool IsRegisteredOperator(const std::string &name);
0059 
0060    // List of registered operators (in alphabetical order)
0061    std::vector<std::string> GetRegisteredOperators();
0062 
0063    // Set the type of the tensor
0064    void RegisterTensorType(const std::string & /*name*/, ETensorType /*type*/);
0065 
0066    // Check if the type of the tensor is registered
0067    bool IsRegisteredTensorType(const std::string & /*name*/);
0068 
0069    // check verbosity
0070    bool Verbose() const {
0071       return fVerbose;
0072    }
0073 
0074    // Get the type of the tensor
0075    ETensorType GetTensorType(const std::string &name);
0076 
0077    // Parse the index'th node from the ONNX graph
0078    std::unique_ptr<ROperator> ParseOperator(const size_t /*index*/, const onnx::GraphProto & /*graphproto*/,
0079                                             const std::vector<size_t> & /*nodes*/, const std::vector<int> & /* children */);
0080 
0081    // check a graph for missing operators
0082    void CheckGraph(const onnx::GraphProto & g, int & level, std::map<std::string, int> & missingOperators);
0083 
0084    // parse the ONNX graph
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    // check the model for missing operators - return false in case some operator implementation is missing
0100    bool CheckModel(std::string filename, bool verbose = false);
0101 
0102    //set external data full path (needed if external data are not stored in the default modelName.onnx.data)
0103    // call this function before parsing
0104    void SetExternalDataFile(const std::string & dataFileName) {
0105       fDataFileName = dataFileName;
0106    }
0107 
0108    ~RModelParser_ONNX();
0109 };
0110 
0111 } // namespace SOFIE
0112 } // namespace Experimental
0113 } // namespace TMVA
0114 
0115 #endif // TMVA_SOFIE_RMODELPARSER_ONNX