diff --git a/include/onnxruntime/core/graph/graph.h b/include/onnxruntime/core/graph/graph.h index 3e31fc5c8d..8a4f1fd08b 100644 --- a/include/onnxruntime/core/graph/graph.h +++ b/include/onnxruntime/core/graph/graph.h @@ -54,9 +54,10 @@ class Node { Construct an EdgeEnd @param node The source node if this is an input edge to the current node, or the destination node if this is an output edge from the current node. - @param node_arg The NodeArg to use for the edge. + @param src_arg_index The node arg index of source node of the edge. + @param dst_arg_index The node arg index of destination node of the edge. */ - EdgeEnd(const Node& node, const NodeArg& node_arg) noexcept; + EdgeEnd(const Node& node, int src_arg_index, int dst_arg_index) noexcept; /** Construct a control edge. @param node The node the edge joins to the current node. @@ -66,13 +67,18 @@ class Node { /** Gets the Node that this EdgeEnd refers to. */ const Node& GetNode() const noexcept; - /** Gets the NodeArg that this edge end refers to. - @returns NodeArg pointer or nullptr if this is a control edge. */ - const NodeArg* GetNodeArg() const noexcept; + /** Gets the source arg index. + @returns the source arg index of <*this> edge.*/ + int GetSrcArgIndex() const; + + /** Gets the destination arg index. + @returns the destination arg index of <*this> edge.*/ + int GetDstArgIndex() const; private: const Node* node_; - const NodeArg* node_arg_; + int src_arg_index_; + int dst_arg_index_; }; /** Gets the Node's NodeIndex. */ @@ -146,7 +152,7 @@ class Node { /** Gets the implicit inputs to this Node. If this Node contains a subgraph, these are the NodeArg's that are implicitly consumed by Nodes within that subgraph. e.g. If and Loop operators.*/ - const std::vector& ImplicitInputDefs() const noexcept { + const std::vector& ImplicitInputDefs() const noexcept { return definitions_.implicit_input_defs; } @@ -160,11 +166,10 @@ class Node { struct EdgeEndCompare { bool operator()(const EdgeEnd& lhs, const EdgeEnd& rhs) const { if (lhs.GetNode().Index() == rhs.GetNode().Index()) { - auto lhs_arg = lhs.GetNodeArg(); - auto rhs_arg = rhs.GetNodeArg(); - std::string lhs_arg_name = lhs_arg == nullptr ? "" : lhs_arg->Name(); - std::string rhs_arg_name = rhs_arg == nullptr ? "" : rhs_arg->Name(); - return lhs_arg_name.compare(rhs_arg_name) < 0; + if (lhs.GetSrcArgIndex() == rhs.GetSrcArgIndex()) { + return lhs.GetDstArgIndex() < rhs.GetDstArgIndex(); + } + return lhs.GetSrcArgIndex() < rhs.GetSrcArgIndex(); } return lhs.GetNode().Index() < rhs.GetNode().Index(); } @@ -305,7 +310,7 @@ class Node { @remarks For example, a subgraph in an 'If' node gets all its input values via this mechanism rather than there being explicit inputs to the 'If' node that are passed to the subgraph. They are pseudo-inputs to this Node as it has an implicit dependency on them. */ - std::vector implicit_input_defs; + std::vector implicit_input_defs; private: ONNXRUNTIME_DISALLOW_COPY_ASSIGNMENT_AND_MOVE(Definitions); @@ -391,7 +396,7 @@ class Node { // OperatorSchema that <*this> node refers to. const ONNX_NAMESPACE::OpSchema* op_ = nullptr; Node::Type node_type_ = Node::Type::Primitive; - + // The function body is owned by graph_ const Function* func_body_ = nullptr; @@ -588,6 +593,12 @@ class Graph { Node& AddNode(const Node& other); /** Remove a Node from this Graph and free it. + The output edges of this specified node MUST have been removed before removing the node. + The input edges of this specified node is removed while removing the node. The process of + removing a node from a graph should be, + 1. Remove out edges of this specified node. + 2. Remove this specified node. + 3. Add new input edges connected with all out nodes. @returns true if the node_index was valid @remarks Do not call AddNode and Remove Node concurrently as they are not thread-safe. */ @@ -596,16 +607,18 @@ class Graph { /** Add an edge between two Nodes. @param src_node_index NodeIndex of source Node that is providing output to the destination Node. @param dst_node_index NodeIndex of destination Node that is receiving input from the source Node. - @param node_arg NodeArg to use for the edge. + @param src_arg_index node arg index of source node. + @param dst_arg_index node arg index of destination node. */ - void AddEdge(NodeIndex src_node_index, NodeIndex dst_node_index, const NodeArg& node_arg); + void AddEdge(NodeIndex src_node_index, NodeIndex dst_node_index, int src_arg_index, int dst_arg_index); /** Remove an edge between two Nodes. @param src_node_index NodeIndex of source Node to remove an output edge from. @param dst_node_index NodeIndex of destination Node to remove an input edge from. - @param node_arg NodeArg that is used by the edge that is being removed. + @param src_arg_index node arg index of source node. + @param dst_arg_index node arg index of destination node. */ - void RemoveEdge(NodeIndex src_node_index, NodeIndex dst_node_index, const NodeArg& node_arg); + void RemoveEdge(NodeIndex src_node_index, NodeIndex dst_node_index, int src_arg_index, int dst_arg_index); /** Add a control edge between two Nodes in this Graph. @@ -779,7 +792,7 @@ class Graph { struct ResolveContext { ResolveContext() = default; - std::unordered_map output_args; + std::unordered_map> output_args; std::unordered_set inputs_and_initializers; std::unordered_set outer_scope_node_args; std::unordered_map node_name_to_index; @@ -798,7 +811,7 @@ class Graph { }; // search this and up through any parent_graph_ instance for a NodeArg - const NodeArg* GetNodeArgIncludingParentGraphs(const std::string& node_arg_name) const; + NodeArg* GetNodeArgIncludingParentGraphs(const std::string& node_arg_name); // Initialize all the graph inputs, initializers and outputs common::Status InitInputsInitializersOutputs(); diff --git a/onnxruntime/core/framework/op_kernel_context_internal.h b/onnxruntime/core/framework/op_kernel_context_internal.h index afce7fc3c1..cc10bca1d8 100644 --- a/onnxruntime/core/framework/op_kernel_context_internal.h +++ b/onnxruntime/core/framework/op_kernel_context_internal.h @@ -17,7 +17,7 @@ class OpKernelContextInternal : public OpKernelContext { explicit OpKernelContextInternal(ExecutionFrame& frame, const OpKernel& kernel, const logging::Logger& logger, - const std::vector& implicit_inputs, + const std::vector& implicit_inputs, const bool& terminate_flag) : OpKernelContext(&frame, &kernel, logger), implicit_inputs_{implicit_inputs}, @@ -51,7 +51,7 @@ class OpKernelContextInternal : public OpKernelContext { const bool& GetTerminateFlag() const noexcept { return terminate_flag_; } private: - const std::vector& implicit_inputs_; + const std::vector& implicit_inputs_; const bool& terminate_flag_; }; diff --git a/onnxruntime/core/framework/session_state_initializer.cc b/onnxruntime/core/framework/session_state_initializer.cc index b438718ae3..b83cdb7341 100644 --- a/onnxruntime/core/framework/session_state_initializer.cc +++ b/onnxruntime/core/framework/session_state_initializer.cc @@ -70,7 +70,7 @@ SessionStateInitializer::SessionStateInitializer(onnxruntime::Graph& graph, common::Status SessionStateInitializer::CreatePlan(const onnxruntime::GraphTransformerManager& graph_transformation_manager, const InsertCastTransformer& insert_cast_transformer, - const std::vector& outer_scope_node_args, + const std::vector& outer_scope_node_args, bool enable_sequential_execution) { ONNXRUNTIME_RETURN_IF_ERROR(TransformGraph(graph_, graph_transformation_manager, execution_providers_, kernel_registry_manager_, diff --git a/onnxruntime/core/framework/session_state_initializer.h b/onnxruntime/core/framework/session_state_initializer.h index c0540e3f37..dd1d6c1dcb 100644 --- a/onnxruntime/core/framework/session_state_initializer.h +++ b/onnxruntime/core/framework/session_state_initializer.h @@ -31,7 +31,7 @@ class SessionStateInitializer { // First perform any transformations and create the execution plan common::Status CreatePlan(const onnxruntime::GraphTransformerManager& graph_transformation_manager, const InsertCastTransformer& insert_cast_transformer, - const std::vector& outer_scope_node_args, + const std::vector& outer_scope_node_args, bool enable_sequential_execution); // initialize tensors, and save. save kernels and input/output node mappings diff --git a/onnxruntime/core/graph/conv_bn_fusion.cc b/onnxruntime/core/graph/conv_bn_fusion.cc index 25a3eb74b9..aaf15ab6ec 100644 --- a/onnxruntime/core/graph/conv_bn_fusion.cc +++ b/onnxruntime/core/graph/conv_bn_fusion.cc @@ -170,14 +170,13 @@ Status ConvBNFusion::Apply(onnxruntime::Graph& graph, bool& modified) const { } } } - removed_nodes.push_back(bn_node.Index()); } for (auto i : removed_nodes) { graph.RemoveNode(i); } - + if (!removed_nodes.empty()) { modified = true; ONNXRUNTIME_RETURN_IF_ERROR(graph.Resolve()); diff --git a/onnxruntime/core/graph/conv_mul_fusion.cc b/onnxruntime/core/graph/conv_mul_fusion.cc index 8a7fc9b51a..152d689547 100644 --- a/onnxruntime/core/graph/conv_mul_fusion.cc +++ b/onnxruntime/core/graph/conv_mul_fusion.cc @@ -126,7 +126,7 @@ Status ConvMulFusion::Apply(onnxruntime::Graph& graph, bool& modified) const { } } } - + removed_nodes.push_back(mul_node.Index()); } diff --git a/onnxruntime/core/graph/graph.cc b/onnxruntime/core/graph/graph.cc index ddbe873d8d..8265bbb70b 100644 --- a/onnxruntime/core/graph/graph.cc +++ b/onnxruntime/core/graph/graph.cc @@ -223,20 +223,26 @@ bool NodeArg::Exists() const noexcept { return exists_; } -Node::EdgeEnd::EdgeEnd(const Node& node, const NodeArg& node_arg) noexcept - : node_(&node), node_arg_(&node_arg) { +Node::EdgeEnd::EdgeEnd(const Node& node, int src_arg_index, int dst_arg_index) noexcept + : node_(&node), + src_arg_index_(src_arg_index), + dst_arg_index_(dst_arg_index) { } Node::EdgeEnd::EdgeEnd(const Node& node) noexcept - : node_(&node), node_arg_(nullptr) { + : EdgeEnd(node, INT_MAX, INT_MAX) { } const Node& Node::EdgeEnd::GetNode() const noexcept { return *node_; } -const NodeArg* Node::EdgeEnd::GetNodeArg() const noexcept { - return node_arg_; +int Node::EdgeEnd::GetSrcArgIndex() const { + return src_arg_index_; +} + +int Node::EdgeEnd::GetDstArgIndex() const { + return dst_arg_index_; } Node::NodeConstIterator::NodeConstIterator(EdgeConstIterator p_iter) { @@ -692,9 +698,9 @@ Graph::Graph(Graph& parent_graph, ONNX_NAMESPACE::GraphProto& subgraph_proto) } Status Graph::VerifyNoDuplicateName() { - const std::unordered_set& inputs_and_initializers = resolve_context_.inputs_and_initializers; - std::unordered_map& output_args = resolve_context_.output_args; - std::unordered_map& node_name_to_index = resolve_context_.node_name_to_index; + auto& inputs_and_initializers = resolve_context_.inputs_and_initializers; + auto& output_args = resolve_context_.output_args; + auto& node_name_to_index = resolve_context_.node_name_to_index; output_args.clear(); node_name_to_index.clear(); @@ -715,7 +721,9 @@ Status Graph::VerifyNoDuplicateName() { node_name_to_index[node_name] = node.Index(); // Verify node outputs' name should be unique. + int output_index = -1; for (const auto* output_def : node.OutputDefs()) { + ++output_index; if (output_def->Exists()) { auto& output_arg_name = output_def->Name(); if (inputs_and_initializers.count(output_arg_name)) { @@ -723,7 +731,7 @@ Status Graph::VerifyNoDuplicateName() { "Error: Duplicate definition of name (" + output_arg_name + ")."); return status; } - auto result = output_args.insert({output_arg_name, &node}); + auto result = output_args.insert({output_arg_name, {&node, output_index}}); if (!result.second) { // Two outputs with same name, so that insertion fails. Status status(ONNXRUNTIME, FAIL, @@ -760,7 +768,7 @@ common::Status Graph::SetOuterScopeNodeArgs(const std::unordered_set& entry) { return entry.first; }); + [](const std::pair>& entry) { return entry.first; }); for (auto* node : resolve_context_.nodes_with_subgraphs) { for (auto& subgraph : node->MutableSubgraphs()) { @@ -773,8 +781,8 @@ common::Status Graph::SetOuterScopeNodeArgs(const std::unordered_setGetNodeArgIncludingParentGraphs(node_arg_name); @@ -783,7 +791,7 @@ const NodeArg* Graph::GetNodeArgIncludingParentGraphs(const std::string& node_ar return node_arg; } -void Graph::AddEdge(NodeIndex src_node_index, NodeIndex dst_node_index, const NodeArg& node_arg) { +void Graph::AddEdge(NodeIndex src_node_index, NodeIndex dst_node_index, int src_arg_slot, int dst_arg_slot) { if (nodes_.size() <= src_node_index || nodes_.size() <= dst_node_index || nullptr == nodes_[src_node_index] || @@ -791,34 +799,47 @@ void Graph::AddEdge(NodeIndex src_node_index, NodeIndex dst_node_index, const No // Invalid node indexes specified. ONNXRUNTIME_THROW("Invalid node indexes specified when adding edge."); } - // Verify whether the node_arg is input of dst and output of src firstly. - bool valid = false; - for (auto arg : nodes_[src_node_index]->OutputDefs()) { - if (arg == &node_arg) { - valid = true; - break; + + NodeArg *src_arg = nullptr, *dst_arg = nullptr; + if (nodes_[src_node_index]->MutableDefinitions().output_defs.size() > src_arg_slot) { + src_arg = nodes_[src_node_index]->MutableDefinitions().output_defs[src_arg_slot]; + } + + if (nullptr == src_arg) { + ONNXRUNTIME_THROW("Invalid source node arg slot specified when adding edge."); + } + + auto& dst_node_defs = nodes_[dst_node_index]->MutableDefinitions(); + NodeArg** dst_arg_pointer = nullptr; + if (dst_node_defs.input_defs.size() > dst_arg_slot) { + dst_arg_pointer = &dst_node_defs.input_defs[dst_arg_slot]; + dst_arg = *dst_arg_pointer; + } else { + auto num_of_explicit_inputs = dst_node_defs.input_defs.size(); + if (num_of_explicit_inputs + dst_node_defs.implicit_input_defs.size() > dst_arg_slot) { + dst_arg_pointer = &dst_node_defs.implicit_input_defs[dst_arg_slot - num_of_explicit_inputs]; + dst_arg = *dst_arg_pointer; } } - ONNXRUNTIME_ENFORCE(valid); - valid = false; - for (auto arg : nodes_[dst_node_index]->InputDefs()) { - if (arg == &node_arg) { - valid = true; - break; + if (nullptr == dst_arg) { + ONNXRUNTIME_THROW("Invalid destination node arg slot specified when adding edge."); + } + + if (src_arg != dst_arg) { + if (src_arg->Type() != dst_arg->Type()) { + // The output type of source node arg does not match the input type of destination node arg. + ONNXRUNTIME_THROW("Argument type mismatch when adding edge."); + } else { + src_arg->UpdateTypeAndShape(*dst_arg); + *dst_arg_pointer = src_arg; } } - for (auto arg : nodes_[dst_node_index]->ImplicitInputDefs()) { - if (arg == &node_arg) { - valid = true; - break; - } - } - ONNXRUNTIME_ENFORCE(valid); - nodes_[dst_node_index]->MutableRelationships().input_edges.insert(Node::EdgeEnd(*nodes_[src_node_index], node_arg)); - nodes_[src_node_index]->MutableRelationships().output_edges.insert(Node::EdgeEnd(*nodes_[dst_node_index], node_arg)); + + nodes_[src_node_index]->MutableRelationships().output_edges.insert(Node::EdgeEnd(*nodes_[dst_node_index], src_arg_slot, dst_arg_slot)); + nodes_[dst_node_index]->MutableRelationships().input_edges.insert(Node::EdgeEnd(*nodes_[src_node_index], src_arg_slot, dst_arg_slot)); } -void Graph::RemoveEdge(NodeIndex src_node_index, NodeIndex dst_node_index, const NodeArg& node_arg) { +void Graph::RemoveEdge(NodeIndex src_node_index, NodeIndex dst_node_index, int src_arg_slot, int dst_arg_slot) { if (nodes_.size() <= src_node_index || nodes_.size() <= dst_node_index || nullptr == nodes_[src_node_index] || @@ -826,31 +847,36 @@ void Graph::RemoveEdge(NodeIndex src_node_index, NodeIndex dst_node_index, const // Invalid node indexes specified. ONNXRUNTIME_THROW("Invalid node indexes specified when removing edge."); } - // Verify whether the node_arg is input of dst and output of src firstly. - bool valid = false; - for (auto arg : nodes_[src_node_index]->OutputDefs()) { - if (arg == &node_arg) { - valid = true; - break; + + const NodeArg *src_arg = nullptr, *dst_arg = nullptr; + if (nodes_[src_node_index]->GetDefinitions().output_defs.size() > src_arg_slot) { + src_arg = nodes_[src_node_index]->GetDefinitions().output_defs[src_arg_slot]; + } + + if (nullptr == src_arg) { + ONNXRUNTIME_THROW("Invalid source node arg slot specified when removing edge."); + } + + auto& dst_node_defs = nodes_[dst_node_index]->GetDefinitions(); + if (dst_node_defs.input_defs.size() > dst_arg_slot) { + dst_arg = dst_node_defs.input_defs[dst_arg_slot]; + } else { + auto num_of_explicit_inputs = dst_node_defs.input_defs.size(); + if (num_of_explicit_inputs + dst_node_defs.implicit_input_defs.size() > dst_arg_slot) { + dst_arg = dst_node_defs.implicit_input_defs[dst_arg_slot - num_of_explicit_inputs]; } } - ONNXRUNTIME_ENFORCE(valid); - valid = false; - for (auto arg : nodes_[dst_node_index]->InputDefs()) { - if (arg == &node_arg) { - valid = true; - break; - } + if (nullptr == dst_arg) { + ONNXRUNTIME_THROW("Invalid destination node arg slot specified when removing edge."); } - for (auto arg : nodes_[dst_node_index]->ImplicitInputDefs()) { - if (arg == &node_arg) { - valid = true; - break; - } + if (src_arg != dst_arg) { + // The edge ends specified by source and destination arg slot are not referring to same node arg. + // It means there was no edge between these two slots before. + ONNXRUNTIME_THROW("Argument type mismatch when removing edge."); } - ONNXRUNTIME_ENFORCE(valid); - nodes_[dst_node_index]->MutableRelationships().input_edges.erase(Node::EdgeEnd(*nodes_[src_node_index], node_arg)); - nodes_[src_node_index]->MutableRelationships().output_edges.erase(Node::EdgeEnd(*nodes_[dst_node_index], node_arg)); + + nodes_[dst_node_index]->MutableRelationships().input_edges.erase(Node::EdgeEnd(*nodes_[src_node_index], src_arg_slot, dst_arg_slot)); + nodes_[src_node_index]->MutableRelationships().output_edges.erase(Node::EdgeEnd(*nodes_[dst_node_index], src_arg_slot, dst_arg_slot)); } GSL_SUPPRESS(es .84) // ignoring return value from unordered_map::insert causes noisy complaint @@ -867,7 +893,7 @@ Status Graph::BuildConnections(std::vector& outer_scope_node_args_c for (auto& node_arg_name : node_args_consumed) { bool node_arg_in_parent_graph = false; - const auto* node_arg = GetNodeArg(node_arg_name); + auto node_arg = GetNodeArg(node_arg_name); if (node_arg == nullptr) { // it's a node arg from outside this graph's scope, so add that to the list we return @@ -897,8 +923,13 @@ Status Graph::BuildConnections(std::vector& outer_scope_node_args_c // add it to the Node's list of implicit inputs auto& implicit_inputs = node->MutableDefinitions().implicit_input_defs; - if (std::find(implicit_inputs.cbegin(), implicit_inputs.cend(), node_arg) == implicit_inputs.cend()) { + int input_slot_index = static_cast(node->GetDefinitions().input_defs.size()); + auto iter = std::find(implicit_inputs.cbegin(), implicit_inputs.cend(), node_arg); + if (implicit_inputs.cend() == iter) { implicit_inputs.push_back(node_arg); + input_slot_index += static_cast(implicit_inputs.size() - 1); + } else { + input_slot_index += static_cast(iter - implicit_inputs.cbegin()); } if (node_arg_in_parent_graph || @@ -914,8 +945,8 @@ Status Graph::BuildConnections(std::vector& outer_scope_node_args_c ONNXRUNTIME_ENFORCE(entry != resolve_context_.output_args.end()); // Create relationship between this node (node), and the node providing the output (output_node). - Node& output_node = *entry->second; - AddEdge(output_node.Index(), node->Index(), *node_arg); + Node& output_node = *entry->second.first; + AddEdge(output_node.Index(), node->Index(), entry->second.second, input_slot_index); inner_nodes.insert(&output_node); } @@ -932,7 +963,9 @@ Status Graph::BuildConnections(std::vector& outer_scope_node_args_c if (input_args.size() > 0) { // This node needs inputs. + int input_slot_index = -1; for (const auto* input_arg : input_args) { + ++input_slot_index; if (!input_arg->Exists()) { // This input could be optional and it does not exist in this case. continue; @@ -953,8 +986,8 @@ Status Graph::BuildConnections(std::vector& outer_scope_node_args_c } // Create relationship between this node (node), and the node providing the output (output_node). - Node& output_node = *output_arg_iter->second; - AddEdge(output_node.Index(), node.Index(), *input_arg); + Node& output_node = *output_arg_iter->second.first; + AddEdge(output_node.Index(), node.Index(), output_arg_iter->second.second, input_slot_index); inner_nodes.insert(&output_node); } @@ -2023,6 +2056,19 @@ Node& Graph::AddNode(const std::string& name, } bool Graph::RemoveNode(NodeIndex p_index) { + auto node = GetNode(p_index); + if (nullptr == node /*|| 0 != node->GetRelationships().output_edges.size()*/) { + // Node should be removed after all out edges are removed. + // TODO: add the check commented out back. + return false; + } + + // Remove all input edges. + auto input_edges = node->GetRelationships().input_edges; + + for (auto& input_edge : input_edges) { + RemoveEdge(input_edge.GetNode().Index(), p_index, input_edge.GetSrcArgIndex(), input_edge.GetDstArgIndex()); + } return ReleaseNode(p_index); } @@ -2437,6 +2483,14 @@ Node& Graph::FuseSubGraph(std::unique_ptr<::onnxruntime::IndexedSubGraph> sub_gr // Remove nodes fused above. auto& sub_graph_ref = function_container_.back()->GetIndexedSubGraph(); for (auto node_index : sub_graph_ref.nodes) { + auto node = GetNode(node_index); + if (nullptr == node) { + continue; + } + auto output_edges = node->GetRelationships().output_edges; + for (auto output_edge : output_edges) { + RemoveEdge(node->Index(), output_edge.GetNode().Index(), output_edge.GetSrcArgIndex(), output_edge.GetDstArgIndex()); + } RemoveNode(node_index); } return fused_node; @@ -2446,6 +2500,10 @@ Status Graph::InlineFunction(Node& node) { // Remove the function node, add the nodes in function's subgraph into the // main graph. const Graph& subgraph = node.GetFunctionBody()->Body(); + auto output_edges = node.GetRelationships().output_edges; + for (auto output_edge : output_edges) { + RemoveEdge(node.Index(), output_edge.GetNode().Index(), output_edge.GetSrcArgIndex(), output_edge.GetDstArgIndex()); + } RemoveNode(node.Index()); for (const auto& subgraph_node : subgraph.Nodes()) { AddNode(subgraph_node); diff --git a/onnxruntime/core/graph/unsqueeze_elimination.cc b/onnxruntime/core/graph/unsqueeze_elimination.cc index 31b0347c24..ab3ee845ec 100644 --- a/onnxruntime/core/graph/unsqueeze_elimination.cc +++ b/onnxruntime/core/graph/unsqueeze_elimination.cc @@ -82,6 +82,7 @@ Status UnsqueezeElimination::Apply(onnxruntime::Graph& graph, bool& modified) co } } } + removed_nodes.push_back(node.Index()); } diff --git a/onnxruntime/test/ir/graph_test.cc b/onnxruntime/test/ir/graph_test.cc index 5947c61aac..09e62d482e 100644 --- a/onnxruntime/test/ir/graph_test.cc +++ b/onnxruntime/test/ir/graph_test.cc @@ -165,8 +165,8 @@ TEST(GraphTraversalTest, ReverseDFS) { EXPECT_TRUE(status.IsOK()) << status.ErrorMessage(); // Remove/Add edge should not ask for resolving again. - graph.RemoveEdge(node_1.Index(), node_3.Index(), *graph.GetNodeArg("node_1_out_1")); - graph.AddEdge(node_1.Index(), node_3.Index(), *graph.GetNodeArg("node_1_out_1")); + graph.RemoveEdge(node_1.Index(), node_3.Index(), 0, 0); + graph.AddEdge(node_1.Index(), node_3.Index(), 0, 0); std::vector from; for (auto& node : graph.Nodes()) {