diff --git a/cgmanifests/cgmanifest.json b/cgmanifests/cgmanifest.json index e043e07ea1..379ff6921c 100644 --- a/cgmanifests/cgmanifest.json +++ b/cgmanifests/cgmanifest.json @@ -46,7 +46,7 @@ "component": { "type": "git", "git": { - "commitHash": "d721d320bd2f66d342d24b71600fe1f5e222e952", + "commitHash": "ffd5f70370642c909222f9a4cae8400023dacbdc", "repositoryUrl": "https://github.com/apache/tvm.git" }, "comments": "needed for TVM EP" diff --git a/cmake/external/tvm.cmake b/cmake/external/tvm.cmake index c1ee5fdde5..3f425a0938 100644 --- a/cmake/external/tvm.cmake +++ b/cmake/external/tvm.cmake @@ -4,7 +4,7 @@ if (onnxruntime_USE_TVM) FetchContent_Declare( tvm GIT_REPOSITORY https://github.com/apache/tvm.git - GIT_TAG d721d320bd2f66d342d24b71600fe1f5e222e952 + GIT_TAG ffd5f70370642c909222f9a4cae8400023dacbdc ) FetchContent_GetProperties(tvm) diff --git a/include/onnxruntime/core/providers/tvm/tvm_provider_factory.h b/include/onnxruntime/core/providers/tvm/tvm_provider_factory.h index 6d2478a948..3dbd270cde 100644 --- a/include/onnxruntime/core/providers/tvm/tvm_provider_factory.h +++ b/include/onnxruntime/core/providers/tvm/tvm_provider_factory.h @@ -10,7 +10,7 @@ extern "C" { #endif -ORT_API_STATUS(OrtSessionOptionsAppendExecutionProvider_Tvm, _In_ OrtSessionOptions* options, _In_ const char* settings); +ORT_API_STATUS(OrtSessionOptionsAppendExecutionProvider_Tvm, _In_ OrtSessionOptions* options, _In_ const char* opt_str); #ifdef __cplusplus } diff --git a/onnxruntime/core/providers/tvm/custom_logging.cc b/onnxruntime/core/providers/tvm/custom_logging.cc index 08053e456a..3140683825 100644 --- a/onnxruntime/core/providers/tvm/custom_logging.cc +++ b/onnxruntime/core/providers/tvm/custom_logging.cc @@ -11,31 +11,33 @@ #include #include + // TODO(agladyshev): Make conditional choice of sep for Windows and UNIX std::string GetFileName(const std::string& file_path, char sep = '/') { - return {std::next(file_path.begin(), file_path.find_last_of(sep) + 1), - file_path.end()}; + return {std::next(file_path.begin(), file_path.find_last_of(sep) + 1), + file_path.end()}; } std::string GetTimedLogMessage(const std::string& file, int lineno, const std::string& message) { - std::stringstream sstream; - std::string file_name = GetFileName(file); - std::time_t t = std::time(nullptr); - sstream << "[" << std::put_time(std::localtime(&t), "%H:%M:%S") << "][TVM] " - << file_name << ":" << lineno << ": " + message; - return sstream.str(); + std::stringstream sstream; + std::string file_name = GetFileName(file); + std::time_t t = std::time(nullptr); + sstream << "[" << std::put_time(std::localtime(&t), "%H:%M:%S") << "][TVM] " + << file_name << ":" << lineno << ": " + message; + return sstream.str(); } namespace tvm { namespace runtime { namespace detail { - void LogFatalImpl(const std::string& file, int lineno, const std::string& message) { - throw std::runtime_error(GetTimedLogMessage(file, lineno, message)); - } + void LogFatalImpl(const std::string& file, int lineno, const std::string& message) { + throw std::runtime_error(GetTimedLogMessage(file, lineno, message)); + } + + void LogMessageImpl(const std::string& file, int lineno, const std::string& message) { + std::cerr << GetTimedLogMessage(file, lineno, message) << std::endl; + } - void LogMessageImpl(const std::string& file, int lineno, const std::string& message) { - std::cerr << GetTimedLogMessage(file, lineno, message) << std::endl; - } } // namespace detail } // namespace runtime } // namespace tvm diff --git a/onnxruntime/core/providers/tvm/tvm_allocator.cc b/onnxruntime/core/providers/tvm/tvm_allocator.cc index 4fadff5c57..ef06e1f59a 100644 --- a/onnxruntime/core/providers/tvm/tvm_allocator.cc +++ b/onnxruntime/core/providers/tvm/tvm_allocator.cc @@ -8,6 +8,7 @@ namespace onnxruntime { +namespace tvm { void* TVMAllocator::Alloc(size_t size) { void* p = nullptr; @@ -24,4 +25,5 @@ void TVMAllocator::Free(void* p) { TVMDeviceFreeDataSpace(ctx, p); } -} // namespace onnxruntime +} // namespace tvm +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/tvm/tvm_allocator.h b/onnxruntime/core/providers/tvm/tvm_allocator.h index 1185402487..50a13f890a 100644 --- a/onnxruntime/core/providers/tvm/tvm_allocator.h +++ b/onnxruntime/core/providers/tvm/tvm_allocator.h @@ -7,7 +7,9 @@ #include "core/framework/allocator.h" #include "tvm_common.h" + namespace onnxruntime { +namespace tvm { #define TVM_ALLOC_ALIGN 128 @@ -22,14 +24,14 @@ class TVMAllocator : public IAllocator { : IAllocator(info) { switch (info.device.Type()) { case OrtDevice::CPU: - ctx = {kDLCPU, info.device.Id()}; - break; + ctx = {kDLCPU, info.device.Id()}; + break; case OrtDevice::GPU: - ctx = {kDLVulkan, info.device.Id()}; - break; + ctx = {kDLVulkan, info.device.Id()}; + break; default: - ORT_NOT_IMPLEMENTED("Unsupported device"); - break; + ORT_NOT_IMPLEMENTED("Unsupported device"); + break; } } @@ -38,5 +40,7 @@ class TVMAllocator : public IAllocator { DLDevice ctx; }; -} // namespace onnxruntime +} // namespace tvm +} // namespace onnxruntime + #endif // TVM_ALLOCATOR diff --git a/onnxruntime/core/providers/tvm/tvm_api.cc b/onnxruntime/core/providers/tvm/tvm_api.cc index f225e06528..ff61c6c43d 100644 --- a/onnxruntime/core/providers/tvm/tvm_api.cc +++ b/onnxruntime/core/providers/tvm/tvm_api.cc @@ -17,16 +17,9 @@ using TvmPackedFunc = ::tvm::PackedFunc; TvmModule TVMCompile(const std::string& onnx_txt, const std::string& model_path, - const std::string& executor, - const std::string& target, - const std::string& target_host, - int opt_level, + const TvmEPOptions& options, int opset, - bool freeze_params, - const std::vector>& input_shapes, - bool nhwc, - const std::string& tuning_logfile, - const std::string& tuning_type) + const TVMTensorShapes& input_shapes) { ::tvm::Array shapes; for (size_t i = 0; i < input_shapes.size(); ++i) @@ -41,19 +34,18 @@ TvmModule TVMCompile(const std::string& onnx_txt, const TvmPackedFunc* compile = ::tvm::runtime::Registry::Get("tvm_onnx_import_and_compile"); ORT_ENFORCE(compile != nullptr, "Unable to retrieve 'tvm_onnx_import_and_compile'."); - TvmModule mod = (*compile)( - TVMByteArray{onnx_txt.data(), onnx_txt.size()}, - model_path, - executor, - target, - target_host, - opt_level, - opset, - freeze_params, - shapes, - nhwc, - tuning_logfile, - tuning_type); + TvmModule mod = (*compile)(TVMByteArray{onnx_txt.data(), onnx_txt.size()}, + model_path, + options.executor, + options.target, + options.target_host, + options.opt_level, + opset, + options.freeze_weights, + shapes, + options.to_nhwc, + options.tuning_file_path, + options.tuning_type); ORT_ENFORCE(mod.get() != nullptr, "Compiled TVM Module is nullptr!"); return mod; } @@ -108,20 +100,19 @@ void TVM_VM_GetOutputs(TvmModule& mod, } void TVMGetOutputShapes(TvmModule& mod, - size_t num_outputs, - std::vector>& output_shapes) + TVMTensorShapes& output_shapes) { - output_shapes.clear(); + size_t size = output_shapes.size(); TvmPackedFunc get_output = mod.GetFunction("get_output", false); - for (size_t i = 0; i < num_outputs; ++i) { + for (size_t i = 0; i < size; ++i) { ::tvm::runtime::NDArray output_array = get_output(i); ::tvm::runtime::ShapeTuple shape_tuple = output_array.Shape(); size_t dims_num = shape_tuple.size(); - std::vector dims; + TensorShapeVector dims; for (size_t j = 0; j < dims_num; ++j) { dims.push_back(int64_t(shape_tuple[j])); } - output_shapes.push_back(dims); + output_shapes[i] = dims; } } diff --git a/onnxruntime/core/providers/tvm/tvm_api.h b/onnxruntime/core/providers/tvm/tvm_api.h index 291da43861..77241def6e 100644 --- a/onnxruntime/core/providers/tvm/tvm_api.h +++ b/onnxruntime/core/providers/tvm/tvm_api.h @@ -9,31 +9,27 @@ #include "tvm_common.h" #include "tvm_defaults.h" +#include "tvm_ep_options.h" + namespace onnxruntime { namespace tvm { - TvmModule TVMCompile(const std::string& onnx_txt, - const std::string& model_path, - const std::string& executor, - const std::string& target, - const std::string& target_host, - int opt_level, - int opset, - bool freeze_params, - const std::vector>& input_shapes, - bool nhwc = false, - const std::string& tuning_logfile = "", - const std::string& tuning_type = std::string(onnxruntime::tvm::default_tuning_type)); - void TVMSetInputs(TvmModule& mod, std::vector& inds, std::vector& inputs); - void TVM_VM_SetInputs(TvmModule& mod, std::vector& inds, std::vector& inputs); - void TVMGetOutputs(TvmModule& mod, std::vector& outputs); - void TVM_VM_GetOutputs(TvmModule& mod, std::vector& outputs); - void TVMGetOutputShapes(TvmModule& mod, - size_t num_outputs, - std::vector>& output_shapes); - void TVMRun(TvmModule& mod); - void TVM_VM_Run(TvmModule& mod); + + TvmModule TVMCompile(const std::string& onnx_txt, + const std::string& model_path, + const TvmEPOptions& options, + int opset, + const TVMTensorShapes& input_shapes); + void TVMSetInputs(TvmModule& mod, std::vector& inds, std::vector& inputs); + void TVM_VM_SetInputs(TvmModule& mod, std::vector& inds, std::vector& inputs); + void TVMGetOutputs(TvmModule& mod, std::vector& outputs); + void TVM_VM_GetOutputs(TvmModule& mod, std::vector& outputs); + void TVMGetOutputShapes(TvmModule& mod, + TVMTensorShapes& output_shapes); + void TVMRun(TvmModule& mod); + void TVM_VM_Run(TvmModule& mod); + } // namespace tvm } // namespace onnxruntime -#endif // TVM_API_H \ No newline at end of file +#endif // TVM_API_H diff --git a/onnxruntime/core/providers/tvm/tvm_common.h b/onnxruntime/core/providers/tvm/tvm_common.h index 5b3a0c4dea..17120eb2a2 100644 --- a/onnxruntime/core/providers/tvm/tvm_common.h +++ b/onnxruntime/core/providers/tvm/tvm_common.h @@ -4,10 +4,20 @@ #ifndef TVM_COMMON_H #define TVM_COMMON_H +#include +#include + #include #include #include -using TvmModule = tvm::runtime::Module; + +namespace onnxruntime { +namespace tvm { + +using TvmModule = ::tvm::runtime::Module; + +} // namespace tvm +} // namespace onnxruntime #endif // TVM_COMMON_H diff --git a/onnxruntime/core/providers/tvm/tvm_compiler.cc b/onnxruntime/core/providers/tvm/tvm_compiler.cc new file mode 100644 index 0000000000..dfbf182506 --- /dev/null +++ b/onnxruntime/core/providers/tvm/tvm_compiler.cc @@ -0,0 +1,38 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include + +#include "tvm_compiler.h" +#include "tvm_api.h" + + +namespace onnxruntime { +namespace tvm { + +TVMCompiler::TVMCompiler(std::string&& onnx_model_str, + const std::string& model_path, + int opset) : +onnx_model_str_(std::move(onnx_model_str)), +model_path_(model_path), +opset_(opset) { +} + +auto TVMCompiler::operator()(const TvmEPOptions& options, + const TVMTensorShapes& input_shapes) -> ModulePtr { + if (mod_) { + return mod_; + } + + mod_ = std::make_shared(); + *mod_ = tvm::TVMCompile(onnx_model_str_, + model_path_, + options, + opset_, + input_shapes); + onnx_model_str_.clear(); + return mod_; +} + +} // namespace tvm +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/tvm/tvm_compiler.h b/onnxruntime/core/providers/tvm/tvm_compiler.h new file mode 100644 index 0000000000..057ed058fd --- /dev/null +++ b/onnxruntime/core/providers/tvm/tvm_compiler.h @@ -0,0 +1,40 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#ifndef TVM_COMPILER_H +#define TVM_COMPILER_H + +#include +#include + +#include "tvm_common.h" +#include "tvm_ep_options.h" + + +namespace onnxruntime { +namespace tvm { + +class TVMCompiler { + using ModulePtr = std::shared_ptr; +public: + TVMCompiler() = delete; + ~TVMCompiler() = default; + + TVMCompiler(std::string&& onnx_model_str, + const std::string& model_path, + int opset); + + ModulePtr operator()(const TvmEPOptions& options, + const TVMTensorShapes& input_shapes); + +private: + ModulePtr mod_; + std::string onnx_model_str_; + std::string model_path_; + int opset_; +}; + +} // namespace tvm +} // namespace onnxruntime + +#endif // TVM_COMPILER_H diff --git a/onnxruntime/core/providers/tvm/tvm_ep_options.cc b/onnxruntime/core/providers/tvm/tvm_ep_options.cc new file mode 100644 index 0000000000..3b20bbc917 --- /dev/null +++ b/onnxruntime/core/providers/tvm/tvm_ep_options.cc @@ -0,0 +1,256 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include +#include + +#include "core/common/common.h" +#include "core/common/cpuid_info.h" +#include "core/framework/provider_options_utils.h" + +#include "tvm_ep_options.h" + + +namespace onnxruntime { +namespace tvm { + +namespace provider_option_names { +constexpr const char* kExecutor = "executor"; +constexpr const char* kTarget = "target"; +constexpr const char* kTargetHost = "target_host"; +constexpr const char* kOptLevel = "opt_level"; +constexpr const char* kFreezeWeights = "freeze_weights"; +constexpr const char* kToNHWC = "to_nhwc"; +constexpr const char* kTuningFilePath = "tuning_file_path"; +constexpr const char* kTuningType = "tuning_type"; +constexpr const char* kInputNames = "input_names"; +constexpr const char* kInputShapes = "input_shapes"; + +static const std::unordered_set valid_keys { + std::string{kExecutor}, + std::string{kTarget}, + std::string{kTargetHost}, + std::string{kOptLevel}, + std::string{kFreezeWeights}, + std::string{kToNHWC}, + std::string{kTuningFilePath}, + std::string{kTuningType}, + std::string{kInputNames}, + std::string{kInputShapes} +}; + +} // namespace provider_option_names + +size_t split(const std::string &src, std::vector &dst, char ch) { + dst.clear(); + + size_t pos = src.find( ch ); + size_t initialPos = 0; + while( pos != std::string::npos ) { + dst.push_back( src.substr( initialPos, pos - initialPos ) ); + initialPos = pos + 1; + + pos = src.find( ch, initialPos ); + } + dst.push_back( src.substr( initialPos, std::min( pos, src.size() ) - initialPos + 1 ) ); + + return dst.size(); +} + +TvmEPOptions TvmEPOptionsHelper::FromOptionsString(const char* opt_str) { + std::string settings{opt_str}; + ProviderOptions options; + if (!settings.empty()) { + const std::string& str = settings; + + // tokenize settings + std::regex reg("\\s*,\\s*"); + std::sregex_token_iterator iter(str.begin(), str.end(), reg, -1); + std::sregex_token_iterator iter_end; + std::vector pairs(iter, iter_end); + + ORT_ENFORCE(pairs.size() > 0); + + for(const auto& pair : pairs) { + auto pos_colon = pair.find(':'); + ORT_ENFORCE(pos_colon != std::string::npos, "Invalid key value pair."); + std::string key = pair.substr(0, pos_colon); + std::string value = pair.substr(pos_colon + 1); + + // trim leading and trailing spaces from key/value + key = whitespace_trimming(key); + value = whitespace_trimming(value); + + // Check keys of obtained options + if (tvm::provider_option_names::valid_keys.count(key) == 0) { + ORT_NOT_IMPLEMENTED("TvmOptions: unknown option (", key, ")"); + } + + options[key] = value; + } + } + + return TvmEPOptionsHelper::FromProviderOptions(options); +} + +std::string TvmEPOptionsHelper::whitespace_trimming(const std::string& str) { + const std::string WHITESPACE = " \n\r\t\f\v"; + size_t start = str.find_first_not_of(WHITESPACE); + if (start == std::string::npos) { + return ""; + } else { + size_t end = str.find_last_not_of(WHITESPACE); + ORT_ENFORCE(end != std::string::npos); + return str.substr(start, end + 1); + } +} + +TvmEPOptions TvmEPOptionsHelper::FromProviderOptions(const ProviderOptions& pr_options) { + TvmEPOptions options{}; + + ORT_THROW_IF_ERROR( + ProviderOptionsParser{} + .AddAssignmentToReference(tvm::provider_option_names::kExecutor, options.executor) + .AddAssignmentToReference(tvm::provider_option_names::kTarget, options.target) + .AddAssignmentToReference(tvm::provider_option_names::kTargetHost, options.target_host) + .AddAssignmentToReference(tvm::provider_option_names::kOptLevel, options.opt_level) + .AddAssignmentToReference(tvm::provider_option_names::kFreezeWeights, options.freeze_weights) + .AddAssignmentToReference(tvm::provider_option_names::kToNHWC, options.to_nhwc) + .AddAssignmentToReference(tvm::provider_option_names::kTuningFilePath, options.tuning_file_path) + .AddAssignmentToReference(tvm::provider_option_names::kTuningType, options.tuning_type) + .AddAssignmentToReference(tvm::provider_option_names::kInputNames, options.input_names_str) + .AddAssignmentToReference(tvm::provider_option_names::kInputShapes, options.input_shapes_str) + .Parse(pr_options)); + + optionsPostprocess(options); + + return options; +} + +void TvmEPOptionsHelper::optionsPostprocess(TvmEPOptions& options) { + setInputShapes(options); + targetPostprocess(options.target); + targetHostPostprocess(options.target, options.target_host); + optLevelPostprocess(options.opt_level); +} + +bool TvmEPOptionsHelper::checkGPUTarget(const std::string& target) { + bool check = ( + target.find("cuda") != std::string::npos || + target.find("opencl") != std::string::npos || + target.find("metal") != std::string::npos || + target.find("vulkan") != std::string::npos + ); + return check; +} + +void TvmEPOptionsHelper::setInputShapes(TvmEPOptions& options) { + if (options.input_names_str.empty() && options.input_shapes_str.empty()) + return; + ORT_ENFORCE(!options.input_names_str.empty() && !options.input_shapes_str.empty(), + "Both provider options \"input_names\" and \"input_shapes\" should be empty or full"); + + std::vector name_set; + std::string trimmed_names = whitespace_trimming(options.input_names_str); + size_t inp_tensors_num = split(trimmed_names, name_set, ' '); + ORT_ENFORCE(inp_tensors_num, "There is no any input tensor names!"); + + std::string trimmed_shapes = whitespace_trimming(options.input_shapes_str); + size_t end_pos = trimmed_shapes.find_last_of(']'); + ORT_ENFORCE(end_pos != std::string::npos, "Invalid string for input shapes. Symbol ] is not found"); + ORT_ENFORCE(end_pos == (trimmed_shapes.size() - 1), + "Invalid string for input shapes. Symbol ] should be last after whitespace trimming"); + + std::vector shape_set; + split(trimmed_shapes, shape_set, ']'); + shape_set.pop_back(); + ORT_ENFORCE( shape_set.size() == inp_tensors_num, + "Number of shapes is not the same as number of input tensor names"); + + for (size_t i = 0; i < inp_tensors_num; ++i) { + size_t pos = shape_set[i].find('['); + ORT_ENFORCE(pos != std::string::npos, "There is no symbol [ as pair for ]"); + std::string numbers = shape_set[i].substr(pos + 1); + std::vector number_set; + ORT_ENFORCE(split(numbers, number_set, ' '), "There is no any number between [ and ] symbols"); + + TensorShapeVector dims; + for(const auto& number : number_set) { + dims.push_back(std::stoi(number)); + } + + options.input_shapes[name_set[i]] = dims; + } +} + +void TvmEPOptionsHelper::targetPostprocess(std::string& target) { + if(target == tvm::cpu_target_str || + target == tvm::llvm_target_str) { + ProcessCPUTarget(target); + } else if(target == tvm::gpu_target_str) { + ProcessGPUTarget(); + } else if(target.empty()) { + ORT_NOT_IMPLEMENTED("target option is empty!"); + } else { + // TODO(vvchernov): extend mechanism of auto-definition of target + // target is gotten from option set up by client + } +} + +void TvmEPOptionsHelper::ProcessCPUTarget(std::string& target) { + const auto& cpu_id_info = CPUIDInfo::GetCPUIDInfo(); + // auto detect from CPU ID + if (cpu_id_info.HasAVX512Skylake()) { + target = tvm::cpu_targets::LLVM_TARGET_SKYLAKE_AVX512; + } else if (cpu_id_info.HasAVX512f()) { + target = tvm::cpu_targets::LLVM_TARGET_AVX512; + } else if (cpu_id_info.HasAVX2()) { + target = tvm::cpu_targets::LLVM_TARGET_AVX2; + } else if (cpu_id_info.HasAVX()) { + target = tvm::cpu_targets::LLVM_TARGET_AVX; + } else { + // TODO(vvchernov): extend mechanism of auto-definition of cpu target + target = tvm::llvm_target_str; + } +} + +void TvmEPOptionsHelper::ProcessGPUTarget() { + ORT_NOT_IMPLEMENTED("GPU target auto-defenition is not implemented now!"); +} + +void TvmEPOptionsHelper::targetHostPostprocess(const std::string& target, std::string& target_host) { + if((target_host == tvm::cpu_target_str || + target_host == tvm::llvm_target_str) && + target_host != target) { + target_host = target; + } else if (target_host.empty()) { + target_host = target; + } else { + // TODO(vvchernov): extend mechanism of auto-definition of target host + // target host is gotten from option set up by client + } +} + +void TvmEPOptionsHelper::optLevelPostprocess(unsigned int& opt_level) { + if(opt_level < 1) { + opt_level = tvm::default_opt_level; + } +} + +std::ostream& operator<<(std::ostream& out, const TvmEPOptions& options) { + out << "TVM EP options:\n" << + "executor type: " << options.executor << "\n" << + "target: " << options.target << "\n" << + "target_host: " << options.target_host << "\n" << + "opt level: " << options.opt_level << "\n" << + "freeze weights: " << options.freeze_weights << "\n" << + "tuning file path: " << options.tuning_file_path << "\n" << + "tuning type: " << options.tuning_type << "\n" << + "convert layout to NHWC: " << options.to_nhwc << "\n" << + "input tensor names: " << options.input_names_str << "\n" << + "input tensor shapes: " << options.input_shapes_str; + return out; +} + +} // namespace tvm +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/tvm/tvm_execution_provider_info.h b/onnxruntime/core/providers/tvm/tvm_ep_options.h similarity index 53% rename from onnxruntime/core/providers/tvm/tvm_execution_provider_info.h rename to onnxruntime/core/providers/tvm/tvm_ep_options.h index b3c8932ab6..4d323698b0 100644 --- a/onnxruntime/core/providers/tvm/tvm_execution_provider_info.h +++ b/onnxruntime/core/providers/tvm/tvm_ep_options.h @@ -1,17 +1,20 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. -#ifndef TVM_EXECUTION_PROVIDER_INFO_H -#define TVM_EXECUTION_PROVIDER_INFO_H +#ifndef TVM_EXECUTION_PROVIDER_OPTIONS_H +#define TVM_EXECUTION_PROVIDER_OPTIONS_H #include #include #include +#include #include "core/framework/provider_options.h" +#include "core/framework/tensor_shape.h" #include "tvm_defaults.h" + namespace onnxruntime { namespace tvm { @@ -22,12 +25,13 @@ const std::string LLVM_TARGET_AVX2 = "llvm -mcpu=core-avx2"; const std::string LLVM_TARGET_SKYLAKE_AVX512 = "llvm -mcpu=skylake-avx512"; const std::string LLVM_TARGET_AVX512 = "llvm -mcpu=skylake-avx512"; } // namespace cpu_targets -} // namespace tvm -using TVMInputShapes = std::unordered_map>; +using TVMTensorShapes = std::vector; +using TVMInputShapes = std::unordered_map; +using InputsInfoMap = std::unordered_map; // Information needed to construct an TVM execution provider. -struct TvmExecutionProviderInfo { +struct TvmEPOptions { std::string executor{tvm::default_executor_type}; std::string target{tvm::default_target_str}; std::string target_host{tvm::default_target_str}; @@ -39,12 +43,30 @@ struct TvmExecutionProviderInfo { std::string input_names_str{""}; std::string input_shapes_str{""}; TVMInputShapes input_shapes{}; - - static std::string whitespace_trimming(const std::string& str); - static TvmExecutionProviderInfo FromProviderOptions(const ProviderOptions& options); - static TvmExecutionProviderInfo FromOptionsString(const char* options); + TVMTensorShapes output_shapes{}; }; +std::ostream& operator<<(std::ostream& out, const TvmEPOptions& options); + +class TvmEPOptionsHelper { +public: + static TvmEPOptions FromOptionsString(const char* options); + static TvmEPOptions FromProviderOptions(const ProviderOptions& options); + static std::string whitespace_trimming(const std::string& str); + + static bool checkGPUTarget(const std::string& target); + +private: + static void optionsPostprocess(TvmEPOptions& options); + static void setInputShapes(TvmEPOptions& options); + static void targetPostprocess(std::string& target); + static void ProcessCPUTarget(std::string& target); + static void ProcessGPUTarget(); + static void targetHostPostprocess(const std::string& target, std::string& target_host); + static void optLevelPostprocess(unsigned int& opt_level); +}; + +} // namespace tvm } // namespace onnxruntime -#endif // TVM_EXECUTION_PROVIDER_INFO_H +#endif // TVM_EXECUTION_PROVIDER_OPTIONS_H diff --git a/onnxruntime/core/providers/tvm/tvm_execution_provider.cc b/onnxruntime/core/providers/tvm/tvm_execution_provider.cc index cafd1561f0..60b6c8469f 100644 --- a/onnxruntime/core/providers/tvm/tvm_execution_provider.cc +++ b/onnxruntime/core/providers/tvm/tvm_execution_provider.cc @@ -10,7 +10,6 @@ #include "core/framework/compute_capability.h" #include "core/platform/env.h" #include "core/graph/model.h" -#include "core/common/cpuid_info.h" #include "tvm_execution_provider.h" #include "xpu_data_transfer.h" @@ -22,239 +21,19 @@ using namespace ONNX_NAMESPACE; namespace onnxruntime { +namespace tvm { // Information to construct kernel function state. struct TVMFuncState { AllocateFunc allocate_func = nullptr; DestroyFunc release_func = nullptr; AllocatorHandle allocator = nullptr; - TvmModule* module = nullptr; - std::function>& input_shapes)> compiler = nullptr; + std::shared_ptr compiler = nullptr; }; -class TVMRunner { - public: - using TVMTensorShape = std::vector; - using TVMTensorShapes = std::vector; - using InputsInfoMap = std::map; - using ORTGraphNodes = std::vector; - - TVMRunner() = delete; - ~TVMRunner() = default; - - TVMRunner(TvmExecutionProvider* ep, - const std::string& name, - const Graph& graph) : - use_vm_(ep->info_.executor == "vm") { - // Extract input shapes - const ORTGraphNodes& all_nodes = graph.GetInputsIncludingInitializers(); - TVMTensorShapes input_shapes; - size_t indx = 0; - if (ep->info_.freeze_weights) { - for (const auto* node : all_nodes) { - const auto& node_name = node->Name(); - if(!graph.IsInitializedTensor(node_name)) { - TVMTensorShape ishape; - if(!ep->info_.input_shapes.empty() && - ep->info_.input_shapes.count(node_name)) { - ishape = ep->info_.input_shapes[node_name]; - inputs_info_[indx] = ishape; - update_output_shapes_ = true; - } else { - getTensorInfo(*node->Shape(), ishape, indx); - } - input_shapes.emplace_back(ishape); - } - ++indx; - } - } else { - for (const auto* node : all_nodes) { - const auto& node_name = node->Name(); - TVMTensorShape ishape; - if(!ep->info_.input_shapes.empty() && - ep->info_.input_shapes.count(node_name)) { - ishape = ep->info_.input_shapes[node_name]; - inputs_info_[indx++] = ishape; - update_output_shapes_ = true; - } else { - getTensorInfo(*node->Shape(), ishape, indx++); - } - if(!graph.IsInitializedTensor(node_name)) { - input_shapes.emplace_back(ishape); - } - } - } - - // Get module from tvm - mod_ = ep->CompileFunc(name, input_shapes); - - // Prepare draft for output tvm tensors - const ORTGraphNodes& ort_outputs_info = graph.GetOutputs(); - size_t num_outputs = ort_outputs_info.size(); - - if (update_output_shapes_) { - if (!use_vm_) { - tvm::TVMGetOutputShapes(*mod_, num_outputs, output_shapes_); - } - } else { - for (auto i = 0u; i < num_outputs; i++) { - TensorShape ort_shape = utils::GetTensorShapeFromTensorShapeProto(*ort_outputs_info[i]->Shape()); - int dims = ort_shape.NumDimensions(); - - TVMTensorShape oshape(dims); - for (int j = 0; j < dims; ++j) { - oshape[j] = int64_t(ort_shape[j]); - } - output_shapes_.emplace_back(oshape); - } - } - - for (auto i = 0u; i < num_outputs; i++) { - DLTensor t; - // Draft for tensor, correct data is defined during inference - t.strides = nullptr; - t.byte_offset = 0; - t.data = nullptr; - if (!(use_vm_ && update_output_shapes_)) { - t.ndim = output_shapes_[i].size(); - t.shape = output_shapes_[i].data(); - } else { - t.ndim = 0; - t.shape = nullptr; - } - - tensors_outputs_.push_back(t); - } - } - - common::Status operator()(FunctionState state, const OrtCustomOpApi* api, OrtKernelContext* context) { - Ort::CustomOpApi ort{*api}; - - size_t num = inputs_info_.size(); - std::vector inds(num); - std::vector dl_tensors_inputs(num); - size_t counter = 0u; - for (auto& info : inputs_info_) { - // TODO(vvchernov): decomposition declaration only available with -std=c++1z or -std=gnu++1z - auto& i = info.first; - auto& shape = info.second; - const OrtValue* input_tensor = ort.KernelContext_GetInput(context, i); - ORT_ENFORCE(input_tensor->IsTensor()); - const Tensor& tensor = input_tensor->Get(); - const OrtDevice& device = tensor.Location().device; - auto tensor_info = ort.GetTensorTypeAndShape(input_tensor); - auto tensor_type = ort.GetTensorElementType(tensor_info); - if (!update_output_shapes_) { - std::vector ort_shape = ort.GetTensorShape(tensor_info); - ORT_ENFORCE(compare_shapes(shape, ort_shape)); - } - ort.ReleaseTensorTypeAndShapeInfo(tensor_info); - - DLTensor t; - t.device = GetDLDevice(device); - t.dtype = GetDataType(tensor_type); - t.strides = nullptr; - t.byte_offset = 0; - t.data = const_cast(ort.GetTensorData(input_tensor)); - t.ndim = shape.size(); - t.shape = shape.data(); - dl_tensors_inputs[counter] = t; - inds[counter++] = i; - } - if (use_vm_) { - tvm::TVM_VM_SetInputs(*mod_, inds, dl_tensors_inputs); - // Infer once for calculating of output shapes - if(!probe_infer_) { - tvm::TVM_VM_Run(*mod_); - size_t num_outputs = tensors_outputs_.size(); - tvm::TVMGetOutputShapes(*mod_, num_outputs, output_shapes_); - for (size_t i = 0; i < num_outputs; ++i) { - tensors_outputs_[i].ndim = output_shapes_[i].size(); - tensors_outputs_[i].shape = output_shapes_[i].data(); - } - probe_infer_ = true; - } - } else { - tvm::TVMSetInputs(*mod_, inds, dl_tensors_inputs); - } - - size_t num_outputs = tensors_outputs_.size(); - for (auto i = 0u; i < num_outputs; i++) { - //setup output tensor property - OrtValue* output_tensor = ort.KernelContext_GetOutput(context, - i, - output_shapes_[i].data(), - output_shapes_[i].size()); - ORT_ENFORCE(output_tensor->IsTensor()); - const Tensor& tensor = output_tensor->Get(); - const OrtDevice& device = tensor.Location().device; - auto tensor_info = ort.GetTensorTypeAndShape(output_tensor); - auto tensor_type = ort.GetTensorElementType(tensor_info); - ort.ReleaseTensorTypeAndShapeInfo(tensor_info); - - tensors_outputs_[i].device = GetDLDevice(device); - tensors_outputs_[i].dtype = GetDataType(tensor_type); - tensors_outputs_[i].data = ort.GetTensorMutableData(output_tensor); - } - - if (use_vm_) { - tvm::TVM_VM_Run(*mod_); - tvm::TVM_VM_GetOutputs(*mod_, tensors_outputs_); - } else { - tvm::TVMRun(*mod_); - tvm::TVMGetOutputs(*mod_, tensors_outputs_); - } - - return Status::OK(); - } - private: - void getTensorInfo(const TensorShapeProto& shape_proto, - TVMTensorShape& ishape, - size_t indx) { - TensorShape ort_shape = utils::GetTensorShapeFromTensorShapeProto(shape_proto); - int dims = ort_shape.NumDimensions(); - - ishape.resize(dims); - for (int j = 0; j < dims; ++j) { - int64_t dim = int64_t(ort_shape[j]); - ORT_ENFORCE(dim > 0, "Input dimension is not positive value (dim = " + std::to_string(dim) + "). " + - "Please use provider options to setup input_names and input_shapes"); - ishape[j] = dim; - } - inputs_info_[indx] = ishape; - } - - bool compare_shapes(const TVMTensorShape& shape1, const TVMTensorShape& shape2) { - size_t size = shape1.size(); - if (shape2.size() == size) { - for (size_t i = 0; i < size; ++i) { - if(shape1[i] != shape2[i]) { - return false; - } - } - } else { - return false; - } - - return true; - } - - private: - TvmModule* mod_; - bool use_vm_ = true; - bool probe_infer_ = false; - InputsInfoMap inputs_info_{}; - bool update_output_shapes_ = false; - TVMTensorShapes output_shapes_; - std::vector tensors_outputs_; -}; - -TvmExecutionProvider::TvmExecutionProvider(const TvmExecutionProviderInfo& info) +TvmExecutionProvider::TvmExecutionProvider(const TvmEPOptions& options) : IExecutionProvider{kTvmExecutionProvider}, - info_{info} { - ProcessInfo(); - + options_{options} { AllocatorCreationInfo default_memory_info = {[](int) { return std::make_unique(); }, @@ -273,10 +52,6 @@ TvmExecutionProvider::TvmExecutionProvider(const TvmExecutionProviderInfo& info) TvmExecutionProvider::~TvmExecutionProvider() {} -AllocatorPtr TvmExecutionProvider::GetAllocator(int id, OrtMemType mem_type) const { - return allocator_; -} - std::vector> TvmExecutionProvider::GetCapability(const GraphViewer& graph_viewer, const std::vector& /*kernel_registries*/) const { @@ -327,8 +102,8 @@ TvmExecutionProvider::GetCapability(const GraphViewer& graph_viewer, } common::Status TvmExecutionProvider::Compile(const std::vector& nodes, - std::vector& node_compute_funcs) { - PrintProviderOptions(); + std::vector& node_compute_funcs) { + printOptions(); for (auto* fused_node : nodes) { auto func_body = fused_node->GetFunctionBody(); if (!func_body) @@ -345,32 +120,28 @@ common::Status TvmExecutionProvider::Compile(const std::vector& nodes, opset->set_domain(kOnnxDomain); opset->set_version(node_graph.DomainToVersionMap().at(kOnnxDomain)); - std::string string_buf; - model_proto.SerializeToString(&string_buf); - buffers_[func_name] = string_buf; - opsets_[func_name] = int(opset->version()); - model_paths_[func_name] = fused_node->ModelPath().ToPathString();; + std::string onnx_model_str; + model_proto.SerializeToString(&onnx_model_str); + compilers_[func_name] = std::make_shared(std::move(onnx_model_str), + fused_node->ModelPath().ToPathString(), + int(opset->version())); + InputsInfoMap all_input_shapes; + auto mod = compileModel(func_name, node_graph, all_input_shapes); + + std::vector output_tensors; + prepareOutputTensors(mod, output_tensors, node_graph.GetOutputs().size()); + + runners_[func_name] = std::make_shared(options_, mod, all_input_shapes, output_tensors); if (dump_subgraphs_) { - std::fstream dump("/tmp/" + fused_node->Name() + ".onnx", + std::fstream dump("/tmp/" + func_name + ".onnx", std::ios::out | std::ios::trunc | std::ios::binary); model_proto.SerializeToOstream(&dump); } - NodeComputeInfo compute_info; - compute_info.create_state_func = std::bind(&TvmExecutionProvider::CreateStateFunc, - this, - std::placeholders::_1, - std::placeholders::_2); - - compute_info.release_state_func = [](FunctionState state) { - if (state) - delete static_cast(state); - }; // TODO(vvchernov): implement ops checking and mechanism of gracefully passing the responsibility to other EPs // if the checking fails due to unsupported op(s) - runners_[func_name] = std::make_shared(this, func_name, node_graph); - compute_info.compute_func = *runners_[func_name].get(); + NodeComputeInfo compute_info = prepareComputeInfo(func_name); node_compute_funcs.push_back(compute_info); } @@ -378,182 +149,156 @@ common::Status TvmExecutionProvider::Compile(const std::vector& nodes, } std::unique_ptr TvmExecutionProvider::GetDataTransfer() const { - if (GPUTargetCheck()) { - return std::make_unique(); - } else if (info_.target.find("llvm") != std::string::npos) { - return std::make_unique(); - } else { - ORT_NOT_IMPLEMENTED("TVM GetDataTransfer is not implemented for target ", info_.target); - } -} - -bool TvmExecutionProvider::GPUTargetCheck() const { //TODO(vvchernov): target or target host? - bool check = ( - info_.target.find("cuda") != std::string::npos || - info_.target.find("opencl") != std::string::npos || - info_.target.find("metal") != std::string::npos || - info_.target.find("vulkan") != std::string::npos - ); - return check; + if (TvmEPOptionsHelper::checkGPUTarget(options_.target)) { + return std::make_unique(); + } else if (options_.target.find("llvm") != std::string::npos) { + return std::make_unique(); + } else { + ORT_NOT_IMPLEMENTED("TVM GetDataTransfer is not implemented for target ", options_.target); + } } -size_t TvmExecutionProvider::split(const std::string &txt, std::vector &strs, char ch) const { - size_t pos = txt.find( ch ); - size_t initialPos = 0; - strs.clear(); - - while( pos != std::string::npos ) { - strs.push_back( txt.substr( initialPos, pos - initialPos ) ); - initialPos = pos + 1; - - pos = txt.find( ch, initialPos ); - } - - strs.push_back( txt.substr( initialPos, std::min( pos, txt.size() ) - initialPos + 1 ) ); - - return strs.size(); +AllocatorPtr TvmExecutionProvider::GetAllocator(int id, OrtMemType mem_type) const { + return allocator_; } -void TvmExecutionProvider::ProcessInfo() { - if(!info_.input_shapes_str.empty()) { - ORT_ENFORCE(!info_.input_names_str.empty(), - "Please insert input tensor names. Input shapes only is invalid case"); - // Parse strings and set to input_shapes map - std::vector tmp_strs; - std::vector names_strs; +void TvmExecutionProvider::printOptions() { + LOGS(*GetLogger(), INFO) << options_; +} - std::string names_str = TvmExecutionProviderInfo::whitespace_trimming(info_.input_names_str); - std::string shapes_str = TvmExecutionProviderInfo::whitespace_trimming(info_.input_shapes_str); +std::shared_ptr TvmExecutionProvider::compileModel(const std::string& func_name, + const Graph& graph, + InputsInfoMap& all_input_shapes) { + all_input_shapes.clear(); - ORT_ENFORCE(split(names_str, names_strs, ' '), "There is no any input tensor names!"); - size_t inp_tensors_num = names_strs.size(); + TVMTensorShapes input_shapes; + if (options_.freeze_weights) { + setInputShapesForFreezedNN(graph, input_shapes, all_input_shapes); + } else { + setInputShapesForUnfreezedNN(graph, input_shapes, all_input_shapes); + } - size_t end_pos = shapes_str.find_last_of(']'); - ORT_ENFORCE(end_pos != std::string::npos, "Invalid string for input shapes. Symbol ] is not found"); - ORT_ENFORCE(end_pos == (shapes_str.size() - 1), - "Invalid string for input shapes. Symbol ] should be last after whitespace trimming"); - split(shapes_str, tmp_strs, ']'); - tmp_strs.pop_back(); - ORT_ENFORCE( tmp_strs.size() == inp_tensors_num, - "Number of shapes is not the same as number of input tensor names"); - for (size_t i = 0; i < inp_tensors_num; ++i) { - size_t pos = tmp_strs[i].find('['); - ORT_ENFORCE(pos != std::string::npos, "There is no symbol [ as pair for ]"); - std::string nums_str = tmp_strs[i].substr(pos + 1); - std::vector nums_strs; - ORT_ENFORCE(split(nums_str, nums_strs, ' '), "There is no any numbers between [ and ] symbols"); - std::vector dims; - for(const auto& num_str : nums_strs) { - dims.push_back(std::stoi(num_str)); - } + std::shared_ptr mod = compilers_[func_name]->operator()(options_, input_shapes); - info_.input_shapes[names_strs[i]] = dims; + return mod; +} + +void TvmExecutionProvider::setInputShapesForFreezedNN(const Graph& graph, + TVMTensorShapes& input_shapes, + InputsInfoMap& all_input_shapes) { + const std::vector& all_nodes = graph.GetInputsIncludingInitializers(); + + size_t indx = 0; + for (const auto* node : all_nodes) { + if(!graph.IsInitializedTensor(node->Name())) { + TensorShapeVector shape = getInputShape(node); + all_input_shapes[indx++] = shape; + input_shapes.emplace_back(shape); } } +} - if(info_.target == tvm::cpu_target_str || - info_.target == tvm::llvm_target_str) { - ProcessCPUTarget(); - } else if(info_.target == tvm::gpu_target_str) { - ProcessGPUTarget(); - } else if(info_.target.empty()) { - ORT_NOT_IMPLEMENTED("target option is empty!"); - } else { - // TODO(vvchernov): extend mechanism of auto-definition of target - // target is gotten from option set up by client - } +void TvmExecutionProvider::setInputShapesForUnfreezedNN(const Graph& graph, + TVMTensorShapes& input_shapes, + InputsInfoMap& all_input_shapes) { + const std::vector& all_nodes = graph.GetInputsIncludingInitializers(); - if((info_.target_host == tvm::cpu_target_str || - info_.target_host == tvm::llvm_target_str) && - info_.target_host != info_.target) { - info_.target_host = info_.target; - } else if (info_.target_host.empty()) { - info_.target_host = info_.target; - } else { - // TODO(vvchernov): extend mechanism of auto-definition of target host - // target host is gotten from option set up by client - } - - if(info_.opt_level < 1) { - info_.opt_level = tvm::default_opt_level; + size_t indx = 0; + for (const auto* node : all_nodes) { + TensorShapeVector shape = getInputShape(node); + all_input_shapes[indx++] = shape; + if(!graph.IsInitializedTensor(node->Name())) { + input_shapes.emplace_back(shape); + } } } -void TvmExecutionProvider::ProcessCPUTarget() { - const auto& cpu_id_info = CPUIDInfo::GetCPUIDInfo(); - // auto detect from CPU ID - if (cpu_id_info.HasAVX512Skylake()) { - info_.target = tvm::cpu_targets::LLVM_TARGET_SKYLAKE_AVX512; - } else if (cpu_id_info.HasAVX512f()) { - info_.target = tvm::cpu_targets::LLVM_TARGET_AVX512; - } else if (cpu_id_info.HasAVX2()) { - info_.target = tvm::cpu_targets::LLVM_TARGET_AVX2; - } else if (cpu_id_info.HasAVX()) { - info_.target = tvm::cpu_targets::LLVM_TARGET_AVX; - } else { - // TODO(vvchernov): extend mechanism of auto-definition of cpu target - info_.target = tvm::llvm_target_str; +TensorShapeVector TvmExecutionProvider::getInputShape(const NodeArg* node) { + TensorShapeVector shape; + const auto& node_name = node->Name(); + if(!options_.input_shapes.empty() && + options_.input_shapes.count(node_name)) { + shape = options_.input_shapes[node_name]; + } else { + shape = convertTensorShape(*node->Shape()); + } + + return shape; +} + +TensorShapeVector TvmExecutionProvider::convertTensorShape(const TensorShapeProto& shape_proto) { + TensorShape ort_shape = utils::GetTensorShapeFromTensorShapeProto(shape_proto); + size_t dims = ort_shape.NumDimensions(); + + TensorShapeVector shape(dims); + for (size_t j = 0; j < dims; ++j) { + int64_t dim = int64_t(ort_shape[j]); + ORT_ENFORCE(dim > 0, "Input dimension is not positive value (dim = " + std::to_string(dim) + "). " + + "Please use provider options to setup input_names and input_shapes"); + shape[j] = dim; + } + + return shape; +} + +void TvmExecutionProvider::prepareOutputTensors(const std::shared_ptr& mod, + std::vector& output_tensors, + size_t num) { + ORT_ENFORCE(mod != nullptr, "TVM module is not compiled"); + output_tensors.clear(); + options_.output_shapes.clear(); + options_.output_shapes.resize(num); + + if (options_.executor != "vm") { + tvm::TVMGetOutputShapes(*mod, options_.output_shapes); + } + + for (auto& output_shape : options_.output_shapes) { + DLTensor t; + // Draft for tensor, correct data is defined during inference + t.strides = nullptr; + t.byte_offset = 0; + t.data = nullptr; + if (options_.executor == "vm") { + t.ndim = 0; + t.shape = nullptr; + } else { + t.ndim = output_shape.size(); + t.shape = output_shape.data(); + } + + output_tensors.push_back(t); } } -void TvmExecutionProvider::ProcessGPUTarget() { - ORT_NOT_IMPLEMENTED("GPU target auto-defenition is not implemented now!"); +NodeComputeInfo TvmExecutionProvider::prepareComputeInfo(const std::string& func_name) { + NodeComputeInfo compute_info; + compute_info.create_state_func = std::bind(&TvmExecutionProvider::createStateFunc, + this, + std::placeholders::_1, + std::placeholders::_2); + + compute_info.release_state_func = [](FunctionState state) { + if (state) + delete static_cast(state); + }; + + compute_info.compute_func = *runners_[func_name].get(); + + return compute_info; } -void TvmExecutionProvider::PrintProviderOptions() const { - LOGS(*GetLogger(), INFO) << "TVM EP options:\n" << - "executor type: " << info_.executor << "\n" << - "target: " << info_.target << "\n" << - "target_host: " << info_.target_host << "\n" << - "opt level: " << info_.opt_level << "\n" << - "freeze weights: " << info_.freeze_weights << "\n" << - "tuning file path: " << info_.tuning_file_path << "\n" << - "tuning type: " << info_.tuning_type << "\n" << - "convert layout to NHWC: " << info_.to_nhwc << "\n" << - "input tensor names: " << info_.input_names_str << "\n" << - "input tensor shapes: " << info_.input_shapes_str; -} - -int TvmExecutionProvider::CreateStateFunc(ComputeContext* context, FunctionState* state) { +int TvmExecutionProvider::createStateFunc(ComputeContext* context, FunctionState* state) { auto* state_ptr = new TVMFuncState(); *state_ptr = {context->allocate_func, - context->release_func, - context->allocator_handle, - nullptr, - std::bind(&TvmExecutionProvider::CompileFunc, - this, - std::placeholders::_1, - std::placeholders::_2)}; + context->release_func, + context->allocator_handle, + compilers_[context->node_name]}; + // TODO(vvchernov): Who and when release state? *state = state_ptr; return 0; } -TvmModule* TvmExecutionProvider::CompileFunc(std::string func_name, - const TVMTensorShapes& input_shapes) { - if (modules_.count(func_name)) { - return modules_[func_name].get(); - } - - TvmModule mod_f = tvm::TVMCompile(buffers_[func_name], - model_paths_[func_name], - info_.executor, - info_.target, - info_.target_host, - info_.opt_level, - opsets_[func_name], - info_.freeze_weights, - input_shapes, - info_.to_nhwc, - info_.tuning_file_path, - info_.tuning_type); - auto module_ptr = std::make_shared(); - *module_ptr = mod_f; - modules_[func_name] = module_ptr; - // Release memory after module generation - buffers_.erase(func_name); - opsets_.erase(func_name); - return modules_[func_name].get(); -} - +} // namespace tvm } // namespace onnxruntime diff --git a/onnxruntime/core/providers/tvm/tvm_execution_provider.h b/onnxruntime/core/providers/tvm/tvm_execution_provider.h index 6a5d2a2b4c..9d891ee292 100644 --- a/onnxruntime/core/providers/tvm/tvm_execution_provider.h +++ b/onnxruntime/core/providers/tvm/tvm_execution_provider.h @@ -13,28 +13,27 @@ #include "core/framework/execution_provider.h" #include "core/platform/ort_mutex.h" -#include "tvm_common.h" -#include "tvm_execution_provider_info.h" +#include "tvm_compiler.h" +#include "tvm_runner.h" + namespace onnxruntime { - + class Graph; + class NodeArg; namespace tvm { + namespace env_vars { static const std::string kDumpSubgraphs = "ORT_TVM_DUMP_SUBGRAPHS"; } // namespace env_vars -} // namespace tvm - -class TVMRunner; class TvmExecutionProvider : public IExecutionProvider { - friend TVMRunner; + using Compiler = tvm::TVMCompiler; + using Compilers = std::unordered_map>; + using Runner = tvm::TVMRunner; + using Runners = std::unordered_map>; - using TVMTensorShape = std::vector; - using TVMTensorShapes = std::vector; - using TVMRunners = std::unordered_map>; - using TVMModules = std::unordered_map>; public: - explicit TvmExecutionProvider(const TvmExecutionProviderInfo& info); + explicit TvmExecutionProvider(const TvmEPOptions& options); virtual ~TvmExecutionProvider(); std::vector> @@ -47,27 +46,27 @@ class TvmExecutionProvider : public IExecutionProvider { AllocatorPtr GetAllocator(int id, OrtMemType mem_type) const override; private: - bool GPUTargetCheck() const; - size_t split(const std::string &txt, std::vector &strs, char ch) const; - void ProcessInfo(); - void ProcessCPUTarget(); - void ProcessGPUTarget(); - void PrintProviderOptions() const; - // Bindings for compute info - int CreateStateFunc(ComputeContext*, FunctionState*); - TvmModule* CompileFunc(std::string func_name, const TVMTensorShapes& input_shapes); + void printOptions(); + std::shared_ptr compileModel(const std::string& func_name, + const Graph& graph, + InputsInfoMap& inputs_info); + void setInputShapesForFreezedNN(const Graph& graph, TVMTensorShapes& input_shapes, InputsInfoMap& all_input_shapes); + void setInputShapesForUnfreezedNN(const Graph& graph, TVMTensorShapes& input_shapes, InputsInfoMap& all_input_shapes); + TensorShapeVector getInputShape(const NodeArg* node); + TensorShapeVector convertTensorShape(const ONNX_NAMESPACE::TensorShapeProto& shape_proto); + void prepareOutputTensors(const std::shared_ptr& mod, std::vector& output_tensors, size_t num); + NodeComputeInfo prepareComputeInfo(const std::string& func_name); + int createStateFunc(ComputeContext*, FunctionState*); private: - TVMRunners runners_; - std::unordered_map buffers_; - std::unordered_map opsets_; - std::unordered_map model_paths_; + TvmEPOptions options_; + Compilers compilers_; + Runners runners_; bool dump_subgraphs_ = false; OrtMutex tvm_mu_; AllocatorPtr allocator_; - TvmExecutionProviderInfo info_; - TVMModules modules_; }; +} // namespace tvm } // namespace onnxruntime #endif // TVM_EXECUTION_PROVIDER_H diff --git a/onnxruntime/core/providers/tvm/tvm_execution_provider_info.cc b/onnxruntime/core/providers/tvm/tvm_execution_provider_info.cc deleted file mode 100644 index 2bb2d1d692..0000000000 --- a/onnxruntime/core/providers/tvm/tvm_execution_provider_info.cc +++ /dev/null @@ -1,111 +0,0 @@ -// Copyright (c) Microsoft Corporation. All rights reserved. -// Licensed under the MIT License. - -#include -#include - -#include "core/common/common.h" -#include "core/framework/provider_options_utils.h" - -#include "tvm_execution_provider_info.h" - - -namespace onnxruntime { -namespace tvm { -namespace provider_option_names { -constexpr const char* kExecutor = "executor"; -constexpr const char* kTarget = "target"; -constexpr const char* kTargetHost = "target_host"; -constexpr const char* kOptLevel = "opt_level"; -constexpr const char* kFreezeWeights = "freeze_weights"; -constexpr const char* kToNHWC = "to_nhwc"; -constexpr const char* kTuningFilePath = "tuning_file_path"; -constexpr const char* kTuningType = "tuning_type"; -constexpr const char* kInputNames = "input_names"; -constexpr const char* kInputShapes = "input_shapes"; - -static const std::unordered_set valid_keys { - std::string{kExecutor}, - std::string{kTarget}, - std::string{kTargetHost}, - std::string{kOptLevel}, - std::string{kFreezeWeights}, - std::string{kToNHWC}, - std::string{kTuningFilePath}, - std::string{kTuningType}, - std::string{kInputNames}, - std::string{kInputShapes} -}; - -} // namespace provider_option_names -} // namespace tvm - -std::string TvmExecutionProviderInfo::whitespace_trimming(const std::string& str) { - const std::string WHITESPACE = " \n\r\t\f\v"; - size_t start = str.find_first_not_of(WHITESPACE); - if (start == std::string::npos) { - return ""; - } else { - size_t end = str.find_last_not_of(WHITESPACE); - ORT_ENFORCE(end != std::string::npos); - return str.substr(start, end + 1); - } -} - -TvmExecutionProviderInfo TvmExecutionProviderInfo::FromProviderOptions(const ProviderOptions& options) { - TvmExecutionProviderInfo info{}; - - ORT_THROW_IF_ERROR( - ProviderOptionsParser{} - .AddAssignmentToReference(tvm::provider_option_names::kExecutor, info.executor) - .AddAssignmentToReference(tvm::provider_option_names::kTarget, info.target) - .AddAssignmentToReference(tvm::provider_option_names::kTargetHost, info.target_host) - .AddAssignmentToReference(tvm::provider_option_names::kOptLevel, info.opt_level) - .AddAssignmentToReference(tvm::provider_option_names::kFreezeWeights, info.freeze_weights) - .AddAssignmentToReference(tvm::provider_option_names::kToNHWC, info.to_nhwc) - .AddAssignmentToReference(tvm::provider_option_names::kTuningFilePath, info.tuning_file_path) - .AddAssignmentToReference(tvm::provider_option_names::kTuningType, info.tuning_type) - .AddAssignmentToReference(tvm::provider_option_names::kInputNames, info.input_names_str) - .AddAssignmentToReference(tvm::provider_option_names::kInputShapes, info.input_shapes_str) - .Parse(options)); - - return info; -} - -TvmExecutionProviderInfo TvmExecutionProviderInfo::FromOptionsString(const char* opt_str) { - std::string settings{opt_str}; - ProviderOptions options; - if (!settings.empty()) { - const std::string& str = settings; - - // tokenize settings - std::regex reg("\\s*,\\s*"); - std::sregex_token_iterator iter(str.begin(), str.end(), reg, -1); - std::sregex_token_iterator iter_end; - std::vector pairs(iter, iter_end); - - ORT_ENFORCE(pairs.size() > 0); - - for(const auto& pair : pairs) { - auto pos_colon = pair.find(':'); - ORT_ENFORCE(pos_colon != std::string::npos, "Invalid key value pair."); - std::string key = pair.substr(0, pos_colon); - std::string value = pair.substr(pos_colon + 1); - - // trim leading and trailing spaces from key/value - key = whitespace_trimming(key); - value = whitespace_trimming(value); - - // Check keys of obtained options - if (tvm::provider_option_names::valid_keys.count(key) == 0) { - ORT_NOT_IMPLEMENTED("TvmOptions: unknown option (", key, ")"); - } - - options[key] = value; - } - } - - return FromProviderOptions(options); -} - -} // namespace onnxruntime diff --git a/onnxruntime/core/providers/tvm/tvm_provider_factory.cc b/onnxruntime/core/providers/tvm/tvm_provider_factory.cc index b63077e3b3..bcfeb637bd 100644 --- a/onnxruntime/core/providers/tvm/tvm_provider_factory.cc +++ b/onnxruntime/core/providers/tvm/tvm_provider_factory.cc @@ -13,32 +13,31 @@ namespace onnxruntime { struct TvmProviderFactory : IExecutionProviderFactory { - TvmProviderFactory(const TvmExecutionProviderInfo& info) : info_{info} {} + TvmProviderFactory(const tvm::TvmEPOptions& options) : options_{options} {} ~TvmProviderFactory() = default; std::unique_ptr CreateProvider() override { - return std::make_unique(info_); + return std::make_unique(options_); } - private: - TvmExecutionProviderInfo info_; +private: + tvm::TvmEPOptions options_; }; -std::shared_ptr CreateExecutionProviderFactory_Tvm(const char* settings) { - TvmExecutionProviderInfo info = TvmExecutionProviderInfo::FromOptionsString(settings); - return std::make_shared(info); +std::shared_ptr CreateExecutionProviderFactory_Tvm(const char* opt_str) { + tvm::TvmEPOptions options = tvm::TvmEPOptionsHelper::FromOptionsString(opt_str); + return std::make_shared(options); } -std::shared_ptr CreateExecutionProviderFactory_Tvm(const TvmExecutionProviderInfo& info) -{ - return std::make_shared(info); +std::shared_ptr CreateExecutionProviderFactory_Tvm(const tvm::TvmEPOptions& options) { + return std::make_shared(options); } } // namespace onnxruntime ORT_API_STATUS_IMPL(OrtSessionOptionsAppendExecutionProvider_Tvm, _In_ OrtSessionOptions* options, - _In_ const char* settings) { - onnxruntime::TvmExecutionProviderInfo info = onnxruntime::TvmExecutionProviderInfo::FromOptionsString(settings); - options->provider_factories.push_back(onnxruntime::CreateExecutionProviderFactory_Tvm(info)); + _In_ const char* opt_str) { + onnxruntime::tvm::TvmEPOptions tvm_options = onnxruntime::tvm::TvmEPOptionsHelper::FromOptionsString(opt_str); + options->provider_factories.push_back(onnxruntime::CreateExecutionProviderFactory_Tvm(tvm_options)); return nullptr; } diff --git a/onnxruntime/core/providers/tvm/tvm_runner.cc b/onnxruntime/core/providers/tvm/tvm_runner.cc new file mode 100644 index 0000000000..117ecea680 --- /dev/null +++ b/onnxruntime/core/providers/tvm/tvm_runner.cc @@ -0,0 +1,26 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/graph/model.h" +#include "core/framework/tensorprotoutils.h" + +#include "tvm_runner.h" + + +using namespace ONNX_NAMESPACE; +namespace onnxruntime { +namespace tvm { + +TVMRunner::TVMRunner(const TvmEPOptions& options, + const std::shared_ptr& mod, + const InputsInfoMap& inputs_info, + const std::vector& output_tensors) { + runner_ = getTVMRunnerImpl(mod, options, inputs_info, output_tensors); +} + +common::Status TVMRunner::operator()(FunctionState state, const OrtCustomOpApi* api, OrtKernelContext* context) { + return runner_->run(api, context); +} + +} // namespace tvm +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/tvm/tvm_runner.h b/onnxruntime/core/providers/tvm/tvm_runner.h new file mode 100644 index 0000000000..85d37ccec1 --- /dev/null +++ b/onnxruntime/core/providers/tvm/tvm_runner.h @@ -0,0 +1,35 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#ifndef TVM_RUNNER_H +#define TVM_RUNNER_H + +#include +#include + +#include "tvm_runner_impl.h" + + +namespace onnxruntime { +namespace tvm { + +class TVMRunner { +public: + TVMRunner() = delete; + virtual ~TVMRunner() = default; + + TVMRunner(const TvmEPOptions& options, + const std::shared_ptr& mod, + const InputsInfoMap& inputs_info, + const std::vector& output_tensor); + + common::Status operator()(FunctionState state, const OrtCustomOpApi* api, OrtKernelContext* context); + +private: + std::shared_ptr runner_; +}; + +} // namespace tvm +} // namespace onnxruntime + +#endif // TVM_TVM_RUNNER_H diff --git a/onnxruntime/core/providers/tvm/tvm_runner_impl.cc b/onnxruntime/core/providers/tvm/tvm_runner_impl.cc new file mode 100644 index 0000000000..bade84b680 --- /dev/null +++ b/onnxruntime/core/providers/tvm/tvm_runner_impl.cc @@ -0,0 +1,165 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#include "core/framework/tensorprotoutils.h" + +#include "tvm_runner_impl.h" +#include "tvm_utils.h" +#include "tvm_api.h" + + +namespace onnxruntime { +namespace tvm { + +/* ------------------------------------ RunnerImplFactory ----------------------------- */ + +std::shared_ptr getTVMRunnerImpl(const std::shared_ptr& mod, + const TvmEPOptions& options, + const InputsInfoMap& inputs_info, + const std::vector output_tensors) { + const std::string& name = options.executor; + if (name == "graph") { + return std::make_shared(mod, inputs_info, options.output_shapes, output_tensors); + } else if (name == "vm") { + return std::make_shared(mod, inputs_info, options.output_shapes, output_tensors); + } + return nullptr; +} + +/* ------------------------------------ RunnerImpl ------------------------------------ */ + +RunnerImpl::RunnerImpl(const std::shared_ptr& mod, + const InputsInfoMap& inputs_info, + const TVMTensorShapes output_shapes, + const std::vector output_tensors) : + mod_(mod), + inputs_info_(inputs_info), + output_shapes_(output_shapes), + output_tensors_(output_tensors) { +} + +void RunnerImpl::convert_input_tensors2dl_tensors(Ort::CustomOpApi& ort, + OrtKernelContext* context, + std::vector& dst, + std::vector& dst_inds) { + size_t num = inputs_info_.size(); + dst.reserve(num); + dst_inds.reserve(num); + for (auto& info : inputs_info_) { + // TODO(vvchernov): decomposition declaration only available with -std=c++1z or -std=gnu++1z + auto& i = info.first; + auto& shape = info.second; + const OrtValue* input_tensor = ort.KernelContext_GetInput(context, i); + ORT_ENFORCE(input_tensor->IsTensor()); + const Tensor& tensor = input_tensor->Get(); + const OrtDevice& device = tensor.Location().device; + auto tensor_info = ort.GetTensorTypeAndShape(input_tensor); + auto tensor_type = ort.GetTensorElementType(tensor_info); + ort.ReleaseTensorTypeAndShapeInfo(tensor_info); + + DLTensor t; + t.device = GetDLDevice(device); + t.dtype = GetDataType(tensor_type); + t.strides = nullptr; + t.byte_offset = 0; + t.data = const_cast(ort.GetTensorData(input_tensor)); + t.ndim = shape.size(); + t.shape = shape.data(); + dst.emplace_back(t); + dst_inds.push_back(i); + } +} + +void RunnerImpl::add_device_type_data2output_tensors(Ort::CustomOpApi& ort, + OrtKernelContext* context) { + size_t num_outputs = output_tensors_.size(); + for (auto i = 0u; i < num_outputs; i++) { + //setup output tensor property + OrtValue* output_tensor = ort.KernelContext_GetOutput(context, + i, + output_shapes_[i].data(), + output_shapes_[i].size()); + ORT_ENFORCE(output_tensor->IsTensor()); + const Tensor& tensor = output_tensor->Get(); + const OrtDevice& device = tensor.Location().device; + auto tensor_info = ort.GetTensorTypeAndShape(output_tensor); + auto tensor_type = ort.GetTensorElementType(tensor_info); + ort.ReleaseTensorTypeAndShapeInfo(tensor_info); + + output_tensors_[i].device = GetDLDevice(device); + output_tensors_[i].dtype = GetDataType(tensor_type); + output_tensors_[i].data = ort.GetTensorMutableData(output_tensor); + } +} + +/* ------------------------------------ GERunnerImpl ------------------------------------ */ + +GERunnerImpl::GERunnerImpl(const std::shared_ptr& mod, + const InputsInfoMap& inputs_info, + const TVMTensorShapes output_shapes, + const std::vector output_tensors) : + RunnerImpl(mod, inputs_info, output_shapes, output_tensors) { +} + +void GERunnerImpl::set_input(Ort::CustomOpApi& ort, OrtKernelContext* context) { + std::vector inds; + std::vector dl_tensors_inputs; + convert_input_tensors2dl_tensors(ort, context, dl_tensors_inputs, inds); + + tvm::TVMSetInputs(*mod_, inds, dl_tensors_inputs); +} + +void GERunnerImpl::connect_output_tensors2ort(Ort::CustomOpApi& ort, OrtKernelContext* context) { + add_device_type_data2output_tensors(ort, context); +} + +void GERunnerImpl::run_and_get_output() { + tvm::TVMRun(*mod_); + tvm::TVMGetOutputs(*mod_, output_tensors_); +} + +/* ------------------------------------ VMRunnerImpl ------------------------------------ */ + +VMRunnerImpl::VMRunnerImpl(const std::shared_ptr& mod, + const InputsInfoMap& inputs_info, + const TVMTensorShapes output_shapes, + const std::vector output_tensors) : + RunnerImpl(mod, inputs_info, output_shapes, output_tensors) { +} + +void VMRunnerImpl::set_input(Ort::CustomOpApi& ort, OrtKernelContext* context) { + std::vector inds; + std::vector dl_tensors_inputs; + convert_input_tensors2dl_tensors(ort, context, dl_tensors_inputs, inds); + + tvm::TVM_VM_SetInputs(*mod_, inds, dl_tensors_inputs); +} + +void VMRunnerImpl::connect_output_tensors2ort(Ort::CustomOpApi& ort, OrtKernelContext* context) { + if(!probe_infer_) { + infer_once_to_get_output_shapes(); + } + + add_device_type_data2output_tensors(ort, context); +} + +void VMRunnerImpl::run_and_get_output() { + tvm::TVM_VM_Run(*mod_); + tvm::TVM_VM_GetOutputs(*mod_, output_tensors_); +} + +void VMRunnerImpl::infer_once_to_get_output_shapes() { + tvm::TVM_VM_Run(*mod_); + size_t num_outputs = output_tensors_.size(); + // TODO(vvchernov): check it + output_shapes_.resize(num_outputs); + tvm::TVMGetOutputShapes(*mod_, output_shapes_); + for (size_t i = 0; i < num_outputs; ++i) { + output_tensors_[i].ndim = output_shapes_[i].size(); + output_tensors_[i].shape = output_shapes_[i].data(); + } + probe_infer_ = true; +} + +} // namespace tvm +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/tvm/tvm_runner_impl.h b/onnxruntime/core/providers/tvm/tvm_runner_impl.h new file mode 100644 index 0000000000..e9104859c7 --- /dev/null +++ b/onnxruntime/core/providers/tvm/tvm_runner_impl.h @@ -0,0 +1,104 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +// Licensed under the MIT License. + +#ifndef TVM_RUNNER_IMPL_H +#define TVM_RUNNER_IMPL_H + +#include +#include +#include + +#include "core/framework/func_api.h" +#include "core/session/onnxruntime_cxx_api.h" + +#include "tvm_common.h" +#include "tvm_ep_options.h" + + +namespace onnxruntime { +namespace tvm { + +class RunnerImpl { +public: + RunnerImpl() = delete; + RunnerImpl(const std::shared_ptr& mod, + const InputsInfoMap& inputs_info, + const TVMTensorShapes output_shapes, + const std::vector tensors_outputs); + virtual ~RunnerImpl() = default; + + virtual common::Status run(const OrtCustomOpApi* api, OrtKernelContext* context) { + Ort::CustomOpApi ort{*api}; + + set_input(ort, context); + connect_output_tensors2ort(ort, context); + run_and_get_output(); + + return Status::OK(); + } + + virtual void set_input(Ort::CustomOpApi& ort, OrtKernelContext* context) = 0; + virtual void connect_output_tensors2ort(Ort::CustomOpApi& ort, OrtKernelContext* context) = 0; + virtual void run_and_get_output() = 0; + +protected: + void convert_input_tensors2dl_tensors(Ort::CustomOpApi& ort, + OrtKernelContext* context, + std::vector& dst, + std::vector& dst_inds); + void add_device_type_data2output_tensors(Ort::CustomOpApi& ort, + OrtKernelContext* context); + +protected: + std::shared_ptr mod_; + InputsInfoMap inputs_info_; + TVMTensorShapes output_shapes_; + std::vector output_tensors_; +}; + + +class GERunnerImpl : public RunnerImpl { +public: + GERunnerImpl() = delete; + GERunnerImpl(const std::shared_ptr& mod, + const InputsInfoMap& inputs_info, + const TVMTensorShapes output_shapes, + const std::vector tensors_outputs); + virtual ~GERunnerImpl() = default; + + virtual void set_input(Ort::CustomOpApi& ort, OrtKernelContext* context) override final; + virtual void connect_output_tensors2ort(Ort::CustomOpApi& ort, OrtKernelContext* context) override final; + virtual void run_and_get_output() override final; +}; + + +class VMRunnerImpl : public RunnerImpl { +public: + VMRunnerImpl() = delete; + VMRunnerImpl(const std::shared_ptr& mod, + const InputsInfoMap& inputs_info, + const TVMTensorShapes output_shapes, + const std::vector tensors_outputs); + virtual ~VMRunnerImpl() = default; + + virtual void set_input(Ort::CustomOpApi& ort, OrtKernelContext* context) override final; + virtual void connect_output_tensors2ort(Ort::CustomOpApi& ort, OrtKernelContext* context) override final; + virtual void run_and_get_output() override final; + +private: + void infer_once_to_get_output_shapes(); + +private: + bool probe_infer_ = false; +}; + + +std::shared_ptr getTVMRunnerImpl(const std::shared_ptr& mod, + const TvmEPOptions& options, + const InputsInfoMap& inputs_info, + const std::vector output_tensors); + +} // namespace tvm +} // namespace onnxruntime + +#endif // TVM_TVM_RUNNER_IMPL_H diff --git a/onnxruntime/core/providers/tvm/tvm_utils.h b/onnxruntime/core/providers/tvm/tvm_utils.h index ab0e8da565..9471afb135 100644 --- a/onnxruntime/core/providers/tvm/tvm_utils.h +++ b/onnxruntime/core/providers/tvm/tvm_utils.h @@ -10,7 +10,9 @@ #include "core/framework/ortdevice.h" #include "core/common/common.h" + namespace onnxruntime { +namespace tvm { inline DLDataType GetDataType(ONNXTensorElementDataType type) { if (type == ONNX_TENSOR_ELEMENT_DATA_TYPE_DOUBLE) { @@ -50,6 +52,7 @@ inline DLDevice GetDLDevice(const OrtDevice& device) { return context; } -} // namespace onnxruntime +} // namespace tvm +} // namespace onnxruntime #endif // TVM_UTILS_H diff --git a/onnxruntime/core/providers/tvm/xpu_data_transfer.cc b/onnxruntime/core/providers/tvm/xpu_data_transfer.cc index 4efb171dda..5247382566 100644 --- a/onnxruntime/core/providers/tvm/xpu_data_transfer.cc +++ b/onnxruntime/core/providers/tvm/xpu_data_transfer.cc @@ -6,7 +6,10 @@ #include "xpu_data_transfer.h" #include "tvm_utils.h" + namespace onnxruntime { +namespace tvm { + XPUDataTransfer::XPUDataTransfer() { } @@ -14,8 +17,8 @@ XPUDataTransfer::~XPUDataTransfer() { } bool XPUDataTransfer::CanCopy(const OrtDevice& src_device, const OrtDevice& dst_device) const { - return (src_device.Type() == OrtDevice::CPU && dst_device.Type() == OrtDevice::CPU) || - (src_device.Type() == OrtDevice::GPU || dst_device.Type() == OrtDevice::GPU); + return (src_device.Type() == OrtDevice::CPU && dst_device.Type() == OrtDevice::CPU) || + (src_device.Type() == OrtDevice::GPU || dst_device.Type() == OrtDevice::GPU); } common::Status XPUDataTransfer::CopyTensor(const Tensor& src, Tensor& dst, int _exec_queue_id) const { @@ -27,11 +30,11 @@ common::Status XPUDataTransfer::CopyTensor(const Tensor& src, Tensor& dst, int _ const OrtDevice& dst_device = dst.Location().device; if ((src_device.Type() == OrtDevice::CPU) && (dst_device.Type() == OrtDevice::CPU)) { - if (src_data == dst_data) { - // no need copying as both pointers are referring to same piece of memory. - return Status::OK(); - } - memcpy(dst_data, src_data, bytes); + if (src_data == dst_data) { + // no need copying as both pointers are referring to same piece of memory. + return Status::OK(); + } + memcpy(dst_data, src_data, bytes); } else { DLTensor tvm_src, tvm_dst; DLDataType dl_type{kDLInt, 8, 1}; @@ -80,4 +83,5 @@ common::Status TvmCPUDataTransfer::CopyTensor(const Tensor& src, Tensor& dst, in return Status::OK(); } -} // namespace onnxruntime +} // namespace tvm +} // namespace onnxruntime diff --git a/onnxruntime/core/providers/tvm/xpu_data_transfer.h b/onnxruntime/core/providers/tvm/xpu_data_transfer.h index f07c117943..0b38f71baa 100644 --- a/onnxruntime/core/providers/tvm/xpu_data_transfer.h +++ b/onnxruntime/core/providers/tvm/xpu_data_transfer.h @@ -7,10 +7,12 @@ #include "core/framework/data_transfer.h" #include "tvm_common.h" + namespace onnxruntime { +namespace tvm { class XPUDataTransfer : public IDataTransfer { - public: +public: XPUDataTransfer(); ~XPUDataTransfer(); @@ -23,7 +25,7 @@ class XPUDataTransfer : public IDataTransfer { }; class TvmCPUDataTransfer : public IDataTransfer { - public: +public: TvmCPUDataTransfer() = default; // Dampen MSVC warning about not fully overriding CopyTensor using IDataTransfer::CopyTensor; @@ -31,5 +33,7 @@ class TvmCPUDataTransfer : public IDataTransfer { common::Status CopyTensor(const Tensor& src, Tensor& dst, int exec_queue_id) const override; }; -} // namespace onnxruntime +} // namespace tvm +} // namespace onnxruntime + #endif // XPU_DATA_TRANSFER diff --git a/onnxruntime/python/onnxruntime_pybind_state.cc b/onnxruntime/python/onnxruntime_pybind_state.cc index ff4b4e697f..fe1a1f0dc5 100644 --- a/onnxruntime/python/onnxruntime_pybind_state.cc +++ b/onnxruntime/python/onnxruntime_pybind_state.cc @@ -649,10 +649,10 @@ std::unique_ptr CreateExecutionProviderInstance( #endif } else if (type == kTvmExecutionProvider) { #if USE_TVM - onnxruntime::TvmExecutionProviderInfo info{}; + onnxruntime::tvm::TvmEPOptions info{}; const auto it = provider_options_map.find(type); if (it != provider_options_map.end()) { - info = onnxruntime::TvmExecutionProviderInfo::FromProviderOptions(it->second); + info = onnxruntime::tvm::TvmEPOptionsHelper::FromProviderOptions(it->second); } return onnxruntime::CreateExecutionProviderFactory_Tvm(info)->CreateProvider(); diff --git a/onnxruntime/python/onnxruntime_pybind_state_common.h b/onnxruntime/python/onnxruntime_pybind_state_common.h index 093fced91c..225f31a1d1 100644 --- a/onnxruntime/python/onnxruntime_pybind_state_common.h +++ b/onnxruntime/python/onnxruntime_pybind_state_common.h @@ -157,7 +157,7 @@ extern std::string nuphar_settings; } // namespace onnxruntime #endif #ifdef USE_TVM -#include "core/providers/tvm/tvm_execution_provider_info.h" +#include "core/providers/tvm/tvm_ep_options.h" #endif #ifdef USE_VITISAI #include "core/providers/vitisai/vitisai_provider_factory.h" @@ -484,7 +484,7 @@ std::shared_ptr CreateExecutionProviderFactory_Dnnl(i std::shared_ptr CreateExecutionProviderFactory_OpenVINO(const OrtOpenVINOProviderOptions* params); std::shared_ptr CreateExecutionProviderFactory_Nuphar(bool, const char*); #ifdef USE_TVM -std::shared_ptr CreateExecutionProviderFactory_Tvm(const TvmExecutionProviderInfo& info); +std::shared_ptr CreateExecutionProviderFactory_Tvm(const tvm::TvmEPOptions& info); std::shared_ptr CreateExecutionProviderFactory_Tvm(const char* params); #endif std::shared_ptr CreateExecutionProviderFactory_VITISAI(const char* backend_type, int device_id,