From 8837b49026cceee4f08f930da9772c8fd4e42c8b Mon Sep 17 00:00:00 2001 From: MediaPipe Team Date: Wed, 27 Sep 2023 16:50:37 -0700 Subject: [PATCH] get_vector_item stream utility function. PiperOrigin-RevId: 568998504 --- mediapipe/framework/api2/stream/BUILD | 31 +++++ .../framework/api2/stream/get_vector_item.h | 66 +++++++++ .../api2/stream/get_vector_item_test.cc | 130 ++++++++++++++++++ 3 files changed, 227 insertions(+) create mode 100644 mediapipe/framework/api2/stream/get_vector_item.h create mode 100644 mediapipe/framework/api2/stream/get_vector_item_test.cc diff --git a/mediapipe/framework/api2/stream/BUILD b/mediapipe/framework/api2/stream/BUILD index 57533745..88607844 100644 --- a/mediapipe/framework/api2/stream/BUILD +++ b/mediapipe/framework/api2/stream/BUILD @@ -64,6 +64,37 @@ cc_test( ], ) +cc_library( + name = "get_vector_item", + hdrs = ["get_vector_item.h"], + deps = [ + "//mediapipe/calculators/core:get_vector_item_calculator", + "//mediapipe/framework/api2:builder", + "//mediapipe/framework/api2:port", + "//mediapipe/framework/formats:classification_cc_proto", + "//mediapipe/framework/formats:landmark_cc_proto", + "//mediapipe/framework/formats:rect_cc_proto", + "@org_tensorflow//tensorflow/lite/c:common", + ], +) + +cc_test( + name = "get_vector_item_test", + srcs = ["get_vector_item_test.cc"], + deps = [ + ":get_vector_item", + "//mediapipe/framework:calculator_framework", + "//mediapipe/framework/api2:builder", + "//mediapipe/framework/formats:classification_cc_proto", + "//mediapipe/framework/formats:landmark_cc_proto", + "//mediapipe/framework/formats:rect_cc_proto", + "//mediapipe/framework/port:gtest", + "//mediapipe/framework/port:gtest_main", + "//mediapipe/framework/port:parse_text_proto", + "//mediapipe/framework/port:status_matchers", + ], +) + cc_library( name = "landmarks_to_detection", srcs = ["landmarks_to_detection.cc"], diff --git a/mediapipe/framework/api2/stream/get_vector_item.h b/mediapipe/framework/api2/stream/get_vector_item.h new file mode 100644 index 00000000..acdddeda --- /dev/null +++ b/mediapipe/framework/api2/stream/get_vector_item.h @@ -0,0 +1,66 @@ +#ifndef MEDIAPIPE_FRAMEWORK_API2_STREAM_GET_VECTOR_ITEM_H_ +#define MEDIAPIPE_FRAMEWORK_API2_STREAM_GET_VECTOR_ITEM_H_ + +#include +#include + +#include "mediapipe/calculators/core/get_vector_item_calculator.h" +#include "mediapipe/framework/api2/builder.h" +#include "mediapipe/framework/api2/port.h" +#include "mediapipe/framework/formats/classification.pb.h" +#include "mediapipe/framework/formats/landmark.pb.h" +#include "mediapipe/framework/formats/rect.pb.h" +#include "tensorflow/lite/c/common.h" + +namespace mediapipe::api2::builder { + +namespace internal_get_vector_item { + +// Helper function that adds a node to a graph, that is capable of getting item +// from a vector of type (T). +template +mediapipe::api2::builder::GenericNode& AddGetVectorItemNode( + mediapipe::api2::builder::Graph& graph) { + if constexpr (std::is_same_v) { + return graph.AddNode("GetNormalizedLandmarkListVectorItemCalculator"); + } else if constexpr (std::is_same_v) { + return graph.AddNode("GetLandmarkListVectorItemCalculator"); + } else if constexpr (std::is_same_v) { + return graph.AddNode("GetClassificationListVectorItemCalculator"); + } else if constexpr (std::is_same_v) { + return graph.AddNode("GetNormalizedRectVectorItemCalculator"); + } else if constexpr (std::is_same_v) { + return graph.AddNode("GetRectVectorItemCalculator"); + } else { + static_assert( + dependent_false::value, + "Get vector item node is not available for the specified type."); + } +} + +} // namespace internal_get_vector_item + +// Gets item from the vector. +// +// Example: +// ``` +// +// Graph graph; +// +// Stream> multi_landmarks = ...; +// Stream landmarks = +// GetItem(multi_landmarks, 0, graph); +// +// ``` +template +Stream GetItem(Stream> items, Stream idx, + mediapipe::api2::builder::Graph& graph) { + auto& getter = internal_get_vector_item::AddGetVectorItemNode(graph); + items.ConnectTo(getter.In("VECTOR")); + idx.ConnectTo(getter.In("INDEX")); + return getter.Out("ITEM").template Cast(); +} + +} // namespace mediapipe::api2::builder + +#endif // MEDIAPIPE_FRAMEWORK_API2_STREAM_GET_VECTOR_ITEM_H_ diff --git a/mediapipe/framework/api2/stream/get_vector_item_test.cc b/mediapipe/framework/api2/stream/get_vector_item_test.cc new file mode 100644 index 00000000..b07e797b --- /dev/null +++ b/mediapipe/framework/api2/stream/get_vector_item_test.cc @@ -0,0 +1,130 @@ +#include "mediapipe/framework/api2/stream/get_vector_item.h" + +#include + +#include "mediapipe/framework/api2/builder.h" +#include "mediapipe/framework/calculator_framework.h" +#include "mediapipe/framework/formats/classification.pb.h" +#include "mediapipe/framework/formats/landmark.pb.h" +#include "mediapipe/framework/formats/rect.pb.h" +#include "mediapipe/framework/port/gmock.h" +#include "mediapipe/framework/port/gtest.h" +#include "mediapipe/framework/port/parse_text_proto.h" +#include "mediapipe/framework/port/status_matchers.h" + +namespace mediapipe::api2::builder { +namespace { + +using ::mediapipe::api2::builder::Graph; + +TEST(GetItem, GetNormalizedLandmarkListVectorItem) { + Graph graph; + Stream> items = + graph.In("ITEMS").Cast>(); + Stream idx = graph.In("IDX").Cast(); + Stream item = GetItem(items, idx, graph); + item.SetName("item"); + EXPECT_THAT(graph.GetConfig(), + EqualsProto(ParseTextProtoOrDie(R"pb( + node { + calculator: "GetNormalizedLandmarkListVectorItemCalculator" + input_stream: "INDEX:__stream_0" + input_stream: "VECTOR:__stream_1" + output_stream: "ITEM:item" + } + input_stream: "IDX:__stream_0" + input_stream: "ITEMS:__stream_1" + )pb"))); + CalculatorGraph calculator_graph; + MP_EXPECT_OK(calculator_graph.Initialize(graph.GetConfig())); +} + +TEST(GetItem, GetLandmarkListVectorItem) { + Graph graph; + Stream> items = + graph.In("ITEMS").Cast>(); + Stream idx = graph.In("IDX").Cast(); + Stream item = GetItem(items, idx, graph); + item.SetName("item"); + EXPECT_THAT(graph.GetConfig(), + EqualsProto(ParseTextProtoOrDie(R"pb( + node { + calculator: "GetLandmarkListVectorItemCalculator" + input_stream: "INDEX:__stream_0" + input_stream: "VECTOR:__stream_1" + output_stream: "ITEM:item" + } + input_stream: "IDX:__stream_0" + input_stream: "ITEMS:__stream_1" + )pb"))); + CalculatorGraph calculator_graph; + MP_EXPECT_OK(calculator_graph.Initialize(graph.GetConfig())); +} + +TEST(GetItem, GetClassificationListVectorItem) { + Graph graph; + Stream> items = + graph.In("ITEMS").Cast>(); + Stream idx = graph.In("IDX").Cast(); + Stream item = GetItem(items, idx, graph); + item.SetName("item"); + EXPECT_THAT(graph.GetConfig(), + EqualsProto(ParseTextProtoOrDie(R"pb( + node { + calculator: "GetClassificationListVectorItemCalculator" + input_stream: "INDEX:__stream_0" + input_stream: "VECTOR:__stream_1" + output_stream: "ITEM:item" + } + input_stream: "IDX:__stream_0" + input_stream: "ITEMS:__stream_1" + )pb"))); + CalculatorGraph calculator_graph; + MP_EXPECT_OK(calculator_graph.Initialize(graph.GetConfig())); +} + +TEST(GetItem, GetNormalizedRectVectorItem) { + Graph graph; + Stream> items = + graph.In("ITEMS").Cast>(); + Stream idx = graph.In("IDX").Cast(); + Stream item = GetItem(items, idx, graph); + item.SetName("item"); + EXPECT_THAT(graph.GetConfig(), + EqualsProto(ParseTextProtoOrDie(R"pb( + node { + calculator: "GetNormalizedRectVectorItemCalculator" + input_stream: "INDEX:__stream_0" + input_stream: "VECTOR:__stream_1" + output_stream: "ITEM:item" + } + input_stream: "IDX:__stream_0" + input_stream: "ITEMS:__stream_1" + )pb"))); + CalculatorGraph calculator_graph; + MP_EXPECT_OK(calculator_graph.Initialize(graph.GetConfig())); +} + +TEST(GetItem, GetRectVectorItem) { + Graph graph; + Stream> items = graph.In("ITEMS").Cast>(); + Stream idx = graph.In("IDX").Cast(); + Stream item = GetItem(items, idx, graph); + item.SetName("item"); + EXPECT_THAT(graph.GetConfig(), + EqualsProto(ParseTextProtoOrDie(R"pb( + node { + calculator: "GetRectVectorItemCalculator" + input_stream: "INDEX:__stream_0" + input_stream: "VECTOR:__stream_1" + output_stream: "ITEM:item" + } + input_stream: "IDX:__stream_0" + input_stream: "ITEMS:__stream_1" + )pb"))); + CalculatorGraph calculator_graph; + MP_EXPECT_OK(calculator_graph.Initialize(graph.GetConfig())); +} + +} // namespace +} // namespace mediapipe::api2::builder