diff --git a/mediapipe/framework/api2/stream/BUILD b/mediapipe/framework/api2/stream/BUILD new file mode 100644 index 00000000..f9f371d2 --- /dev/null +++ b/mediapipe/framework/api2/stream/BUILD @@ -0,0 +1,14 @@ +package(default_visibility = ["//visibility:public"]) + +licenses(["notice"]) + +cc_library( + name = "loopback", + hdrs = ["loopback.h"], + deps = [ + "//mediapipe/calculators/core:previous_loopback_calculator", + "//mediapipe/framework/api2:builder", + "//mediapipe/framework/api2:port", + ], + alwayslink = 1, +) diff --git a/mediapipe/framework/api2/stream/loopback.h b/mediapipe/framework/api2/stream/loopback.h new file mode 100644 index 00000000..3ad2f0a2 --- /dev/null +++ b/mediapipe/framework/api2/stream/loopback.h @@ -0,0 +1,55 @@ +#ifndef MEDIAPIPE_FRAMEWORK_API2_STREAM_LOOPBACK_H_ +#define MEDIAPIPE_FRAMEWORK_API2_STREAM_LOOPBACK_H_ + +#include +#include + +#include "mediapipe/framework/api2/builder.h" +#include "mediapipe/framework/api2/port.h" + +namespace mediapipe::api2::builder { + +// Returns a pair of two values: +// - A stream with loopback data. Such stream, for each new packet in @tick +// stream, provides a packet previously calculated within the graph. +// - A function to define/set loopback data producing stream. +// NOTE: +// * function must be called and only once, otherwise graph validation will +// fail. +// * calling function after graph is destroyed results in undefined behavior +// +// The function wraps `PreviousLoopbackCalculator` into a convenience function +// and allows graph input to be processed together with some previous output. +// +// ------- +// +// Example: +// +// ``` +// +// Graph graph; +// Stream<...> tick = ...; // E.g. main input can surve as a tick. +// auto [prev_data, set_loopback_fn] = GetLoopbackData(tick, graph); +// ... +// Stream data = ...; +// set_loopback_fn(data); +// +// ``` +template +std::pair, std::function)>> GetLoopbackData( + Stream tick, mediapipe::api2::builder::Graph& graph) { + auto& prev = graph.AddNode("PreviousLoopbackCalculator"); + tick.ConnectTo(prev.In("MAIN")); + return {prev.Out("PREV_LOOP").template Cast(), + [prev_ptr = &prev](Stream data) { + // TODO: input stream info must be specified, but + // builder api doesn't support it at the moment. As a workaround, + // input stream info is added by GraphBuilder as a graph building + // post processing step. + data.ConnectTo(prev_ptr->In("LOOP")); + }}; +} + +} // namespace mediapipe::api2::builder + +#endif // MEDIAPIPE_FRAMEWORK_API2_STREAM_LOOPBACK_H_ diff --git a/mediapipe/framework/api2/stream/loopback_test.cc b/mediapipe/framework/api2/stream/loopback_test.cc new file mode 100644 index 00000000..8b5694db --- /dev/null +++ b/mediapipe/framework/api2/stream/loopback_test.cc @@ -0,0 +1,55 @@ +#include "mediapipe/framework/api2/stream/loopback.h" + +#include "mediapipe/framework/api2/builder.h" +#include "mediapipe/framework/api2/node.h" +#include "mediapipe/framework/api2/port.h" +#include "mediapipe/framework/port/gmock.h" +#include "mediapipe/framework/port/gtest.h" +#include "mediapipe/framework/port/parse_text_proto.h" + +namespace mediapipe::api2::builder { +namespace { + +class TestDataProducer : public NodeIntf { + public: + static constexpr Input kLoopbackData{"LOOPBACK_DATA"}; + static constexpr Output kProducedData{"PRODUCED_DATA"}; + MEDIAPIPE_NODE_INTERFACE(TestDataProducer, kLoopbackData, kProducedData); +}; + +TEST(LoopbackTest, GetLoopbackData) { + Graph graph; + + Stream tick = graph.In("TICK").Cast(); + + auto [data, set_loopback_data_fn] = GetLoopbackData(tick, graph); + + auto& producer = graph.AddNode(); + data.ConnectTo(producer[TestDataProducer::kLoopbackData]); + Stream data_to_loopback(producer[TestDataProducer::kProducedData]); + + set_loopback_data_fn(data_to_loopback); + + // PreviousLoopbackCalculator configuration is incorrect here and should be + // updated when corresponding b/175887687 is fixed. + // Use mediapipe::aimatter::GraphBuilder to fix back edges in the graph. + EXPECT_THAT(graph.GetConfig(), + testing::EqualsProto( + mediapipe::ParseTextProtoOrDie(R"pb( + node { + calculator: "PreviousLoopbackCalculator" + input_stream: "LOOP:__stream_2" + input_stream: "MAIN:__stream_0" + output_stream: "PREV_LOOP:__stream_1" + } + node { + calculator: "TestDataProducer" + input_stream: "LOOPBACK_DATA:__stream_1" + output_stream: "PRODUCED_DATA:__stream_2" + } + input_stream: "TICK:__stream_0" + )pb"))); +} + +} // namespace +} // namespace mediapipe::api2::builder