diff --git a/onnxruntime/core/providers/common.h b/onnxruntime/core/providers/common.h index fca366b295..77a5b468c4 100644 --- a/onnxruntime/core/providers/common.h +++ b/onnxruntime/core/providers/common.h @@ -137,9 +137,9 @@ inline Status ComputePadAndOutputShape(const int64_t in_dim, return Status::OK(); } -template -inline bool Contains(const Map& map, const Key& key) { - return map.find(key) != map.end(); +template +inline bool Contains(const AssociativeContainer& container, const Key& key) { + return container.find(key) != container.end(); } // Note: This helper function will not have overflow protection diff --git a/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/helper.cc b/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/helper.cc index 4646d1cb08..c15241ebc6 100644 --- a/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/helper.cc +++ b/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/helper.cc @@ -367,13 +367,9 @@ void GetFlattenOutputShape(const Node& node, const Shape& input_shape, int32_t& dim_2 = std::accumulate(input_shape.cbegin() + axis, input_shape.cend(), 1, std::multiplies()); } -bool IsValidSupportedNodesGroup(const std::vector& supported_node_group, const GraphViewer& graph_viewer) { - if (supported_node_group.empty()) - return false; - - if (supported_node_group.size() == 1) { - const auto& node_indices = graph_viewer.GetNodesInTopologicalOrder(); - const auto* node(graph_viewer.GetNode(node_indices[supported_node_group[0]])); +bool IsValidSupportedNodeGroup(const std::vector& supported_node_partition) { + if (supported_node_partition.size() == 1) { + const auto* node = supported_node_partition[0]; const auto& op = node->OpType(); // It is not worth it to perform a single Reshape/Flatten/Identity operator // which is only copying the data in NNAPI @@ -440,9 +436,9 @@ bool IsNodeSupported(const Node& node, const GraphViewer& graph_viewer, const Op return op_support_checker->IsOpSupported(graph_viewer.GetAllInitializedTensors(), node, params); } -bool IsNodeSupportedInternal(const Node& node, const GraphViewer& graph_viewer, - const OpSupportCheckParams& params, - const std::unordered_set& node_outputs_in_group) { +bool IsNodeSupportedInGroup(const Node& node, const GraphViewer& graph_viewer, + const OpSupportCheckParams& params, + const std::unordered_set& node_outputs_in_group) { if (!IsNodeSupported(node, graph_viewer, params)) return false; @@ -459,7 +455,7 @@ bool IsInputSupported(const NodeArg& input, const std::string& parent_name) { // We do not support input with no shape if (!shape_proto) { LOGS_DEFAULT(VERBOSE) << "Input [" << input_name << "] of [" << parent_name - << "] has not shape"; + << "] has no shape"; return false; } @@ -474,61 +470,6 @@ bool IsInputSupported(const NodeArg& input, const std::string& parent_name) { return true; } -std::vector> GetSupportedNodes(const GraphViewer& graph_viewer, const OpSupportCheckParams& params) { - std::vector> supported_node_groups; - if (params.android_feature_level < ORT_NNAPI_MIN_API_LEVEL) { - LOGS_DEFAULT(WARNING) << "All ops will fallback to CPU EP, because system NNAPI feature level [" - << params.android_feature_level - << "] is lower than minimal supported NNAPI API feature level [" - << ORT_NNAPI_MIN_API_LEVEL - << "] of this build for NNAPI"; - return supported_node_groups; - } - - // Disable NNAPI if the graph has input with dynamic shape - for (const auto* input : graph_viewer.GetInputs()) { - if (!IsInputSupported(*input, "graph")) { - return supported_node_groups; - } - } - - // This holds the supported node's topological index - std::vector supported_node_group; - // This holds the NodeIndex of the nodes in the above group - std::unordered_set node_outputs_in_group; - const auto& node_indices = graph_viewer.GetNodesInTopologicalOrder(); - for (size_t i = 0; i < node_indices.size(); i++) { - const auto* node(graph_viewer.GetNode(node_indices[i])); - bool supported = IsNodeSupportedInternal(*node, graph_viewer, params, node_outputs_in_group); - LOGS_DEFAULT(VERBOSE) << "Operator type: [" << node->OpType() - << "] index: [" << i - << "] name: [" << node->Name() - << "] supported: [" << supported - << "]"; - if (supported) { - supported_node_group.push_back(i); - - // We want to put all the output names of nodes in the current group for easy query - // See IsInternalQuantizationSupported() - for (const auto* output : node->OutputDefs()) { - node_outputs_in_group.insert(output->Name()); - } - } else { - if (IsValidSupportedNodesGroup(supported_node_group, graph_viewer)) { - supported_node_groups.push_back(supported_node_group); - } - - supported_node_group.clear(); - node_outputs_in_group.clear(); - } - } - - if (IsValidSupportedNodesGroup(supported_node_group, graph_viewer)) - supported_node_groups.push_back(supported_node_group); - - return supported_node_groups; -} - std::string Shape2String(const std::vector& shape) { std::ostringstream os; os << "[ "; diff --git a/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/helper.h b/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/helper.h index 71ff61d3e7..f73bf7f2d2 100644 --- a/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/helper.h +++ b/onnxruntime/core/providers/nnapi/nnapi_builtin/builders/helper.h @@ -126,8 +126,17 @@ void GetFlattenOutputShape(const Node& node, const Shape& input_shape, int32_t& // If a node is supported by NNAPI bool IsNodeSupported(const Node& node, const GraphViewer& graph_viewer, const OpSupportCheckParams& params); -// Get a list of groups of supported nodes, each group represents a subgraph supported by NNAPI EP -std::vector> GetSupportedNodes(const GraphViewer& graph_viewer, const OpSupportCheckParams& params); +// If a node is supported by NNAPI in a partition node group +// `node_outputs_in_group` is the set of the output names of the nodes added to this group so far +bool IsNodeSupportedInGroup(const Node& node, const GraphViewer& graph_viewer, + const OpSupportCheckParams& params, + const std::unordered_set& node_outputs_in_group); + +// If a graph input is supported by NNAPI +bool IsInputSupported(const NodeArg& input, const std::string& parent_name); + +// If an NNAPI partition node group is valid +bool IsValidSupportedNodeGroup(const std::vector& supported_node_group); // Get string representation of a Shape std::string Shape2String(const std::vector& shape); diff --git a/onnxruntime/core/providers/nnapi/nnapi_builtin/nnapi_execution_provider.cc b/onnxruntime/core/providers/nnapi/nnapi_builtin/nnapi_execution_provider.cc index 6a049ccb24..da17d487dc 100644 --- a/onnxruntime/core/providers/nnapi/nnapi_builtin/nnapi_execution_provider.cc +++ b/onnxruntime/core/providers/nnapi/nnapi_builtin/nnapi_execution_provider.cc @@ -1,28 +1,36 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. -#include "nnapi_execution_provider.h" +#include "core/providers/nnapi/nnapi_builtin/nnapi_execution_provider.h" -#include "builders/helper.h" -#include "builders/op_support_checker.h" #include "core/framework/allocatormgr.h" #include "core/framework/compute_capability.h" #include "core/graph/graph_viewer.h" +#include "core/providers/common.h" +#include "core/providers/nnapi/nnapi_builtin/builders/helper.h" +#include "core/providers/nnapi/nnapi_builtin/builders/op_support_checker.h" +#include "core/providers/nnapi/nnapi_builtin/nnapi_lib/nnapi_implementation.h" +#include "core/providers/partitioning_utils.h" #include "core/session/onnxruntime_cxx_api.h" -#include "nnapi_lib/nnapi_implementation.h" #ifdef __ANDROID__ -#include "model.h" -#include "builders/model_builder.h" +#include "core/providers/nnapi/nnapi_builtin/builders/model_builder.h" +#include "core/providers/nnapi/nnapi_builtin/model.h" #endif namespace onnxruntime { constexpr const char* NNAPI = "Nnapi"; +constexpr std::array kDefaultPartitioningStopOps{ + "NonMaxSuppression", +}; + NnapiExecutionProvider::NnapiExecutionProvider(uint32_t nnapi_flags) - : IExecutionProvider{onnxruntime::kNnapiExecutionProvider}, - nnapi_flags_(nnapi_flags) { + : IExecutionProvider{onnxruntime::kNnapiExecutionProvider, true}, + nnapi_flags_(nnapi_flags), + // TODO make this configurable + partitioning_stop_ops_(kDefaultPartitioningStopOps.begin(), kDefaultPartitioningStopOps.end()) { AllocatorCreationInfo device_info( [](int) { return std::make_unique(OrtMemoryInfo(NNAPI, OrtAllocatorType::OrtDeviceAllocator)); @@ -42,165 +50,104 @@ NnapiExecutionProvider::NnapiExecutionProvider(uint32_t nnapi_flags) NnapiExecutionProvider::~NnapiExecutionProvider() {} std::vector> -NnapiExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph_view, +NnapiExecutionProvider::GetCapability(const onnxruntime::GraphViewer& graph_viewer, const std::vector& /*kernel_registries*/) const { std::vector> result; // TODO: Task 812756: NNAPI EP, add support for subgraph (If and Loop operators) - if (graph_view.IsSubgraph()) { + if (graph_viewer.IsSubgraph()) { return result; } - std::unordered_set all_node_inputs; - for (const auto& node : graph_view.Nodes()) { - for (auto* input : node.InputDefs()) { - all_node_inputs.insert(input->Name()); - } - } - // We need to get the Android system API level to ensure the GetCapability giving the correct result // based on the system. // If we are actually running on Android system, we can get the API level by querying the system // However, since we also allow the NNAPI EP run GetCapability for model conversion on a non-Android system, - // since we cannot get the runtime system API level, we have to specify it using complie definition. - int32_t android_feature_level; + // since we cannot get the runtime system API level, we have to specify it using compile definition. + static const int32_t android_feature_level = []() { #ifdef __ANDROID__ - const auto* _nnapi = NnApiImplementation(); - android_feature_level = _nnapi->nnapi_runtime_feature_level; + const auto* nnapi = NnApiImplementation(); + return nnapi->nnapi_runtime_feature_level; #else - android_feature_level = ORT_NNAPI_MAX_SUPPORTED_API_LEVEL; + return ORT_NNAPI_MAX_SUPPORTED_API_LEVEL; #endif + }(); - nnapi::OpSupportCheckParams params{ + const nnapi::OpSupportCheckParams params{ android_feature_level, !!(nnapi_flags_ & NNAPI_FLAG_USE_NCHW), }; - const auto supported_nodes_vector = GetSupportedNodes(graph_view, params); - size_t num_of_supported_nodes = 0; - - // Find inputs, initializers and outputs for each supported subgraph - const std::vector& node_index = graph_view.GetNodesInTopologicalOrder(); - const auto& graph_outputs = graph_view.GetOutputs(); - for (const auto& group : supported_nodes_vector) { - if (group.empty()) - continue; - - num_of_supported_nodes += group.size(); - LOGS_DEFAULT(VERBOSE) << "NnapiExecutionProvider::GetCapability, current supported node group size: " - << group.size(); - - std::unordered_set node_set; - node_set.reserve(group.size()); - for (const auto& index : group) { - node_set.insert(node_index[index]); - } - - std::unique_ptr sub_graph = std::make_unique(); - // Find inputs and outputs of the subgraph - std::unordered_map fused_inputs, fused_outputs, fused_outputs_to_add; - std::unordered_set erased; - int input_order = 0; - int output_order = 0; - - for (const auto& index : group) { - sub_graph->nodes.push_back(node_index[index]); - const auto* node = graph_view.GetNode(node_index[index]); - - for (const auto* input : node->InputDefs()) { - const auto it = fused_outputs.find(input); - if (it != fused_outputs.end()) { - fused_outputs.erase(it); - erased.insert(input); - } - //only when input is neither in output list nor erased list, add the input to input list - else if (erased.find(input) == erased.end()) { - fused_inputs[input] = input_order++; - } - } - - // For output searching, there is a special case: - // If certain output is used more than once, - // if the output is connected to nodes that don't belong to the subgraph, the output need to be added - // to the output list - - std::unordered_set processed_outputs; - for (auto it = node->OutputEdgesBegin(), end = node->OutputEdgesEnd(); it != end; ++it) { - const auto node_idx = it->GetNode().Index(); - const auto* output = node->OutputDefs()[it->GetSrcArgIndex()]; - - if (node_set.find(node_idx) != node_set.end()) { - const auto iter = fused_inputs.find(output); - if (iter != fused_inputs.end()) { - fused_inputs.erase(iter); - erased.insert(output); - } else if (erased.find(output) == erased.end()) { - fused_outputs[output] = output_order++; - } - } else { - fused_outputs_to_add[output] = output_order++; - } - - processed_outputs.insert(output); - } - - for (const auto* output : node->OutputDefs()) { - if (processed_outputs.find(output) != processed_outputs.end()) - continue; - - const auto iter = fused_inputs.find(output); - if (iter != fused_inputs.end()) { - fused_inputs.erase(iter); - erased.insert(output); - } - // only when output is neither in input list nor erased list, add the output to output list - else if (erased.find(output) == erased.end() && output->Exists()) { - fused_outputs[output] = output_order++; - } - } - } - - fused_outputs.insert(fused_outputs_to_add.begin(), fused_outputs_to_add.end()); - // Sort inputs and outputs by the order they were added - std::multimap inputs, outputs; - - for (auto it = fused_inputs.begin(), end = fused_inputs.end(); it != end; ++it) { - inputs.insert(std::pair(it->second, it->first)); - } - - for (auto it = fused_outputs.begin(), end = fused_outputs.end(); it != end; ++it) { - if (all_node_inputs.find(it->first->Name()) != all_node_inputs.end()) { - outputs.insert(std::pair(it->second, it->first)); - } else if (std::find(graph_outputs.begin(), graph_outputs.end(), it->first) != graph_outputs.end()) { - outputs.insert(std::pair(it->second, it->first)); - } - } - - // Assign inputs and outputs to subgraph's meta_def - auto meta_def = std::make_unique<::onnxruntime::IndexedSubGraph::MetaDef>(); - meta_def->name = "NNAPI_" + std::to_string(metadef_id_++); - meta_def->domain = kMSDomain; - - for (const auto& input : inputs) { - meta_def->inputs.push_back(input.second->Name()); - } - - for (const auto& output : outputs) { - meta_def->outputs.push_back(output.second->Name()); - } - - // meta_def->status = ONNX_NAMESPACE::EXPERIMENTAL; - meta_def->since_version = 1; - sub_graph->SetMetaDef(std::move(meta_def)); - - result.push_back(std::make_unique(std::move(sub_graph))); + if (params.android_feature_level < ORT_NNAPI_MIN_API_LEVEL) { + LOGS_DEFAULT(WARNING) << "All ops will fallback to CPU EP, because system NNAPI feature level [" + << params.android_feature_level + << "] is lower than minimal supported NNAPI API feature level [" + << ORT_NNAPI_MIN_API_LEVEL + << "] of this build for NNAPI"; + return result; } - auto num_of_partitions = result.size(); + // Disable NNAPI if the graph has any unsupported inputs + for (const auto* input : graph_viewer.GetInputs()) { + if (!nnapi::IsInputSupported(*input, "graph")) { + return result; + } + } + + const auto excluded_nodes = utils::CreateExcludedNodeSet(graph_viewer, partitioning_stop_ops_); + const bool check_excluded_nodes = !excluded_nodes.empty(); + + std::unordered_set node_outputs_in_current_group{}; + + const auto is_node_supported = [&](const Node& node) -> bool { + const bool excluded = check_excluded_nodes && Contains(excluded_nodes, &node); + const bool supported = !excluded && + nnapi::IsNodeSupportedInGroup(node, graph_viewer, params, + node_outputs_in_current_group); + LOGS_DEFAULT(VERBOSE) << "Operator type: [" << node.OpType() + << "] index: [" << node.Index() + << "] name: [" << node.Name() + << "] supported: [" << supported + << "]"; + + if (supported) { + // We want to save all the output names of nodes in the current group for easy query + // See nnapi::IsNodeSupportedInGroup() + for (const auto* output : node.OutputDefs()) { + node_outputs_in_current_group.insert(output->Name()); + } + } + + return supported; + }; + + const auto on_group_closed = [&](const std::vector& group) -> bool { + // reset per-partition node group tracking + node_outputs_in_current_group.clear(); + return nnapi::IsValidSupportedNodeGroup(group); + }; + + const auto gen_metadef_name = [&]() { + uint64_t model_hash; + int metadef_id = GenerateMetaDefId(graph_viewer, model_hash); + return MakeString(NNAPI, "_", model_hash, "_", metadef_id); + }; + + result = utils::CreateSupportedPartitions(graph_viewer, is_node_supported, on_group_closed, + gen_metadef_name, NNAPI); + + const auto num_of_partitions = result.size(); + const auto num_of_supported_nodes = std::transform_reduce( + result.begin(), result.end(), + size_t{0}, std::plus<>{}, + [](const auto& partition) -> size_t { + return partition && partition->sub_graph ? partition->sub_graph->nodes.size() : 0; + }); + const auto summary_msg = MakeString( "NnapiExecutionProvider::GetCapability,", " number of partitions supported by NNAPI: ", num_of_partitions, - " number of nodes in the graph: ", graph_view.NumberOfNodes(), + " number of nodes in the graph: ", graph_viewer.NumberOfNodes(), " number of nodes supported by NNAPI: ", num_of_supported_nodes); // If the graph is partitioned in multiple subgraphs, and this may impact performance, diff --git a/onnxruntime/core/providers/nnapi/nnapi_builtin/nnapi_execution_provider.h b/onnxruntime/core/providers/nnapi/nnapi_builtin/nnapi_execution_provider.h index da283dcc56..a60bb0021e 100644 --- a/onnxruntime/core/providers/nnapi/nnapi_builtin/nnapi_execution_provider.h +++ b/onnxruntime/core/providers/nnapi/nnapi_builtin/nnapi_execution_provider.h @@ -31,13 +31,12 @@ class NnapiExecutionProvider : public IExecutionProvider { uint32_t GetNNAPIFlags() const { return nnapi_flags_; } private: - // unique counter to name each fused kernel across the entire model - mutable int metadef_id_{0}; - // The bit flags which define bool options for NNAPI EP, bits are defined as // NNAPIFlags in include/onnxruntime/core/providers/nnapi/nnapi_provider_factory.h const uint32_t nnapi_flags_; + const std::unordered_set partitioning_stop_ops_; + #ifdef __ANDROID__ std::unordered_map> nnapi_models_; #endif diff --git a/onnxruntime/core/providers/partitioning_utils.cc b/onnxruntime/core/providers/partitioning_utils.cc index bd85850aff..402513b445 100644 --- a/onnxruntime/core/providers/partitioning_utils.cc +++ b/onnxruntime/core/providers/partitioning_utils.cc @@ -3,97 +3,208 @@ #include "core/providers/partitioning_utils.h" +#include +#include +#include #include #include "core/framework/compute_capability.h" #include "core/framework/execution_provider.h" -#include "core/graph/model.h" +#include "core/graph/graph_viewer.h" +#include "core/providers/common.h" namespace onnxruntime { namespace utils { // internal helpers namespace { -// Kahn's topological sort with awareness of whether a node should be in a supported or unsupported partition -std::vector PartitionAwareTopoSort(const GraphViewer& graph_viewer, - const std::unordered_set& supported_nodes) { - std::queue supported_to_visit, unsupported_to_visit; - std::unordered_map in_degree; - std::vector topo_order; +#ifndef NDEBUG +std::string NodeDebugString(const Node& node) { + std::ostringstream oss; + oss << node.Index() << " '" << node.Name() << "'(" << node.OpType() << ")"; + return oss.str(); +} - auto num_nodes = graph_viewer.NumberOfNodes(); - topo_order.reserve(num_nodes); - in_degree.reserve(num_nodes); +template +std::string NodeGroupDebugString(const Container& group, bool show_all = false) { + static_assert(std::is_same_v); - auto add_to_visit = [&](const Node& node) { - if (supported_nodes.count(&node)) { - supported_to_visit.push(&node); - } else { - unsupported_to_visit.push(&node); + if (group.empty()) { + return ""; + } + + std::ostringstream oss; + oss << "<" << group.size() << (group.size() == 1 ? " node> " : " nodes> "); + if (show_all) { + auto node_it = group.begin(); + oss << NodeDebugString(**(node_it++)); + while (node_it != group.end()) { + oss << ", " << NodeDebugString(**(node_it++)); + } + } else { + const Node& start_node = *group.front(); + const Node& end_node = *group.back(); + oss << NodeDebugString(start_node) << " to " << NodeDebugString(end_node); + } + + return oss.str(); +} +#endif + +/** +Create partition node groups. + +A partition node group (a.k.a. a group) contains supported nodes that will run in a partition. + +All nodes in a group can be run together. This means that two nodes with an intervening unsupported node cannot be in +the same group. On the other hand, nodes within the same group do not necessarily have to be connected. + +The partitioning algorithm attempts to form the largest possible groups in a greedy fashion. It is a variant of Kahn's +topological sort algorithm that forms the group(s) as it goes. + +Conceptually, we consider nodes in a sequence of waves starting from the root nodes. One wave produces at most one +group. A wave flows over nodes in topological order, adding supported nodes to the current group, and stops at the +border of the current group. The next wave starts where the previous wave stopped. + +When generating the topological ordering, we maintain a set of nodes that have no inputs produced by unprocessed nodes. +From this set, we select the next node to process. + +When selecting the next node to process, we first take: +- a supported node (which will be part of the group) +- an unsupported node that does not consume an output of any node in the group + +The remaining unsupported nodes mark the border of the current group so they will be processed later when we consider +the next group. + +@param graph_viewer GraphViewer that IExecutionProvider::GetCapability is called with. +@param is_node_supported_fn Callback to check whether a node is supported. +@param on_group_closed_fn Callback to indicate a completed partition node group. +@param debug_output Print diagnostic output about the partitions and reasons for partition breaks. + No-op in a release build. +@return The partition node groups. +*/ +std::vector> CreateSupportedPartitionNodeGroups( + const GraphViewer& graph_viewer, + const IsNodeSupportedFn& is_node_supported_fn, + const OnGroupClosedFn& on_group_closed_fn, + bool debug_output) { +#ifdef NDEBUG + ORT_UNUSED_PARAMETER(debug_output); +#endif + + ORT_ENFORCE(is_node_supported_fn, "Node support test is required."); + + std::vector> supported_groups{}; + + // number of inputs from unprocessed nodes (in-degree) per node + std::unordered_map in_degree{}; + // nodes that are ready to process + std::deque nodes_to_process{}; + // nodes that will be processed when considering the next partition node group + std::deque nodes_to_process_with_next_group{}; + + // initialize in-degrees and find root nodes + for (const auto& node : graph_viewer.Nodes()) { + const auto node_input_edge_count = node.GetInputEdgesCount(); + in_degree.insert({node.Index(), node_input_edge_count}); + if (node_input_edge_count == 0) { + nodes_to_process.push_back(&node); + } + } + + std::vector supported_group{}; + // the partition node group's border is the aggregate of its nodes' output nodes + std::unordered_set supported_group_border{}; + + auto close_group = [&]() { + if (!supported_group.empty()) { +#ifndef NDEBUG + if (debug_output) { + LOGS_DEFAULT(VERBOSE) << "New partition node group.\n" + << "Unsupported nodes on group border: " + << NodeGroupDebugString(nodes_to_process_with_next_group, true) << "\n" + << "Nodes in group: " << NodeGroupDebugString(supported_group); + } +#endif + + // if no on_group_closed_fn callback was given, keep the partition + // otherwise, let the callback determine whether to keep it + const bool keep_partition = !on_group_closed_fn || on_group_closed_fn(supported_group); + + if (keep_partition) { + supported_groups.emplace_back(std::move(supported_group)); + } +#ifndef NDEBUG + else { + LOGS_DEFAULT_IF(debug_output, VERBOSE) << "Discarded partition node group."; + } +#endif + + supported_group.clear(); + supported_group_border.clear(); } }; - // find root nodes - for (auto& node : graph_viewer.Nodes()) { - size_t input_edge_count = node.GetInputEdgesCount(); - in_degree.insert({node.Index(), input_edge_count}); - if (input_edge_count == 0) { - add_to_visit(node); - } - } - - // prefer unsupported nodes first. this will increase the number of inputs potentially available to the first - // partition handled by this EP. - bool processing_supported_nodes = false; - - while (!supported_to_visit.empty() || !unsupported_to_visit.empty()) { - const Node* current = nullptr; - - // see if we need to flip - if ((processing_supported_nodes && supported_to_visit.empty()) || - (!processing_supported_nodes && unsupported_to_visit.empty())) { - processing_supported_nodes = !processing_supported_nodes; + while (!nodes_to_process.empty() || !nodes_to_process_with_next_group.empty()) { + if (nodes_to_process.empty()) { + // we have processed all the nodes that we can while building this partition node group, start a new one + close_group(); + nodes_to_process.swap(nodes_to_process_with_next_group); continue; } - // get next node from same partition - if (processing_supported_nodes) { - current = supported_to_visit.front(); - supported_to_visit.pop(); - } else { - current = unsupported_to_visit.front(); - unsupported_to_visit.pop(); + const Node& node = *nodes_to_process.front(); + nodes_to_process.pop_front(); + + const bool is_node_supported = is_node_supported_fn(node); + + if (!is_node_supported && Contains(supported_group_border, &node)) { + // an unsupported node on the border will be processed after the current partition node group + nodes_to_process_with_next_group.push_back(&node); + continue; } - // when in_degree is zero all the inputs to the node are available - for (auto node_it = current->OutputNodesBegin(), end = current->OutputNodesEnd(); node_it != end; ++node_it) { - in_degree[node_it->Index()]--; + if (is_node_supported) { + // add node to the partition node group + supported_group.push_back(&node); - if (in_degree[node_it->Index()] == 0) { - add_to_visit(*node_it); - } + // remove node from the border and add its outputs to the border + supported_group_border.erase(&node); + + std::for_each( + node.OutputNodesBegin(), node.OutputNodesEnd(), + [&supported_group_border](const Node& output) { + supported_group_border.insert(&output); + }); } - topo_order.push_back(&*current); + // adjust in-degrees of the node outputs and add any new nodes to process + std::for_each( + node.OutputNodesBegin(), node.OutputNodesEnd(), + [&](const Node& output) { + auto& output_node_in_degree = in_degree[output.Index()]; + --output_node_in_degree; + + if (output_node_in_degree == 0) { + nodes_to_process.push_back(&output); + } + }); } - // check we didn't break something - ORT_ENFORCE(graph_viewer.NumberOfNodes() == static_cast(topo_order.size()), - "Partition aware topological sort has produced invalid output."); + close_group(); - return topo_order; + return supported_groups; } +} // namespace std::unordered_set CreateExcludedNodeSet(const GraphViewer& graph_viewer, const std::unordered_set& stop_ops) { std::unordered_set excluded_nodes; - const auto end_stop_ops = stop_ops.cend(); for (const NodeIndex node_index : graph_viewer.GetNodesInTopologicalOrder()) { const Node& node = *graph_viewer.GetNode(node_index); - if (excluded_nodes.find(&node) == excluded_nodes.cend() && - stop_ops.find(node.OpType()) != end_stop_ops) { + if (!Contains(excluded_nodes, &node) && Contains(stop_ops, node.OpType())) { excluded_nodes.insert(&node); // add all the downstream nodes @@ -116,11 +227,9 @@ std::unordered_set CreateExcludedNodeSet(const GraphViewer& graph_v return excluded_nodes; } -} // namespace - std::unique_ptr MakeComputeCapability(const GraphViewer& graph_viewer, const std::vector& group, - const std::function& generate_metadef_name, + const GenerateMetadefNameFn& generate_metadef_name, const std::string& execution_provider_name) { std::unordered_set node_set; node_set.reserve(group.size()); @@ -142,8 +251,8 @@ std::unique_ptr MakeComputeCapability(const GraphViewer& grap for (const auto* input : node->InputDefs()) { // if the node input was not produced by this subgraph, add it to the subgraph inputs. - if (node_outputs.count(input) == 0) { - if (subgraph_inputs.count(input) == 0) { + if (!Contains(node_outputs, input)) { + if (!Contains(subgraph_inputs, input)) { subgraph_inputs.insert(input); ordered_subgraph_inputs.push_back(input); } @@ -154,7 +263,7 @@ std::unique_ptr MakeComputeCapability(const GraphViewer& grap for (const auto* output_def : output_defs) { node_outputs.insert(output_def); // if output is overall graph output we need to produce it. - if (graph_outputs.count(output_def) != 0) { + if (Contains(graph_outputs, output_def)) { ordered_subgraph_outputs.push_back(output_def); } } @@ -162,9 +271,9 @@ std::unique_ptr MakeComputeCapability(const GraphViewer& grap // if output connects to a node not in this subgraph we need to add it // unless it was already added as an overall graph output, for (auto it = node->OutputEdgesBegin(), end = node->OutputEdgesEnd(); it != end; ++it) { - if (node_set.count(&it->GetNode()) == 0) { + if (!Contains(node_set, &it->GetNode())) { const auto* output_def = output_defs[it->GetSrcArgIndex()]; - if (subgraph_outputs.count(output_def) == 0 && graph_outputs.count(output_def) == 0) { + if (!Contains(subgraph_outputs, output_def) && !Contains(graph_outputs, output_def)) { subgraph_outputs.insert(output_def); ordered_subgraph_outputs.push_back(output_def); } @@ -194,143 +303,51 @@ std::unique_ptr MakeComputeCapability(const GraphViewer& grap std::vector> CreateSupportedPartitions(const GraphViewer& graph_viewer, - const std::unordered_set& supported_nodes, - const std::unordered_set& stop_ops, - const std::function& generate_metadef_name, + const IsNodeSupportedFn& is_node_supported_fn, + const OnGroupClosedFn& on_partition_closed_fn, + const GenerateMetadefNameFn& generate_metadef_name_fn, const std::string& execution_provider_name, bool debug_output) { - // find any nodes we need to exclude - std::unordered_set excluded_nodes = CreateExcludedNodeSet(graph_viewer, stop_ops); + const auto groups = CreateSupportedPartitionNodeGroups(graph_viewer, + is_node_supported_fn, + on_partition_closed_fn, + debug_output); -#ifndef NDEBUG - auto node_str = [](const Node& node) { - std::ostringstream oss; - oss << node.Index() << " '" << node.Name() << "'(" << node.OpType() << ")"; - return oss.str(); - }; + std::vector> partitions{}; + partitions.reserve(groups.size()); - auto group_str = [&node_str](const std::vector& group) { - const Node& start_node = *group.front(); - const Node& end_node = *group.back(); - std::ostringstream oss; - oss << node_str(start_node) << " to " << node_str(end_node) << "\n"; - return oss.str(); - }; -#endif + std::transform( + groups.begin(), groups.end(), + std::back_inserter(partitions), + [&](const auto& supported_partition) { + return MakeComputeCapability(graph_viewer, supported_partition, generate_metadef_name_fn, + execution_provider_name); + }); - // partition aware sort. this groups all the nodes we can and can't handle - const std::vector new_order = PartitionAwareTopoSort(graph_viewer, supported_nodes); - - // create groups using the new sort order - auto cur_topo_node = new_order.cbegin(); - auto end_topo_nodes = new_order.cend(); - - std::queue nodes_to_process; // supported nodes to process - std::unordered_set processed_nodes; // supported nodes we have processed - std::map> node_groups; - std::vector cur_group; - - bool check_excluded_nodes = !excluded_nodes.empty(); - const auto excluded_nodes_end = excluded_nodes.cend(); - - while (cur_topo_node != end_topo_nodes) { - const Node* node = *cur_topo_node; - ++cur_topo_node; - - if (processed_nodes.find(node) != processed_nodes.cend()) { - continue; - } - - if (check_excluded_nodes && excluded_nodes.find(node) != excluded_nodes_end) { - processed_nodes.insert(node); - continue; - } - - bool supported = supported_nodes.count(node) != 0; - bool in_partition = !cur_group.empty(); - - // check if end of a partition. - if (in_partition && !supported) { -#ifndef NDEBUG - if (debug_output) { - LOGS_DEFAULT(VERBOSE) << "New partition due to " << node_str(*node) - << ". Nodes in old partition: " << cur_group.size() << "\n"; - LOGS_DEFAULT(VERBOSE) << group_str(cur_group) << "\n"; - } -#else - ORT_UNUSED_PARAMETER(debug_output); -#endif - node_groups.insert({cur_group.front()->Index(), std::move(cur_group)}); - } - - // add the node and any connected downstream nodes that we can handle if supported. - // if not mark as processed so we know its inputs are available - if (supported) { - nodes_to_process.push(node); - } else { - processed_nodes.insert(node); - } - - while (!nodes_to_process.empty()) { - node = nodes_to_process.front(); - nodes_to_process.pop(); - - if (processed_nodes.find(node) == processed_nodes.cend()) { - // add to partition if all inputs available - bool inputs_available = true; - for (auto cur = node->InputNodesBegin(), end = node->InputNodesEnd(); cur != end; ++cur) { - if (processed_nodes.find(&*cur) == processed_nodes.cend()) { - inputs_available = false; - break; - } - } - - if (inputs_available) { - cur_group.push_back(node); - processed_nodes.insert(node); - - for (auto cur = node->OutputNodesBegin(), end = node->OutputNodesEnd(); cur != end; ++cur) { - const Node& downstream_node = *cur; - - // nodes will get added to the queue once per input from a supported node. - // we need this to happen as they can't be added to the group until all inputs are known to be available. - if (supported_nodes.count(&downstream_node) != 0) { - nodes_to_process.push(&downstream_node); - } - } - } else { - // we need all other nodes providing input to this node to have been processed - // before it can be added to cur_group. - // - // e.g. given A B with a topological order of A, B, C. - // \ / - // C - // - // When we process A we add C via the output edge to nodes_to_process. After we finish with A we look at C - // as the next node in nodes_to_process, but the input from B is missing. - // There are no more entries in nodes_to_process so we move to the next node in the topological order and - // process B. Again C is added to nodes_to_process via the output edge. After we finish with B we look at C - // again as the next node in nodes_to_process. - // Now all the inputs are available and C is added to the current group. - } - } - } - } - - if (!cur_group.empty()) { - node_groups.insert({cur_group.front()->Index(), std::move(cur_group)}); - } - - // create ComputeCapability instances - std::vector> results; - results.reserve(node_groups.size()); - - for (const auto& idx_to_group : node_groups) { - results.push_back( - MakeComputeCapability(graph_viewer, idx_to_group.second, generate_metadef_name, execution_provider_name)); - } - - return results; + return partitions; } + +std::vector> +CreateSupportedPartitions(const GraphViewer& graph_viewer, + const std::unordered_set& supported_nodes, + const std::unordered_set& stop_ops, + const GenerateMetadefNameFn& generate_metadef_name_fn, + const std::string& execution_provider_name, + bool debug_output) { + const auto excluded_nodes = CreateExcludedNodeSet(graph_viewer, stop_ops); + const bool check_excluded_nodes = !excluded_nodes.empty(); + + return CreateSupportedPartitions( + graph_viewer, + [&](const Node& node) -> bool { + const bool is_excluded = check_excluded_nodes && Contains(excluded_nodes, &node); + return !is_excluded && Contains(supported_nodes, &node); + }, + {}, + generate_metadef_name_fn, + execution_provider_name, + debug_output); +} + } // namespace utils } // namespace onnxruntime diff --git a/onnxruntime/core/providers/partitioning_utils.h b/onnxruntime/core/providers/partitioning_utils.h index e739fafe2a..2d5dd5cf98 100644 --- a/onnxruntime/core/providers/partitioning_utils.h +++ b/onnxruntime/core/providers/partitioning_utils.h @@ -18,27 +18,73 @@ class Node; namespace utils { -/** -Create the supported partitions for the execution provider. +/** +Called to check whether a node is supported. + +@param node The node to check. + +@return Whether the node is supported. +*/ +using IsNodeSupportedFn = std::function; + +/** +Called to indicate a completed partition node group. +The partition is kept or discarded based on the return value. + +@param group The partition node group. + +@return Whether to keep the partition. +*/ +using OnGroupClosedFn = std::function& group)>; + +/** +Called to create a metadef name. +Most likely should call IExecutionProvider::GenerateMetaDefId. +See onnxruntime/test/providers/internal_testing/internal_testing_execution_provider.cc for example usage. + +@return The metadef name. +*/ +using GenerateMetadefNameFn = std::function; + +/** +Create the supported partitions for the execution provider. @param graph_viewer GraphViewer that IExecutionProvider::GetCapability is called with. -@param supported_nodes Set of nodes that the execution provider wants to handle. -@param stop_ops Set of operator names at which we stop considering nodes for assignment to this execution provider. -@param generate_metadef_name Functor to create the name for the MetaDef. - Most likely should call IExecutionProvider::GenerateMetaDefId. - See onnxruntime/test/providers/internal_testing/internal_testing_execution_provider.cc - for example usage. +@param is_node_supported_fn Callback to check whether a node is supported. +@param on_group_closed_fn Callback to indicate a completed partition node group. +@param generate_metadef_name_fn Callback to create the name for the MetaDef. @param execution_provider_name Name of execution provider creating the ComputeCapability instance. -@param debug_output Print diagnostic output about the partitions and reasons for partition breaks. +@param debug_output Print diagnostic output about the partitions and reasons for partition breaks. No-op in a release build. -@returns ComputeCapability instances for all partitions assigned to the execution provider. +@returns ComputeCapability instances for all partitions assigned to the execution provider. +*/ +std::vector> +CreateSupportedPartitions(const GraphViewer& graph_viewer, + const IsNodeSupportedFn& is_node_supported_fn, + const OnGroupClosedFn& on_group_closed_fn, + const GenerateMetadefNameFn& generate_metadef_name_fn, + const std::string& execution_provider_name, + bool debug_output = false); + +/** +Create the supported partitions for the execution provider. + +@param graph_viewer GraphViewer that IExecutionProvider::GetCapability is called with. +@param supported_nodes Set of nodes that the execution provider wants to handle. +@param stop_ops Set of operator names at which we stop considering nodes for assignment to this execution provider. +@param generate_metadef_name Functor to create the name for the MetaDef. +@param execution_provider_name Name of execution provider creating the ComputeCapability instance. +@param debug_output Print diagnostic output about the partitions and reasons for partition breaks. + No-op in a release build. + +@returns ComputeCapability instances for all partitions assigned to the execution provider. */ std::vector> CreateSupportedPartitions( const GraphViewer& graph_viewer, const std::unordered_set& supported_nodes, const std::unordered_set& stop_ops, - const std::function& generate_metadef_name, + const GenerateMetadefNameFn& generate_metadef_name, const std::string& execution_provider_name, bool debug_output = false); @@ -49,9 +95,6 @@ Will automatically determine the inputs and outputs required. @param graph_viewer GraphViewer that IExecutionProvider::GetCapability is called with. @param group Group of nodes to include in the ComputeCapability instance. @param generate_metadef_name Functor to create the name for the MetaDef. - Most likely should call IExecutionProvider::GenerateMetaDefId. - See onnxruntime/test/providers/internal_testing/internal_testing_execution_provider.cc for - example usage. @param execution_provider_name Name of execution provider creating the ComputeCapability instance. @returns New ComputeCapability instance. @@ -61,7 +104,19 @@ Will automatically determine the inputs and outputs required. */ std::unique_ptr MakeComputeCapability(const GraphViewer& graph_viewer, const std::vector& group, - const std::function& generate_metadef_name, + const GenerateMetadefNameFn& generate_metadef_name, const std::string& execution_provider_name); + +/** +Create the set of nodes to exclude based on a set of stop ops. +Stop op nodes and nodes downstream from them will be excluded. + +@param graph_viewer GraphViewer with the nodes to consider. +@param stop_ops The set of stop ops. + +@return The set of excluded nodes. +*/ +std::unordered_set CreateExcludedNodeSet(const GraphViewer& graph_viewer, + const std::unordered_set& stop_ops); } // namespace utils } // namespace onnxruntime diff --git a/onnxruntime/test/testdata/ep_partitioning_tests.py b/onnxruntime/test/testdata/ep_partitioning_tests.py index 23bad4cdcf..fe68f6cc5e 100644 --- a/onnxruntime/test/testdata/ep_partitioning_tests.py +++ b/onnxruntime/test/testdata/ep_partitioning_tests.py @@ -25,21 +25,9 @@ def create_model_1(): # Naively creating groups based on iterating this order and whether a node is supported gives the following groups # (a1), (s1), (a2), (s2), (a3, a4). This is similar to what most EPs do currently. # - # If we also consider downstream nodes with all inputs available when adding via the topological sort we get two - # less groups as s2 gets added with s1. - # (a1), (s1, s2), (a2, a3, a4) - # - # If the EP handles Sub that's fine. If the EP handles Add that's not. - # - # Finally, if we do a partition aware sort, and prefer unhandled nodes first to maximize the inputs that would be - # available each time we go to the EP, we can choose either of the root nodes (a1 or s1) to start at. - # - # If the EP is handling Sub we would start with a1 and get the same groups as above - which is perfectly fine as - # there's a single partition with (s1, s2) run on the EP. - # - # If the EP is handling Add we would start with s1 (due to preferring unhandled nodes first) and get the following - # groups, which also achieves a single partition of the handled nodes. - # (s2, s2), (a1, a2, a3, a4) + # To improve on that, we may consider all reachable supported nodes when iterating the topological ordering. + # In this model, regardless of whether Add or Sub is supported, we get the groups (s1, s2), (a1, a2, a3, a4). + # One of those groups is the resulting partition. # # So if this model is loaded in a partitioning test, there should only be one partition running on the EP regardless # of whether Add or Sub is supported by it.