From 943445fba84ec55d1833d769e1dc73f806f6444b Mon Sep 17 00:00:00 2001 From: MediaPipe Team Date: Wed, 7 Jun 2023 20:01:56 -0700 Subject: [PATCH] Update base audio/vision tasks api to suuport proto3 graph options. PiperOrigin-RevId: 538661975 --- mediapipe/tasks/cc/audio/core/BUILD | 1 + mediapipe/tasks/cc/audio/core/audio_task_api_factory.h | 10 +++------- mediapipe/tasks/cc/core/task_api_factory.h | 1 - mediapipe/tasks/cc/vision/core/BUILD | 2 ++ .../tasks/cc/vision/core/vision_task_api_factory.h | 10 +++------- 5 files changed, 9 insertions(+), 15 deletions(-) diff --git a/mediapipe/tasks/cc/audio/core/BUILD b/mediapipe/tasks/cc/audio/core/BUILD index 8372b1aa..4f821f6d 100644 --- a/mediapipe/tasks/cc/audio/core/BUILD +++ b/mediapipe/tasks/cc/audio/core/BUILD @@ -43,6 +43,7 @@ cc_library( ":base_audio_task_api", "//mediapipe/calculators/core:flow_limiter_calculator", "//mediapipe/framework:calculator_cc_proto", + "//mediapipe/tasks/cc/core:task_api_factory", "@com_google_absl//absl/status", "@com_google_absl//absl/status:statusor", "@com_google_absl//absl/strings", diff --git a/mediapipe/tasks/cc/audio/core/audio_task_api_factory.h b/mediapipe/tasks/cc/audio/core/audio_task_api_factory.h index bdac1cad..901419a5 100644 --- a/mediapipe/tasks/cc/audio/core/audio_task_api_factory.h +++ b/mediapipe/tasks/cc/audio/core/audio_task_api_factory.h @@ -27,6 +27,7 @@ limitations under the License. #include "absl/strings/str_cat.h" #include "mediapipe/framework/calculator.pb.h" #include "mediapipe/tasks/cc/audio/core/base_audio_task_api.h" +#include "mediapipe/tasks/cc/core/task_api_factory.h" #include "tensorflow/lite/core/api/op_resolver.h" namespace mediapipe { @@ -60,13 +61,8 @@ class AudioTaskApiFactory { "Task graph config should only contain one task subgraph node.", MediaPipeTasksStatus::kInvalidTaskGraphConfigError); } else { - if (!node.options().HasExtension(Options::ext)) { - return CreateStatusWithPayload( - absl::StatusCode::kInvalidArgument, - absl::StrCat(node.calculator(), - " is missing the required task options field."), - MediaPipeTasksStatus::kInvalidTaskGraphConfigError); - } + MP_RETURN_IF_ERROR( + tasks::core::TaskApiFactory::CheckHasValidOptions(node)); found_task_subgraph = true; } } diff --git a/mediapipe/tasks/cc/core/task_api_factory.h b/mediapipe/tasks/cc/core/task_api_factory.h index 6c55a656..6f604dd4 100644 --- a/mediapipe/tasks/cc/core/task_api_factory.h +++ b/mediapipe/tasks/cc/core/task_api_factory.h @@ -81,7 +81,6 @@ class TaskApiFactory { return std::make_unique(std::move(runner)); } - private: template static absl::Status CheckHasValidOptions( const CalculatorGraphConfig::Node& node) { diff --git a/mediapipe/tasks/cc/vision/core/BUILD b/mediapipe/tasks/cc/vision/core/BUILD index 59fb4562..6bcf2f5d 100644 --- a/mediapipe/tasks/cc/vision/core/BUILD +++ b/mediapipe/tasks/cc/vision/core/BUILD @@ -43,6 +43,7 @@ cc_library( "//mediapipe/framework/formats:rect_cc_proto", "//mediapipe/tasks/cc/components/containers:rect", "//mediapipe/tasks/cc/core:base_task_api", + "//mediapipe/tasks/cc/core:task_api_factory", "//mediapipe/tasks/cc/core:task_runner", "//mediapipe/tasks/cc/vision/utils:image_tensor_specs", "@com_google_absl//absl/status", @@ -58,6 +59,7 @@ cc_library( ":base_vision_task_api", "//mediapipe/calculators/core:flow_limiter_calculator", "//mediapipe/framework:calculator_cc_proto", + "//mediapipe/tasks/cc/core:task_api_factory", "@com_google_absl//absl/status", "@com_google_absl//absl/status:statusor", "@com_google_absl//absl/strings", diff --git a/mediapipe/tasks/cc/vision/core/vision_task_api_factory.h b/mediapipe/tasks/cc/vision/core/vision_task_api_factory.h index c68e432c..48fc3384 100644 --- a/mediapipe/tasks/cc/vision/core/vision_task_api_factory.h +++ b/mediapipe/tasks/cc/vision/core/vision_task_api_factory.h @@ -26,6 +26,7 @@ limitations under the License. #include "absl/status/statusor.h" #include "absl/strings/str_cat.h" #include "mediapipe/framework/calculator.pb.h" +#include "mediapipe/tasks/cc/core/task_api_factory.h" #include "mediapipe/tasks/cc/vision/core/base_vision_task_api.h" #include "tensorflow/lite/core/api/op_resolver.h" @@ -60,13 +61,8 @@ class VisionTaskApiFactory { "Task graph config should only contain one task subgraph node.", MediaPipeTasksStatus::kInvalidTaskGraphConfigError); } else { - if (!node.options().HasExtension(Options::ext)) { - return CreateStatusWithPayload( - absl::StatusCode::kInvalidArgument, - absl::StrCat(node.calculator(), - " is missing the required task options field."), - MediaPipeTasksStatus::kInvalidTaskGraphConfigError); - } + MP_RETURN_IF_ERROR( + tasks::core::TaskApiFactory::CheckHasValidOptions(node)); found_task_subgraph = true; } }