Project import generated by Copybara.
GitOrigin-RevId: ff83882955f1a1e2a043ff4e71278be9d7217bbe
This commit is contained in:
@@ -20,6 +20,7 @@ load(
|
||||
"mediapipe_binary_graph",
|
||||
)
|
||||
load("//mediapipe/framework:mediapipe_cc_test.bzl", "mediapipe_cc_test")
|
||||
load("@bazel_skylib//:bzl_library.bzl", "bzl_library")
|
||||
|
||||
licenses(["notice"])
|
||||
|
||||
@@ -29,6 +30,30 @@ exports_files([
|
||||
"simple_subgraph_template.cc",
|
||||
])
|
||||
|
||||
bzl_library(
|
||||
name = "mediapipe_graph_bzl",
|
||||
srcs = [
|
||||
"mediapipe_graph.bzl",
|
||||
],
|
||||
visibility = ["//visibility:public"],
|
||||
deps = [
|
||||
":build_defs_bzl",
|
||||
"//mediapipe/framework:encode_binary_proto",
|
||||
"//mediapipe/framework:transitive_protos_bzl",
|
||||
"//mediapipe/framework/deps:expand_template_bzl",
|
||||
],
|
||||
)
|
||||
|
||||
bzl_library(
|
||||
name = "build_defs_bzl",
|
||||
srcs = [
|
||||
"build_defs.bzl",
|
||||
],
|
||||
visibility = [
|
||||
"//mediapipe/framework:__subpackages__",
|
||||
],
|
||||
)
|
||||
|
||||
cc_library(
|
||||
name = "text_to_binary_graph",
|
||||
srcs = ["text_to_binary_graph.cc"],
|
||||
@@ -744,5 +769,7 @@ cc_test(
|
||||
|
||||
exports_files(
|
||||
["build_defs.bzl"],
|
||||
visibility = ["//mediapipe/framework:__subpackages__"],
|
||||
visibility = [
|
||||
"//mediapipe/framework:__subpackages__",
|
||||
],
|
||||
)
|
||||
|
||||
@@ -19,6 +19,7 @@
|
||||
#include "mediapipe/framework/tool/sink.h"
|
||||
|
||||
#include <memory>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "absl/strings/str_cat.h"
|
||||
@@ -168,8 +169,19 @@ void AddMultiStreamCallback(
|
||||
std::function<void(const std::vector<Packet>&)> callback,
|
||||
CalculatorGraphConfig* config,
|
||||
std::pair<std::string, Packet>* side_packet) {
|
||||
std::map<std::string, Packet> side_packets;
|
||||
AddMultiStreamCallback(streams, callback, config, &side_packets,
|
||||
/*observe_timestamp_bounds=*/false);
|
||||
*side_packet = *side_packets.begin();
|
||||
}
|
||||
|
||||
void AddMultiStreamCallback(
|
||||
const std::vector<std::string>& streams,
|
||||
std::function<void(const std::vector<Packet>&)> callback,
|
||||
CalculatorGraphConfig* config, std::map<std::string, Packet>* side_packets,
|
||||
bool observe_timestamp_bounds) {
|
||||
CHECK(config);
|
||||
CHECK(side_packet);
|
||||
CHECK(side_packets);
|
||||
CalculatorGraphConfig::Node* sink_node = config->add_node();
|
||||
const std::string name = GetUnusedNodeName(
|
||||
*config, absl::StrCat("multi_callback_", absl::StrJoin(streams, "_")));
|
||||
@@ -179,15 +191,23 @@ void AddMultiStreamCallback(
|
||||
sink_node->add_input_stream(stream_name);
|
||||
}
|
||||
|
||||
if (observe_timestamp_bounds) {
|
||||
const std::string observe_ts_bounds_packet_name = GetUnusedSidePacketName(
|
||||
*config, absl::StrCat(name, "_observe_ts_bounds"));
|
||||
sink_node->add_input_side_packet(absl::StrCat(
|
||||
"OBSERVE_TIMESTAMP_BOUNDS:", observe_ts_bounds_packet_name));
|
||||
InsertIfNotPresent(side_packets, observe_ts_bounds_packet_name,
|
||||
MakePacket<bool>(true));
|
||||
}
|
||||
const std::string input_side_packet_name =
|
||||
GetUnusedSidePacketName(*config, absl::StrCat(name, "_callback"));
|
||||
side_packet->first = input_side_packet_name;
|
||||
sink_node->add_input_side_packet(
|
||||
absl::StrCat("VECTOR_CALLBACK:", input_side_packet_name));
|
||||
|
||||
side_packet->second =
|
||||
InsertIfNotPresent(
|
||||
side_packets, input_side_packet_name,
|
||||
MakePacket<std::function<void(const std::vector<Packet>&)>>(
|
||||
std::move(callback));
|
||||
std::move(callback)));
|
||||
}
|
||||
|
||||
void AddCallbackWithHeaderCalculator(const std::string& stream_name,
|
||||
@@ -240,6 +260,10 @@ absl::Status CallbackCalculator::GetContract(CalculatorContract* cc) {
|
||||
return mediapipe::InvalidArgumentErrorBuilder(MEDIAPIPE_LOC)
|
||||
<< "InputSidePackets must use tags.";
|
||||
}
|
||||
if (cc->InputSidePackets().HasTag("OBSERVE_TIMESTAMP_BOUNDS")) {
|
||||
cc->InputSidePackets().Tag("OBSERVE_TIMESTAMP_BOUNDS").Set<bool>();
|
||||
cc->SetProcessTimestampBounds(true);
|
||||
}
|
||||
|
||||
int count = allow_multiple_streams ? cc->Inputs().NumEntries("") : 1;
|
||||
for (int i = 0; i < count; ++i) {
|
||||
@@ -266,6 +290,12 @@ absl::Status CallbackCalculator::Open(CalculatorContext* cc) {
|
||||
return mediapipe::InvalidArgumentErrorBuilder(MEDIAPIPE_LOC)
|
||||
<< "missing callback.";
|
||||
}
|
||||
if (cc->InputSidePackets().HasTag("OBSERVE_TIMESTAMP_BOUNDS") &&
|
||||
!cc->InputSidePackets().Tag("OBSERVE_TIMESTAMP_BOUNDS").Get<bool>()) {
|
||||
return mediapipe::InvalidArgumentErrorBuilder(MEDIAPIPE_LOC)
|
||||
<< "The value of the OBSERVE_TIMESTAMP_BOUNDS input side packet "
|
||||
"must be set to true";
|
||||
}
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
|
||||
@@ -115,6 +115,12 @@ void AddMultiStreamCallback(
|
||||
std::function<void(const std::vector<Packet>&)> callback,
|
||||
CalculatorGraphConfig* config, std::pair<std::string, Packet>* side_packet);
|
||||
|
||||
void AddMultiStreamCallback(
|
||||
const std::vector<std::string>& streams,
|
||||
std::function<void(const std::vector<Packet>&)> callback,
|
||||
CalculatorGraphConfig* config, std::map<std::string, Packet>* side_packets,
|
||||
bool observe_timestamp_bounds = false);
|
||||
|
||||
// Add a CallbackWithHeaderCalculator to intercept packets sent on
|
||||
// stream stream_name, and the header packet on stream stream_header.
|
||||
// The input side packet with the produced name callback_side_packet_name
|
||||
|
||||
@@ -146,5 +146,63 @@ TEST(CallbackTest, TestAddMultiStreamCallback) {
|
||||
EXPECT_THAT(sums, testing::ElementsAre(15, 7, 9));
|
||||
}
|
||||
|
||||
class TimestampBoundTestCalculator : public CalculatorBase {
|
||||
public:
|
||||
static absl::Status GetContract(CalculatorContract* cc) {
|
||||
cc->Outputs().Index(0).Set<int>();
|
||||
cc->Outputs().Index(1).Set<int>();
|
||||
return absl::OkStatus();
|
||||
}
|
||||
absl::Status Open(CalculatorContext* cc) final { return absl::OkStatus(); }
|
||||
absl::Status Process(CalculatorContext* cc) final {
|
||||
if (count_ % 5 == 0) {
|
||||
cc->Outputs().Index(0).SetNextTimestampBound(Timestamp(count_ + 1));
|
||||
cc->Outputs().Index(1).SetNextTimestampBound(Timestamp(count_ + 1));
|
||||
}
|
||||
++count_;
|
||||
if (count_ == 13) {
|
||||
return tool::StatusStop();
|
||||
}
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
private:
|
||||
int count_ = 1;
|
||||
};
|
||||
REGISTER_CALCULATOR(TimestampBoundTestCalculator);
|
||||
|
||||
TEST(CallbackTest, TestAddMultiStreamCallbackWithTimestampNotification) {
|
||||
std::string config_str = R"(
|
||||
node {
|
||||
calculator: "TimestampBoundTestCalculator"
|
||||
output_stream: "foo"
|
||||
output_stream: "bar"
|
||||
}
|
||||
)";
|
||||
CalculatorGraphConfig graph_config =
|
||||
mediapipe::ParseTextProtoOrDie<CalculatorGraphConfig>(config_str);
|
||||
|
||||
std::vector<int> sums;
|
||||
|
||||
std::map<std::string, Packet> side_packets;
|
||||
tool::AddMultiStreamCallback(
|
||||
{"foo", "bar"},
|
||||
[&sums](const std::vector<Packet>& packets) {
|
||||
Packet foo_p = packets[0];
|
||||
Packet bar_p = packets[1];
|
||||
ASSERT_TRUE(foo_p.IsEmpty() && bar_p.IsEmpty());
|
||||
int foo = foo_p.Timestamp().Value();
|
||||
int bar = bar_p.Timestamp().Value();
|
||||
sums.push_back(foo + bar);
|
||||
},
|
||||
&graph_config, &side_packets, true);
|
||||
|
||||
CalculatorGraph graph(graph_config);
|
||||
MP_ASSERT_OK(graph.StartRun(side_packets));
|
||||
MP_ASSERT_OK(graph.WaitUntilDone());
|
||||
|
||||
EXPECT_THAT(sums, testing::ElementsAre(10, 20));
|
||||
}
|
||||
|
||||
} // namespace
|
||||
} // namespace mediapipe
|
||||
|
||||
Reference in New Issue
Block a user