From 9f3ae4279f7fd0c6af9e6607e8a94964f697a12a Mon Sep 17 00:00:00 2001 From: Scott McKay Date: Fri, 18 Jan 2019 04:51:23 +1000 Subject: [PATCH] Handle copy to/from non-CPU devices across control flow nodes (#339) --- include/onnxruntime/core/graph/graph.h | 10 +- onnxruntime/core/framework/session_state.cc | 45 ++- onnxruntime/core/framework/session_state.h | 2 +- .../framework/session_state_initializer.cc | 77 +++- .../framework/session_state_initializer.h | 6 +- .../core/framework/transformer_memcpy.cc | 7 +- onnxruntime/core/framework/utils.cc | 349 +++++++++++++++++- onnxruntime/core/framework/utils.h | 53 ++- onnxruntime/core/graph/graph.cc | 2 +- .../core/providers/cpu/controlflow/if.cc | 13 +- .../core/providers/cpu/controlflow/loop.cc | 18 +- .../providers/cpu/controlflow/scan_utils.cc | 9 +- onnxruntime/core/session/IOBinding.cc | 90 +---- onnxruntime/core/session/IOBinding.h | 5 - onnxruntime/core/session/inference_session.cc | 244 +----------- .../test/providers/cpu/controlflow/if_test.cc | 26 +- .../providers/cpu/controlflow/loop_test.cc | 27 +- .../providers/cpu/controlflow/scan_test.cc | 27 +- .../test/providers/cpu/rnn/rnn_op_test.cc | 13 +- .../test/providers/provider_test_utils.cc | 261 +++++++------ .../test/providers/provider_test_utils.h | 14 +- 21 files changed, 788 insertions(+), 510 deletions(-) diff --git a/include/onnxruntime/core/graph/graph.h b/include/onnxruntime/core/graph/graph.h index 92ec7e7a5b..15ff88f73c 100644 --- a/include/onnxruntime/core/graph/graph.h +++ b/include/onnxruntime/core/graph/graph.h @@ -265,6 +265,14 @@ class Node { */ Graph* GetMutableGraphAttribute(const std::string& attr_name); + /** Gets a map of attribute name to the mutable Graph instances for all subgraphs of the Node. + @returns Map of the attribute name that defines the subgraph to the subgraph's Graph instance. + nullptr if the Node has no subgraphs. + */ + const std::unordered_map>& GetAttributeNameToMutableSubgraphMap() { + return attr_to_subgraph_map_; + } + /** Gets the execution ProviderType that this node will be executed by. */ ProviderType GetExecutionProviderType() const noexcept; @@ -420,7 +428,7 @@ class Node { Graph* graph_; // Map of attribute name to the Graph instance created from the GraphProto attribute - std::unordered_map attr_to_subgraph_map_; + std::unordered_map> attr_to_subgraph_map_; // Graph instances for subgraphs that are owned by this Node std::vector> subgraphs_; diff --git a/onnxruntime/core/framework/session_state.cc b/onnxruntime/core/framework/session_state.cc index 4c28418d62..70618f74d4 100644 --- a/onnxruntime/core/framework/session_state.cc +++ b/onnxruntime/core/framework/session_state.cc @@ -7,6 +7,7 @@ #include "core/common/logging/logging.h" #include "core/framework/op_kernel.h" +#include "core/framework/utils.h" using namespace ::onnxruntime::common; namespace onnxruntime { @@ -109,8 +110,48 @@ bool SessionState::GetEnableMemoryPattern() const { return enable_mem_pattern_; } -void SessionState::AddInputNameToNodeInfoMapping(const std::string& input_name, const NodeInfo& node_info) { - input_names_to_nodeinfo_mapping_[input_name].push_back(node_info); +common::Status SessionState::AddInputNameToNodeInfoMapping(const std::string& input_name, const NodeInfo& node_info) { + auto status = Status::OK(); + + // in the future we could support multiple nodes on difference devices using an input, however right now + // the logic in utils::CopyOneInputAcrossDevices only checks the first entry. + // Instead of failing silently and adding extra entries that will be ignored, check if the required provider + // is the same for any duplicate entries. If it differs we can't run the model. + + auto& entries = input_names_to_nodeinfo_mapping_[input_name]; + + if (entries.empty()) { + entries.push_back(node_info); + } else { + const auto& existing_entry = entries.front(); + + // if index == max it's an entry for an implicit input to a subgraph or unused graph input. + // we want to prefer the entry for explicit usage in this graph, as the implicit usage in a + // subgraph will be handled by the subgraph's SessionState. + if (node_info.index == std::numeric_limits::max()) { + // ignore and preserve existing value + } else if (existing_entry.index == std::numeric_limits::max()) { + // replace existing entry that is for an implicit input with new entry for explicit usage in this graph + entries[0] = node_info; + } else { + // if the providers match we can add the new entry for completeness (it will be ignored in + // utils::CopyOneInputAcrossDevices though). + // if they don't, we are broken. + const auto& current_provider = utils::GetNodeInputProviderType(entries[0]); + const auto& new_provider = utils::GetNodeInputProviderType(node_info); + + if (current_provider == new_provider) { + entries.push_back(node_info); + } else { + ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, + "Using an input in multiple nodes on different devices is not supported currently. Input:", + input_name, " is used by node ", existing_entry.p_node->Name(), " (", current_provider, + ") and node ", node_info.p_node->Name(), " (", new_provider, ")."); + } + } + } + + return status; } common::Status SessionState::GetInputNodeInfo(const std::string& input_name, std::vector& node_info_vec) const { diff --git a/onnxruntime/core/framework/session_state.h b/onnxruntime/core/framework/session_state.h index a855e0741a..a380b71ce2 100644 --- a/onnxruntime/core/framework/session_state.h +++ b/onnxruntime/core/framework/session_state.h @@ -136,7 +136,7 @@ class SessionState { }; using NameNodeInfoMapType = std::unordered_map>; - void AddInputNameToNodeInfoMapping(const std::string& input_name, const NodeInfo& node_info); + common::Status AddInputNameToNodeInfoMapping(const std::string& input_name, const NodeInfo& node_info); common::Status GetInputNodeInfo(const std::string& input_name, std::vector& node_info_vec) const; const NameNodeInfoMapType& GetInputNodeInfoMap() const; diff --git a/onnxruntime/core/framework/session_state_initializer.cc b/onnxruntime/core/framework/session_state_initializer.cc index 2fbae5a81f..7f786179ea 100644 --- a/onnxruntime/core/framework/session_state_initializer.cc +++ b/onnxruntime/core/framework/session_state_initializer.cc @@ -4,6 +4,7 @@ #include "core/framework/session_state_initializer.h" #include +#include #include "core/common/common.h" #include "core/common/logging/logging.h" @@ -48,18 +49,18 @@ static common::Status SaveKernels(const ExecutionProviders& execution_providers, static common::Status SaveInputOutputNamesToNodeMapping(const onnxruntime::Graph& graph, const KernelRegistryManager& custom_registry_manager, - SessionState& session_state); + SessionState& session_state, + const std::vector* implicit_inputs); SessionStateInitializer::SessionStateInitializer(onnxruntime::Graph& graph, SessionState& session_state, const ExecutionProviders& providers, - KernelRegistryManager& kernel_registry_manager, - const logging::Logger& logger) + KernelRegistryManager& kernel_registry_manager) : graph_{graph}, session_state_{session_state}, execution_providers_{providers}, kernel_registry_manager_{kernel_registry_manager}, - logger_{logger} { + logger_{session_state.Logger()} { } common::Status SessionStateInitializer::CreatePlan(const std::vector& outer_scope_node_args, @@ -104,7 +105,8 @@ common::Status SessionStateInitializer::CreatePlan(const std::vector& } common::Status SessionStateInitializer::InitializeAndSave(bool enable_memory_pattern, - std::map& weights_buffers) { + std::map& weights_buffers, + const std::vector* implicit_inputs) { const auto* exec_plan_ptr = session_state_.GetExecutionPlan(); ORT_ENFORCE(exec_plan_ptr, "Execution plan was not found in SessionState. CreatePlan must be called first."); @@ -123,7 +125,8 @@ common::Status SessionStateInitializer::InitializeAndSave(bool enable_memory_pat graph_.CleanAllInitializedTensors(); // remove weights from the graph now to save memory ORT_RETURN_IF_ERROR(SaveKernels(execution_providers_, session_state_, kernel_registry_manager_, logger_)); - ORT_RETURN_IF_ERROR(SaveInputOutputNamesToNodeMapping(graph_, kernel_registry_manager_, session_state_)); + ORT_RETURN_IF_ERROR(SaveInputOutputNamesToNodeMapping(graph_, kernel_registry_manager_, session_state_, + implicit_inputs)); return Status::OK(); } @@ -421,8 +424,9 @@ common::Status SaveKernels(const ExecutionProviders& execution_providers, return Status::OK(); } +template // T is const NodeArg or NodeArg static bool IsArgNameInInputsOutputs(const std::string& name, - const std::vector& graph_args) { + const std::vector& graph_args) { auto it = std::find_if(std::begin(graph_args), std::end(graph_args), [&name](const onnxruntime::NodeArg* arg) { return arg->Name() == name; }); @@ -431,11 +435,20 @@ static bool IsArgNameInInputsOutputs(const std::string& name, common::Status SaveInputOutputNamesToNodeMapping(const onnxruntime::Graph& graph, const KernelRegistryManager& custom_registry_manager, - SessionState& session_state) { + SessionState& session_state, + const std::vector* implicit_inputs) { auto& graph_inputs = graph.GetInputsIncludingInitializers(); auto& graph_outputs = graph.GetOutputs(); + if (implicit_inputs && implicit_inputs->empty()) { + implicit_inputs = nullptr; + } + for (auto& node : graph.Nodes()) { + // note that KernelCreateInfo may not exist for custom kernel + const KernelCreateInfo* kci = nullptr; + custom_registry_manager.SearchKernelRegistry(node, &kci); + ORT_RETURN_IF_ERROR( onnxruntime::Node::ForEachWithIndex( node.InputDefs(), @@ -444,10 +457,6 @@ common::Status SaveInputOutputNamesToNodeMapping(const onnxruntime::Graph& graph return Status::OK(); } - // note that KernelCreateInfo may not exist for custom kernel - const KernelCreateInfo* kci = nullptr; - custom_registry_manager.SearchKernelRegistry(node, &kci); - SessionState::NodeInfo node_info(index, &node, kci); if (IsArgNameInInputsOutputs(arg.Name(), graph_inputs)) { @@ -455,6 +464,13 @@ common::Status SaveInputOutputNamesToNodeMapping(const onnxruntime::Graph& graph return Status::OK(); } + if (implicit_inputs) { + if (IsArgNameInInputsOutputs(arg.Name(), *implicit_inputs)) { + session_state.AddInputNameToNodeInfoMapping(arg.Name(), node_info); + return Status::OK(); + } + } + if (IsArgNameInInputsOutputs(arg.Name(), graph_outputs)) { session_state.AddOutputNameToNodeInfoMapping(arg.Name(), node_info); return Status::OK(); @@ -462,6 +478,43 @@ common::Status SaveInputOutputNamesToNodeMapping(const onnxruntime::Graph& graph return Status::OK(); })); + + // implicit inputs to a node could come directly from a feed, so we need to make sure they have an entry too + const auto& node_implicit_inputs = node.ImplicitInputDefs(); + if (!node_implicit_inputs.empty()) { + // nested subgraph. for now map them to this node (which will be CPU based as all the control flow nodes + // are currently CPU based and they're the only ones that have implicit inputs) as the inputs will be passed as a + // feed when executing the subgraph and need to be in the mapping. + // in the future we want to recurse and find where the implicit input is actually used to try and avoid a + // copy to/from CPU to go through the control flow nodes where possible/applicable. + // the processing for the subgraph where the implicit input is consumed will do the real check on whether any + // copy to a different device is required + SessionState::NodeInfo node_info(std::numeric_limits::max(), &node, kci); + for (const auto& input_def : node_implicit_inputs) { + session_state.AddInputNameToNodeInfoMapping(input_def->Name(), node_info); + } + } + } + + // It's possible (although assumably rare) for a graph to have inputs that aren't used. one reasonable occurrence + // is in the Loop subgraph where the value of the condition used to decide whether to continue looping is passed in. + // The condition evaluated to 'true' given the subgraph is being executed, so it's of dubious value as an input. + // Similar is the current iteration number which may or may not be needed by the Loop subgraph. + // In order to handle those, create a dummy entry in the input name to node info mapping so that + // utils::CopyOneInputAcrossDevices is happy. + + auto& input_map = session_state.GetInputNodeInfoMap(); + auto end_map = input_map.cend(); + SessionState::NodeInfo empty_node_info(std::numeric_limits::max(), nullptr, nullptr); + + for (const auto& graph_input : graph_inputs) { + const auto& name = graph_input->Name(); + if (input_map.find(name) == end_map) { + // dummy entry for an input that we didn't find a use of in the graph. warn about it in case that's a bug. + // utils::CopyOneInputAcrossDevices will use the input MLValue as is given we don't believe it's used anywhere. + LOGS(session_state.Logger(), WARNING) << "Graph input with name " << name << " is not associated with a node. "; + session_state.AddInputNameToNodeInfoMapping(name, empty_node_info); + } } return Status::OK(); diff --git a/onnxruntime/core/framework/session_state_initializer.h b/onnxruntime/core/framework/session_state_initializer.h index fa8e4b5f98..da2681596b 100644 --- a/onnxruntime/core/framework/session_state_initializer.h +++ b/onnxruntime/core/framework/session_state_initializer.h @@ -25,8 +25,7 @@ class SessionStateInitializer { SessionStateInitializer(onnxruntime::Graph& graph, SessionState& session_state, const ExecutionProviders& providers, - KernelRegistryManager& kernel_registry_manager, - const logging::Logger& logger); + KernelRegistryManager& kernel_registry_manager); // First perform any transformations and create the execution plan common::Status CreatePlan(const std::vector& outer_scope_node_args, @@ -35,7 +34,8 @@ class SessionStateInitializer { // initialize tensors, and save. save kernels and input/output node mappings // @param enable_memory_pattern common::Status InitializeAndSave(bool enable_memory_pattern, - std::map& weights_buffers); + std::map& weights_buffers, + const std::vector* implicit_inputs = nullptr); private: onnxruntime::Graph& graph_; diff --git a/onnxruntime/core/framework/transformer_memcpy.cc b/onnxruntime/core/framework/transformer_memcpy.cc index 8c7fe38636..a7ae9e883a 100644 --- a/onnxruntime/core/framework/transformer_memcpy.cc +++ b/onnxruntime/core/framework/transformer_memcpy.cc @@ -81,6 +81,10 @@ void TransformerMemcpyImpl::ProcessDefs(onnxruntime::Node& node, const KernelReg return Status::OK(); }) .IsOK()); + + // we don't need to handle implicit input here as provider_ is never kCpuExecutionProvider, all control flow + // nodes are CPU based, and only control flow nodes have implicit inputs. + auto& output_defs = node.MutableOutputDefs(); for (size_t i = 0; i < output_defs.size(); ++i) { auto arg = output_defs[i]; @@ -94,7 +98,8 @@ void TransformerMemcpyImpl::ProcessDefs(onnxruntime::Node& node, const KernelReg } } else { // TODO: copy between devices? i.e. multiple GPUs - if (node.GetExecutionProviderType() != onnxruntime::kCpuExecutionProvider && !node.GetExecutionProviderType().empty()) { + if (node.GetExecutionProviderType() != onnxruntime::kCpuExecutionProvider && + !node.GetExecutionProviderType().empty()) { ORT_THROW("Execution type '", node.GetExecutionProviderType(), "' doesn't support memcpy "); } diff --git a/onnxruntime/core/framework/utils.cc b/onnxruntime/core/framework/utils.cc index 4bb3a192f5..9136f87593 100644 --- a/onnxruntime/core/framework/utils.cc +++ b/onnxruntime/core/framework/utils.cc @@ -5,11 +5,14 @@ #include "core/graph/graph_viewer.h" +#include "core/framework/execution_frame.h" #include "core/framework/execution_providers.h" #include "core/framework/kernel_def_builder.h" #include "core/framework/kernel_registry_manager.h" -#include "core/framework/op_kernel.h" +#include "core/framework/op_kernel_context_internal.h" +#include "core/framework/parallel_executor.h" #include "core/framework/session_state.h" +#include "core/framework/sequential_executor.h" namespace onnxruntime { namespace utils { @@ -39,5 +42,349 @@ AllocatorPtr GetAllocator(const SessionState& session_state, const OrtAllocatorI return GetAllocator(session_state.GetExecutionProviders(), allocator_info); } +common::Status AllocateHelper(const IExecutionProvider& execution_provider, + int device_id, + const Tensor& fetched_tensor, + MLValue& output_mlvalue) { + auto allocator = execution_provider.GetAllocator(device_id, OrtMemTypeDefault); + if (!allocator) { + return Status(common::ONNXRUNTIME, common::FAIL, "invalid allocator"); + } + + void* buffer = nullptr; + if (fetched_tensor.Size() != 0) { + buffer = allocator->Alloc(fetched_tensor.Size()); + if (!buffer) { + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Failed to allocate buffer. Execution provider type=", + execution_provider.Type()); + } + } + + std::unique_ptr p_tensor = std::make_unique(fetched_tensor.DataType(), + fetched_tensor.Shape(), + buffer, + allocator->Info(), + allocator); + output_mlvalue.Init(p_tensor.release(), + DataTypeImpl::GetType(), + DataTypeImpl::GetType()->GetDeleteFunc()); + + return Status::OK(); +} + +const std::string& GetNodeInputProviderType(const SessionState::NodeInfo& info) { + // the input index will be std::numeric_limits::max() if it's an implicit input to a control flow node. + // the input will be processed fully when executing the subgraph that consumes the implicit input. + bool implicit_input = info.index == std::numeric_limits::max(); + + // node may declare input_mem_type to be on CPU explicitly + // skip implicit inputs as they don't have a valid 'index' value + bool node_input_on_cpu = !implicit_input && + info.kci && MemTypeOnCpuExplicitly(info.kci->kernel_def->InputMemoryType(info.index)); + + // need a std::string that doesn't go away for kCpuExecutionProvider so we can return a reference. + static const std::string cpu_execution_provider{onnxruntime::kCpuExecutionProvider}; + + auto& required_provider_type = node_input_on_cpu ? cpu_execution_provider + : info.p_node->GetExecutionProviderType(); + + return required_provider_type; +} + +// TODO should we handle the case of one input name feeding 2 nodes placed on different devices? +common::Status CopyOneInputAcrossDevices(const SessionState& session_state, + const std::string& input_name, + const MLValue& orig_mlvalue, + MLValue& new_mlvalue) { + //TODO: make it configurable + const int target_device_id = 0; + std::vector node_info_vec; + ORT_RETURN_IF_ERROR(session_state.GetInputNodeInfo(input_name, node_info_vec)); + + auto& exec_providers = session_state.GetExecutionProviders(); + + // currently we only support one device per input. see SessionState::AddInputNameToNodeInfoMapping for more + // info on the logic to create the node_info_vec. + // for (auto& node_info : node_info_vec) { + auto& node_info = node_info_vec.front(); + + if (node_info.p_node == nullptr) { + // dummy entry for an input that we didn't find a use of in the graph. + // use the input as is given we don't believe it's actually needed. + new_mlvalue = orig_mlvalue; + return Status::OK(); + } + + if (!orig_mlvalue.IsTensor()) { + // copying not supported for non-tensor types + new_mlvalue = orig_mlvalue; + return Status::OK(); + } + + auto& required_provider_type = GetNodeInputProviderType(node_info); + auto& input_tensor = orig_mlvalue.Get(); + auto& input_tensor_loc = input_tensor.Location(); + + auto* p_input_provider = exec_providers.Get(input_tensor_loc); + if (!p_input_provider) { + p_input_provider = exec_providers.Get(onnxruntime::kCpuExecutionProvider); + ORT_ENFORCE(p_input_provider); + } + + auto input_provider_type = p_input_provider->Type(); + if (input_provider_type == required_provider_type && input_tensor_loc.mem_type == OrtMemTypeDefault) { + new_mlvalue = orig_mlvalue; + return Status::OK(); + } + + // If a node requires input on cpu and input tensor is allocated with pinned memory allocator, don't do copy + if (required_provider_type == onnxruntime::kCpuExecutionProvider && + (input_tensor_loc.mem_type == OrtMemTypeCPU || + input_tensor_loc.mem_type == OrtMemTypeCPUOutput)) { + new_mlvalue = orig_mlvalue; + return Status::OK(); + } + + auto* required_provider = exec_providers.Get(required_provider_type); + ORT_ENFORCE(required_provider); + ORT_RETURN_IF_ERROR(utils::AllocateHelper(*required_provider, target_device_id, input_tensor, new_mlvalue)); + + auto* new_tensor = new_mlvalue.GetMutable(); + + // our CPU exec provider doesn't support copy from GPU->CPU + if (required_provider_type != onnxruntime::kCpuExecutionProvider) { + ORT_RETURN_IF_ERROR(required_provider->CopyTensor(input_tensor, *new_tensor)); + } else { + ORT_RETURN_IF_ERROR(p_input_provider->CopyTensor(input_tensor, *new_tensor)); + } + + // } loop of node_info_vec + + return Status::OK(); +} + +// copies inputs across devices only if required +common::Status CopyInputsAcrossDevices(const SessionState& session_state, + const NameMLValMap& orig_feeds, + NameMLValMap& new_feeds) { + for (auto& pair : orig_feeds) { + MLValue new_mlvalue; + auto& input_name = pair.first; + auto& orig_mlvalue = pair.second; + ORT_RETURN_IF_ERROR(CopyOneInputAcrossDevices(session_state, input_name, orig_mlvalue, new_mlvalue)); + new_feeds[input_name] = new_mlvalue; + } + + return Status::OK(); +} + +static std::pair Contains(const std::vector& output_names, + const std::string& name) { + auto it = std::find(std::begin(output_names), std::end(output_names), name); + if (it == output_names.end()) { + return {false, 0}; + } + return {true, it - output_names.begin()}; +} + +// ensures pre-allocated outputs match the node providers. +common::Status MatchOutputsWithProviders(const SessionState& session_state, + const std::vector& output_names, + std::vector& fetches, + std::vector& new_fetches) { + const auto& execution_providers = session_state.GetExecutionProviders(); + + if (fetches.empty()) { + fetches.resize(output_names.size()); + } + + new_fetches.resize(output_names.size()); + + std::set seen_outputs; + auto p_graph = session_state.GetGraphViewer(); + ORT_ENFORCE(p_graph); + + std::pair found; + for (auto& node : p_graph->Nodes()) { // TODO optimize this + if (seen_outputs.size() == fetches.size()) { + break; + } + + for (auto* arg : node.OutputDefs()) { + if (!arg->Exists() || + arg->Name().empty() || + !(found = Contains(output_names, arg->Name())).first) { + continue; + } + + seen_outputs.insert(arg->Name()); + size_t idx = found.second; + MLValue orig_mlvalue = fetches[idx]; + + if (orig_mlvalue.IsAllocated()) { + if (!orig_mlvalue.IsTensor()) { + new_fetches[idx] = fetches[idx]; + continue; + } + + auto& node_provider_type = node.GetExecutionProviderType(); + auto& orig_tensor = orig_mlvalue.Get(); + auto& orig_tensor_loc = orig_tensor.Location(); + auto* tensor_provider = execution_providers.Get(orig_tensor_loc); + if (!tensor_provider) { + tensor_provider = execution_providers.Get(onnxruntime::kCpuExecutionProvider); + } + + auto tensor_provider_type = tensor_provider->Type(); + if (node_provider_type == tensor_provider_type) { + new_fetches[idx] = fetches[idx]; + continue; + } + + // leave the new_fetches[idx] as it is since it'll get allocated on the appropriate + // provider by the op kernel context when requested. + continue; + + } else { + new_fetches[idx] = fetches[idx]; + continue; + } + } + } + + // If we've already seen all the outputs requested just return. + if (seen_outputs.size() == output_names.size()) { + return Status::OK(); + } + + // Handle the case when a constant is an output but has been folded into a weight + // and hence it doesn't show up in any of the OutputDefs before. + // assume that the weight has already been placed in the appropriate device before + auto& defs = p_graph->GetOutputs(); + auto& mlvalue_name_idx_map{session_state.GetMLValueNameIdxMap()}; + auto& weights = session_state.GetInitializedTensors(); + + for (auto& one_def : defs) { + if (!one_def->Exists() || + one_def->Name().empty() || + seen_outputs.count(one_def->Name()) || + !(found = Contains(output_names, one_def->Name())).first) { + continue; + } + + auto& def_name = one_def->Name(); + size_t idx = found.second; + int mlvalue_idx; + ORT_RETURN_IF_ERROR(mlvalue_name_idx_map.GetIdx(def_name, mlvalue_idx)); + if (!weights.count(mlvalue_idx)) { + LOGS(session_state.Logger(), INFO) << "Output with name " << def_name << " is not a weight."; + continue; + } + + seen_outputs.insert(def_name); + const auto& weight = weights.at(mlvalue_idx); + new_fetches[idx] = weight; + } + + if (seen_outputs.size() != output_names.size()) // make sure we've seen all outputs + return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "output size mismatch, expected ", output_names.size(), + " got ", seen_outputs.size()); + + return Status::OK(); +} + +// copies outputs across devices only if required +common::Status CopyOutputsAcrossDevices(const SessionState& session_state, + std::vector& fetches, + std::vector& user_fetches) { + auto& execution_providers = session_state.GetExecutionProviders(); + + for (size_t idx = 0, end = fetches.size(); idx < end; ++idx) { + auto& fetched_mlvalue = fetches[idx]; + if (!fetched_mlvalue.IsTensor()) { + user_fetches[idx] = fetched_mlvalue; + continue; + } + + auto& fetched_tensor = fetched_mlvalue.Get(); + auto& fetched_tensor_location = fetched_tensor.Location(); + auto* p_fetched_provider = execution_providers.Get(fetched_tensor_location); + if (!p_fetched_provider) { + p_fetched_provider = execution_providers.Get(onnxruntime::kCpuExecutionProvider); + ORT_ENFORCE(p_fetched_provider); + } + + auto fetched_provider_type = p_fetched_provider->Type(); + auto& output_mlvalue = user_fetches[idx]; + + if (!output_mlvalue.IsAllocated()) { + if (fetched_provider_type != onnxruntime::kCpuExecutionProvider) { + ORT_RETURN_IF_ERROR(utils::AllocateHelper(*execution_providers.Get(onnxruntime::kCpuExecutionProvider), + 0, + fetched_tensor, + output_mlvalue)); + } else { + user_fetches[idx] = fetched_mlvalue; + continue; + } + } + + Tensor* p_output_tensor = output_mlvalue.GetMutable(); + auto& output_tensor_loc = p_output_tensor->Location(); + auto* p_output_provider = execution_providers.Get(output_tensor_loc); + if (!p_output_provider) { + p_output_provider = execution_providers.Get(onnxruntime::kCpuExecutionProvider); + ORT_ENFORCE(p_output_provider); + } + + auto output_provider_type = p_output_provider->Type(); + + if (output_provider_type == fetched_provider_type || fetched_tensor_location.mem_type == OrtMemTypeCPUOutput) { + user_fetches[idx] = fetched_mlvalue; + continue; + } + + // our CPU exec provider doesn't support copy from GPU->CPU + if (fetched_provider_type != onnxruntime::kCpuExecutionProvider) { + ORT_RETURN_IF_ERROR(p_fetched_provider->CopyTensor(fetched_tensor, *p_output_tensor)); + } else { + ORT_RETURN_IF_ERROR(p_output_provider->CopyTensor(fetched_tensor, *p_output_tensor)); + } + } + + return Status::OK(); +} + +common::Status ExecuteGraph(const SessionState& session_state, + const NameMLValMap& feeds, + const std::vector& output_names, + std::vector& fetches, + bool sequential_execution, + const bool& terminate_flag, + const logging::Logger& logger) { + // TODO: Would be better to check upfront whether there was a need to copy inputs/outputs across devices, + // especially when a subgraph is repeatedly executed in a Scan or Loop node. If we checked once and no copy was + // needed we can skip everything here apart from the Execute call. + + NameMLValMap device_feeds; + ORT_RETURN_IF_ERROR(utils::CopyInputsAcrossDevices(session_state, feeds, device_feeds)); + + std::vector device_fetches; + ORT_RETURN_IF_ERROR(utils::MatchOutputsWithProviders(session_state, output_names, fetches, device_fetches)); + + std::unique_ptr p_exec; + + if (sequential_execution) { + p_exec = std::unique_ptr(new SequentialExecutor(terminate_flag)); + } else { + p_exec = std::unique_ptr(new ParallelExecutor(session_state, terminate_flag)); + } + + ORT_RETURN_IF_ERROR(p_exec->Execute(session_state, device_feeds, output_names, device_fetches, logger)); + ORT_RETURN_IF_ERROR(utils::CopyOutputsAcrossDevices(session_state, device_fetches, fetches)); + + return Status::OK(); +} + } // namespace utils } // namespace onnxruntime diff --git a/onnxruntime/core/framework/utils.h b/onnxruntime/core/framework/utils.h index af9fc2166e..5ff5d69b0d 100644 --- a/onnxruntime/core/framework/utils.h +++ b/onnxruntime/core/framework/utils.h @@ -6,17 +6,18 @@ #include "core/graph/basic_types.h" #include "core/framework/allocator.h" #include "core/framework/data_types.h" - -namespace onnxruntime { -class Node; -class Graph; -} // namespace onnxruntime +#include "core/framework/framework_common.h" +#include "core/framework/session_state.h" namespace onnxruntime { class ExecutionProviders; +class Graph; class KernelDef; class KernelRegistryManager; -class SessionState; +class IExecutionProvider; +class MLValue; +class Node; +class Tensor; namespace logging { class Logger; @@ -26,14 +27,42 @@ namespace utils { const KernelDef* GetKernelDef(const KernelRegistryManager& kernel_registry, const onnxruntime::Node& node); -const KernelDef* GetKernelDef(const onnxruntime::Graph& graph, - const KernelRegistryManager& kernel_registry, - const onnxruntime::NodeIndex node_id); - AllocatorPtr GetAllocator(const ExecutionProviders& exec_providers, const OrtAllocatorInfo& allocator_info); -AllocatorPtr GetAllocator(const SessionState& session_state, - const OrtAllocatorInfo& allocator_info); +AllocatorPtr GetAllocator(const SessionState& session_state, const OrtAllocatorInfo& allocator_info); + +common::Status AllocateHelper(const IExecutionProvider& execution_provider, + int device_id, + const Tensor& fetched_tensor, + MLValue& output_mlvalue); + +const std::string& GetNodeInputProviderType(const SessionState::NodeInfo& info); + +common::Status CopyOneInputAcrossDevices(const SessionState& session_state, + const std::string& input_name, + const MLValue& orig_mlvalue, + MLValue& new_mlvalue); + +common::Status CopyInputsAcrossDevices(const SessionState& session_state, + const NameMLValMap& orig_feeds, + NameMLValMap& new_feeds); + +common::Status MatchOutputsWithProviders(const SessionState& session_state, + const std::vector& output_names, + std::vector& fetches, + std::vector& new_fetches); + +common::Status CopyOutputsAcrossDevices(const SessionState& session_state, + std::vector& fetches, + std::vector& user_fetches); + +common::Status ExecuteGraph(const SessionState& session_state, + const NameMLValMap& feeds, + const std::vector& output_names, + std::vector& fetches, + bool sequential_execution, + const bool& terminate_flag, + const logging::Logger& logger); #define DispatchOnTensorType(tensor_type, function, ...) \ if (tensor_type == DataTypeImpl::GetType()) \ diff --git a/onnxruntime/core/graph/graph.cc b/onnxruntime/core/graph/graph.cc index 37a0ff354a..9c23dadd52 100644 --- a/onnxruntime/core/graph/graph.cc +++ b/onnxruntime/core/graph/graph.cc @@ -401,7 +401,7 @@ void Node::CreateSubgraph(const std::string& attr_name) { if (attr != attributes_.cend() && attr->second.has_g()) { GraphProto& mutable_graph = *attr->second.mutable_g(); std::unique_ptr subgraph{new Graph(*graph_, mutable_graph)}; - attr_to_subgraph_map_[attr_name] = subgraph.get(); + attr_to_subgraph_map_.insert({std::string{attr_name}, gsl::not_null{subgraph.get()}}); subgraphs_.push_back(std::move(subgraph)); } } diff --git a/onnxruntime/core/providers/cpu/controlflow/if.cc b/onnxruntime/core/providers/cpu/controlflow/if.cc index febad17178..e742e22636 100644 --- a/onnxruntime/core/providers/cpu/controlflow/if.cc +++ b/onnxruntime/core/providers/cpu/controlflow/if.cc @@ -7,9 +7,9 @@ #include "core/framework/op_kernel_context_internal.h" #include "core/framework/sequential_executor.h" #include "core/framework/session_state.h" +#include "core/framework/utils.h" #include "core/framework/tensorprotoutils.h" -// #include "core/providers/cpu/tensor/utils.h" using namespace ONNX_NAMESPACE; using namespace onnxruntime::common; @@ -120,7 +120,7 @@ Status IfImpl::Initialize() { if (num_subgraph_outputs != num_outputs_) { return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "'If' node has ", num_outputs_, - " outputs which doesn't match the subgraph's ", num_subgraph_outputs, " outputs."); + " outputs which doesn't match the subgraph's ", num_subgraph_outputs, " outputs."); } subgraph_output_names_.reserve(num_subgraph_outputs); @@ -145,7 +145,7 @@ Status IfImpl::AllocateOutputTensors() { auto* graph_output_shape = graph_output->Shape(); if (!graph_output_shape) { return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Subgraph must have the shape set for all outputs but ", - graph_output->Name(), " did not."); + graph_output->Name(), " did not."); } TensorShape output_shape{onnxruntime::utils::GetTensorShapeFromTensorShapeProto(*graph_output_shape)}; @@ -182,7 +182,7 @@ Status IfImpl::Execute() { // pass in implicit inputs as feeds. for (auto& entry : implicit_inputs_) { ORT_ENFORCE(entry.second, "All implicit inputs should have MLValue instances by now. ", - entry.first, " did not."); + entry.first, " did not."); // prune to values that are in this subgraph as the implicit inputs cover both 'then' and 'else' subgraphs. // alternatively we could track implicit inputs on a per-attribute basis in the node, but that @@ -192,7 +192,6 @@ Status IfImpl::Execute() { feeds[entry.first] = *entry.second; } } - std::vector fetches; fetches.reserve(num_outputs_); @@ -200,8 +199,8 @@ Status IfImpl::Execute() { fetches.push_back(outputs_[i].second); } - SequentialExecutor executor{context_.GetTerminateFlag()}; - status = executor.Execute(session_state_, feeds, subgraph_output_names_, fetches, context_.Logger()); + status = utils::ExecuteGraph(session_state_, feeds, subgraph_output_names_, fetches, /*sequential_execution*/ true, + context_.GetTerminateFlag(), context_.Logger()); ORT_RETURN_IF_ERROR(status); for (int i = 0; i < num_outputs_; ++i) { diff --git a/onnxruntime/core/providers/cpu/controlflow/loop.cc b/onnxruntime/core/providers/cpu/controlflow/loop.cc index 454100421d..7e62437baa 100644 --- a/onnxruntime/core/providers/cpu/controlflow/loop.cc +++ b/onnxruntime/core/providers/cpu/controlflow/loop.cc @@ -16,7 +16,7 @@ #include "core/framework/sequential_executor.h" #include "core/framework/session_state.h" #include "core/framework/tensorprotoutils.h" - +#include "core/framework/utils.h" #include "core/providers/cpu/tensor/utils.h" #include "gsl/gsl_algorithm" @@ -187,8 +187,8 @@ Status LoopImpl::Initialize() { // validate that the subgraph has that many inputs. if (num_subgraph_inputs_ != subgraph_inputs.size()) { return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, - "Graph in 'body' attribute of Loop should have ", - num_subgraph_inputs_, " inputs. Found:", subgraph_.GetInputs().size()); + "Graph in 'body' attribute of Loop should have ", + num_subgraph_inputs_, " inputs. Found:", subgraph_.GetInputs().size()); } auto& subgraph_outputs = subgraph_.GetOutputs(); @@ -197,8 +197,8 @@ Status LoopImpl::Initialize() { // check num outputs are correct. the 'cond' output from the subgraph is not a Loop output, so diff is 1 if (num_subgraph_outputs - 1 != num_outputs_) { return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "'Loop' node has ", num_outputs_, - " outputs so the subgraph requires ", num_outputs_ + 1, - " but has ", num_subgraph_outputs); + " outputs so the subgraph requires ", num_outputs_ + 1, + " but has ", num_subgraph_outputs); } AllocatorPtr allocator; @@ -242,7 +242,7 @@ NameMLValMap LoopImpl::CreateInitialFeeds() { // pass in implicit inputs as feeds. for (auto& entry : implicit_inputs_) { ORT_ENFORCE(entry.second, "All implicit inputs should have MLValue instances by now. ", - entry.first, " did not."); + entry.first, " did not."); feeds[entry.first] = *entry.second; } @@ -290,7 +290,7 @@ Status LoopImpl::ConcatenateLoopOutput(std::vector& per_iteration_outpu // sanity check if (bytes_per_iteration != iteration_data.Size()) { return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "Inconsistent shape in loop output for output ", output_index, - " Expected:", per_iteration_shape, " Got:", iteration_data.Shape()); + " Expected:", per_iteration_shape, " Got:", iteration_data.Shape()); } auto num_bytes = iteration_data.Size(); @@ -316,8 +316,8 @@ Status LoopImpl::Execute() { fetches.clear(); } - SequentialExecutor executor{context_.GetTerminateFlag()}; - status = executor.Execute(session_state_, feeds, subgraph_output_names_, fetches, context_.Logger()); + status = utils::ExecuteGraph(session_state_, feeds, subgraph_output_names_, fetches, /*sequential_execution*/ true, + context_.GetTerminateFlag(), context_.Logger()); ORT_RETURN_IF_ERROR(status); condition_mlvalue_ = fetches[0]; diff --git a/onnxruntime/core/providers/cpu/controlflow/scan_utils.cc b/onnxruntime/core/providers/cpu/controlflow/scan_utils.cc index d341314006..f4039cfda2 100644 --- a/onnxruntime/core/providers/cpu/controlflow/scan_utils.cc +++ b/onnxruntime/core/providers/cpu/controlflow/scan_utils.cc @@ -16,6 +16,7 @@ #include "core/framework/op_kernel_context_internal.h" #include "core/framework/sequential_executor.h" #include "core/framework/tensorprotoutils.h" +#include "core/framework/utils.h" #ifdef _MSC_VER #pragma warning(pop) @@ -165,8 +166,12 @@ Status IterateSequence(OpKernelContextInternal& context, // Many of the other pieces are constant across usages. // Not sure how best to handle the memory pattern side of things though. // For now just making it work. Optimization and refinement will follow. - SequentialExecutor executor{context.GetTerminateFlag()}; - status = executor.Execute(session_state, feeds, subgraph_output_names, fetches, context.Logger()); + //SequentialExecutor executor{context.GetTerminateFlag()}; + //status = executor.Execute(session_state, feeds, subgraph_output_names, fetches, context.Logger()); + //ORT_RETURN_IF_ERROR(status); + + status = utils::ExecuteGraph(session_state, feeds, subgraph_output_names, fetches, /*sequential_execution*/ true, + context.GetTerminateFlag(), context.Logger()); ORT_RETURN_IF_ERROR(status); // cycle the LoopStateVariable input/output in preparation for the next iteration diff --git a/onnxruntime/core/session/IOBinding.cc b/onnxruntime/core/session/IOBinding.cc index cd24146d71..a1dd01f02d 100644 --- a/onnxruntime/core/session/IOBinding.cc +++ b/onnxruntime/core/session/IOBinding.cc @@ -5,6 +5,7 @@ #include "core/common/logging/logging.h" #include "core/framework/session_state.h" #include "core/framework/op_kernel.h" +#include "core/framework/utils.h" namespace onnxruntime { IOBinding::IOBinding(const SessionState& session_state) : session_state_(session_state) { @@ -17,98 +18,11 @@ common::Status IOBinding::BindInput(const std::string& name, const MLValue& ml_v } MLValue new_mlvalue; - ORT_RETURN_IF_ERROR(CopyOneInputAcrossDevices(session_state_, name, ml_value, new_mlvalue)); + ORT_RETURN_IF_ERROR(utils::CopyOneInputAcrossDevices(session_state_, name, ml_value, new_mlvalue)); feeds_[name] = new_mlvalue; return Status::OK(); } -static common::Status AllocateHelper(const SessionState& session_state, - int id, onnxruntime::ProviderType provider_type, - const MLValue& fetched_mlvalue, - MLValue& output_mlvalue) { - auto* p_provider = session_state.GetExecutionProviders().Get(provider_type); - ORT_ENFORCE(p_provider); - auto allocator = p_provider->GetAllocator(id, OrtMemTypeDefault); - ORT_ENFORCE(allocator != nullptr); - auto& fetched_tensor = fetched_mlvalue.Get(); - void* buffer = allocator->Alloc(fetched_tensor.Size()); - ORT_ENFORCE(buffer); - std::unique_ptr p_tensor = std::make_unique(fetched_tensor.DataType(), - fetched_tensor.Shape(), - buffer, - allocator->Info(), - allocator); - output_mlvalue.Init(p_tensor.release(), - DataTypeImpl::GetType(), - DataTypeImpl::GetType()->GetDeleteFunc()); - - return Status::OK(); -} - -// TODO should we handle the case of one input name feeding 2 nodes placed on different -// devices. -common::Status IOBinding::CopyOneInputAcrossDevices(const SessionState& session_state, - const std::string& input_name, - const MLValue& orig_mlvalue, - MLValue& new_mlvalue) { - //TODO: make it configurable - const int target_device_id = 0; - std::vector node_info_vec; - ORT_RETURN_IF_ERROR(session_state.GetInputNodeInfo(input_name, node_info_vec)); - - for (auto& node_info : node_info_vec) { - size_t index = node_info.index; - auto& node = *node_info.p_node; - const KernelCreateInfo* kci = node_info.kci; - - // node may declare input_mem_type to be on CPU explicitly - bool node_input_on_cpu = kci && MemTypeOnCpuExplicitly(kci->kernel_def->InputMemoryType(index)); - auto& required_provider_type = node_input_on_cpu ? onnxruntime::kCpuExecutionProvider : node.GetExecutionProviderType(); - if (!orig_mlvalue.IsTensor()) { - // copying not supported for non-tensor types - new_mlvalue = orig_mlvalue; - return Status::OK(); - } - auto& input_tensor = orig_mlvalue.Get(); - auto& input_tensor_loc = input_tensor.Location(); - auto& exec_providers = session_state.GetExecutionProviders(); - - auto* p_input_provider = exec_providers.Get(input_tensor_loc); - if (!p_input_provider) { - p_input_provider = exec_providers.Get(onnxruntime::kCpuExecutionProvider); - ORT_ENFORCE(p_input_provider); - } - - auto input_provider_type = p_input_provider->Type(); - if (input_provider_type == required_provider_type && input_tensor_loc.mem_type == OrtMemTypeDefault) { - new_mlvalue = orig_mlvalue; - return Status::OK(); - } - - //If node require input on cpu and input tensor is allocated with pinned memory allocator, don't do copy - if (node_input_on_cpu && (input_tensor_loc.mem_type == OrtMemTypeCPU || input_tensor_loc.mem_type == OrtMemTypeCPUOutput)) { - new_mlvalue = orig_mlvalue; - return Status::OK(); - } - - auto* node_provider = exec_providers.Get(required_provider_type); - ORT_ENFORCE(node_provider); - ORT_RETURN_IF_ERROR(AllocateHelper(session_state, target_device_id, required_provider_type, orig_mlvalue, new_mlvalue)); - auto* new_tensor = new_mlvalue.GetMutable(); - auto* node_exec_provider = exec_providers.Get(required_provider_type); - ORT_ENFORCE(node_exec_provider); - - // our CPU exec provider doesn't support copy from GPU->CPU - if (required_provider_type != onnxruntime::kCpuExecutionProvider) { - ORT_RETURN_IF_ERROR(node_exec_provider->CopyTensor(input_tensor, *new_tensor)); - } else { - ORT_RETURN_IF_ERROR(p_input_provider->CopyTensor(input_tensor, *new_tensor)); - } - } - - return Status::OK(); -} - static common::Status SyncProviders(const SessionState::NameNodeInfoMapType& node_info_map, const SessionState& session_state) { std::set providers; diff --git a/onnxruntime/core/session/IOBinding.h b/onnxruntime/core/session/IOBinding.h index 06e1c29748..18ab77a7cb 100644 --- a/onnxruntime/core/session/IOBinding.h +++ b/onnxruntime/core/session/IOBinding.h @@ -86,11 +86,6 @@ class IOBinding { std::vector output_names_; std::vector outputs_; - static common::Status CopyOneInputAcrossDevices(const SessionState& session_state, - const std::string& input_name, - const MLValue& orig_mlvalue, - MLValue& new_mlvalue); - ORT_DISALLOW_COPY_ASSIGNMENT_AND_MOVE(IOBinding); }; } // namespace onnxruntime diff --git a/onnxruntime/core/session/inference_session.cc b/onnxruntime/core/session/inference_session.cc index cc90cd5e57..2e2cdd7204 100644 --- a/onnxruntime/core/session/inference_session.cc +++ b/onnxruntime/core/session/inference_session.cc @@ -333,16 +333,18 @@ class InferenceSession::Impl { // create SessionState for executing subgraph subgraph_info.session_state = std::make_unique(execution_providers_); subgraph_info.session_state->SetProfiler(session_profiler_); + subgraph_info.session_state->SetLogger(*session_logger_); // setup everything required to execute the subgraph and save it in subgraph_session_state SessionStateInitializer initializer{*subgraph, *subgraph_info.session_state, - execution_providers_, kernel_registry_manager_, *session_logger_}; + execution_providers_, kernel_registry_manager_}; ORT_RETURN_IF_ERROR(initializer.CreatePlan(node.ImplicitInputDefs(), session_options_.enable_sequential_execution)); ORT_RETURN_IF_ERROR(initializer.InitializeAndSave(session_state_.GetEnableMemoryPattern(), - subgraph_info.weights_buffers)); + subgraph_info.weights_buffers, + &node.ImplicitInputDefs())); // add the subgraph SessionState instance to the parent graph SessionState so it can be retrieved // by Compute() via OpKernelContextInternal. @@ -404,7 +406,7 @@ class InferenceSession::Impl { insert_cast_transformer_.AddKernelRegistries(kernel_registry_manager_.GetAllKernelRegistries()); SessionStateInitializer session_initializer{graph, session_state_, execution_providers_, - kernel_registry_manager_, *session_logger_}; + kernel_registry_manager_}; // apply any transformations to the main graph and any subgraphs ORT_RETURN_IF_ERROR(TransformGraph(graph, graph_transformation_mgr_, @@ -594,210 +596,6 @@ class InferenceSession::Impl { return common::Status::OK(); } - // copies inputs across devices only if required - common::Status CopyInputsAcrossDevices(const SessionState& session_state, - const NameMLValMap& orig_feeds, - NameMLValMap& new_feeds) { - for (auto& pair : orig_feeds) { - MLValue new_mlvalue; - auto& input_name = pair.first; - auto& orig_mlvalue = pair.second; - ORT_RETURN_IF_ERROR(IOBinding::CopyOneInputAcrossDevices(session_state, - input_name, - orig_mlvalue, - new_mlvalue)); - new_feeds[input_name] = new_mlvalue; - } - return Status::OK(); - } - - // ensures pre-allocated outputs match the node providers. - common::Status MatchOutputsWithProviders(const std::vector& output_names, - std::vector& fetches, - std::vector& new_fetches) { - if (fetches.empty()) { - fetches.resize(output_names.size()); - } - new_fetches.resize(output_names.size()); - - std::set seen_outputs; - auto p_graph = session_state_.GetGraphViewer(); - ORT_ENFORCE(p_graph); - - std::pair found; - for (auto& node : p_graph->Nodes()) { // TODO optimize this - if (seen_outputs.size() == fetches.size()) { - break; - } - for (auto* arg : node.OutputDefs()) { - if (!arg->Exists() || - arg->Name().empty() || - !(found = Contains(output_names, arg->Name())).first) { - continue; - } - - seen_outputs.insert(arg->Name()); - size_t idx = found.second; - MLValue orig_mlvalue = fetches[idx]; - if (orig_mlvalue.IsAllocated()) { - if (!orig_mlvalue.IsTensor()) { - new_fetches[idx] = fetches[idx]; - continue; - } - - auto& node_provider_type = node.GetExecutionProviderType(); - auto& orig_tensor = orig_mlvalue.Get(); - auto& orig_tensor_loc = orig_tensor.Location(); - auto* tensor_provider = execution_providers_.Get(orig_tensor_loc); - if (!tensor_provider) { - tensor_provider = execution_providers_.Get(onnxruntime::kCpuExecutionProvider); - } - - auto tensor_provider_type = tensor_provider->Type(); - if (node_provider_type == tensor_provider_type) { - new_fetches[idx] = fetches[idx]; - continue; - } - // leave the new_fetches[idx] as it is since it'll get allocated on the appropriate - // provider by the op kernel context when requested. - continue; - - } else { - new_fetches[idx] = fetches[idx]; - continue; - } - } - } - - // If we've already seen all the outputs requested just return. - if (seen_outputs.size() == output_names.size()) { - return Status::OK(); - } - - // Handle the case when a constant is an output but has been folded into a weight - // and hence it doesn't show up in any of the OutputDefs before. - // assume that the weight has already been placed in the appropriate device before - auto& defs = p_graph->GetOutputs(); - auto& mlvalue_name_idx_map{session_state_.GetMLValueNameIdxMap()}; - auto& weights = session_state_.GetInitializedTensors(); - - for (auto& one_def : defs) { - if (!one_def->Exists() || - one_def->Name().empty() || - seen_outputs.count(one_def->Name()) || - !(found = Contains(output_names, one_def->Name())).first) { - continue; - } - - auto& def_name = one_def->Name(); - size_t idx = found.second; - int mlvalue_idx; - ORT_RETURN_IF_ERROR(mlvalue_name_idx_map.GetIdx(def_name, mlvalue_idx)); - if (!weights.count(mlvalue_idx)) { - LOGS(*session_logger_, INFO) << "Output with name " << def_name << " is not a weight."; - continue; - } - seen_outputs.insert(def_name); - const auto& weight = weights.at(mlvalue_idx); - new_fetches[idx] = weight; - } - - if (seen_outputs.size() != output_names.size()) // make sure we've seen all outputs - return ORT_MAKE_STATUS(ONNXRUNTIME, FAIL, "output size mismatch, expected ", output_names.size(), - " got ", seen_outputs.size()); - - return Status::OK(); - } - - common::Status AllocateHelper(onnxruntime::ProviderType provider_type, - int device_id, - const Tensor& fetched_tensor, - MLValue& output_mlvalue) { - auto* p_provider = execution_providers_.Get(provider_type); - if (!p_provider) - return Status(common::ONNXRUNTIME, common::INVALID_ARGUMENT, "invalid provider_type"); - - auto allocator = p_provider->GetAllocator(device_id, OrtMemTypeDefault); - if (!allocator) - return Status(common::ONNXRUNTIME, common::FAIL, "invalid allocator"); - - void* buffer = nullptr; - if (fetched_tensor.Shape().Size() != 0) { - buffer = allocator->Alloc(fetched_tensor.DataType()->Size() * fetched_tensor.Shape().Size()); - if (!buffer) - return Status(common::ONNXRUNTIME, common::FAIL, "invalid buffer"); - } - - std::unique_ptr p_tensor = std::make_unique(fetched_tensor.DataType(), - fetched_tensor.Shape(), - buffer, - allocator->Info(), - allocator); - output_mlvalue.Init(p_tensor.release(), - DataTypeImpl::GetType(), - DataTypeImpl::GetType()->GetDeleteFunc()); - - return Status::OK(); - } - - // copies outputs across devices only if required - common::Status CopyOutputsAcrossDevices(std::vector& fetches, - std::vector& user_fetches) { - for (size_t idx = 0, end = fetches.size(); idx < end; ++idx) { - auto& fetched_mlvalue = fetches[idx]; - if (!fetched_mlvalue.IsTensor()) { - user_fetches[idx] = fetched_mlvalue; - continue; - } - - auto& fetched_tensor = fetched_mlvalue.Get(); - auto& fetched_tensor_location = fetched_tensor.Location(); - auto* p_fetched_provider = execution_providers_.Get(fetched_tensor_location); - if (!p_fetched_provider) { - p_fetched_provider = execution_providers_.Get(onnxruntime::kCpuExecutionProvider); - ORT_ENFORCE(p_fetched_provider); - } - - auto fetched_provider_type = p_fetched_provider->Type(); - - auto& output_mlvalue = user_fetches[idx]; - if (!output_mlvalue.IsAllocated()) { - if (fetched_provider_type != onnxruntime::kCpuExecutionProvider) { - ORT_RETURN_IF_ERROR(AllocateHelper(onnxruntime::kCpuExecutionProvider, 0, - fetched_tensor, - output_mlvalue)); - } else { - user_fetches[idx] = fetched_mlvalue; - continue; - } - } - - Tensor* p_output_tensor = output_mlvalue.GetMutable(); - auto& output_tensor_loc = p_output_tensor->Location(); - auto* p_output_provider = execution_providers_.Get(output_tensor_loc); - if (!p_output_provider) { - p_output_provider = execution_providers_.Get(onnxruntime::kCpuExecutionProvider); - ORT_ENFORCE(p_output_provider); - } - - auto output_provider_type = p_output_provider->Type(); - - if (output_provider_type == fetched_provider_type || fetched_tensor_location.mem_type == OrtMemTypeCPUOutput) { - user_fetches[idx] = fetched_mlvalue; - continue; - } - - // our CPU exec provider doesn't support copy from GPU->CPU - if (fetched_provider_type != onnxruntime::kCpuExecutionProvider) { - ORT_RETURN_IF_ERROR(p_fetched_provider->CopyTensor(fetched_tensor, *p_output_tensor)); - } else { - ORT_RETURN_IF_ERROR(p_output_provider->CopyTensor(fetched_tensor, *p_output_tensor)); - } - } - - return Status::OK(); - } - Status Run(const RunOptions& run_options, const NameMLValMap& feeds, const std::vector& output_names, @@ -834,28 +632,13 @@ class InferenceSession::Impl { // info all execution providers InferenceSession:Run started // TODO: only call OnRunStart for all providers in-use - for (auto& xp : execution_providers_) + for (auto& xp : execution_providers_) { ORT_CHECK_AND_SET_RETVAL(xp->OnRunStart()); - - NameMLValMap copied_feeds; - ORT_CHECK_AND_SET_RETVAL(CopyInputsAcrossDevices(session_state_, feeds, copied_feeds)); - - std::vector new_fetches; - ORT_CHECK_AND_SET_RETVAL(MatchOutputsWithProviders(output_names, *p_fetches, new_fetches)); - - std::unique_ptr p_exec; - - if (retval.IsOK()) { - if (session_options_.enable_sequential_execution) { - p_exec = std::unique_ptr(new SequentialExecutor(run_options.terminate)); - } else { - p_exec = std::unique_ptr(new ParallelExecutor(session_state_, run_options.terminate)); - } } - ORT_CHECK_AND_SET_RETVAL(p_exec->Execute(session_state_, copied_feeds, output_names, new_fetches, run_logger)); - ORT_CHECK_AND_SET_RETVAL(CopyOutputsAcrossDevices(new_fetches, *p_fetches)); - + ORT_CHECK_AND_SET_RETVAL( + utils::ExecuteGraph(session_state_, feeds, output_names, *p_fetches, + session_options_.enable_sequential_execution, run_options.terminate, run_logger)); } catch (const std::exception& e) { retval = Status(common::ONNXRUNTIME, common::FAIL, e.what()); } catch (...) { @@ -956,15 +739,6 @@ class InferenceSession::Impl { } private: - static std::pair Contains(const std::vector& output_names, - const std::string& name) { - auto it = std::find(std::begin(output_names), std::end(output_names), name); - if (it == output_names.end()) { - return {false, 0}; - } - return {true, it - output_names.begin()}; - } - bool HasLocalSchema() const { return !custom_schema_registries_.empty(); } diff --git a/onnxruntime/test/providers/cpu/controlflow/if_test.cc b/onnxruntime/test/providers/cpu/controlflow/if_test.cc index f7334d9570..fa8178caf6 100644 --- a/onnxruntime/test/providers/cpu/controlflow/if_test.cc +++ b/onnxruntime/test/providers/cpu/controlflow/if_test.cc @@ -9,6 +9,8 @@ #include "test/providers/provider_test_utils.h" #include "core/session/inference_session.h" +#include "test/util/include/default_providers.h" + using namespace ONNX_NAMESPACE; namespace onnxruntime { @@ -18,6 +20,7 @@ struct RunOptions { bool include_dim_values_in_main_graph = false; int symbolic_dim_value_in_main_graph = -1; bool include_dim_values_in_subgraph = true; + bool mixed_execution_providers = false; }; static const ONNX_NAMESPACE::GraphProto CreateSubgraph(bool then_branch, const RunOptions& options); @@ -86,8 +89,6 @@ class IfOpTester : public OpTester { // add Identity node so if_graph_input_0 comes from graph inputs { - MTypeProto map_type; - inputs = {if_input}; outputs = {&graph.GetOrCreateNodeArg("if_input_0", if_input->TypeAsProto())}; graph.AddNode("identity", "Identity", "Pass if input through from graph inputs.", inputs, outputs); @@ -201,7 +202,17 @@ void RunTest(bool condition_value, test.AddOutput("if_out_0", output_shape, {11.f}); } - test.Run(expect_result, failure_message); + if (options.mixed_execution_providers) { + // we want the CUDA provider to be first, and the CPU provider second. all except the Scannode should run on + // CUDA given that, which creates the scenario where we need to copy to/from CPU to execute the Scan node correctly. + std::vector> execution_providers; + execution_providers.push_back(DefaultCudaExecutionProvider()); + execution_providers.push_back(DefaultCpuExecutionProvider()); + + test.Run(expect_result, failure_message, {}, nullptr, &execution_providers); + } else { + test.Run(expect_result, failure_message); + } } TEST(If, ShapeInMainGraph_NoShapeInSubgraph_True) { @@ -236,6 +247,14 @@ TEST(If, NoShapeInMainGraph_ShapeInSubgraph_False) { RunTest(false, options); } +#ifdef USE_CUDA +TEST(If, MixedExecutionProviders) { + RunOptions options{}; + options.mixed_execution_providers = true; + RunTest(true, options); +} +#endif // USE_CUDA + /* These tests require subgraphs with nodes that support symbolic dimensions. 'Add' does not. @@ -260,5 +279,6 @@ TEST(If, SymbolicShapeInMainGraph_NoShapeInSubgraph_False) { RunTest(false, options); } */ + } // namespace test } // namespace onnxruntime diff --git a/onnxruntime/test/providers/cpu/controlflow/loop_test.cc b/onnxruntime/test/providers/cpu/controlflow/loop_test.cc index 98e2ddf98b..7f536a288b 100644 --- a/onnxruntime/test/providers/cpu/controlflow/loop_test.cc +++ b/onnxruntime/test/providers/cpu/controlflow/loop_test.cc @@ -7,9 +7,11 @@ #include "core/common/logging/logging.h" #include "core/framework/session_state.h" -#include "test/providers/provider_test_utils.h" #include "core/session/inference_session.h" +#include "test/providers/provider_test_utils.h" +#include "test/util/include/default_providers.h" + using namespace ONNX_NAMESPACE; namespace onnxruntime { @@ -19,6 +21,7 @@ struct RunOptions { bool include_dim_values_in_main_graph = true; bool include_dim_values_in_subgraph = false; bool include_types_in_subgraph = false; + bool mixed_execution_providers = false; }; static const ONNX_NAMESPACE::GraphProto CreateSubgraph(const RunOptions& options); @@ -308,7 +311,17 @@ void RunTest(int64_t max_iterations, test.AddOutput("loop_var_1_final", loop_var_1_final_shape, loop_var_1_final); test.AddOutput("loop_out_0_final", loop_out_0_final_shape, loop_out_0_final); - test.Run(expect_result, failure_message); + if (options.mixed_execution_providers) { + // we want the CUDA provider to be first, and the CPU provider second. all except the Loop node should run on + // CUDA given that, which creates the scenario where we need to copy to/from CPU to execute the Loop node correctly. + std::vector> execution_providers; + execution_providers.push_back(DefaultCudaExecutionProvider()); + execution_providers.push_back(DefaultCpuExecutionProvider()); + + test.Run(expect_result, failure_message, {}, nullptr, &execution_providers); + } else { + test.Run(expect_result, failure_message); + } } // exit due to hitting condition that the sum is < kSumMax which is 8 @@ -474,5 +487,15 @@ TEST(Loop, InfiniteLoopTermination) { terminator_thread.join(); } +#ifdef USE_CUDA +// test that when part of the subgraph run on CUDA it executes successfully +TEST(Loop, MixedExecutionProviders) { + RunOptions options{}; + options.mixed_execution_providers = true; + + ExitDueToCond(options); +} +#endif + } // namespace test } // namespace onnxruntime diff --git a/onnxruntime/test/providers/cpu/controlflow/scan_test.cc b/onnxruntime/test/providers/cpu/controlflow/scan_test.cc index bce792bd8c..f017a1e282 100644 --- a/onnxruntime/test/providers/cpu/controlflow/scan_test.cc +++ b/onnxruntime/test/providers/cpu/controlflow/scan_test.cc @@ -4,9 +4,11 @@ #include "gtest/gtest.h" #include "gmock/gmock.h" #include "core/framework/session_state.h" -#include "test/providers/provider_test_utils.h" #include "core/session/inference_session.h" +#include "test/providers/provider_test_utils.h" +#include "test/util/include/default_providers.h" + using namespace ONNX_NAMESPACE; namespace onnxruntime { @@ -20,6 +22,7 @@ struct RunOptions { bool include_outer_scope_add = false; bool scalar_loop_state_value = false; bool add_bad_shape = false; + bool mixed_execution_providers = false; }; static void CreateSubgraph(Graph& graph, RunOptions& options, const std::string& failure_message = ""); @@ -368,7 +371,17 @@ static void RunTest_v9(const std::string test_name, int64_t sequence_len, int64_ test.AddOutput("scan_output_2", output_shape, output_2); test.AddOutput("scan_output_3", output_shape, output_3); - test.Run(expect_result, failure_message); + if (options.mixed_execution_providers) { + // we want the CUDA provider to be first, and the CPU provider second. all except the Scannode should run on + // CUDA given that, which creates the scenario where we need to copy to/from CPU to execute the Scan node correctly. + std::vector> execution_providers; + execution_providers.push_back(DefaultCudaExecutionProvider()); + execution_providers.push_back(DefaultCpuExecutionProvider()); + + test.Run(expect_result, failure_message, {}, nullptr, &execution_providers); + } else { + test.Run(expect_result, failure_message); + } } static void ShortSequenceOneInBatchOneLoopStateVar(const RunOptions& options, const std::string& expected_error = "") { @@ -1018,5 +1031,15 @@ void UnknownDimInSubgraphOutput(bool is_v8) { TEST_8_AND_9(UnknownDimInSubgraphOutput); +#ifdef USE_CUDA +TEST(Scan, MixedExecutionProviders) { + RunOptions options{}; + options.is_v8 = false; + options.mixed_execution_providers = true; + + ShortSequenceOneInBatchOneLoopStateVar(options); +} +#endif + } // namespace test } // namespace onnxruntime diff --git a/onnxruntime/test/providers/cpu/rnn/rnn_op_test.cc b/onnxruntime/test/providers/cpu/rnn/rnn_op_test.cc index 43e3aa4976..a7412d32e2 100644 --- a/onnxruntime/test/providers/cpu/rnn/rnn_op_test.cc +++ b/onnxruntime/test/providers/cpu/rnn/rnn_op_test.cc @@ -26,7 +26,7 @@ void MemoryLayoutTransposeRNNInputCNTKToONNXRuntime(const T* X_data_cntk, T* X_d // onnxruntime takes output of shape [seq_length, num_directions, batch_size, hidden_size) template void MemoryLayoutTransposeRNNOutputCNTKToONNXRuntime(const T* X_data_cntk, T* X_data_onnx, - int64_t seq_length, int64_t num_directions, int64_t batch_size, int64_t hidden_size) { + int64_t seq_length, int64_t num_directions, int64_t batch_size, int64_t hidden_size) { for (int seq = 0; seq < seq_length; seq++) { for (int dir = 0; dir < num_directions; dir++) { for (int batch = 0; batch < batch_size; batch++) { @@ -110,7 +110,7 @@ TEST(RNNTest, RNN_bidirectional_bias_initial_zigged_batch) { 0.0F, 0.0F, 0.0F, 0.0F, 0.0F, 0.0F}); std::vector Y_data(seq_length * num_directions * batch_size * hidden_size); MemoryLayoutTransposeRNNOutputCNTKToONNXRuntime(&Y_data_in_batchs[0], &Y_data[0], - seq_length, num_directions, batch_size, hidden_size); + seq_length, num_directions, batch_size, hidden_size); test.AddOutput("Y", Y_dims, Y_data); std::vector Y_h_dims{num_directions, batch_size, hidden_size}; @@ -187,7 +187,7 @@ TEST(RNNTest, RNN_bidirectional_zigged_batch) { 0.0F, 0.0F, 0.0F, 0.0F, 0.0F, 0.0F}); std::vector Y_data(seq_length * num_directions * batch_size * hidden_size); MemoryLayoutTransposeRNNOutputCNTKToONNXRuntime(&Y_data_in_batchs[0], &Y_data[0], - seq_length, num_directions, batch_size, hidden_size); + seq_length, num_directions, batch_size, hidden_size); test.AddOutput("Y", Y_dims, Y_data); std::vector Y_h_dims{num_directions, batch_size, hidden_size}; @@ -261,7 +261,7 @@ TEST(RNNTest, RNN_reverse_direction_zigged_batch) { 0.0F, 0.0F, 0.0F}); std::vector Y_data(seq_length * num_directions * batch_size * hidden_size); MemoryLayoutTransposeRNNOutputCNTKToONNXRuntime(&Y_data_in_batchs[0], &Y_data[0], - seq_length, num_directions, batch_size, hidden_size); + seq_length, num_directions, batch_size, hidden_size); test.AddOutput("Y", Y_dims, Y_data); std::vector Y_h_dims{num_directions, batch_size, hidden_size}; @@ -335,7 +335,7 @@ TEST(RNNTest, RNN_forward_direction_zigged_batch) { 0.0F, 0.0F, 0.0F}); std::vector Y_data(seq_length * num_directions * batch_size * hidden_size); MemoryLayoutTransposeRNNOutputCNTKToONNXRuntime(&Y_data_in_batchs[0], &Y_data[0], - seq_length, num_directions, batch_size, hidden_size); + seq_length, num_directions, batch_size, hidden_size); test.AddOutput("Y", Y_dims, Y_data); std::vector Y_h_dims{num_directions, batch_size, hidden_size}; @@ -682,7 +682,8 @@ TEST(RNNTest, RNN_invalid_sequence_lens) { std::vector Y_h_data{0.f, 0.f, 0.f, 0.f, 0.f, 0.f}; test.AddOutput("Y_h", Y_h_dims, Y_h_data); - test.Run(OpTester::ExpectResult::kExpectFailure, error_msg); + // the CUDA RNN version allows the invalid sequence lengths, so disable testing on CUDA + test.Run(OpTester::ExpectResult::kExpectFailure, error_msg, {kCudaExecutionProvider}); }; // should batch batch_size to be valid diff --git a/onnxruntime/test/providers/provider_test_utils.cc b/onnxruntime/test/providers/provider_test_utils.cc index a4e5080f5e..4cc35d1ef9 100644 --- a/onnxruntime/test/providers/provider_test_utils.cc +++ b/onnxruntime/test/providers/provider_test_utils.cc @@ -240,10 +240,100 @@ std::unique_ptr OpTester::BuildGraph() { return p_model; } +void OpTester::ExecuteModel(Model& model, + InferenceSession& session_object, + ExpectResult expect_result, + const std::string& expected_failure_string, + const RunOptions* run_options, + std::unordered_map feeds, + std::vector output_names, + const std::string& provider_type) { + std::stringstream s1; + model.ToProto().SerializeToOstream(&s1); + auto status = session_object.Load(s1); + EXPECT_TRUE(status.IsOK()) << status.ErrorMessage(); + if (!status.IsOK()) { + LOGS_DEFAULT(ERROR) << "Load failed with status: " << status.ErrorMessage(); + return; + } + + status = session_object.Initialize(); + if (!status.IsOK()) { + if (expect_result == ExpectResult::kExpectFailure) { + EXPECT_TRUE(!status.IsOK()); + EXPECT_THAT(status.ErrorMessage(), testing::HasSubstr(expected_failure_string)); + } else { + LOGS_DEFAULT(ERROR) << "Initialize failed with status: " << status.ErrorMessage(); + EXPECT_TRUE(status.IsOK()) << status.ErrorMessage(); + } + } + + if (!status.IsOK()) { + return; + } + + RunOptions default_run_options{}; + default_run_options.run_tag = op_; + default_run_options.run_log_verbosity_level = 1; + + std::vector fetches; + status = session_object.Run(run_options ? *run_options : default_run_options, feeds, output_names, &fetches); + if (status.IsOK()) { + EXPECT_TRUE(expect_result == ExpectResult::kExpectSuccess); + if (expect_result == ExpectResult::kExpectFailure) { + return; + } + } else { + if (expect_result == ExpectResult::kExpectFailure) { + EXPECT_THAT(status.ErrorMessage(), testing::HasSubstr(expected_failure_string)); + } else { + LOGS_DEFAULT(ERROR) << "Run failed with status: " << status.ErrorMessage(); + EXPECT_TRUE(status.IsOK()) << status.ErrorMessage(); + } + return; + } + + // Verify the outputs + // Todo: support check output with map/sequence/.... + size_t idx = 0; + for (auto& expected_data : output_data_) { + MLValue& mlvalue = fetches[idx]; + if (mlvalue.Fence()) + mlvalue.Fence()->BeforeUsingAsInput(onnxruntime::kCpuExecutionProvider, 0); + + if (expected_data.def_.Exists()) { // optional outputs won't exist + if (expected_data.data_.IsTensor()) { + // verify output shape inference when input defs have shape + if (add_shape_to_tensor_data_) { + auto out_shape_proto = expected_data.def_.Shape(); + EXPECT_TRUE(out_shape_proto != nullptr); + auto inferred_dims = utils::GetTensorShapeFromTensorShapeProto(*out_shape_proto); + const auto& expected_shape = expected_data.data_.Get().Shape(); + EXPECT_TRUE(inferred_dims.size() == expected_shape.NumDimensions()); + for (int d = 0; d < inferred_dims.size(); ++d) { + // check equal unless the input involved a symbolic dimension + if (inferred_dims[d] != -1) + EXPECT_EQ(expected_shape[d], inferred_dims[d]) << "Output idx = " << idx << " dim = " << d; + } + } + Check(expected_data, mlvalue.Get(), provider_type); + } else { + Check(expected_data, mlvalue, provider_type); + } + ++idx; + + // skip missing trailing optional outputs + if (idx == fetches.size()) + break; + } + } +} + void OpTester::Run(ExpectResult expect_result, const std::string& expected_failure_string, const std::unordered_set& excluded_provider_types, - const RunOptions* run_options) { + const RunOptions* run_options, + std::vector>* execution_providers) { try { #ifndef NDEBUG run_called_ = true; @@ -297,135 +387,74 @@ void OpTester::Run(ExpectResult expect_result, bool has_run = false; - for (const std::string& provider_type : all_provider_types) { - if (excluded_provider_types.count(provider_type) > 0) - continue; - + if (execution_providers) { InferenceSession session_object{so}; - for (auto& custom_session_registry : custom_session_registries_) - session_object.RegisterCustomRegistry(custom_session_registry); + ASSERT_TRUE(!execution_providers->empty()) << "Empty execution providers vector."; + std::string provider_types; - std::unique_ptr execution_provider; - if (provider_type == onnxruntime::kCpuExecutionProvider) - execution_provider = DefaultCpuExecutionProvider(); - else if (provider_type == onnxruntime::kCudaExecutionProvider) - execution_provider = DefaultCudaExecutionProvider(); - else if (provider_type == onnxruntime::kMklDnnExecutionProvider) - execution_provider = DefaultMkldnnExecutionProvider(); - else if (provider_type == onnxruntime::kNupharExecutionProvider) - execution_provider = DefaultNupharExecutionProvider(); - else if (provider_type == onnxruntime::kBrainSliceExecutionProvider) - execution_provider = DefaultBrainSliceExecutionProvider(); - // skip if execution provider is disabled - if (execution_provider == nullptr) - continue; + for (auto& entry : *execution_providers) { + provider_types += entry->Type() + ":"; + EXPECT_TRUE(session_object.RegisterExecutionProvider(std::move(entry)).IsOK()); + } - bool valid = true; - - // set execution provider for all nodes in the graph - for (auto& node : graph.Nodes()) { - if (node.OpType() == kConstant) + ExecuteModel(*p_model, session_object, expect_result, expected_failure_string, run_options, + feeds, output_names, provider_types); + } else { + for (const std::string& provider_type : all_provider_types) { + if (excluded_provider_types.count(provider_type) > 0) continue; - //if node is not registered for the provider, skip - node.SetExecutionProviderType(provider_type); - auto reg = execution_provider->GetKernelRegistry(); - const KernelCreateInfo* kci = reg->TryFindKernel(node, execution_provider->Type()); - if (!kci) { - valid = false; - break; - } - } + InferenceSession session_object{so}; - if (!valid) - continue; + for (auto& custom_session_registry : custom_session_registries_) + session_object.RegisterCustomRegistry(custom_session_registry); - has_run = true; + std::unique_ptr execution_provider; + if (provider_type == onnxruntime::kCpuExecutionProvider) + execution_provider = DefaultCpuExecutionProvider(); + else if (provider_type == onnxruntime::kCudaExecutionProvider) + execution_provider = DefaultCudaExecutionProvider(); + else if (provider_type == onnxruntime::kMklDnnExecutionProvider) + execution_provider = DefaultMkldnnExecutionProvider(); + else if (provider_type == onnxruntime::kNupharExecutionProvider) + execution_provider = DefaultNupharExecutionProvider(); + else if (provider_type == onnxruntime::kBrainSliceExecutionProvider) + execution_provider = DefaultBrainSliceExecutionProvider(); + // skip if execution provider is disabled + if (execution_provider == nullptr) + continue; - EXPECT_TRUE(session_object.RegisterExecutionProvider(std::move(execution_provider)).IsOK()); + bool valid = true; - std::stringstream s1; - p_model->ToProto().SerializeToOstream(&s1); - status = session_object.Load(s1); - EXPECT_TRUE(status.IsOK()) << status.ErrorMessage(); - if (!status.IsOK()) { - LOGS_DEFAULT(ERROR) << "Load failed with status: " << status.ErrorMessage(); - return; - } + // set execution provider for all nodes in the graph + for (auto& node : graph.Nodes()) { + if (node.OpType() == kConstant) + continue; - status = session_object.Initialize(); - if (!status.IsOK()) { - if (expect_result == ExpectResult::kExpectFailure) { - EXPECT_TRUE(!status.IsOK()); - EXPECT_THAT(status.ErrorMessage(), testing::HasSubstr(expected_failure_string)); - } else { - LOGS_DEFAULT(ERROR) << "Initialize failed with status: " << status.ErrorMessage(); - EXPECT_TRUE(status.IsOK()) << status.ErrorMessage(); - } - } - if (!status.IsOK()) { - return; - } - - RunOptions default_run_options{}; - default_run_options.run_tag = op_; - default_run_options.run_log_verbosity_level = 1; - - std::vector fetches; - status = session_object.Run(run_options ? *run_options : default_run_options, feeds, output_names, &fetches); - if (status.IsOK()) { - EXPECT_TRUE(expect_result == ExpectResult::kExpectSuccess); - if (expect_result == ExpectResult::kExpectFailure) { - return; - } - } else { - if (expect_result == ExpectResult::kExpectFailure) { - EXPECT_THAT(status.ErrorMessage(), testing::HasSubstr(expected_failure_string)); - } else { - LOGS_DEFAULT(ERROR) << "Run failed with status: " << status.ErrorMessage(); - EXPECT_TRUE(status.IsOK()) << status.ErrorMessage(); - } - return; - } - - // Verify the outputs - // Todo: support check output with map/sequence/.... - size_t idx = 0; - for (auto& expected_data : output_data_) { - MLValue& mlvalue = fetches[idx]; - if (mlvalue.Fence()) - mlvalue.Fence()->BeforeUsingAsInput(onnxruntime::kCpuExecutionProvider, 0); - - if (expected_data.def_.Exists()) { // optional outputs won't exist - if (expected_data.data_.IsTensor()) { - // verify output shape inference when input defs have shape - if (add_shape_to_tensor_data_) { - auto out_shape_proto = expected_data.def_.Shape(); - EXPECT_TRUE(out_shape_proto != nullptr); - auto inferred_dims = utils::GetTensorShapeFromTensorShapeProto(*out_shape_proto); - const auto& expected_shape = expected_data.data_.Get().Shape(); - EXPECT_TRUE(inferred_dims.size() == expected_shape.NumDimensions()); - for (int d = 0; d < inferred_dims.size(); ++d) { - // check equal unless the input involved a symbolic dimension - if (inferred_dims[d] != -1) - EXPECT_EQ(expected_shape[d], inferred_dims[d]) << "Output idx = " << idx << " dim = " << d; - } - } - Check(expected_data, mlvalue.Get(), provider_type); - } else { - Check(expected_data, mlvalue, provider_type); - } - ++idx; - - // skip missing trailing optional outputs - if (idx == fetches.size()) + //if node is not registered for the provider, skip + node.SetExecutionProviderType(provider_type); + auto reg = execution_provider->GetKernelRegistry(); + const KernelCreateInfo* kci = reg->TryFindKernel(node, execution_provider->Type()); + if (!kci) { + valid = false; break; + } } - } - } - EXPECT_TRUE(has_run) << "No registered execution providers were able to run the model."; + if (!valid) + continue; + + has_run = true; + + EXPECT_TRUE(session_object.RegisterExecutionProvider(std::move(execution_provider)).IsOK()); + + ExecuteModel(*p_model, session_object, expect_result, expected_failure_string, run_options, + feeds, output_names, provider_type); + } + + EXPECT_TRUE(has_run) << "No registered execution providers were able to run the model."; + } } catch (const std::exception& ex) { std::cerr << ex.what(); // rethrow as some tests for error handling expect this diff --git a/onnxruntime/test/providers/provider_test_utils.h b/onnxruntime/test/providers/provider_test_utils.h index b2056b856a..fe356b29fe 100644 --- a/onnxruntime/test/providers/provider_test_utils.h +++ b/onnxruntime/test/providers/provider_test_utils.h @@ -22,6 +22,8 @@ #include namespace onnxruntime { +class InferenceSession; + namespace test { // unfortunately std::optional is in C++17 so use a miniversion of it template @@ -238,7 +240,8 @@ class OpTester { void Run(ExpectResult expect_result = ExpectResult::kExpectSuccess, const std::string& expected_failure_string = "", const std::unordered_set& excluded_provider_types = {}, - const RunOptions* run_options = nullptr); + const RunOptions* run_options = nullptr, + std::vector>* execution_providers = nullptr); struct Data { onnxruntime::NodeArg def_; @@ -307,6 +310,15 @@ class OpTester { } } + void ExecuteModel(Model& model, + InferenceSession& session_object, + ExpectResult expect_result, + const std::string& expected_failure_string, + const RunOptions* run_options, + std::unordered_map feeds, + std::vector output_names, + const std::string& provider_type); + const char* domain_; int opset_version_; bool add_shape_to_tensor_data_ = true;