Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
105 changes: 89 additions & 16 deletions onnxruntime/core/graph/function_utils.cc
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
// Copyright (c) Microsoft Corporation. All rights reserved.
// Licensed under the MIT License.
#include <algorithm>
#include <sstream>
#include <unordered_map>

#include "core/graph/function_utils.h"
Expand All @@ -15,6 +17,33 @@
using string = std::string;
using namespace ONNX_NAMESPACE;

namespace {

bool AttributeContainsSubgraphs(const AttributeProto& attr) {
return attr.has_g() || attr.graphs_size() > 0;
}

std::string FormatAttributeReferenceChain(const std::vector<std::string>& active_attr_refs, const std::string& next_attr_ref) {
std::ostringstream stream;

for (size_t i = 0; i < active_attr_refs.size(); ++i) {
if (i != 0) {
stream << " -> ";
}

stream << active_attr_refs[i];
}

if (!active_attr_refs.empty()) {
stream << " -> ";
}

stream << next_attr_ref;
return stream.str();
}

} // namespace

// Utilify function to get the imported version of domain from opset imports
// Returns -1 if requested domain is not found in the opset_imports
template <template <typename, typename> class M>
Expand Down Expand Up @@ -365,6 +394,7 @@
std::string prefix_;
const onnxruntime::NodeAttributes& attr_map_;
std::vector<InlinedHashMap<std::string, std::string>> rename_scopes_;
std::vector<std::string> active_attr_refs_;

Check warning on line 397 in onnxruntime/core/graph/function_utils.cc

View workflow job for this annotation

GitHub Actions / Optional Lint C++

[cpplint] reported by reviewdog 🐶 Add #include <vector> for vector<> [build/include_what_you_use] [4] Raw Output: onnxruntime/core/graph/function_utils.cc:397: Add #include <vector> for vector<> [build/include_what_you_use] [4]

Inliner(const std::string& prefix, const onnxruntime::NodeAttributes& attr_map) : prefix_(prefix),
attr_map_(attr_map) {
Expand Down Expand Up @@ -432,7 +462,7 @@
}

// Process a node:
void transform(NodeProto& n) {
Status transform(NodeProto& n) {
if (!n.name().empty())
n.set_name(prefix_ + "_" + n.name());

Expand All @@ -446,11 +476,22 @@
const auto attr_map_end = attr_map_.cend();
for (auto attr_iter = attributes.begin(); attr_iter != attributes.end();) {
auto& attr = *attr_iter;
bool tracked_attr_ref = false;
if (!attr.ref_attr_name().empty()) {
// Attribute-references must be replaced by the corresponding attribute-value in the call-node
// if the call-node contains the attribute. Otherwise, this attribute must be removed.
auto entry = attr_map_.find(attr.ref_attr_name());
if (entry != attr_map_end) {
if (AttributeContainsSubgraphs(entry->second)) {
const auto active_attr_ref = std::find(active_attr_refs_.cbegin(), active_attr_refs_.cend(), attr.ref_attr_name());
ORT_RETURN_IF(active_attr_ref != active_attr_refs_.cend(),
"Function attribute graph expansion is recursive: ",
FormatAttributeReferenceChain(active_attr_refs_, attr.ref_attr_name()));

active_attr_refs_.push_back(attr.ref_attr_name());
tracked_attr_ref = true;
}

// Copy value of attribute, but retain original name:
std::string name = attr.name();
attr = entry->second;
Expand All @@ -462,56 +503,88 @@
}
// Subgraphs must be recursively processed.
if (attr.has_g()) {
transform(*attr.mutable_g());
auto status = transform(*attr.mutable_g());
if (!status.IsOK()) {
if (tracked_attr_ref) {
active_attr_refs_.pop_back();
}

return status;
}
}
for (auto& graph : *attr.mutable_graphs()) {
auto status = transform(graph);
if (!status.IsOK()) {
if (tracked_attr_ref) {
active_attr_refs_.pop_back();
}

return status;
}
}

if (tracked_attr_ref) {
active_attr_refs_.pop_back();
}
for (auto& graph : *attr.mutable_graphs())
transform(graph);

++attr_iter;
}

return Status::OK();
}

// Process a sub-graph, contained as an attribute in a control-flow op node.
void transform(GraphProto& graph) {
Status transform(GraphProto& graph) {
rename_scopes_.emplace_back();
for (auto& x : *graph.mutable_input())
make_unique(*x.mutable_name());
for (auto& init : *graph.mutable_initializer())
make_unique(*init.mutable_name());
for (auto& y : *graph.mutable_output())
make_unique(*y.mutable_name());
for (auto& n : *graph.mutable_node())
transform(n);
for (auto& n : *graph.mutable_node()) {
auto status = transform(n);
if (!status.IsOK()) {
rename_scopes_.pop_back();
return status;
}
}
rename_scopes_.pop_back();

return Status::OK();
}

public:
// The main specialization method: specialize a FunctionProto for a particular call-site.
static void specialize(const NodeProto& callnode, FunctionProto& callee, const onnxruntime::NodeAttributes& attr_map,
const std::string& unique_prefix) {
static Status specialize(const NodeProto& callnode, FunctionProto& callee, const onnxruntime::NodeAttributes& attr_map,
const std::string& unique_prefix) {
Inliner inliner(unique_prefix, attr_map);

inliner.bind<false>(*callee.mutable_input(), callnode.input());
inliner.bind<true>(*callee.mutable_output(), callnode.output());

for (auto& n : *callee.mutable_node())
inliner.transform(n);
for (auto& n : *callee.mutable_node()) {
ORT_RETURN_IF_ERROR(inliner.transform(n));
}

return Status::OK();
}
};

void Specialize(ONNX_NAMESPACE::FunctionProto& called_function, const ONNX_NAMESPACE::NodeProto& calling_node,
const onnxruntime::NodeAttributes& attr_map, const std::string& unique_prefix) {
Inliner::specialize(calling_node, called_function, attr_map, unique_prefix);
Status Specialize(ONNX_NAMESPACE::FunctionProto& called_function, const ONNX_NAMESPACE::NodeProto& calling_node,
const onnxruntime::NodeAttributes& attr_map, const std::string& unique_prefix) {
return Inliner::specialize(calling_node, called_function, attr_map, unique_prefix);
}

void Specialize(ONNX_NAMESPACE::FunctionProto& called_function, const Node& calling_node, const std::string& unique_prefix) {
Status Specialize(ONNX_NAMESPACE::FunctionProto& called_function, const Node& calling_node, const std::string& unique_prefix) {

Check warning on line 579 in onnxruntime/core/graph/function_utils.cc

View workflow job for this annotation

GitHub Actions / Optional Lint C++

[cpplint] reported by reviewdog 🐶 Add #include <string> for string [build/include_what_you_use] [4] Raw Output: onnxruntime/core/graph/function_utils.cc:579: Add #include <string> for string [build/include_what_you_use] [4]
ONNX_NAMESPACE::NodeProto calling_node_proto;
calling_node.ToProto(calling_node_proto);

onnxruntime::NodeAttributes attr_map = calling_node.GetAttributes();
for (auto& attribute_proto : called_function.attribute_proto()) {
ORT_IGNORE_RETURN_VALUE(attr_map.emplace(attribute_proto.name(), attribute_proto));
}
Specialize(called_function, calling_node_proto, attr_map, unique_prefix);
return Specialize(called_function, calling_node_proto, attr_map, unique_prefix);
}

} // namespace function_utils
Expand Down
6 changes: 3 additions & 3 deletions onnxruntime/core/graph/function_utils.h
Original file line number Diff line number Diff line change
Expand Up @@ -63,10 +63,10 @@ inline std::string GetFunctionIdentifier(std::string_view function_domain, std::
return id;
}

void Specialize(ONNX_NAMESPACE::FunctionProto& called_function, const ONNX_NAMESPACE::NodeProto& calling_node,
const onnxruntime::NodeAttributes& attr_map, const std::string& unique_prefix);
Status Specialize(ONNX_NAMESPACE::FunctionProto& called_function, const ONNX_NAMESPACE::NodeProto& calling_node,
const onnxruntime::NodeAttributes& attr_map, const std::string& unique_prefix);

void Specialize(ONNX_NAMESPACE::FunctionProto& called_function, const Node& calling_node, const std::string& unique_prefix);
Status Specialize(ONNX_NAMESPACE::FunctionProto& called_function, const Node& calling_node, const std::string& unique_prefix);

} // namespace function_utils

Expand Down
2 changes: 1 addition & 1 deletion onnxruntime/core/graph/graph.cc
Original file line number Diff line number Diff line change
Expand Up @@ -6584,7 +6584,7 @@ Status Graph::InlineFunction(Node& callnode) {
ORT_ENFORCE(callnode.TryGetFunctionProto(inlined_fp), "Node has no function body and cannot be inlined.");

// Make all the names unique and resolve nested graphs inputs to the outer scope.
function_utils::Specialize(inlined_fp, callnode, uniq_identifier);
ORT_RETURN_IF_ERROR(function_utils::Specialize(inlined_fp, callnode, uniq_identifier));

// In this case, global Resolve() will take care of everything.
ORT_RETURN_IF_ERROR(InlineFunctionProto(inlined_fp, parent_annotation));
Expand Down
Loading
Loading