mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
Add C++ Session ctor taking model bytes and OrtPrepackedWeightsContainer (#12333)
This commit is contained in:
parent
df8dd41a8e
commit
d5a1c01b38
3 changed files with 246 additions and 226 deletions
|
|
@ -29,14 +29,14 @@
|
|||
#endif
|
||||
|
||||
/** \brief All C++ Onnxruntime APIs are defined inside this namespace
|
||||
*
|
||||
*/
|
||||
*
|
||||
*/
|
||||
namespace Ort {
|
||||
|
||||
/** \brief All C++ methods that can fail will throw an exception of this type
|
||||
*
|
||||
* If <tt>ORT_NO_EXCEPTIONS</tt> is defined, then any error will result in a call to abort()
|
||||
*/
|
||||
*
|
||||
* If <tt>ORT_NO_EXCEPTIONS</tt> is defined, then any error will result in a call to abort()
|
||||
*/
|
||||
struct Exception : std::exception {
|
||||
Exception(std::string&& string, OrtErrorCode code) : message_{std::move(string)}, code_{code} {}
|
||||
|
||||
|
|
@ -121,44 +121,44 @@ ORT_DEFINE_RELEASE(ArenaCfg);
|
|||
#undef ORT_DEFINE_RELEASE
|
||||
|
||||
/** \brief IEEE 754 half-precision floating point data type
|
||||
* \details It is necessary for type dispatching to make use of C++ API
|
||||
* The type is implicitly convertible to/from uint16_t.
|
||||
* The size of the structure should align with uint16_t and one can freely cast
|
||||
* uint16_t buffers to/from Ort::Float16_t to feed and retrieve data.
|
||||
*
|
||||
* Generally, you can feed any of your types as float16/blfoat16 data to create a tensor
|
||||
* on top of it, providing it can form a continuous buffer with 16-bit elements with no padding.
|
||||
* And you can also feed a array of uint16_t elements directly. For example,
|
||||
*
|
||||
* \code{.unparsed}
|
||||
* uint16_t values[] = { 15360, 16384, 16896, 17408, 17664};
|
||||
* constexpr size_t values_length = sizeof(values) / sizeof(values[0]);
|
||||
* std::vector<int64_t> dims = {values_length}; // one dimensional example
|
||||
* Ort::MemoryInfo info("Cpu", OrtDeviceAllocator, 0, OrtMemTypeDefault);
|
||||
* // Note we are passing bytes count in this api, not number of elements -> sizeof(values)
|
||||
* auto float16_tensor = Ort::Value::CreateTensor(info, values, sizeof(values),
|
||||
* dims.data(), dims.size(), ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT16);
|
||||
* \endcode
|
||||
*
|
||||
* Here is another example, a little bit more elaborate. Let's assume that you use your own float16 type and you want to use
|
||||
* a templated version of the API above so the type is automatically set based on your type. You will need to supply an extra
|
||||
* template specialization.
|
||||
*
|
||||
* \code{.unparsed}
|
||||
* namespace yours { struct half {}; } // assume this is your type, define this:
|
||||
* namespace Ort {
|
||||
* template<>
|
||||
* struct TypeToTensorType<yours::half> { static constexpr ONNXTensorElementDataType type = ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT16; };
|
||||
* } //namespace Ort
|
||||
*
|
||||
* std::vector<yours::half> values;
|
||||
* std::vector<int64_t> dims = {values.size()}; // one dimensional example
|
||||
* Ort::MemoryInfo info("Cpu", OrtDeviceAllocator, 0, OrtMemTypeDefault);
|
||||
* // Here we are passing element count -> values.size()
|
||||
* auto float16_tensor = Ort::Value::CreateTensor<yours::half>(info, values.data(), values.size(), dims.data(), dims.size());
|
||||
*
|
||||
* \endcode
|
||||
*/
|
||||
* \details It is necessary for type dispatching to make use of C++ API
|
||||
* The type is implicitly convertible to/from uint16_t.
|
||||
* The size of the structure should align with uint16_t and one can freely cast
|
||||
* uint16_t buffers to/from Ort::Float16_t to feed and retrieve data.
|
||||
*
|
||||
* Generally, you can feed any of your types as float16/blfoat16 data to create a tensor
|
||||
* on top of it, providing it can form a continuous buffer with 16-bit elements with no padding.
|
||||
* And you can also feed a array of uint16_t elements directly. For example,
|
||||
*
|
||||
* \code{.unparsed}
|
||||
* uint16_t values[] = { 15360, 16384, 16896, 17408, 17664};
|
||||
* constexpr size_t values_length = sizeof(values) / sizeof(values[0]);
|
||||
* std::vector<int64_t> dims = {values_length}; // one dimensional example
|
||||
* Ort::MemoryInfo info("Cpu", OrtDeviceAllocator, 0, OrtMemTypeDefault);
|
||||
* // Note we are passing bytes count in this api, not number of elements -> sizeof(values)
|
||||
* auto float16_tensor = Ort::Value::CreateTensor(info, values, sizeof(values),
|
||||
* dims.data(), dims.size(), ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT16);
|
||||
* \endcode
|
||||
*
|
||||
* Here is another example, a little bit more elaborate. Let's assume that you use your own float16 type and you want to use
|
||||
* a templated version of the API above so the type is automatically set based on your type. You will need to supply an extra
|
||||
* template specialization.
|
||||
*
|
||||
* \code{.unparsed}
|
||||
* namespace yours { struct half {}; } // assume this is your type, define this:
|
||||
* namespace Ort {
|
||||
* template<>
|
||||
* struct TypeToTensorType<yours::half> { static constexpr ONNXTensorElementDataType type = ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT16; };
|
||||
* } //namespace Ort
|
||||
*
|
||||
* std::vector<yours::half> values;
|
||||
* std::vector<int64_t> dims = {values.size()}; // one dimensional example
|
||||
* Ort::MemoryInfo info("Cpu", OrtDeviceAllocator, 0, OrtMemTypeDefault);
|
||||
* // Here we are passing element count -> values.size()
|
||||
* auto float16_tensor = Ort::Value::CreateTensor<yours::half>(info, values.data(), values.size(), dims.data(), dims.size());
|
||||
*
|
||||
* \endcode
|
||||
*/
|
||||
struct Float16_t {
|
||||
uint16_t value;
|
||||
constexpr Float16_t() noexcept : value(0) {}
|
||||
|
|
@ -171,13 +171,13 @@ struct Float16_t {
|
|||
static_assert(sizeof(Float16_t) == sizeof(uint16_t), "Sizes must match");
|
||||
|
||||
/** \brief bfloat16 (Brain Floating Point) data type
|
||||
* \details It is necessary for type dispatching to make use of C++ API
|
||||
* The type is implicitly convertible to/from uint16_t.
|
||||
* The size of the structure should align with uint16_t and one can freely cast
|
||||
* uint16_t buffers to/from Ort::BFloat16_t to feed and retrieve data.
|
||||
*
|
||||
* See also code examples for Float16_t above.
|
||||
*/
|
||||
* \details It is necessary for type dispatching to make use of C++ API
|
||||
* The type is implicitly convertible to/from uint16_t.
|
||||
* The size of the structure should align with uint16_t and one can freely cast
|
||||
* uint16_t buffers to/from Ort::BFloat16_t to feed and retrieve data.
|
||||
*
|
||||
* See also code examples for Float16_t above.
|
||||
*/
|
||||
struct BFloat16_t {
|
||||
uint16_t value;
|
||||
constexpr BFloat16_t() noexcept : value(0) {}
|
||||
|
|
@ -190,12 +190,12 @@ struct BFloat16_t {
|
|||
static_assert(sizeof(BFloat16_t) == sizeof(uint16_t), "Sizes must match");
|
||||
|
||||
/** \brief Used internally by the C++ API. C++ wrapper types inherit from this
|
||||
*
|
||||
* This is a zero cost abstraction to wrap the C API objects and delete them on destruction.
|
||||
* There is a secondary class 'Unowned<T>' that is used to prevent deletion on destruction (Used for return types that are
|
||||
* not owned by the caller)
|
||||
*
|
||||
*/
|
||||
*
|
||||
* This is a zero cost abstraction to wrap the C API objects and delete them on destruction.
|
||||
* There is a secondary class 'Unowned<T>' that is used to prevent deletion on destruction (Used for return types that are
|
||||
* not owned by the caller)
|
||||
*
|
||||
*/
|
||||
template <typename T>
|
||||
struct Base {
|
||||
using contained_type = T;
|
||||
|
|
@ -233,9 +233,9 @@ struct Base {
|
|||
};
|
||||
|
||||
/** \brief Wraps an object that inherits from Ort::Base and stops it from deleting the contained pointer on destruction
|
||||
*
|
||||
* This has the effect of making it not own the memory held by Ort::Base.
|
||||
*/
|
||||
*
|
||||
* This has the effect of making it not own the memory held by Ort::Base.
|
||||
*/
|
||||
template <typename T>
|
||||
struct Unowned : T {
|
||||
Unowned(typename T::contained_type* p) : T{p} {}
|
||||
|
|
@ -256,7 +256,9 @@ struct AllocatedFree {
|
|||
OrtAllocator* allocator_;
|
||||
explicit AllocatedFree(OrtAllocator* allocator)
|
||||
: allocator_(allocator) {}
|
||||
void operator()(void* ptr) const { if(ptr) allocator_->Free(allocator_, ptr); }
|
||||
void operator()(void* ptr) const {
|
||||
if (ptr) allocator_->Free(allocator_, ptr);
|
||||
}
|
||||
};
|
||||
} // namespace detail
|
||||
|
||||
|
|
@ -267,10 +269,10 @@ struct AllocatedFree {
|
|||
using AllocatedStringPtr = std::unique_ptr<char, detail::AllocatedFree>;
|
||||
|
||||
/** \brief The Env (Environment)
|
||||
*
|
||||
* The Env holds the logging state used by all other objects.
|
||||
* <b>Note:</b> One Env must be created before using any other Onnxruntime functionality
|
||||
*/
|
||||
*
|
||||
* The Env holds the logging state used by all other objects.
|
||||
* <b>Note:</b> One Env must be created before using any other Onnxruntime functionality
|
||||
*/
|
||||
struct Env : Base<OrtEnv> {
|
||||
explicit Env(std::nullptr_t) {} ///< Create an empty Env object, must be assigned a valid one to be used
|
||||
|
||||
|
|
@ -297,8 +299,8 @@ struct Env : Base<OrtEnv> {
|
|||
};
|
||||
|
||||
/** \brief Custom Op Domain
|
||||
*
|
||||
*/
|
||||
*
|
||||
*/
|
||||
struct CustomOpDomain : Base<OrtCustomOpDomain> {
|
||||
explicit CustomOpDomain(std::nullptr_t) {} ///< Create an empty CustomOpDomain object, must be assigned a valid one to be used
|
||||
|
||||
|
|
@ -324,23 +326,23 @@ struct RunOptions : Base<OrtRunOptions> {
|
|||
RunOptions& AddConfigEntry(const char* config_key, const char* config_value); ///< Wraps OrtApi::AddRunConfigEntry
|
||||
|
||||
/** \brief Terminates all currently executing Session::Run calls that were made using this RunOptions instance
|
||||
*
|
||||
* If a currently executing session needs to be force terminated, this can be called from another thread to force it to fail with an error
|
||||
* Wraps OrtApi::RunOptionsSetTerminate
|
||||
*/
|
||||
*
|
||||
* If a currently executing session needs to be force terminated, this can be called from another thread to force it to fail with an error
|
||||
* Wraps OrtApi::RunOptionsSetTerminate
|
||||
*/
|
||||
RunOptions& SetTerminate();
|
||||
|
||||
/** \brief Clears the terminate flag so this RunOptions instance can be used in a new Session::Run call without it instantly terminating
|
||||
*
|
||||
* Wraps OrtApi::RunOptionsUnsetTerminate
|
||||
*/
|
||||
*
|
||||
* Wraps OrtApi::RunOptionsUnsetTerminate
|
||||
*/
|
||||
RunOptions& UnsetTerminate();
|
||||
};
|
||||
|
||||
/** \brief Options object used when creating a new Session object
|
||||
*
|
||||
* Wraps ::OrtSessionOptions object and methods
|
||||
*/
|
||||
*
|
||||
* Wraps ::OrtSessionOptions object and methods
|
||||
*/
|
||||
struct SessionOptions : Base<OrtSessionOptions> {
|
||||
explicit SessionOptions(std::nullptr_t) {} ///< Create an empty SessionOptions object, must be assigned a valid one to be used
|
||||
SessionOptions(); ///< Wraps OrtApi::CreateSessionOptions
|
||||
|
|
@ -395,149 +397,151 @@ struct SessionOptions : Base<OrtSessionOptions> {
|
|||
};
|
||||
|
||||
/** \brief Wrapper around ::OrtModelMetadata
|
||||
*
|
||||
*/
|
||||
*
|
||||
*/
|
||||
struct ModelMetadata : Base<OrtModelMetadata> {
|
||||
explicit ModelMetadata(std::nullptr_t) {} ///< Create an empty ModelMetadata object, must be assigned a valid one to be used
|
||||
explicit ModelMetadata(OrtModelMetadata* p) : Base<OrtModelMetadata>{p} {} ///< Used for interop with the C API
|
||||
|
||||
/** \deprecated use GetProducerNameAllocated()
|
||||
* [[deprecated]]
|
||||
* This interface produces a pointer that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
char* GetProducerName(OrtAllocator* allocator) const; ///< Wraps OrtApi::ModelMetadataGetProducerName
|
||||
* [[deprecated]]
|
||||
* This interface produces a pointer that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
char* GetProducerName(OrtAllocator* allocator) const; ///< Wraps OrtApi::ModelMetadataGetProducerName
|
||||
|
||||
/** \brief Returns a copy of the producer name.
|
||||
*
|
||||
* \param allocator to allocate memory for the copy of the name returned
|
||||
* \return a instance of smart pointer that would deallocate the buffer when out of scope.
|
||||
* The OrtAllocator instances must be valid at the point of memory release.
|
||||
*/
|
||||
*
|
||||
* \param allocator to allocate memory for the copy of the name returned
|
||||
* \return a instance of smart pointer that would deallocate the buffer when out of scope.
|
||||
* The OrtAllocator instances must be valid at the point of memory release.
|
||||
*/
|
||||
AllocatedStringPtr GetProducerNameAllocated(OrtAllocator* allocator) const; ///< Wraps OrtApi::ModelMetadataGetProducerName
|
||||
|
||||
/** \deprecated use GetGraphNameAllocated()
|
||||
* [[deprecated]]
|
||||
* This interface produces a pointer that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
char* GetGraphName(OrtAllocator* allocator) const; ///< Wraps OrtApi::ModelMetadataGetGraphName
|
||||
* [[deprecated]]
|
||||
* This interface produces a pointer that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
char* GetGraphName(OrtAllocator* allocator) const; ///< Wraps OrtApi::ModelMetadataGetGraphName
|
||||
|
||||
/** \brief Returns a copy of the graph name.
|
||||
*
|
||||
* \param allocator to allocate memory for the copy of the name returned
|
||||
* \return a instance of smart pointer that would deallocate the buffer when out of scope.
|
||||
* The OrtAllocator instances must be valid at the point of memory release.
|
||||
*/
|
||||
*
|
||||
* \param allocator to allocate memory for the copy of the name returned
|
||||
* \return a instance of smart pointer that would deallocate the buffer when out of scope.
|
||||
* The OrtAllocator instances must be valid at the point of memory release.
|
||||
*/
|
||||
AllocatedStringPtr GetGraphNameAllocated(OrtAllocator* allocator) const; ///< Wraps OrtApi::ModelMetadataGetGraphName
|
||||
|
||||
/** \deprecated use GetDomainAllocated()
|
||||
* [[deprecated]]
|
||||
* This interface produces a pointer that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
char* GetDomain(OrtAllocator* allocator) const; ///< Wraps OrtApi::ModelMetadataGetDomain
|
||||
* [[deprecated]]
|
||||
* This interface produces a pointer that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
char* GetDomain(OrtAllocator* allocator) const; ///< Wraps OrtApi::ModelMetadataGetDomain
|
||||
|
||||
/** \brief Returns a copy of the domain name.
|
||||
*
|
||||
* \param allocator to allocate memory for the copy of the name returned
|
||||
* \return a instance of smart pointer that would deallocate the buffer when out of scope.
|
||||
* The OrtAllocator instances must be valid at the point of memory release.
|
||||
*/
|
||||
*
|
||||
* \param allocator to allocate memory for the copy of the name returned
|
||||
* \return a instance of smart pointer that would deallocate the buffer when out of scope.
|
||||
* The OrtAllocator instances must be valid at the point of memory release.
|
||||
*/
|
||||
AllocatedStringPtr GetDomainAllocated(OrtAllocator* allocator) const; ///< Wraps OrtApi::ModelMetadataGetDomain
|
||||
|
||||
/** \deprecated use GetDescriptionAllocated()
|
||||
* [[deprecated]]
|
||||
* This interface produces a pointer that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
char* GetDescription(OrtAllocator* allocator) const; ///< Wraps OrtApi::ModelMetadataGetDescription
|
||||
* [[deprecated]]
|
||||
* This interface produces a pointer that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
char* GetDescription(OrtAllocator* allocator) const; ///< Wraps OrtApi::ModelMetadataGetDescription
|
||||
|
||||
/** \brief Returns a copy of the description.
|
||||
*
|
||||
* \param allocator to allocate memory for the copy of the string returned
|
||||
* \return a instance of smart pointer that would deallocate the buffer when out of scope.
|
||||
* The OrtAllocator instances must be valid at the point of memory release.
|
||||
*/
|
||||
*
|
||||
* \param allocator to allocate memory for the copy of the string returned
|
||||
* \return a instance of smart pointer that would deallocate the buffer when out of scope.
|
||||
* The OrtAllocator instances must be valid at the point of memory release.
|
||||
*/
|
||||
AllocatedStringPtr GetDescriptionAllocated(OrtAllocator* allocator) const; ///< Wraps OrtApi::ModelMetadataGetDescription
|
||||
|
||||
/** \deprecated use GetGraphDescriptionAllocated()
|
||||
* [[deprecated]]
|
||||
* This interface produces a pointer that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
char* GetGraphDescription(OrtAllocator* allocator) const; ///< Wraps OrtApi::ModelMetadataGetGraphDescription
|
||||
* [[deprecated]]
|
||||
* This interface produces a pointer that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
char* GetGraphDescription(OrtAllocator* allocator) const; ///< Wraps OrtApi::ModelMetadataGetGraphDescription
|
||||
|
||||
/** \brief Returns a copy of the graph description.
|
||||
*
|
||||
* \param allocator to allocate memory for the copy of the string returned
|
||||
* \return a instance of smart pointer that would deallocate the buffer when out of scope.
|
||||
* The OrtAllocator instances must be valid at the point of memory release.
|
||||
*/
|
||||
AllocatedStringPtr GetGraphDescriptionAllocated(OrtAllocator* allocator) const; ///< Wraps OrtApi::ModelMetadataGetGraphDescription
|
||||
*
|
||||
* \param allocator to allocate memory for the copy of the string returned
|
||||
* \return a instance of smart pointer that would deallocate the buffer when out of scope.
|
||||
* The OrtAllocator instances must be valid at the point of memory release.
|
||||
*/
|
||||
AllocatedStringPtr GetGraphDescriptionAllocated(OrtAllocator* allocator) const; ///< Wraps OrtApi::ModelMetadataGetGraphDescription
|
||||
|
||||
/** \deprecated use GetCustomMetadataMapKeysAllocated()
|
||||
* [[deprecated]]
|
||||
* This interface produces multiple pointers that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
* [[deprecated]]
|
||||
* This interface produces multiple pointers that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
char** GetCustomMetadataMapKeys(OrtAllocator* allocator, _Out_ int64_t& num_keys) const; ///< Wraps OrtApi::ModelMetadataGetCustomMetadataMapKeys
|
||||
|
||||
std::vector<AllocatedStringPtr> GetCustomMetadataMapKeysAllocated(OrtAllocator* allocator) const; ///< Wraps OrtApi::ModelMetadataGetCustomMetadataMapKeys
|
||||
|
||||
/** \deprecated use LookupCustomMetadataMapAllocated()
|
||||
* [[deprecated]]
|
||||
* This interface produces a pointer that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
char* LookupCustomMetadataMap(const char* key, OrtAllocator* allocator) const; ///< Wraps OrtApi::ModelMetadataLookupCustomMetadataMap
|
||||
* [[deprecated]]
|
||||
* This interface produces a pointer that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
char* LookupCustomMetadataMap(const char* key, OrtAllocator* allocator) const; ///< Wraps OrtApi::ModelMetadataLookupCustomMetadataMap
|
||||
|
||||
/** \brief Looks up a value by a key in the Custom Metadata map
|
||||
*
|
||||
* \param zero terminated string key to lookup
|
||||
* \param allocator to allocate memory for the copy of the string returned
|
||||
* \return a instance of smart pointer that would deallocate the buffer when out of scope.
|
||||
* maybe nullptr if key is not found.
|
||||
*
|
||||
* The OrtAllocator instances must be valid at the point of memory release.
|
||||
*/
|
||||
*
|
||||
* \param zero terminated string key to lookup
|
||||
* \param allocator to allocate memory for the copy of the string returned
|
||||
* \return a instance of smart pointer that would deallocate the buffer when out of scope.
|
||||
* maybe nullptr if key is not found.
|
||||
*
|
||||
* The OrtAllocator instances must be valid at the point of memory release.
|
||||
*/
|
||||
AllocatedStringPtr LookupCustomMetadataMapAllocated(const char* key, OrtAllocator* allocator) const; ///< Wraps OrtApi::ModelMetadataLookupCustomMetadataMap
|
||||
|
||||
int64_t GetVersion() const; ///< Wraps OrtApi::ModelMetadataGetVersion
|
||||
int64_t GetVersion() const; ///< Wraps OrtApi::ModelMetadataGetVersion
|
||||
};
|
||||
|
||||
/** \brief Wrapper around ::OrtSession
|
||||
*
|
||||
*/
|
||||
*
|
||||
*/
|
||||
struct Session : Base<OrtSession> {
|
||||
explicit Session(std::nullptr_t) {} ///< Create an empty Session object, must be assigned a valid one to be used
|
||||
Session(Env& env, const ORTCHAR_T* model_path, const SessionOptions& options); ///< Wraps OrtApi::CreateSession
|
||||
Session(Env& env, const ORTCHAR_T* model_path, const SessionOptions& options, OrtPrepackedWeightsContainer* prepacked_weights_container); ///< Wraps OrtApi::CreateSessionWithPrepackedWeightsContainer
|
||||
Session(Env& env, const void* model_data, size_t model_data_length, const SessionOptions& options); ///< Wraps OrtApi::CreateSessionFromArray
|
||||
Session(Env& env, const void* model_data, size_t model_data_length, const SessionOptions& options,
|
||||
OrtPrepackedWeightsContainer* prepacked_weights_container); ///< Wraps OrtApi::CreateSessionFromArrayWithPrepackedWeightsContainer
|
||||
|
||||
/** \brief Run the model returning results in an Ort allocated vector.
|
||||
*
|
||||
* Wraps OrtApi::Run
|
||||
*
|
||||
* The caller provides a list of inputs and a list of the desired outputs to return.
|
||||
*
|
||||
* See the output logs for more information on warnings/errors that occur while processing the model.
|
||||
* Common errors are.. (TODO)
|
||||
*
|
||||
* \param[in] run_options
|
||||
* \param[in] input_names Array of null terminated strings of length input_count that is the list of input names
|
||||
* \param[in] input_values Array of Value objects of length input_count that is the list of input values
|
||||
* \param[in] input_count Number of inputs (the size of the input_names & input_values arrays)
|
||||
* \param[in] output_names Array of C style strings of length output_count that is the list of output names
|
||||
* \param[in] output_count Number of outputs (the size of the output_names array)
|
||||
* \return A std::vector of Value objects that directly maps to the output_count (eg. output_name[0] is the first entry of the returned vector)
|
||||
*/
|
||||
*
|
||||
* Wraps OrtApi::Run
|
||||
*
|
||||
* The caller provides a list of inputs and a list of the desired outputs to return.
|
||||
*
|
||||
* See the output logs for more information on warnings/errors that occur while processing the model.
|
||||
* Common errors are.. (TODO)
|
||||
*
|
||||
* \param[in] run_options
|
||||
* \param[in] input_names Array of null terminated strings of length input_count that is the list of input names
|
||||
* \param[in] input_values Array of Value objects of length input_count that is the list of input values
|
||||
* \param[in] input_count Number of inputs (the size of the input_names & input_values arrays)
|
||||
* \param[in] output_names Array of C style strings of length output_count that is the list of output names
|
||||
* \param[in] output_count Number of outputs (the size of the output_names array)
|
||||
* \return A std::vector of Value objects that directly maps to the output_count (eg. output_name[0] is the first entry of the returned vector)
|
||||
*/
|
||||
std::vector<Value> Run(const RunOptions& run_options, const char* const* input_names, const Value* input_values, size_t input_count,
|
||||
const char* const* output_names, size_t output_count);
|
||||
|
||||
/** \brief Run the model returning results in user provided outputs
|
||||
* Same as Run(const RunOptions&, const char* const*, const Value*, size_t,const char* const*, size_t)
|
||||
*/
|
||||
* Same as Run(const RunOptions&, const char* const*, const Value*, size_t,const char* const*, size_t)
|
||||
*/
|
||||
void Run(const RunOptions& run_options, const char* const* input_names, const Value* input_values, size_t input_count,
|
||||
const char* const* output_names, Value* output_values, size_t output_count);
|
||||
|
||||
|
|
@ -548,66 +552,66 @@ struct Session : Base<OrtSession> {
|
|||
size_t GetOverridableInitializerCount() const; ///< Returns the number of inputs that have defaults that can be overridden
|
||||
|
||||
/** \deprecated use GetInputNameAllocated()
|
||||
* [[deprecated]]
|
||||
* This interface produces a pointer that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
* [[deprecated]]
|
||||
* This interface produces a pointer that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
char* GetInputName(size_t index, OrtAllocator* allocator) const; ///< Wraps OrtApi::SessionGetInputName
|
||||
|
||||
/** \brief Returns a copy of input name at the specified index.
|
||||
*
|
||||
* \param index must less than the value returned by GetInputCount()
|
||||
* \param allocator to allocate memory for the copy of the name returned
|
||||
* \return a instance of smart pointer that would deallocate the buffer when out of scope.
|
||||
* The OrtAllocator instances must be valid at the point of memory release.
|
||||
*/
|
||||
*
|
||||
* \param index must less than the value returned by GetInputCount()
|
||||
* \param allocator to allocate memory for the copy of the name returned
|
||||
* \return a instance of smart pointer that would deallocate the buffer when out of scope.
|
||||
* The OrtAllocator instances must be valid at the point of memory release.
|
||||
*/
|
||||
AllocatedStringPtr GetInputNameAllocated(size_t index, OrtAllocator* allocator) const;
|
||||
|
||||
/** \deprecated use GetOutputNameAllocated()
|
||||
* [[deprecated]]
|
||||
* This interface produces a pointer that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
* [[deprecated]]
|
||||
* This interface produces a pointer that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
char* GetOutputName(size_t index, OrtAllocator* allocator) const; ///< Wraps OrtApi::SessionGetOutputName
|
||||
|
||||
/** \brief Returns a copy of output name at then specified index.
|
||||
*
|
||||
* \param index must less than the value returned by GetOutputCount()
|
||||
* \param allocator to allocate memory for the copy of the name returned
|
||||
* \return a instance of smart pointer that would deallocate the buffer when out of scope.
|
||||
* The OrtAllocator instances must be valid at the point of memory release.
|
||||
*/
|
||||
*
|
||||
* \param index must less than the value returned by GetOutputCount()
|
||||
* \param allocator to allocate memory for the copy of the name returned
|
||||
* \return a instance of smart pointer that would deallocate the buffer when out of scope.
|
||||
* The OrtAllocator instances must be valid at the point of memory release.
|
||||
*/
|
||||
AllocatedStringPtr GetOutputNameAllocated(size_t index, OrtAllocator* allocator) const;
|
||||
|
||||
/** \deprecated use GetOverridableInitializerNameAllocated()
|
||||
* [[deprecated]]
|
||||
* This interface produces a pointer that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
* [[deprecated]]
|
||||
* This interface produces a pointer that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
char* GetOverridableInitializerName(size_t index, OrtAllocator* allocator) const; ///< Wraps OrtApi::SessionGetOverridableInitializerName
|
||||
|
||||
/** \brief Returns a copy of the overridable initializer name at then specified index.
|
||||
*
|
||||
* \param index must less than the value returned by GetOverridableInitializerCount()
|
||||
* \param allocator to allocate memory for the copy of the name returned
|
||||
* \return a instance of smart pointer that would deallocate the buffer when out of scope.
|
||||
* The OrtAllocator instances must be valid at the point of memory release.
|
||||
*/
|
||||
*
|
||||
* \param index must less than the value returned by GetOverridableInitializerCount()
|
||||
* \param allocator to allocate memory for the copy of the name returned
|
||||
* \return a instance of smart pointer that would deallocate the buffer when out of scope.
|
||||
* The OrtAllocator instances must be valid at the point of memory release.
|
||||
*/
|
||||
AllocatedStringPtr GetOverridableInitializerNameAllocated(size_t index, OrtAllocator* allocator) const; ///< Wraps OrtApi::SessionGetOverridableInitializerName
|
||||
|
||||
/** \deprecated use EndProfilingAllocated()
|
||||
* [[deprecated]]
|
||||
* This interface produces a pointer that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
* [[deprecated]]
|
||||
* This interface produces a pointer that must be released
|
||||
* by the specified allocator and is often leaked. Not exception safe.
|
||||
*/
|
||||
char* EndProfiling(OrtAllocator* allocator) const; ///< Wraps OrtApi::SessionEndProfiling
|
||||
|
||||
/** \brief Returns a copy of the profiling file name.
|
||||
*
|
||||
* \param allocator to allocate memory for the copy of the string returned
|
||||
* \return a instance of smart pointer that would deallocate the buffer when out of scope.
|
||||
* The OrtAllocator instances must be valid at the point of memory release.
|
||||
*/
|
||||
*
|
||||
* \param allocator to allocate memory for the copy of the string returned
|
||||
* \return a instance of smart pointer that would deallocate the buffer when out of scope.
|
||||
* The OrtAllocator instances must be valid at the point of memory release.
|
||||
*/
|
||||
AllocatedStringPtr EndProfilingAllocated(OrtAllocator* allocator) const; ///< Wraps OrtApi::SessionEndProfiling
|
||||
uint64_t GetProfilingStartTimeNs() const; ///< Wraps OrtApi::SessionGetProfilingStartTimeNs
|
||||
ModelMetadata GetModelMetadata() const; ///< Wraps OrtApi::SessionGetModelMetadata
|
||||
|
|
@ -618,8 +622,8 @@ struct Session : Base<OrtSession> {
|
|||
};
|
||||
|
||||
/** \brief Wrapper around ::OrtTensorTypeAndShapeInfo
|
||||
*
|
||||
*/
|
||||
*
|
||||
*/
|
||||
struct TensorTypeAndShapeInfo : Base<OrtTensorTypeAndShapeInfo> {
|
||||
explicit TensorTypeAndShapeInfo(std::nullptr_t) {} ///< Create an empty TensorTypeAndShapeInfo object, must be assigned a valid one to be used
|
||||
explicit TensorTypeAndShapeInfo(OrtTensorTypeAndShapeInfo* p) : Base<OrtTensorTypeAndShapeInfo>{p} {} ///< Used for interop with the C API
|
||||
|
|
@ -635,8 +639,8 @@ struct TensorTypeAndShapeInfo : Base<OrtTensorTypeAndShapeInfo> {
|
|||
};
|
||||
|
||||
/** \brief Wrapper around ::OrtSequenceTypeInfo
|
||||
*
|
||||
*/
|
||||
*
|
||||
*/
|
||||
struct SequenceTypeInfo : Base<OrtSequenceTypeInfo> {
|
||||
explicit SequenceTypeInfo(std::nullptr_t) {} ///< Create an empty SequenceTypeInfo object, must be assigned a valid one to be used
|
||||
explicit SequenceTypeInfo(OrtSequenceTypeInfo* p) : Base<OrtSequenceTypeInfo>{p} {} ///< Used for interop with the C API
|
||||
|
|
@ -645,8 +649,8 @@ struct SequenceTypeInfo : Base<OrtSequenceTypeInfo> {
|
|||
};
|
||||
|
||||
/** \brief Wrapper around ::OrtMapTypeInfo
|
||||
*
|
||||
*/
|
||||
*
|
||||
*/
|
||||
struct MapTypeInfo : Base<OrtMapTypeInfo> {
|
||||
explicit MapTypeInfo(std::nullptr_t) {} ///< Create an empty MapTypeInfo object, must be assigned a valid one to be used
|
||||
explicit MapTypeInfo(OrtMapTypeInfo* p) : Base<OrtMapTypeInfo>{p} {} ///< Used for interop with the C API
|
||||
|
|
@ -1096,19 +1100,19 @@ struct IoBinding : public Base<OrtIoBinding> {
|
|||
};
|
||||
|
||||
/*! \struct Ort::ArenaCfg
|
||||
* \brief it is a structure that represents the configuration of an arena based allocator
|
||||
* \details Please see docs/C_API.md for details
|
||||
*/
|
||||
* \brief it is a structure that represents the configuration of an arena based allocator
|
||||
* \details Please see docs/C_API.md for details
|
||||
*/
|
||||
struct ArenaCfg : Base<OrtArenaCfg> {
|
||||
explicit ArenaCfg(std::nullptr_t) {} ///< Create an empty ArenaCfg object, must be assigned a valid one to be used
|
||||
/**
|
||||
* Wraps OrtApi::CreateArenaCfg
|
||||
* \param max_mem - use 0 to allow ORT to choose the default
|
||||
* \param arena_extend_strategy - use -1 to allow ORT to choose the default, 0 = kNextPowerOfTwo, 1 = kSameAsRequested
|
||||
* \param initial_chunk_size_bytes - use -1 to allow ORT to choose the default
|
||||
* \param max_dead_bytes_per_chunk - use -1 to allow ORT to choose the default
|
||||
* See docs/C_API.md for details on what the following parameters mean and how to choose these values
|
||||
*/
|
||||
* Wraps OrtApi::CreateArenaCfg
|
||||
* \param max_mem - use 0 to allow ORT to choose the default
|
||||
* \param arena_extend_strategy - use -1 to allow ORT to choose the default, 0 = kNextPowerOfTwo, 1 = kSameAsRequested
|
||||
* \param initial_chunk_size_bytes - use -1 to allow ORT to choose the default
|
||||
* \param max_dead_bytes_per_chunk - use -1 to allow ORT to choose the default
|
||||
* See docs/C_API.md for details on what the following parameters mean and how to choose these values
|
||||
*/
|
||||
ArenaCfg(size_t max_mem, int arena_extend_strategy, int initial_chunk_size_bytes, int max_dead_bytes_per_chunk);
|
||||
};
|
||||
|
||||
|
|
|
|||
|
|
@ -605,6 +605,12 @@ inline Session::Session(Env& env, const void* model_data, size_t model_data_leng
|
|||
ThrowOnError(GetApi().CreateSessionFromArray(env, model_data, model_data_length, options, &p_));
|
||||
}
|
||||
|
||||
inline Session::Session(Env& env, const void* model_data, size_t model_data_length,
|
||||
const SessionOptions& options, OrtPrepackedWeightsContainer* prepacked_weights_container) {
|
||||
ThrowOnError(GetApi().CreateSessionFromArrayWithPrepackedWeightsContainer(env, model_data, model_data_length, options,
|
||||
prepacked_weights_container, &p_));
|
||||
}
|
||||
|
||||
inline std::vector<Value> Session::Run(const RunOptions& run_options, const char* const* input_names, const Value* input_values, size_t input_count,
|
||||
const char* const* output_names, size_t output_names_count) {
|
||||
std::vector<Ort::Value> output_values;
|
||||
|
|
@ -1363,7 +1369,7 @@ inline OrtKernelInfo* CustomOpApi::CopyKernelInfo(_In_ const OrtKernelInfo* info
|
|||
|
||||
inline void CustomOpApi::ReleaseKernelInfo(_Frees_ptr_opt_ OrtKernelInfo* info_copy) {
|
||||
api_.ReleaseKernelInfo(info_copy);
|
||||
}
|
||||
}
|
||||
|
||||
inline SessionOptions& SessionOptions::DisablePerSessionThreads() {
|
||||
ThrowOnError(GetApi().DisablePerSessionThreads(p_));
|
||||
|
|
|
|||
|
|
@ -1781,7 +1781,7 @@ TEST(CApiTest, TestSharingOfInitializerAndItsPrepackedVersion) {
|
|||
|
||||
auto default_allocator = std::make_unique<MockedOrtAllocator>();
|
||||
|
||||
// create session 1
|
||||
// create session 1 (using model path)
|
||||
Ort::Session session1(*ort_env, MATMUL_MODEL_URI, session_options, prepacked_weights_container);
|
||||
RunSession<float>(default_allocator.get(),
|
||||
session1,
|
||||
|
|
@ -1791,8 +1791,18 @@ TEST(CApiTest, TestSharingOfInitializerAndItsPrepackedVersion) {
|
|||
expected_values_y,
|
||||
nullptr);
|
||||
|
||||
// create session 2
|
||||
Ort::Session session2(*ort_env, MATMUL_MODEL_URI, session_options, prepacked_weights_container);
|
||||
// create session 2 (using model bytes)
|
||||
std::ifstream model_file_stream(MATMUL_MODEL_URI, std::ios::in | std::ios::binary);
|
||||
ASSERT_TRUE(model_file_stream.good());
|
||||
|
||||
model_file_stream.seekg(0, std::ios::end);
|
||||
size_t size = model_file_stream.tellg();
|
||||
model_file_stream.seekg(0, std::ios::beg);
|
||||
std::vector<char> file_contents(size, 0);
|
||||
model_file_stream.read(&file_contents[0], size);
|
||||
model_file_stream.close();
|
||||
|
||||
Ort::Session session2(*ort_env, file_contents.data(), size, session_options, prepacked_weights_container);
|
||||
RunSession<float>(default_allocator.get(),
|
||||
session2,
|
||||
inputs,
|
||||
|
|
|
|||
Loading…
Reference in a new issue