[ Web Proxy ]
URL:
Viewing: https://raw.githubusercontent.com/cardboardcode/onnx_runtime_cpp/develop/src/OrtSessionHandler.cpp [Back]  [Original]

/**
 * @file    OrtSessionHandler.cpp
 *
 * @author  btran
 *
 */

#include 
#include 

#if ENABLE_TENSORRT
#include 
#endif

#include 

#include 
#include 
#include 
#include 

namespace
{
std::string toString(const ONNXTensorElementDataType dataType)
{
    switch (dataType) {
        case ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT: {
            return "float";
        }
        case ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT8: {
            return "uint8_t";
        }
        case ONNX_TENSOR_ELEMENT_DATA_TYPE_INT8: {
            return "int8_t";
        }
        case ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT16: {
            return "uint16_t";
        }
        case ONNX_TENSOR_ELEMENT_DATA_TYPE_INT16: {
            return "int16_t";
        }
        case ONNX_TENSOR_ELEMENT_DATA_TYPE_INT32: {
            return "int32_t";
        }
        case ONNX_TENSOR_ELEMENT_DATA_TYPE_INT64: {
            return "int64_t";
        }
        case ONNX_TENSOR_ELEMENT_DATA_TYPE_STRING: {
            return "string";
        }
        case ONNX_TENSOR_ELEMENT_DATA_TYPE_BOOL: {
            return "bool";
        }
        case ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT16: {
            return "float16";
        }
        case ONNX_TENSOR_ELEMENT_DATA_TYPE_DOUBLE: {
            return "double";
        }
        case ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT32: {
            return "uint32_t";
        }
        case ONNX_TENSOR_ELEMENT_DATA_TYPE_UINT64: {
            return "uint64_t";
        }
        case ONNX_TENSOR_ELEMENT_DATA_TYPE_COMPLEX64: {
            return "complex with float32 real and imaginary components";
        }
        case ONNX_TENSOR_ELEMENT_DATA_TYPE_COMPLEX128: {
            return "complex with float64 real and imaginary components";
        }
        case ONNX_TENSOR_ELEMENT_DATA_TYPE_BFLOAT16: {
            return "complex with float64 real and imaginary components";
        }
        default:
            return "undefined";
    }
}
}  // namespace

namespace Ort
{
//-----------------------------------------------------------------------------//
// OrtSessionHandlerIml Definition
//-----------------------------------------------------------------------------//

class OrtSessionHandler::OrtSessionHandlerIml
{
 public:
    OrtSessionHandlerIml(const std::string& modelPath,         //
                         const std::optional& gpuIdx,  //
                         const std::optional& inputShapes);
    ~OrtSessionHandlerIml();

    std::vector operator()(const std::vector& inputData) const;

    void updateInputShapes(const std::vector& inputShapes)
    {
        if (inputShapes.size() != m_numInputs) {
            DEBUG_LOG("inputShapes must be of size: %d", m_numInputs);
            return;
        }
        m_inputShapes = inputShapes;

        for (int i = 0; i < m_numInputs; i++) {
            const auto& curInputShape = m_inputShapes[i];
            m_inputTensorSizes[i] =
                std::accumulate(std::begin(curInputShape), std::end(curInputShape), 1, std::multiplies());
        }
    }

 private:
    void initSession();
    void initModelInfo();

 private:
    std::string m_modelPath;

    mutable Ort::Session m_session;
    Ort::Env m_env;
    Ort::AllocatorWithDefaultOptions m_ortAllocator;

    std::optional m_gpuIdx;

    std::vector m_inputShapes;
    std::vector m_outputShapes;

    std::vector m_inputTensorSizes;
    std::vector m_outputTensorSizes;

    uint8_t m_numInputs;
    uint8_t m_numOutputs;

    std::vector m_inputNodeNames;
    std::vector m_outputNodeNames;

    bool m_inputShapesProvided = false;
};

//-----------------------------------------------------------------------------//
// OrtSessionHandler
//-----------------------------------------------------------------------------//

OrtSessionHandler::OrtSessionHandler(const std::string& modelPath,         //
                                     const std::optional& gpuIdx,  //
                                     const std::optional& inputShapes)
    : m_piml(std::make_unique(modelPath,  //
                                                    gpuIdx,     //
                                                    inputShapes))
{
}

OrtSessionHandler::~OrtSessionHandler() = default;

std::vector
OrtSessionHandler::operator()(const std::vector& inputImgData) const
{
    return this->m_piml->operator()(inputImgData);
}

//-----------------------------------------------------------------------------//
// piml class implementation
//-----------------------------------------------------------------------------//

OrtSessionHandler::OrtSessionHandlerIml::OrtSessionHandlerIml(
    const std::string& modelPath,         //
    const std::optional& gpuIdx,  //
    const std::optional& inputShapes)
    : m_modelPath(modelPath)
    , m_session(nullptr)
    , m_env(nullptr)
    , m_ortAllocator()
    , m_gpuIdx(gpuIdx)
    , m_inputShapes()
    , m_outputShapes()
    , m_numInputs(0)
    , m_numOutputs(0)
    , m_inputNodeNames()
    , m_outputNodeNames()
{
    this->initSession();

    if (inputShapes.has_value()) {
        m_inputShapesProvided = true;
        m_inputShapes = inputShapes.value();
    }

    this->initModelInfo();
}

OrtSessionHandler::OrtSessionHandlerIml::~OrtSessionHandlerIml()
{
    for (auto& elem : this->m_inputNodeNames) {
        free(elem);
        elem = nullptr;
    }
    this->m_inputNodeNames.clear();

    for (auto& elem : this->m_outputNodeNames) {
        free(elem);
        elem = nullptr;
    }
    this->m_outputNodeNames.clear();
}

void OrtSessionHandler::OrtSessionHandlerIml::initSession()
{
#if ENABLE_DEBUG
    m_env = Ort::Env(ORT_LOGGING_LEVEL_WARNING, "test");
#else
    m_env = Ort::Env(ORT_LOGGING_LEVEL_ERROR, "test");
#endif
    Ort::SessionOptions sessionOptions;

    sessionOptions.SetIntraOpNumThreads(1);
    // tensorrt options can be customized into sessionOptions
    // https://onnxruntime.ai/docs/execution-providers/TensorRT-ExecutionProvider.html

#if ENABLE_GPU
    if (m_gpuIdx.has_value()) {
        Ort::ThrowOnError(OrtSessionOptionsAppendExecutionProvider_CUDA(sessionOptions, m_gpuIdx.value()));
#if ENABLE_TENSORRT
        Ort::ThrowOnError(OrtSessionOptionsAppendExecutionProvider_Tensorrt(sessionOptions, m_gpuIdx.value()));
#endif
    }
#endif

    sessionOptions.SetGraphOptimizationLevel(GraphOptimizationLevel::ORT_ENABLE_ALL);
    m_session = Ort::Session(m_env, m_modelPath.c_str(), sessionOptions);
    m_numInputs = m_session.GetInputCount();
    DEBUG_LOG("Model number of inputs: %d\n", m_numInputs);

    m_inputNodeNames.reserve(m_numInputs);
    m_inputTensorSizes.reserve(m_numInputs);

    m_numOutputs = m_session.GetOutputCount();
    DEBUG_LOG("Model number of outputs: %d\n", m_numOutputs);

    m_outputNodeNames.reserve(m_numOutputs);
    m_outputTensorSizes.reserve(m_numOutputs);
}

void OrtSessionHandler::OrtSessionHandlerIml::initModelInfo()
{
    for (int i = 0; i < m_numInputs; i++) {
        if (!m_inputShapesProvided) {
            Ort::TypeInfo typeInfo = m_session.GetInputTypeInfo(i);
            auto tensorInfo = typeInfo.GetTensorTypeAndShapeInfo();

            m_inputShapes.emplace_back(tensorInfo.GetShape());
        }

        const auto& curInputShape = m_inputShapes[i];

        m_inputTensorSizes.emplace_back(
            std::accumulate(std::begin(curInputShape), std::end(curInputShape), 1, std::multiplies()));

#if ORT_API_VERSION > 12
        m_inputNodeNames.emplace_back(strdup(m_session.GetInputNameAllocated(i, m_ortAllocator).get()));
#else
        char* inputName = m_session.GetInputName(i, m_ortAllocator);
        m_inputNodeNames.emplace_back(strdup(inputName));
        m_ortAllocator.Free(inputName);
#endif
    }

    {
#if ENABLE_DEBUG
        std::stringstream ssInputs;
        ssInputs 

Web Proxy Viewer  |  New URL  |  Original Page