diff --git a/mediapipe/tasks/web/vision/core/BUILD b/mediapipe/tasks/web/vision/core/BUILD index 7ab822b7..8c405ae6 100644 --- a/mediapipe/tasks/web/vision/core/BUILD +++ b/mediapipe/tasks/web/vision/core/BUILD @@ -1,11 +1,26 @@ # This package contains options shared by all MediaPipe Tasks for Web. -load("//mediapipe/framework/port:build_config.bzl", "mediapipe_ts_library") +load("//mediapipe/framework/port:build_config.bzl", "mediapipe_ts_declaration", "mediapipe_ts_library") package(default_visibility = ["//mediapipe/tasks:internal"]) +mediapipe_ts_declaration( + name = "vision_task_options", + srcs = ["vision_task_options.d.ts"], + deps = [ + "//mediapipe/tasks/web/core", + ], +) + mediapipe_ts_library( - name = "running_mode", - srcs = ["running_mode.ts"], - deps = ["//mediapipe/tasks/cc/core/proto:base_options_jspb_proto"], + name = "vision_task_runner", + srcs = ["vision_task_runner.ts"], + deps = [ + ":vision_task_options", + "//mediapipe/tasks/cc/core/proto:base_options_jspb_proto", + "//mediapipe/tasks/web/components/processors:base_options", + "//mediapipe/tasks/web/core", + "//mediapipe/tasks/web/core:task_runner", + "//mediapipe/web/graph_runner:wasm_mediapipe_lib_ts", + ], ) diff --git a/mediapipe/tasks/web/vision/core/running_mode.ts b/mediapipe/tasks/web/vision/core/vision_task_options.d.ts similarity index 58% rename from mediapipe/tasks/web/vision/core/running_mode.ts rename to mediapipe/tasks/web/vision/core/vision_task_options.d.ts index 1e9b1b9a..8b9562e4 100644 --- a/mediapipe/tasks/web/vision/core/running_mode.ts +++ b/mediapipe/tasks/web/vision/core/vision_task_options.d.ts @@ -14,23 +14,26 @@ * limitations under the License. */ -import {BaseOptions as BaseOptionsProto} from '../../../../tasks/cc/core/proto/base_options_pb'; +import {BaseOptions} from '../../../../tasks/web/core/base_options'; /** - * The running mode of a task. + * The two running modes of a video task. * 1) The image mode for processing single image inputs. * 2) The video mode for processing decoded frames of a video. */ export type RunningMode = 'image'|'video'; -/** Configues the `useStreamMode` option . */ -export function configureRunningMode( - options: {runningMode?: RunningMode}, - proto?: BaseOptionsProto): BaseOptionsProto { - proto = proto ?? new BaseOptionsProto(); - if ('runningMode' in options) { - const useStreamMode = options.runningMode === 'video'; - proto.setUseStreamMode(useStreamMode); - } - return proto; + +/** The options for configuring a MediaPipe vision task. */ +export declare interface VisionTaskOptions { + /** Options to configure the loading of the model assets. */ + baseOptions?: BaseOptions; + + /** + * The running mode of the task. Default to the image mode. + * Vision tasks have two running modes: + * 1) The image mode for processing single image inputs. + * 2) The video mode for processing decoded frames of a video. + */ + runningMode?: RunningMode; } diff --git a/mediapipe/tasks/web/vision/core/vision_task_runner.ts b/mediapipe/tasks/web/vision/core/vision_task_runner.ts new file mode 100644 index 00000000..372ce9ba --- /dev/null +++ b/mediapipe/tasks/web/vision/core/vision_task_runner.ts @@ -0,0 +1,66 @@ +/** + * Copyright 2022 The MediaPipe Authors. All Rights Reserved. + * + * Licensed under the Apache License, Version 2.0 (the "License"); + * you may not use this file except in compliance with the License. + * You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +import {BaseOptions as BaseOptionsProto} from '../../../../tasks/cc/core/proto/base_options_pb'; +import {convertBaseOptionsToProto} from '../../../../tasks/web/components/processors/base_options'; +import {TaskRunner} from '../../../../tasks/web/core/task_runner'; +import {ImageSource} from '../../../../web/graph_runner/wasm_mediapipe_lib'; + +import {VisionTaskOptions} from './vision_task_options'; + +/** Base class for all MediaPipe Vision Tasks. */ +export abstract class VisionTaskRunner extends TaskRunner { + protected abstract baseOptions?: BaseOptionsProto|undefined; + + /** Configures the shared options of a vision task. */ + async setOptions(options: VisionTaskOptions): Promise { + this.baseOptions = this.baseOptions ?? new BaseOptionsProto(); + if (options.baseOptions) { + this.baseOptions = await convertBaseOptionsToProto( + options.baseOptions, this.baseOptions); + } + if ('runningMode' in options) { + const useStreamMode = + !!options.runningMode && options.runningMode !== 'image'; + this.baseOptions.setUseStreamMode(useStreamMode); + } + } + + /** Sends an image packet to the graph and awaits results. */ + protected abstract process(input: ImageSource, timestamp: number): T; + + /** Sends a single image to the graph and awaits results. */ + protected processImageData(image: ImageSource): T { + if (!!this.baseOptions?.getUseStreamMode()) { + throw new Error( + 'Task is not initialized with image mode. ' + + '\'runningMode\' must be set to \'image\'.'); + } + return this.process(image, performance.now()); + } + + /** Sends a single video frame to the graph and awaits results. */ + protected processVideoData(imageFrame: ImageSource, timestamp: number): T { + if (!this.baseOptions?.getUseStreamMode()) { + throw new Error( + 'Task is not initialized with video mode. ' + + '\'runningMode\' must be set to \'video\'.'); + } + return this.process(imageFrame, timestamp); + } +} + + diff --git a/mediapipe/tasks/web/vision/gesture_recognizer/BUILD b/mediapipe/tasks/web/vision/gesture_recognizer/BUILD index d67974a1..f2b66823 100644 --- a/mediapipe/tasks/web/vision/gesture_recognizer/BUILD +++ b/mediapipe/tasks/web/vision/gesture_recognizer/BUILD @@ -19,6 +19,7 @@ mediapipe_ts_library( "//mediapipe/framework/formats:classification_jspb_proto", "//mediapipe/framework/formats:landmark_jspb_proto", "//mediapipe/framework/formats:rect_jspb_proto", + "//mediapipe/tasks/cc/core/proto:base_options_jspb_proto", "//mediapipe/tasks/cc/vision/gesture_recognizer/proto:gesture_classifier_graph_options_jspb_proto", "//mediapipe/tasks/cc/vision/gesture_recognizer/proto:gesture_recognizer_graph_options_jspb_proto", "//mediapipe/tasks/cc/vision/gesture_recognizer/proto:hand_gesture_recognizer_graph_options_jspb_proto", @@ -27,11 +28,10 @@ mediapipe_ts_library( "//mediapipe/tasks/cc/vision/hand_landmarker/proto:hand_landmarks_detector_graph_options_jspb_proto", "//mediapipe/tasks/web/components/containers:category", "//mediapipe/tasks/web/components/containers:landmark", - "//mediapipe/tasks/web/components/processors:base_options", "//mediapipe/tasks/web/components/processors:classifier_options", "//mediapipe/tasks/web/core", "//mediapipe/tasks/web/core:classifier_options", - "//mediapipe/tasks/web/core:task_runner", + "//mediapipe/tasks/web/vision/core:vision_task_runner", "//mediapipe/web/graph_runner:wasm_mediapipe_lib_ts", ], ) @@ -47,5 +47,6 @@ mediapipe_ts_declaration( "//mediapipe/tasks/web/components/containers:landmark", "//mediapipe/tasks/web/core", "//mediapipe/tasks/web/core:classifier_options", + "//mediapipe/tasks/web/vision/core:vision_task_options", ], ) diff --git a/mediapipe/tasks/web/vision/gesture_recognizer/gesture_recognizer.ts b/mediapipe/tasks/web/vision/gesture_recognizer/gesture_recognizer.ts index 6c8072ff..8e745534 100644 --- a/mediapipe/tasks/web/vision/gesture_recognizer/gesture_recognizer.ts +++ b/mediapipe/tasks/web/vision/gesture_recognizer/gesture_recognizer.ts @@ -19,6 +19,7 @@ import {CalculatorOptions} from '../../../../framework/calculator_options_pb'; import {ClassificationList} from '../../../../framework/formats/classification_pb'; import {LandmarkList, NormalizedLandmarkList} from '../../../../framework/formats/landmark_pb'; import {NormalizedRect} from '../../../../framework/formats/rect_pb'; +import {BaseOptions as BaseOptionsProto} from '../../../../tasks/cc/core/proto/base_options_pb'; import {GestureClassifierGraphOptions} from '../../../../tasks/cc/vision/gesture_recognizer/proto/gesture_classifier_graph_options_pb'; import {GestureRecognizerGraphOptions} from '../../../../tasks/cc/vision/gesture_recognizer/proto/gesture_recognizer_graph_options_pb'; import {HandGestureRecognizerGraphOptions} from '../../../../tasks/cc/vision/gesture_recognizer/proto/hand_gesture_recognizer_graph_options_pb'; @@ -27,10 +28,9 @@ import {HandLandmarkerGraphOptions} from '../../../../tasks/cc/vision/hand_landm import {HandLandmarksDetectorGraphOptions} from '../../../../tasks/cc/vision/hand_landmarker/proto/hand_landmarks_detector_graph_options_pb'; import {Category} from '../../../../tasks/web/components/containers/category'; import {Landmark} from '../../../../tasks/web/components/containers/landmark'; -import {convertBaseOptionsToProto} from '../../../../tasks/web/components/processors/base_options'; import {convertClassifierOptionsToProto} from '../../../../tasks/web/components/processors/classifier_options'; -import {TaskRunner} from '../../../../tasks/web/core/task_runner'; import {WasmLoaderOptions} from '../../../../tasks/web/core/wasm_loader_options'; +import {VisionTaskRunner} from '../../../../tasks/web/vision/core/vision_task_runner'; import {createMediaPipeLib, FileLocator, ImageSource, WasmModule} from '../../../../web/graph_runner/wasm_mediapipe_lib'; // Placeholder for internal dependency on trusted resource url @@ -64,7 +64,8 @@ FULL_IMAGE_RECT.setWidth(1); FULL_IMAGE_RECT.setHeight(1); /** Performs hand gesture recognition on images. */ -export class GestureRecognizer extends TaskRunner { +export class GestureRecognizer extends + VisionTaskRunner { private gestures: Category[][] = []; private landmarks: Landmark[][] = []; private worldLandmarks: Landmark[][] = []; @@ -156,10 +157,14 @@ export class GestureRecognizer extends TaskRunner { this.handGestureRecognizerGraphOptions); this.initDefaults(); + } - // Disables the automatic render-to-screen code, which allows for pure - // CPU processing. - this.setAutoRenderToScreen(false); + protected override get baseOptions(): BaseOptionsProto|undefined { + return this.options.getBaseOptions(); + } + + protected override set baseOptions(proto: BaseOptionsProto|undefined) { + this.options.setBaseOptions(proto); } /** @@ -171,12 +176,8 @@ export class GestureRecognizer extends TaskRunner { * * @param options The options for the gesture recognizer. */ - async setOptions(options: GestureRecognizerOptions): Promise { - if (options.baseOptions) { - const baseOptionsProto = await convertBaseOptionsToProto( - options.baseOptions, this.options.getBaseOptions()); - this.options.setBaseOptions(baseOptionsProto); - } + override async setOptions(options: GestureRecognizerOptions): Promise { + await super.setOptions(options); if ('numHands' in options) { this.handDetectorGraphOptions.setNumHands( @@ -233,12 +234,27 @@ export class GestureRecognizer extends TaskRunner { /** * Performs gesture recognition on the provided single image and waits * synchronously for the response. - * @param imageSource An image source to process. - * @param timestamp The timestamp of the current frame, in ms. If not - * provided, defaults to `performance.now()`. + * @param image A single image to process. * @return The detected gestures. */ - recognize(imageSource: ImageSource, timestamp: number = performance.now()): + recognize(image: ImageSource): GestureRecognizerResult { + return this.processImageData(image); + } + + /** + * Performs gesture recognition on the provided video frame and waits + * synchronously for the response. + * @param videoFrame A video frame to process. + * @param timestamp The timestamp of the current frame, in ms. + * @return The detected gestures. + */ + recognizeForVideo(videoFrame: ImageSource, timestamp: number): + GestureRecognizerResult { + return this.processVideoData(videoFrame, timestamp); + } + + /** Runs the gesture recognition and blocks on the response. */ + protected override process(imageSource: ImageSource, timestamp: number): GestureRecognizerResult { this.gestures = []; this.landmarks = []; diff --git a/mediapipe/tasks/web/vision/gesture_recognizer/gesture_recognizer_options.d.ts b/mediapipe/tasks/web/vision/gesture_recognizer/gesture_recognizer_options.d.ts index 45601a74..dd8fc954 100644 --- a/mediapipe/tasks/web/vision/gesture_recognizer/gesture_recognizer_options.d.ts +++ b/mediapipe/tasks/web/vision/gesture_recognizer/gesture_recognizer_options.d.ts @@ -14,14 +14,11 @@ * limitations under the License. */ -import {BaseOptions} from '../../../../tasks/web/core/base_options'; import {ClassifierOptions} from '../../../../tasks/web/core/classifier_options'; +import {VisionTaskOptions} from '../../../../tasks/web/vision/core/vision_task_options'; /** Options to configure the MediaPipe Gesture Recognizer Task */ -export declare interface GestureRecognizerOptions { - /** Options to configure the loading of the model assets. */ - baseOptions?: BaseOptions; - +export declare interface GestureRecognizerOptions extends VisionTaskOptions { /** * The maximum number of hands can be detected by the GestureRecognizer. * Defaults to 1. diff --git a/mediapipe/tasks/web/vision/hand_landmarker/BUILD b/mediapipe/tasks/web/vision/hand_landmarker/BUILD index 25c70e0a..36f1d7eb 100644 --- a/mediapipe/tasks/web/vision/hand_landmarker/BUILD +++ b/mediapipe/tasks/web/vision/hand_landmarker/BUILD @@ -19,14 +19,14 @@ mediapipe_ts_library( "//mediapipe/framework/formats:classification_jspb_proto", "//mediapipe/framework/formats:landmark_jspb_proto", "//mediapipe/framework/formats:rect_jspb_proto", + "//mediapipe/tasks/cc/core/proto:base_options_jspb_proto", "//mediapipe/tasks/cc/vision/hand_detector/proto:hand_detector_graph_options_jspb_proto", "//mediapipe/tasks/cc/vision/hand_landmarker/proto:hand_landmarker_graph_options_jspb_proto", "//mediapipe/tasks/cc/vision/hand_landmarker/proto:hand_landmarks_detector_graph_options_jspb_proto", "//mediapipe/tasks/web/components/containers:category", "//mediapipe/tasks/web/components/containers:landmark", - "//mediapipe/tasks/web/components/processors:base_options", "//mediapipe/tasks/web/core", - "//mediapipe/tasks/web/core:task_runner", + "//mediapipe/tasks/web/vision/core:vision_task_runner", "//mediapipe/web/graph_runner:wasm_mediapipe_lib_ts", ], ) @@ -41,5 +41,6 @@ mediapipe_ts_declaration( "//mediapipe/tasks/web/components/containers:category", "//mediapipe/tasks/web/components/containers:landmark", "//mediapipe/tasks/web/core", + "//mediapipe/tasks/web/vision/core:vision_task_options", ], ) diff --git a/mediapipe/tasks/web/vision/hand_landmarker/hand_landmarker.ts b/mediapipe/tasks/web/vision/hand_landmarker/hand_landmarker.ts index af10305b..0aba5c82 100644 --- a/mediapipe/tasks/web/vision/hand_landmarker/hand_landmarker.ts +++ b/mediapipe/tasks/web/vision/hand_landmarker/hand_landmarker.ts @@ -19,14 +19,14 @@ import {CalculatorOptions} from '../../../../framework/calculator_options_pb'; import {ClassificationList} from '../../../../framework/formats/classification_pb'; import {LandmarkList, NormalizedLandmarkList} from '../../../../framework/formats/landmark_pb'; import {NormalizedRect} from '../../../../framework/formats/rect_pb'; +import {BaseOptions as BaseOptionsProto} from '../../../../tasks/cc/core/proto/base_options_pb'; import {HandDetectorGraphOptions} from '../../../../tasks/cc/vision/hand_detector/proto/hand_detector_graph_options_pb'; import {HandLandmarkerGraphOptions} from '../../../../tasks/cc/vision/hand_landmarker/proto/hand_landmarker_graph_options_pb'; import {HandLandmarksDetectorGraphOptions} from '../../../../tasks/cc/vision/hand_landmarker/proto/hand_landmarks_detector_graph_options_pb'; import {Category} from '../../../../tasks/web/components/containers/category'; import {Landmark} from '../../../../tasks/web/components/containers/landmark'; -import {convertBaseOptionsToProto} from '../../../../tasks/web/components/processors/base_options'; -import {TaskRunner} from '../../../../tasks/web/core/task_runner'; import {WasmLoaderOptions} from '../../../../tasks/web/core/wasm_loader_options'; +import {VisionTaskRunner} from '../../../../tasks/web/vision/core/vision_task_runner'; import {createMediaPipeLib, FileLocator, ImageSource, WasmModule} from '../../../../web/graph_runner/wasm_mediapipe_lib'; // Placeholder for internal dependency on trusted resource url @@ -58,7 +58,7 @@ FULL_IMAGE_RECT.setWidth(1); FULL_IMAGE_RECT.setHeight(1); /** Performs hand landmarks detection on images. */ -export class HandLandmarker extends TaskRunner { +export class HandLandmarker extends VisionTaskRunner { private landmarks: Landmark[][] = []; private worldLandmarks: Landmark[][] = []; private handednesses: Category[][] = []; @@ -138,10 +138,14 @@ export class HandLandmarker extends TaskRunner { this.options.setHandDetectorGraphOptions(this.handDetectorGraphOptions); this.initDefaults(); + } - // Disables the automatic render-to-screen code, which allows for pure - // CPU processing. - this.setAutoRenderToScreen(false); + protected override get baseOptions(): BaseOptionsProto|undefined { + return this.options.getBaseOptions(); + } + + protected override set baseOptions(proto: BaseOptionsProto|undefined) { + this.options.setBaseOptions(proto); } /** @@ -153,12 +157,8 @@ export class HandLandmarker extends TaskRunner { * * @param options The options for the hand landmarker. */ - async setOptions(options: HandLandmarkerOptions): Promise { - if (options.baseOptions) { - const baseOptionsProto = await convertBaseOptionsToProto( - options.baseOptions, this.options.getBaseOptions()); - this.options.setBaseOptions(baseOptionsProto); - } + override async setOptions(options: HandLandmarkerOptions): Promise { + await super.setOptions(options); // Configure hand detector options. if ('numHands' in options) { @@ -186,12 +186,27 @@ export class HandLandmarker extends TaskRunner { /** * Performs hand landmarks detection on the provided single image and waits * synchronously for the response. - * @param imageSource An image source to process. - * @param timestamp The timestamp of the current frame, in ms. If not - * provided, defaults to `performance.now()`. + * @param image An image to process. * @return The detected hand landmarks. */ - detect(imageSource: ImageSource, timestamp: number = performance.now()): + detect(image: ImageSource): HandLandmarkerResult { + return this.processImageData(image); + } + + /** + * Performs hand landmarks detection on the provided video frame and waits + * synchronously for the response. + * @param videoFrame A video frame to process. + * @param timestamp The timestamp of the current frame, in ms. + * @return The detected hand landmarks. + */ + detectForVideo(videoFrame: ImageSource, timestamp: number): + HandLandmarkerResult { + return this.processVideoData(videoFrame, timestamp); + } + + /** Runs the hand landmarker graph and blocks on the response. */ + protected override process(imageSource: ImageSource, timestamp: number): HandLandmarkerResult { this.landmarks = []; this.worldLandmarks = []; diff --git a/mediapipe/tasks/web/vision/hand_landmarker/hand_landmarker_options.d.ts b/mediapipe/tasks/web/vision/hand_landmarker/hand_landmarker_options.d.ts index 53ad9440..fe79b708 100644 --- a/mediapipe/tasks/web/vision/hand_landmarker/hand_landmarker_options.d.ts +++ b/mediapipe/tasks/web/vision/hand_landmarker/hand_landmarker_options.d.ts @@ -14,13 +14,10 @@ * limitations under the License. */ -import {BaseOptions} from '../../../../tasks/web/core/base_options'; +import {VisionTaskOptions} from '../../../../tasks/web/vision/core/vision_task_options'; /** Options to configure the MediaPipe HandLandmarker Task */ -export declare interface HandLandmarkerOptions { - /** Options to configure the loading of the model assets. */ - baseOptions?: BaseOptions; - +export declare interface HandLandmarkerOptions extends VisionTaskOptions { /** * The maximum number of hands can be detected by the HandLandmarker. * Defaults to 1. diff --git a/mediapipe/tasks/web/vision/image_classifier/BUILD b/mediapipe/tasks/web/vision/image_classifier/BUILD index 8506f357..e7e83033 100644 --- a/mediapipe/tasks/web/vision/image_classifier/BUILD +++ b/mediapipe/tasks/web/vision/image_classifier/BUILD @@ -16,15 +16,15 @@ mediapipe_ts_library( "//mediapipe/framework:calculator_jspb_proto", "//mediapipe/framework:calculator_options_jspb_proto", "//mediapipe/tasks/cc/components/containers/proto:classifications_jspb_proto", + "//mediapipe/tasks/cc/core/proto:base_options_jspb_proto", "//mediapipe/tasks/cc/vision/image_classifier/proto:image_classifier_graph_options_jspb_proto", "//mediapipe/tasks/web/components/containers:category", "//mediapipe/tasks/web/components/containers:classification_result", - "//mediapipe/tasks/web/components/processors:base_options", "//mediapipe/tasks/web/components/processors:classifier_options", "//mediapipe/tasks/web/components/processors:classifier_result", "//mediapipe/tasks/web/core", "//mediapipe/tasks/web/core:classifier_options", - "//mediapipe/tasks/web/core:task_runner", + "//mediapipe/tasks/web/vision/core:vision_task_runner", "//mediapipe/web/graph_runner:wasm_mediapipe_lib_ts", ], ) @@ -39,5 +39,6 @@ mediapipe_ts_declaration( "//mediapipe/tasks/web/components/containers:category", "//mediapipe/tasks/web/components/containers:classification_result", "//mediapipe/tasks/web/core:classifier_options", + "//mediapipe/tasks/web/vision/core:vision_task_options", ], ) diff --git a/mediapipe/tasks/web/vision/image_classifier/image_classifier.ts b/mediapipe/tasks/web/vision/image_classifier/image_classifier.ts index 5d60e4a2..0011e9c5 100644 --- a/mediapipe/tasks/web/vision/image_classifier/image_classifier.ts +++ b/mediapipe/tasks/web/vision/image_classifier/image_classifier.ts @@ -17,12 +17,12 @@ import {CalculatorGraphConfig} from '../../../../framework/calculator_pb'; import {CalculatorOptions} from '../../../../framework/calculator_options_pb'; import {ClassificationResult} from '../../../../tasks/cc/components/containers/proto/classifications_pb'; +import {BaseOptions as BaseOptionsProto} from '../../../../tasks/cc/core/proto/base_options_pb'; import {ImageClassifierGraphOptions} from '../../../../tasks/cc/vision/image_classifier/proto/image_classifier_graph_options_pb'; -import {convertBaseOptionsToProto} from '../../../../tasks/web/components/processors/base_options'; import {convertClassifierOptionsToProto} from '../../../../tasks/web/components/processors/classifier_options'; import {convertFromClassificationResultProto} from '../../../../tasks/web/components/processors/classifier_result'; -import {TaskRunner} from '../../../../tasks/web/core/task_runner'; import {WasmLoaderOptions} from '../../../../tasks/web/core/wasm_loader_options'; +import {VisionTaskRunner} from '../../../../tasks/web/vision/core/vision_task_runner'; import {createMediaPipeLib, FileLocator, ImageSource} from '../../../../web/graph_runner/wasm_mediapipe_lib'; // Placeholder for internal dependency on trusted resource url @@ -42,7 +42,7 @@ export {ImageSource}; // Used in the public API // tslint:disable:jspb-use-builder-pattern /** Performs classification on images. */ -export class ImageClassifier extends TaskRunner { +export class ImageClassifier extends VisionTaskRunner { private classificationResult: ImageClassifierResult = {classifications: []}; private readonly options = new ImageClassifierGraphOptions(); @@ -105,6 +105,14 @@ export class ImageClassifier extends TaskRunner { wasmLoaderOptions, new Uint8Array(graphData)); } + protected override get baseOptions(): BaseOptionsProto|undefined { + return this.options.getBaseOptions(); + } + + protected override set baseOptions(proto: BaseOptionsProto|undefined) { + this.options.setBaseOptions(proto); + } + /** * Sets new options for the image classifier. * @@ -114,28 +122,39 @@ export class ImageClassifier extends TaskRunner { * * @param options The options for the image classifier. */ - async setOptions(options: ImageClassifierOptions): Promise { - if (options.baseOptions) { - const baseOptionsProto = await convertBaseOptionsToProto( - options.baseOptions, this.options.getBaseOptions()); - this.options.setBaseOptions(baseOptionsProto); - } - + override async setOptions(options: ImageClassifierOptions): Promise { + await super.setOptions(options); this.options.setClassifierOptions(convertClassifierOptionsToProto( options, this.options.getClassifierOptions())); this.refreshGraph(); } /** - * Performs image classification on the provided image and waits synchronously - * for the response. + * Performs image classification on the provided single image and waits + * synchronously for the response. * - * @param imageSource An image source to process. - * @param timestamp The timestamp of the current frame, in ms. If not - * provided, defaults to `performance.now()`. + * @param image An image to process. * @return The classification result of the image */ - classify(imageSource: ImageSource, timestamp?: number): + classify(image: ImageSource): ImageClassifierResult { + return this.processImageData(image); + } + + /** + * Performs image classification on the provided video frame and waits + * synchronously for the response. + * + * @param videoFrame A video frame to process. + * @param timestamp The timestamp of the current frame, in ms. + * @return The classification result of the image + */ + classifyForVideo(videoFrame: ImageSource, timestamp: number): + ImageClassifierResult { + return this.processVideoData(videoFrame, timestamp); + } + + /** Runs the image classification graph and blocks on the response. */ + protected override process(imageSource: ImageSource, timestamp: number): ImageClassifierResult { // Get classification result by running our MediaPipe graph. this.classificationResult = {classifications: []}; diff --git a/mediapipe/tasks/web/vision/image_classifier/image_classifier_options.d.ts b/mediapipe/tasks/web/vision/image_classifier/image_classifier_options.d.ts index a5f5c238..c1141d28 100644 --- a/mediapipe/tasks/web/vision/image_classifier/image_classifier_options.d.ts +++ b/mediapipe/tasks/web/vision/image_classifier/image_classifier_options.d.ts @@ -14,4 +14,9 @@ * limitations under the License. */ -export {ClassifierOptions as ImageClassifierOptions} from '../../../../tasks/web/core/classifier_options'; +import {ClassifierOptions} from '../../../../tasks/web/core/classifier_options'; +import {VisionTaskOptions} from '../../../../tasks/web/vision/core/vision_task_options'; + +/** Ooptions to configure the image classifier task. */ +export declare interface ImageClassifierOptions extends ClassifierOptions, + VisionTaskOptions {} diff --git a/mediapipe/tasks/web/vision/image_embedder/BUILD b/mediapipe/tasks/web/vision/image_embedder/BUILD index 13ff2e4d..ce1c2570 100644 --- a/mediapipe/tasks/web/vision/image_embedder/BUILD +++ b/mediapipe/tasks/web/vision/image_embedder/BUILD @@ -16,15 +16,15 @@ mediapipe_ts_library( "//mediapipe/framework:calculator_jspb_proto", "//mediapipe/framework:calculator_options_jspb_proto", "//mediapipe/tasks/cc/components/containers/proto:embeddings_jspb_proto", + "//mediapipe/tasks/cc/core/proto:base_options_jspb_proto", "//mediapipe/tasks/cc/vision/image_embedder/proto:image_embedder_graph_options_jspb_proto", "//mediapipe/tasks/web/components/containers:embedding_result", - "//mediapipe/tasks/web/components/processors:base_options", "//mediapipe/tasks/web/components/processors:embedder_options", "//mediapipe/tasks/web/components/processors:embedder_result", "//mediapipe/tasks/web/core", "//mediapipe/tasks/web/core:embedder_options", - "//mediapipe/tasks/web/core:task_runner", - "//mediapipe/tasks/web/vision/core:running_mode", + "//mediapipe/tasks/web/vision/core:vision_task_options", + "//mediapipe/tasks/web/vision/core:vision_task_runner", "//mediapipe/web/graph_runner:wasm_mediapipe_lib_ts", ], ) @@ -39,6 +39,6 @@ mediapipe_ts_declaration( "//mediapipe/tasks/web/components/containers:embedding_result", "//mediapipe/tasks/web/core", "//mediapipe/tasks/web/core:embedder_options", - "//mediapipe/tasks/web/vision/core:running_mode", + "//mediapipe/tasks/web/vision/core:vision_task_options", ], ) diff --git a/mediapipe/tasks/web/vision/image_embedder/image_embedder.ts b/mediapipe/tasks/web/vision/image_embedder/image_embedder.ts index 91d9b511..d17bc72f 100644 --- a/mediapipe/tasks/web/vision/image_embedder/image_embedder.ts +++ b/mediapipe/tasks/web/vision/image_embedder/image_embedder.ts @@ -17,13 +17,12 @@ import {CalculatorGraphConfig} from '../../../../framework/calculator_pb'; import {CalculatorOptions} from '../../../../framework/calculator_options_pb'; import {EmbeddingResult} from '../../../../tasks/cc/components/containers/proto/embeddings_pb'; +import {BaseOptions as BaseOptionsProto} from '../../../../tasks/cc/core/proto/base_options_pb'; import {ImageEmbedderGraphOptions} from '../../../../tasks/cc/vision/image_embedder/proto/image_embedder_graph_options_pb'; -import {convertBaseOptionsToProto} from '../../../../tasks/web/components/processors/base_options'; import {convertEmbedderOptionsToProto} from '../../../../tasks/web/components/processors/embedder_options'; import {convertFromEmbeddingResultProto} from '../../../../tasks/web/components/processors/embedder_result'; -import {TaskRunner} from '../../../../tasks/web/core/task_runner'; import {WasmLoaderOptions} from '../../../../tasks/web/core/wasm_loader_options'; -import {configureRunningMode} from '../../../../tasks/web/vision/core/running_mode'; +import {VisionTaskRunner} from '../../../../tasks/web/vision/core/vision_task_runner'; import {createMediaPipeLib, FileLocator, ImageSource} from '../../../../web/graph_runner/wasm_mediapipe_lib'; // Placeholder for internal dependency on trusted resource url @@ -43,7 +42,7 @@ export * from './image_embedder_result'; export {ImageSource}; // Used in the public API /** Performs embedding extraction on images. */ -export class ImageEmbedder extends TaskRunner { +export class ImageEmbedder extends VisionTaskRunner { private readonly options = new ImageEmbedderGraphOptions(); private embeddings: ImageEmbedderResult = {embeddings: []}; @@ -105,6 +104,14 @@ export class ImageEmbedder extends TaskRunner { wasmLoaderOptions, new Uint8Array(graphData)); } + protected override get baseOptions(): BaseOptionsProto|undefined { + return this.options.getBaseOptions(); + } + + protected override set baseOptions(proto: BaseOptionsProto|undefined) { + this.options.setBaseOptions(proto); + } + /** * Sets new options for the image embedder. * @@ -114,24 +121,16 @@ export class ImageEmbedder extends TaskRunner { * * @param options The options for the image embedder. */ - async setOptions(options: ImageEmbedderOptions): Promise { - let baseOptionsProto = this.options.getBaseOptions(); - if (options.baseOptions) { - baseOptionsProto = await convertBaseOptionsToProto( - options.baseOptions, baseOptionsProto); - } - baseOptionsProto = configureRunningMode(options, baseOptionsProto); - this.options.setBaseOptions(baseOptionsProto); - + override async setOptions(options: ImageEmbedderOptions): Promise { + await super.setOptions(options); this.options.setEmbedderOptions(convertEmbedderOptionsToProto( options, this.options.getEmbedderOptions())); - this.refreshGraph(); } /** - * Performs embedding extraction on the provided image and waits synchronously - * for the response. + * Performs embedding extraction on the provided single image and waits + * synchronously for the response. * * Only use this method when the `useStreamMode` option is not set or * expliclity set to `false`. @@ -140,12 +139,7 @@ export class ImageEmbedder extends TaskRunner { * @return The classification result of the image */ embed(image: ImageSource): ImageEmbedderResult { - if (!!this.options.getBaseOptions()?.getUseStreamMode()) { - throw new Error( - 'Task is not initialized with image mode. ' + - '\'runningMode\' must be set to \'image\'.'); - } - return this.performEmbeddingExtraction(image, performance.now()); + return this.processImageData(image); } /** @@ -160,16 +154,11 @@ export class ImageEmbedder extends TaskRunner { */ embedForVideo(imageFrame: ImageSource, timestamp: number): ImageEmbedderResult { - if (!this.options.getBaseOptions()?.getUseStreamMode()) { - throw new Error( - 'Task is not initialized with video mode. ' + - '\'runningMode\' must be set to \'video\' or \'live_stream\'.'); - } - return this.performEmbeddingExtraction(imageFrame, timestamp); + return this.processVideoData(imageFrame, timestamp); } - /** Runs the embedding extractio and blocks on the response. */ - private performEmbeddingExtraction(image: ImageSource, timestamp: number): + /** Runs the embedding extraction and blocks on the response. */ + protected process(image: ImageSource, timestamp: number): ImageEmbedderResult { // Get embeddings by running our MediaPipe graph. this.addGpuBufferAsImageToStream( diff --git a/mediapipe/tasks/web/vision/image_embedder/image_embedder_options.d.ts b/mediapipe/tasks/web/vision/image_embedder/image_embedder_options.d.ts index 4d795d0d..10000825 100644 --- a/mediapipe/tasks/web/vision/image_embedder/image_embedder_options.d.ts +++ b/mediapipe/tasks/web/vision/image_embedder/image_embedder_options.d.ts @@ -15,17 +15,8 @@ */ import {EmbedderOptions} from '../../../../tasks/web/core/embedder_options'; -import {RunningMode} from '../../../../tasks/web/vision/core/running_mode'; +import {VisionTaskOptions} from '../../../../tasks/web/vision/core/vision_task_options'; /** The options for configuring a MediaPipe image embedder task. */ -export declare interface ImageEmbedderOptions extends EmbedderOptions { - /** - * The running mode of the task. Default to the image mode. - * Image embedder has three running modes: - * 1) The image mode for embedding image on single image inputs. - * 2) The video mode for embedding image on the decoded frames of a video. - * 3) The live stream mode for embedding image on the live stream of input - * data, such as from camera. - */ - runningMode?: RunningMode; -} +export declare interface ImageEmbedderOptions extends EmbedderOptions, + VisionTaskOptions {} diff --git a/mediapipe/tasks/web/vision/object_detector/BUILD b/mediapipe/tasks/web/vision/object_detector/BUILD index a74dc921..0975a9fd 100644 --- a/mediapipe/tasks/web/vision/object_detector/BUILD +++ b/mediapipe/tasks/web/vision/object_detector/BUILD @@ -17,11 +17,11 @@ mediapipe_ts_library( "//mediapipe/framework:calculator_jspb_proto", "//mediapipe/framework:calculator_options_jspb_proto", "//mediapipe/framework/formats:detection_jspb_proto", + "//mediapipe/tasks/cc/core/proto:base_options_jspb_proto", "//mediapipe/tasks/cc/vision/object_detector/proto:object_detector_options_jspb_proto", "//mediapipe/tasks/web/components/containers:category", - "//mediapipe/tasks/web/components/processors:base_options", "//mediapipe/tasks/web/core", - "//mediapipe/tasks/web/core:task_runner", + "//mediapipe/tasks/web/vision/core:vision_task_runner", "//mediapipe/web/graph_runner:wasm_mediapipe_lib_ts", ], ) @@ -35,5 +35,6 @@ mediapipe_ts_declaration( deps = [ "//mediapipe/tasks/web/components/containers:category", "//mediapipe/tasks/web/core", + "//mediapipe/tasks/web/vision/core:vision_task_options", ], ) diff --git a/mediapipe/tasks/web/vision/object_detector/object_detector.ts b/mediapipe/tasks/web/vision/object_detector/object_detector.ts index e17a4202..e6cbd862 100644 --- a/mediapipe/tasks/web/vision/object_detector/object_detector.ts +++ b/mediapipe/tasks/web/vision/object_detector/object_detector.ts @@ -17,10 +17,10 @@ import {CalculatorGraphConfig} from '../../../../framework/calculator_pb'; import {CalculatorOptions} from '../../../../framework/calculator_options_pb'; import {Detection as DetectionProto} from '../../../../framework/formats/detection_pb'; +import {BaseOptions as BaseOptionsProto} from '../../../../tasks/cc/core/proto/base_options_pb'; import {ObjectDetectorOptions as ObjectDetectorOptionsProto} from '../../../../tasks/cc/vision/object_detector/proto/object_detector_options_pb'; -import {convertBaseOptionsToProto} from '../../../../tasks/web/components/processors/base_options'; -import {TaskRunner} from '../../../../tasks/web/core/task_runner'; import {WasmLoaderOptions} from '../../../../tasks/web/core/wasm_loader_options'; +import {VisionTaskRunner} from '../../../../tasks/web/vision/core/vision_task_runner'; import {createMediaPipeLib, FileLocator, ImageSource} from '../../../../web/graph_runner/wasm_mediapipe_lib'; // Placeholder for internal dependency on trusted resource url @@ -41,7 +41,7 @@ export {ImageSource}; // Used in the public API // tslint:disable:jspb-use-builder-pattern /** Performs object detection on images. */ -export class ObjectDetector extends TaskRunner { +export class ObjectDetector extends VisionTaskRunner { private detections: Detection[] = []; private readonly options = new ObjectDetectorOptionsProto(); @@ -103,6 +103,14 @@ export class ObjectDetector extends TaskRunner { wasmLoaderOptions, new Uint8Array(graphData)); } + protected override get baseOptions(): BaseOptionsProto|undefined { + return this.options.getBaseOptions(); + } + + protected override set baseOptions(proto: BaseOptionsProto|undefined) { + this.options.setBaseOptions(proto); + } + /** * Sets new options for the object detector. * @@ -112,12 +120,8 @@ export class ObjectDetector extends TaskRunner { * * @param options The options for the object detector. */ - async setOptions(options: ObjectDetectorOptions): Promise { - if (options.baseOptions) { - const baseOptionsProto = await convertBaseOptionsToProto( - options.baseOptions, this.options.getBaseOptions()); - this.options.setBaseOptions(baseOptionsProto); - } + override async setOptions(options: ObjectDetectorOptions): Promise { + await super.setOptions(options); // Note that we have to support both JSPB and ProtobufJS, hence we // have to expliclity clear the values instead of setting them to @@ -158,12 +162,27 @@ export class ObjectDetector extends TaskRunner { /** * Performs object detection on the provided single image and waits * synchronously for the response. - * @param imageSource An image source to process. - * @param timestamp The timestamp of the current frame, in ms. If not - * provided, defaults to `performance.now()`. + * @param image An image to process. * @return The list of detected objects */ - detect(imageSource: ImageSource, timestamp?: number): Detection[] { + detect(image: ImageSource): Detection[] { + return this.processImageData(image); + } + + /** + * Performs object detection on the provided vidoe frame and waits + * synchronously for the response. + * @param videoFrame A video frame to process. + * @param timestamp The timestamp of the current frame, in ms. + * @return The list of detected objects + */ + detectForVideo(videoFrame: ImageSource, timestamp: number): Detection[] { + return this.processVideoData(videoFrame, timestamp); + } + + /** Runs the object detector graph and blocks on the response. */ + protected override process(imageSource: ImageSource, timestamp: number): + Detection[] { // Get detections by running our MediaPipe graph. this.detections = []; this.addGpuBufferAsImageToStream( diff --git a/mediapipe/tasks/web/vision/object_detector/object_detector_options.d.ts b/mediapipe/tasks/web/vision/object_detector/object_detector_options.d.ts index eec12cf1..1d20ce1e 100644 --- a/mediapipe/tasks/web/vision/object_detector/object_detector_options.d.ts +++ b/mediapipe/tasks/web/vision/object_detector/object_detector_options.d.ts @@ -14,13 +14,10 @@ * limitations under the License. */ -import {BaseOptions} from '../../../../tasks/web/core/base_options'; +import {VisionTaskOptions} from '../../../../tasks/web/vision/core/vision_task_options'; /** Options to configure the MediaPipe Object Detector Task */ -export interface ObjectDetectorOptions { - /** Options to configure the loading of the model assets. */ - baseOptions?: BaseOptions; - +export interface ObjectDetectorOptions extends VisionTaskOptions { /** * The locale to use for display names specified through the TFLite Model * Metadata, if any. Defaults to English.