diff --git a/onnxruntime/core/providers/cuda/math/topk.cc b/onnxruntime/core/providers/cuda/math/topk.cc index 8722513ee9..6b3794c367 100644 --- a/onnxruntime/core/providers/cuda/math/topk.cc +++ b/onnxruntime/core/providers/cuda/math/topk.cc @@ -75,7 +75,7 @@ Status TopK::ComputeInternal(OpKernelContext* ctx) const { auto elem_nums = tensor_X->Shape().AsShapeVector(); auto dimension = elem_nums[axis]; - for (auto i = static_cast(elem_nums.size()) - 2; i >= 0; --i) { + for (auto i = static_cast(elem_nums.size()) - 2; i >= 0; --i) { elem_nums[i] *= elem_nums[i + 1]; } diff --git a/onnxruntime/core/session/standalone_op_invoker.cc b/onnxruntime/core/session/standalone_op_invoker.cc index a30cb915e1..59678f84ee 100644 --- a/onnxruntime/core/session/standalone_op_invoker.cc +++ b/onnxruntime/core/session/standalone_op_invoker.cc @@ -7,6 +7,11 @@ #include "core/session/ort_apis.h" #include +#if defined(_MSC_VER) && !defined(__clang__) +//disabling warning on calling of raw "delete" operator +#pragma warning(disable : 26400) +#endif + #ifdef ORT_MINIMAL_BUILD ORT_API_STATUS_IMPL(OrtApis::CreateOpAttr, @@ -72,20 +77,13 @@ ORT_API(void, OrtApis::ReleaseKernelInfo, _Frees_ptr_opt_ OrtKernelInfo*) { namespace onnxruntime { namespace standalone { -void ReleaseNode(onnxruntime::Node* node) { - if (node) { - for (auto* input_arg : node->InputDefs()) { - delete input_arg; - } - for (auto* output_arg : node->OutputDefs()) { - delete output_arg; - } - delete node; - } -} +using NodePtr = std::unique_ptr; -using NodePtr = std::unique_ptr; -using StandAloneNodesMap = InlinedHashMap; +using ArgPtr = std::unique_ptr; +using ArgPtrs = onnxruntime::InlinedVector; + +using NodeResource = std::pair; +using NodeResourceMap = InlinedHashMap; class NodeRepo { public: @@ -94,13 +92,12 @@ class NodeRepo { return node_repo; } - onnxruntime::Status AddNode(const onnxruntime::OpKernel* kernel, NodePtr node_ptr) { + onnxruntime::Status AddNode(const onnxruntime::OpKernel* kernel, NodePtr&& node_ptr, ArgPtrs&& args) { std::lock_guard guard(mutex_); - auto iter = node_map_.find(kernel); - if (iter != node_map_.end()) { - return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "kernel already mapped to existing node"); + auto ret = resource_map_.try_emplace(kernel, NodeResource{std::move(node_ptr), std::move(args)}); + if (!ret.second) { + return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "kernel already mapped to existing node"); } - node_map_.insert({kernel, std::move(node_ptr)}); return Status::OK(); } @@ -111,11 +108,11 @@ class NodeRepo { size_t expect_output_count{}; { std::lock_guard guard(mutex_); - auto iter = node_map_.find(kernel); - if (iter == node_map_.end()) { + auto iter = resource_map_.find(kernel); + if (iter == resource_map_.end()) { return ORT_MAKE_STATUS(ONNXRUNTIME, INVALID_ARGUMENT, "matching node is missing"); } - auto* node = iter->second.get(); + auto* node = iter->second.first.get(); expect_input_count = node->InputDefs().size(); expect_output_count = node->OutputDefs().size(); } @@ -134,7 +131,7 @@ class NodeRepo { void RemoveNode(const onnxruntime::OpKernel* kernel) { std::lock_guard guard(mutex_); - node_map_.erase(kernel); + resource_map_.erase(kernel); } private: @@ -142,7 +139,7 @@ class NodeRepo { ~NodeRepo() = default; std::mutex mutex_; - StandAloneNodesMap node_map_; + NodeResourceMap resource_map_; }; // For invoking kernels without a graph @@ -377,18 +374,21 @@ onnxruntime::Status CreateOp(const OrtKernelInfo* info, &kernel_create_info); ORT_RETURN_IF_ERROR(status); + ArgPtrs arg_ptrs; std::vector input_args; - for (int i = 0; i < input_count; ++i) { - auto arg_ptr = std::make_unique(std::to_string(i), nullptr); - input_args.push_back(arg_ptr.release()); - } std::vector output_args; - for (int i = 0; i < output_count; ++i) { - auto arg_ptr = std::make_unique(std::to_string(i), nullptr); - output_args.push_back(arg_ptr.release()); + + for (int i = 0; i < input_count; ++i) { + arg_ptrs.push_back(std::make_unique(std::to_string(i), nullptr)); + input_args.push_back(arg_ptrs.back().get()); } - auto tmp_node_holder = std::make_unique(std::string("standalone_") + op_name, op_name, "", input_args, output_args, nullptr, domain); - NodePtr node_ptr(tmp_node_holder.release(), ReleaseNode); + + for (int i = 0; i < output_count; ++i) { + arg_ptrs.push_back(std::make_unique(std::to_string(i), nullptr)); + output_args.push_back(arg_ptrs.back().get()); + } + + NodePtr node_ptr = std::make_unique(std::string("standalone_") + op_name, op_name, "", input_args, output_args, nullptr, domain); for (int i = 0; i < attr_count; ++i) { auto attr_proto = reinterpret_cast(attr_values[i]); node_ptr->AddAttributeProto(*attr_proto); @@ -409,7 +409,7 @@ onnxruntime::Status CreateOp(const OrtKernelInfo* info, static FuncManager kFuncMgr; status = kernel_create_info->kernel_create_func(kFuncMgr, tmp_kernel_info, op_kernel); ORT_RETURN_IF_ERROR(status); - status = NodeRepo::GetInstance().AddNode(op_kernel.get(), std::move(node_ptr)); + status = NodeRepo::GetInstance().AddNode(op_kernel.get(), std::move(node_ptr), std::move(arg_ptrs)); ORT_RETURN_IF_ERROR(status); *op = reinterpret_cast(op_kernel.release()); return status;