Project import generated by Copybara.
GitOrigin-RevId: 5aca6b3f07b67e09988a901f50f595ca5f566e67
This commit is contained in:
@@ -177,7 +177,7 @@ cc_library(
|
||||
deps = [
|
||||
"//mediapipe/framework:packet",
|
||||
"//mediapipe/framework/port:statusor",
|
||||
"@org_tensorflow//tensorflow/core:protos_all_cc",
|
||||
"@org_tensorflow//tensorflow/core:protos_all",
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
@@ -110,7 +110,7 @@ def mediapipe_simple_subgraph(
|
||||
testonly: pass 1 if the graph is to be used only for tests.
|
||||
**kwargs: Remaining keyword args, forwarded to cc_library.
|
||||
"""
|
||||
graph_base_name = graph.replace(":", "/").split("/")[-1].rsplit(".", 1)[0]
|
||||
graph_base_name = name
|
||||
mediapipe_binary_graph(
|
||||
name = name + "_graph",
|
||||
graph = graph,
|
||||
|
||||
@@ -52,6 +52,31 @@ namespace tool {
|
||||
return ::mediapipe::OkStatus();
|
||||
}
|
||||
|
||||
// Returns subgraph streams not requested by a subgraph-node.
|
||||
::mediapipe::Status FindIgnoredStreams(
|
||||
const proto_ns::RepeatedPtrField<ProtoString>& src_streams,
|
||||
const proto_ns::RepeatedPtrField<ProtoString>& dst_streams,
|
||||
std::set<std::string>* result) {
|
||||
ASSIGN_OR_RETURN(auto src_map, tool::TagMap::Create(src_streams));
|
||||
ASSIGN_OR_RETURN(auto dst_map, tool::TagMap::Create(dst_streams));
|
||||
std::set_difference(src_map->Names().begin(), src_map->Names().end(),
|
||||
dst_map->Names().begin(), dst_map->Names().end(),
|
||||
std::inserter(*result, result->begin()));
|
||||
return ::mediapipe::OkStatus();
|
||||
}
|
||||
|
||||
// Removes subgraph streams not requested by a subgraph-node.
|
||||
::mediapipe::Status RemoveIgnoredStreams(
|
||||
proto_ns::RepeatedPtrField<ProtoString>* streams,
|
||||
const std::set<std::string>& missing_streams) {
|
||||
for (int i = streams->size() - 1; i >= 0; --i) {
|
||||
if (missing_streams.count(streams->Get(i)) > 0) {
|
||||
streams->DeleteSubrange(i, 1);
|
||||
}
|
||||
}
|
||||
return ::mediapipe::OkStatus();
|
||||
}
|
||||
|
||||
::mediapipe::Status TransformNames(
|
||||
CalculatorGraphConfig* config,
|
||||
const std::function<std::string(absl::string_view)>& transform) {
|
||||
@@ -190,6 +215,14 @@ static ::mediapipe::Status PrefixNames(int subgraph_index,
|
||||
.SetPrepend()
|
||||
<< "while processing the output side packets of subgraph node "
|
||||
<< subgraph_node.calculator() << ": ";
|
||||
std::set<std::string> ignored_input_streams;
|
||||
MP_RETURN_IF_ERROR(FindIgnoredStreams(subgraph_config->input_stream(),
|
||||
subgraph_node.input_stream(),
|
||||
&ignored_input_streams));
|
||||
std::set<std::string> ignored_input_side_packets;
|
||||
MP_RETURN_IF_ERROR(FindIgnoredStreams(subgraph_config->input_side_packet(),
|
||||
subgraph_node.input_side_packet(),
|
||||
&ignored_input_side_packets));
|
||||
std::map<std::string, std::string>* name_map;
|
||||
auto replace_names = [&name_map](absl::string_view s) {
|
||||
std::string original(s);
|
||||
@@ -207,6 +240,12 @@ static ::mediapipe::Status PrefixNames(int subgraph_index,
|
||||
TransformStreamNames(node.mutable_input_side_packet(), replace_names));
|
||||
MP_RETURN_IF_ERROR(
|
||||
TransformStreamNames(node.mutable_output_side_packet(), replace_names));
|
||||
|
||||
// Remove input streams and side packets ignored by the subgraph-node.
|
||||
MP_RETURN_IF_ERROR(RemoveIgnoredStreams(node.mutable_input_stream(),
|
||||
ignored_input_streams));
|
||||
MP_RETURN_IF_ERROR(RemoveIgnoredStreams(node.mutable_input_side_packet(),
|
||||
ignored_input_side_packets));
|
||||
}
|
||||
name_map = &side_packet_map;
|
||||
for (auto& generator : *subgraph_config->mutable_packet_generator()) {
|
||||
|
||||
Reference in New Issue
Block a user