From 61996332ad8934e179defb3173d669fd210b2dfc Mon Sep 17 00:00:00 2001 From: chenduan-amd Date: Tue, 24 Sep 2024 16:37:05 -0500 Subject: [PATCH] [VitisAI] support run_options in vitisai EP end (#22029) ### Description add OnRunStart() method for Vitis AI execution provider ### Motivation and Context To dynamically obtain some runtime parameters during execution, use run_options within the Vitis AI execution provider (EP). --- .../core/providers/vitisai/imp/global_api.cc | 13 +++++++++++++ .../vitisai/include/vaip/global_api.h | 4 ++++ .../vitisai/vitisai_execution_provider.cc | 18 ++++++++++++++++++ .../vitisai/vitisai_execution_provider.h | 2 +- 4 files changed, 36 insertions(+), 1 deletion(-) diff --git a/onnxruntime/core/providers/vitisai/imp/global_api.cc b/onnxruntime/core/providers/vitisai/imp/global_api.cc index 6f8e153cd1..de16bff7a4 100644 --- a/onnxruntime/core/providers/vitisai/imp/global_api.cc +++ b/onnxruntime/core/providers/vitisai/imp/global_api.cc @@ -49,6 +49,9 @@ struct OrtVitisAIEpAPI { void (*create_ep_context_nodes)( const std::vector>& eps, vaip_core::DllSafe>* ret_value) = nullptr; + int (*vitisai_ep_on_run_start)( + const std::vector>& eps, const void* state, + vaip_core::DllSafe (*get_config_entry)(const void* state, const char* entry_name)) = nullptr; void Ensure() { if (handle_) return; @@ -73,6 +76,7 @@ struct OrtVitisAIEpAPI { std::ignore = env.GetSymbolFromLibrary(handle_, "vaip_get_version", (void**)&vaip_get_version); ORT_THROW_IF_ERROR(env.GetSymbolFromLibrary(handle_, "create_ep_context_nodes", (void**)&create_ep_context_nodes)); + ORT_THROW_IF_ERROR(env.GetSymbolFromLibrary(handle_, "vitisai_ep_on_run_start", (void**)&vitisai_ep_on_run_start)); } private: @@ -105,6 +109,15 @@ std::optional> create_ep_context_nodes( return std::nullopt; } +int vitisai_ep_on_run_start( + const std::vector>& eps, const void* state, + vaip_core::DllSafe (*get_config_entry)(const void* state, const char* entry_name)) { + if (s_library_vitisaiep.vitisai_ep_on_run_start) { + return s_library_vitisaiep.vitisai_ep_on_run_start(eps, state, get_config_entry); + } + return 100; +} + struct MyCustomOpKernel : OpKernel { MyCustomOpKernel(const OpKernelInfo& info, const OrtCustomOp& op) : OpKernel(info), op_(op) { op_kernel_ = diff --git a/onnxruntime/core/providers/vitisai/include/vaip/global_api.h b/onnxruntime/core/providers/vitisai/include/vaip/global_api.h index ec2b98e5b6..1a90f4c7fd 100644 --- a/onnxruntime/core/providers/vitisai/include/vaip/global_api.h +++ b/onnxruntime/core/providers/vitisai/include/vaip/global_api.h @@ -16,3 +16,7 @@ std::shared_ptr get_kernel_registry_vitisaiep(); const std::vector& get_domains_vitisaiep(); std::optional> create_ep_context_nodes( const std::vector>& eps); + +int vitisai_ep_on_run_start( + const std::vector>& eps, const void* state, + vaip_core::DllSafe (*get_config_entry)(const void* state, const char* entry_name)); diff --git a/onnxruntime/core/providers/vitisai/vitisai_execution_provider.cc b/onnxruntime/core/providers/vitisai/vitisai_execution_provider.cc index 57c3e21b70..09b115b4a5 100644 --- a/onnxruntime/core/providers/vitisai/vitisai_execution_provider.cc +++ b/onnxruntime/core/providers/vitisai/vitisai_execution_provider.cc @@ -97,4 +97,22 @@ common::Status VitisAIExecutionProvider::Compile(const std::vector ep_context_node_ptrs; + auto get_config_entry = [](const void* state, const char* entry_name) -> vaip_core::DllSafe { + const onnxruntime::RunOptions& run_options = *static_cast(state); + auto ret = run_options.GetConfigOptions().GetConfigEntry(std::string(entry_name)); + if (ret) { + return vaip_core::DllSafe(new std::string(ret.value())); + } else { + return {}; + }; + }; + auto error_code = vitisai_ep_on_run_start(**execution_providers_, (const void*)&run_options, get_config_entry); + if (error_code) { + return Status(onnxruntime::common::ONNXRUNTIME, onnxruntime::common::StatusCode::FAIL, std::to_string(error_code)); + } + return Status::OK(); +} + } // namespace onnxruntime diff --git a/onnxruntime/core/providers/vitisai/vitisai_execution_provider.h b/onnxruntime/core/providers/vitisai/vitisai_execution_provider.h index 24692dd45c..05d2a97681 100644 --- a/onnxruntime/core/providers/vitisai/vitisai_execution_provider.h +++ b/onnxruntime/core/providers/vitisai/vitisai_execution_provider.h @@ -31,7 +31,7 @@ class VitisAIExecutionProvider : public IExecutionProvider { const IKernelLookup& /*kernel_lookup*/) const override; int GetDeviceId() const { return 0; } - + common::Status OnRunStart(const onnxruntime::RunOptions& /*run_options*/) override; common::Status Compile(const std::vector& fused_nodes_and_graphs, std::vector& node_compute_funcs) override; std::shared_ptr GetKernelRegistry() const override;