Project import generated by Copybara.

GitOrigin-RevId: 5aca6b3f07b67e09988a901f50f595ca5f566e67
This commit is contained in:
MediaPipe Team
2019-11-15 13:10:50 -08:00
committed by jqtang
parent d030c13931
commit 9437483827
116 changed files with 6284 additions and 391 deletions
+1 -1
View File
@@ -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",
],
)
+1 -1
View File
@@ -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()) {