/**
* @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