diff --git a/mediapipe/calculators/tensorflow/BUILD b/mediapipe/calculators/tensorflow/BUILD index 995ba6a7..5f5f5165 100644 --- a/mediapipe/calculators/tensorflow/BUILD +++ b/mediapipe/calculators/tensorflow/BUILD @@ -379,7 +379,6 @@ cc_library( "//mediapipe/util/sequence:media_sequence", "//mediapipe/util/sequence:media_sequence_util", "@com_google_absl//absl/container:flat_hash_map", - "@com_google_absl//absl/log", "@com_google_absl//absl/status", "@com_google_absl//absl/strings", "@org_tensorflow//tensorflow/core:protos_all_cc", diff --git a/mediapipe/calculators/tensorflow/pack_media_sequence_calculator.cc b/mediapipe/calculators/tensorflow/pack_media_sequence_calculator.cc index fcb640ff..d8702914 100644 --- a/mediapipe/calculators/tensorflow/pack_media_sequence_calculator.cc +++ b/mediapipe/calculators/tensorflow/pack_media_sequence_calculator.cc @@ -287,6 +287,13 @@ class PackMediaSequenceCalculator : public CalculatorBase { mpms::ClearClipLabelString(key, sequence_.get()); mpms::ClearClipLabelConfidence(key, sequence_.get()); } + if (absl::StartsWith(tag, kFloatContextFeaturePrefixTag)) { + const std::string& key = + tag.substr(sizeof(kFloatContextFeaturePrefixTag) / + sizeof(*kFloatContextFeaturePrefixTag) - + 1); + mpms::ClearContextFeatureFloats(key, sequence_.get()); + } if (absl::StartsWith(tag, kIntsContextFeaturePrefixTag)) { const std::string& key = tag.substr(sizeof(kIntsContextFeaturePrefixTag) / @@ -536,9 +543,10 @@ class PackMediaSequenceCalculator : public CalculatorBase { sizeof(*kFloatContextFeaturePrefixTag) - 1); RET_CHECK_EQ(cc->InputTimestamp(), Timestamp::PostStream()); - mpms::SetContextFeatureFloats( - key, cc->Inputs().Tag(tag).Get>(), - sequence_.get()); + for (const auto& value : + cc->Inputs().Tag(tag).Get>()) { + mpms::AddContextFeatureFloats(key, value, sequence_.get()); + } } if (absl::StartsWith(tag, kIntsContextFeaturePrefixTag) && !cc->Inputs().Tag(tag).IsEmpty()) { diff --git a/mediapipe/calculators/tensorflow/pack_media_sequence_calculator_test.cc b/mediapipe/calculators/tensorflow/pack_media_sequence_calculator_test.cc index 277e2d63..d9dc56e9 100644 --- a/mediapipe/calculators/tensorflow/pack_media_sequence_calculator_test.cc +++ b/mediapipe/calculators/tensorflow/pack_media_sequence_calculator_test.cc @@ -13,6 +13,7 @@ // limitations under the License. #include +#include #include #include @@ -455,6 +456,83 @@ TEST_F(PackMediaSequenceCalculatorTest, PacksTwoContextFloatLists) { testing::ElementsAre(4, 4)); } +TEST_F(PackMediaSequenceCalculatorTest, ReplaceTwoContextFloatLists) { + SetUpCalculator( + /*input_streams=*/{"FLOAT_CONTEXT_FEATURE_TEST:test", + "FLOAT_CONTEXT_FEATURE_OTHER:test2"}, + /*features=*/{}, + /*output_only_if_all_present=*/false, /*replace_instead_of_append=*/true); + auto input_sequence = std::make_unique(); + mpms::SetContextFeatureFloats("TEST", {2, 3}, input_sequence.get()); + mpms::SetContextFeatureFloats("OTHER", {2, 4}, input_sequence.get()); + + const std::vector vf_1 = {5, 6}; + runner_->MutableInputs() + ->Tag(kFloatContextFeatureTestTag) + .packets.push_back( + MakePacket>(vf_1).At(Timestamp::PostStream())); + const std::vector vf_2 = {7, 8}; + runner_->MutableInputs() + ->Tag(kFloatContextFeatureOtherTag) + .packets.push_back( + MakePacket>(vf_2).At(Timestamp::PostStream())); + + runner_->MutableSidePackets()->Tag(kSequenceExampleTag) = + Adopt(input_sequence.release()); + + MP_ASSERT_OK(runner_->Run()); + + const std::vector& output_packets = + runner_->Outputs().Tag(kSequenceExampleTag).packets; + ASSERT_EQ(1, output_packets.size()); + const tf::SequenceExample& output_sequence = + output_packets[0].Get(); + + ASSERT_THAT(mpms::GetContextFeatureFloats("TEST", output_sequence), + testing::ElementsAre(5, 6)); + ASSERT_THAT(mpms::GetContextFeatureFloats("OTHER", output_sequence), + testing::ElementsAre(7, 8)); +} + +TEST_F(PackMediaSequenceCalculatorTest, AppendTwoContextFloatLists) { + SetUpCalculator( + /*input_streams=*/{"FLOAT_CONTEXT_FEATURE_TEST:test", + "FLOAT_CONTEXT_FEATURE_OTHER:test2"}, + /*features=*/{}, + /*output_only_if_all_present=*/false, + /*replace_instead_of_append=*/false); + auto input_sequence = std::make_unique(); + mpms::SetContextFeatureFloats("TEST", {2, 3}, input_sequence.get()); + mpms::SetContextFeatureFloats("OTHER", {2, 4}, input_sequence.get()); + + const std::vector vf_1 = {5, 6}; + runner_->MutableInputs() + ->Tag(kFloatContextFeatureTestTag) + .packets.push_back( + MakePacket>(vf_1).At(Timestamp::PostStream())); + const std::vector vf_2 = {7, 8}; + runner_->MutableInputs() + ->Tag(kFloatContextFeatureOtherTag) + .packets.push_back( + MakePacket>(vf_2).At(Timestamp::PostStream())); + + runner_->MutableSidePackets()->Tag(kSequenceExampleTag) = + Adopt(input_sequence.release()); + + MP_ASSERT_OK(runner_->Run()); + + const std::vector& output_packets = + runner_->Outputs().Tag(kSequenceExampleTag).packets; + ASSERT_EQ(1, output_packets.size()); + const tf::SequenceExample& output_sequence = + output_packets[0].Get(); + + EXPECT_THAT(mpms::GetContextFeatureFloats("TEST", output_sequence), + testing::ElementsAre(2, 3, 5, 6)); + EXPECT_THAT(mpms::GetContextFeatureFloats("OTHER", output_sequence), + testing::ElementsAre(2, 4, 7, 8)); +} + TEST_F(PackMediaSequenceCalculatorTest, PackTwoContextIntLists) { SetUpCalculator( /*input_streams=*/{"INTS_CONTEXT_FEATURE_TEST:test",