diff --git a/mediapipe/modules/objectron/calculators/filter_detection_calculator.cc b/mediapipe/modules/objectron/calculators/filter_detection_calculator.cc index db0f2748..29f4c79d 100644 --- a/mediapipe/modules/objectron/calculators/filter_detection_calculator.cc +++ b/mediapipe/modules/objectron/calculators/filter_detection_calculator.cc @@ -37,6 +37,7 @@ constexpr char kDetectionTag[] = "DETECTION"; constexpr char kDetectionsTag[] = "DETECTIONS"; constexpr char kLabelsTag[] = "LABELS"; constexpr char kLabelsCsvTag[] = "LABELS_CSV"; +constexpr char kLabelMapTag[] = "LABEL_MAP"; using mediapipe::RE2; using Detections = std::vector; @@ -151,6 +152,11 @@ absl::Status FilterDetectionCalculator::GetContract(CalculatorContract* cc) { if (cc->InputSidePackets().HasTag(kLabelsCsvTag)) { cc->InputSidePackets().Tag(kLabelsCsvTag).Set(); } + if (cc->InputSidePackets().HasTag(kLabelMapTag)) { + cc->InputSidePackets() + .Tag(kLabelMapTag) + .Set>>(); + } return absl::OkStatus(); } @@ -158,7 +164,8 @@ absl::Status FilterDetectionCalculator::Open(CalculatorContext* cc) { cc->SetOffset(TimestampDiff(0)); options_ = cc->Options(); limit_labels_ = cc->InputSidePackets().HasTag(kLabelsTag) || - cc->InputSidePackets().HasTag(kLabelsCsvTag); + cc->InputSidePackets().HasTag(kLabelsCsvTag) || + cc->InputSidePackets().HasTag(kLabelMapTag); if (limit_labels_) { Strings allowlist_labels; if (cc->InputSidePackets().HasTag(kLabelsCsvTag)) { @@ -168,8 +175,16 @@ absl::Status FilterDetectionCalculator::Open(CalculatorContext* cc) { for (auto& e : allowlist_labels) { absl::StripAsciiWhitespace(&e); } - } else { + } else if (cc->InputSidePackets().HasTag(kLabelsTag)) { allowlist_labels = cc->InputSidePackets().Tag(kLabelsTag).Get(); + } else if (cc->InputSidePackets().HasTag(kLabelMapTag)) { + auto label_map = cc->InputSidePackets() + .Tag(kLabelMapTag) + .Get>>() + .get(); + for (const auto& [_, v] : *label_map) { + allowlist_labels.push_back(v); + } } allowed_labels_.insert(allowlist_labels.begin(), allowlist_labels.end()); } diff --git a/mediapipe/modules/objectron/calculators/filter_detection_calculator_test.cc b/mediapipe/modules/objectron/calculators/filter_detection_calculator_test.cc index 958fe4c5..10e750a4 100644 --- a/mediapipe/modules/objectron/calculators/filter_detection_calculator_test.cc +++ b/mediapipe/modules/objectron/calculators/filter_detection_calculator_test.cc @@ -67,5 +67,68 @@ TEST(FilterDetectionCalculatorTest, DetectionFilterTest) { )); } +TEST(FilterDetectionCalculatorTest, DetectionFilterLabelMapTest) { + auto runner = std::make_unique( + ParseTextProtoOrDie(R"pb( + calculator: "FilterDetectionCalculator" + input_stream: "DETECTION:input" + input_side_packet: "LABEL_MAP:input_map" + output_stream: "DETECTION:output" + options { + [mediapipe.FilterDetectionCalculatorOptions.ext]: { min_score: 0.6 } + } + )pb")); + + runner->MutableInputs()->Tag("DETECTION").packets = { + MakePacket(ParseTextProtoOrDie(R"pb( + label: "a" + label: "b" + label: "c" + label: "d" + score: 1 + score: 0.8 + score: 0.3 + score: 0.9 + )pb")) + .At(Timestamp(20)), + MakePacket(ParseTextProtoOrDie(R"pb( + label: "a" + label: "b" + label: "c" + label: "e" + score: 0.6 + score: 0.4 + score: 0.2 + score: 0.7 + )pb")) + .At(Timestamp(40)), + }; + + auto label_map = std::make_unique>(); + (*label_map)[0] = "a"; + (*label_map)[1] = "b"; + (*label_map)[2] = "c"; + runner->MutableSidePackets()->Tag("LABEL_MAP") = + AdoptAsUniquePtr(label_map.release()); + + // Run graph. + MP_ASSERT_OK(runner->Run()); + + // Check output. + EXPECT_THAT( + runner->Outputs().Tag("DETECTION").packets, + ElementsAre(PacketContainsTimestampAndPayload( + Eq(Timestamp(20)), + EqualsProto(R"pb( + label: "a" label: "b" score: 1 score: 0.8 + )pb")), // Packet 1 at timestamp 20. + PacketContainsTimestampAndPayload( + Eq(Timestamp(40)), + EqualsProto(R"pb( + label: "a" score: 0.6 + )pb")) // Packet 2 at timestamp 40. + )); +} + } // namespace } // namespace mediapipe