Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
c688862570 | ||
|
|
4a20e9909d | ||
|
|
7fb37c80e8 | ||
|
|
c6c80c3745 | ||
|
|
cc6a2f7af6 |
@@ -32,6 +32,9 @@ build:macos --copt=-w
|
||||
# Sets the default Apple platform to macOS.
|
||||
build --apple_platform_type=macos
|
||||
|
||||
# Compile ObjC++ files with C++17
|
||||
build --per_file_copt=.*\.mm\$@-std=c++17
|
||||
|
||||
# Allow debugging with XCODE
|
||||
build --apple_generate_dsym
|
||||
|
||||
@@ -88,6 +91,10 @@ build:darwin_x86_64 --apple_platform_type=macos
|
||||
build:darwin_x86_64 --macos_minimum_os=10.12
|
||||
build:darwin_x86_64 --cpu=darwin_x86_64
|
||||
|
||||
build:darwin_arm64 --apple_platform_type=macos
|
||||
build:darwin_arm64 --macos_minimum_os=10.16
|
||||
build:darwin_arm64 --cpu=darwin_arm64
|
||||
|
||||
# This bazelrc file is meant to be written by a setup script.
|
||||
try-import %workspace%/.configure.bazelrc
|
||||
|
||||
|
||||
+1
-1
@@ -1 +1 @@
|
||||
4.2.1
|
||||
5.2.0
|
||||
|
||||
@@ -10,5 +10,3 @@ For questions on how to work with MediaPipe, or support for problems that are no
|
||||
|
||||
If you are reporting a vulnerability, please use the [dedicated reporting process](https://github.com/google/mediapipe/security).
|
||||
|
||||
For high-level discussions about MediaPipe, please post to discuss@mediapipe.org, for questions about the development or internal workings of MediaPipe, or if you would like to know how to contribute to MediaPipe, please post to developers@mediapipe.org.
|
||||
|
||||
|
||||
@@ -15,4 +15,4 @@
|
||||
|
||||
# A list of assignees
|
||||
assignees:
|
||||
- sgowroji
|
||||
- sureshdagooglecom
|
||||
|
||||
+5
-3
@@ -12,7 +12,7 @@
|
||||
# See the License for the specific language governing permissions and
|
||||
# limitations under the License.
|
||||
|
||||
FROM ubuntu:18.04
|
||||
FROM ubuntu:20.04
|
||||
|
||||
MAINTAINER <[email protected]>
|
||||
|
||||
@@ -42,6 +42,8 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
|
||||
software-properties-common && \
|
||||
add-apt-repository -y ppa:openjdk-r/ppa && \
|
||||
apt-get update && apt-get install -y openjdk-8-jdk && \
|
||||
apt-get install -y mesa-common-dev libegl1-mesa-dev libgles2-mesa-dev && \
|
||||
apt-get install -y mesa-utils && \
|
||||
apt-get clean && \
|
||||
rm -rf /var/lib/apt/lists/*
|
||||
|
||||
@@ -50,13 +52,13 @@ RUN pip3 install --upgrade setuptools
|
||||
RUN pip3 install wheel
|
||||
RUN pip3 install future
|
||||
RUN pip3 install six==1.14.0
|
||||
RUN pip3 install tensorflow==1.14.0
|
||||
RUN pip3 install tensorflow==2.2.0
|
||||
RUN pip3 install tf_slim
|
||||
|
||||
RUN ln -s /usr/bin/python3 /usr/bin/python
|
||||
|
||||
# Install bazel
|
||||
ARG BAZEL_VERSION=4.2.1
|
||||
ARG BAZEL_VERSION=5.2.0
|
||||
RUN mkdir /bazel && \
|
||||
wget --no-check-certificate -O /bazel/installer.sh "https://github.com/bazelbuild/bazel/releases/download/${BAZEL_VERSION}/b\
|
||||
azel-${BAZEL_VERSION}-installer-linux-x86_64.sh" && \
|
||||
|
||||
@@ -136,8 +136,8 @@ run code search using
|
||||
|
||||
## Community
|
||||
|
||||
* [Awesome MediaPipe](https://mediapipe.org) - A curated list of awesome
|
||||
MediaPipe related frameworks, libraries and software
|
||||
* [Awesome MediaPipe](https://mediapipe.page.link/awesome-mediapipe) - A
|
||||
curated list of awesome MediaPipe related frameworks, libraries and software
|
||||
* [Slack community](https://mediapipe.page.link/joinslack) for MediaPipe users
|
||||
* [Discuss](https://groups.google.com/forum/#!forum/mediapipe) - General
|
||||
community discussion around MediaPipe
|
||||
|
||||
@@ -35,8 +35,9 @@ http_archive(
|
||||
|
||||
http_archive(
|
||||
name = "rules_cc",
|
||||
strip_prefix = "rules_cc-main",
|
||||
urls = ["https://github.com/bazelbuild/rules_cc/archive/main.zip"],
|
||||
strip_prefix = "rules_cc-2f8c04c04462ab83c545ab14c0da68c3b4c96191",
|
||||
# The commit can be updated if the build passes. Last updated 6/23/22.
|
||||
urls = ["https://github.com/bazelbuild/rules_cc/archive/2f8c04c04462ab83c545ab14c0da68c3b4c96191.zip"],
|
||||
)
|
||||
|
||||
http_archive(
|
||||
@@ -61,11 +62,12 @@ http_archive(
|
||||
sha256 = "de682ea824bfffba05b4e33b67431c247397d6175962534305136aa06f92e049",
|
||||
)
|
||||
|
||||
# Google Benchmark library.
|
||||
# Google Benchmark library v1.6.1 released on 2022-01-10.
|
||||
http_archive(
|
||||
name = "com_google_benchmark",
|
||||
urls = ["https://github.com/google/benchmark/archive/main.zip"],
|
||||
strip_prefix = "benchmark-main",
|
||||
urls = ["https://github.com/google/benchmark/archive/refs/tags/v1.6.1.tar.gz"],
|
||||
strip_prefix = "benchmark-1.6.1",
|
||||
sha256 = "6132883bc8c9b0df5375b16ab520fac1a85dc9e4cf5be59480448ece74b278d4",
|
||||
build_file = "@//third_party:benchmark.BUILD",
|
||||
)
|
||||
|
||||
@@ -201,7 +203,10 @@ new_local_repository(
|
||||
new_local_repository(
|
||||
name = "macos_opencv",
|
||||
build_file = "@//third_party:opencv_macos.BUILD",
|
||||
path = "/usr/local/opt/opencv@3",
|
||||
# For local MacOS builds, the path should point to an opencv@3 installation.
|
||||
# If you edit the path here, you will also need to update the corresponding
|
||||
# prefix in "opencv_macos.BUILD".
|
||||
path = "/usr/local",
|
||||
)
|
||||
|
||||
new_local_repository(
|
||||
@@ -373,9 +378,9 @@ http_archive(
|
||||
)
|
||||
|
||||
# Tensorflow repo should always go after the other external dependencies.
|
||||
# 2021-12-02
|
||||
_TENSORFLOW_GIT_COMMIT = "18a1dc0ba806dc023808531f0373d9ec068e64bf"
|
||||
_TENSORFLOW_SHA256 = "85b90416f7a11339327777bccd634de00ca0de2cf334f5f0727edcb11ff9289a"
|
||||
# 2022-02-15
|
||||
_TENSORFLOW_GIT_COMMIT = "a3419acc751dfc19caf4d34a1594e1f76810ec58"
|
||||
_TENSORFLOW_SHA256 = "b95b2a83632d4055742ae1a2dcc96b45da6c12a339462dbc76c8bca505308e3a"
|
||||
http_archive(
|
||||
name = "org_tensorflow",
|
||||
urls = [
|
||||
@@ -383,7 +388,6 @@ http_archive(
|
||||
],
|
||||
patches = [
|
||||
"@//third_party:org_tensorflow_compatibility_fixes.diff",
|
||||
"@//third_party:org_tensorflow_objc_cxx17.diff",
|
||||
# Diff is generated with a script, don't update it manually.
|
||||
"@//third_party:org_tensorflow_custom_ops.diff",
|
||||
],
|
||||
|
||||
@@ -109,7 +109,7 @@ for app in ${apps}; do
|
||||
if [[ ${category} != "shoe" ]]; then
|
||||
bazel_flags_extended+=(--define ${category}=true)
|
||||
fi
|
||||
bazel "${bazel_flags_extended[@]}"
|
||||
bazelisk "${bazel_flags_extended[@]}"
|
||||
cp -f "${bin}" "${apk}"
|
||||
fi
|
||||
apks+=(${apk})
|
||||
@@ -120,7 +120,7 @@ for app in ${apps}; do
|
||||
if [[ ${app_name} == "templatematchingcpu" ]]; then
|
||||
switch_to_opencv_4
|
||||
fi
|
||||
bazel "${bazel_flags[@]}"
|
||||
bazelisk "${bazel_flags[@]}"
|
||||
cp -f "${bin}" "${apk}"
|
||||
if [[ ${app_name} == "templatematchingcpu" ]]; then
|
||||
switch_to_opencv_3
|
||||
|
||||
@@ -83,7 +83,7 @@ for app in ${apps}; do
|
||||
bazel_flags=("${default_bazel_flags[@]}")
|
||||
bazel_flags+=(${target})
|
||||
|
||||
bazel "${bazel_flags[@]}"
|
||||
bazelisk "${bazel_flags[@]}"
|
||||
cp -f "${bin_dir}/${app}/"*"_cpu" "${out_dir}"
|
||||
fi
|
||||
if [[ $build_only == false ]]; then
|
||||
|
||||
@@ -71,7 +71,7 @@ for app in ${apps}; do
|
||||
bazel_flags+=(--linkopt=-s)
|
||||
fi
|
||||
|
||||
bazel "${bazel_flags[@]}"
|
||||
bazelisk "${bazel_flags[@]}"
|
||||
cp -f "${bin_dir}/${app}/"*".ipa" "${out_dir}"
|
||||
fi
|
||||
done
|
||||
|
||||
@@ -169,7 +169,7 @@ behavior depending on resource constraints.
|
||||
|
||||
[`CalculatorBase`]: https://github.com/google/mediapipe/tree/master/mediapipe/framework/calculator_base.h
|
||||
[`DefaultInputStreamHandler`]: https://github.com/google/mediapipe/tree/master/mediapipe/framework/stream_handler/default_input_stream_handler.h
|
||||
[`SyncSetInputStreamHandler`]: https://github.com/google/mediapipe/tree/master/mediapipe/framework/stream_handler/sync_set_input_stream_handler.h
|
||||
[`ImmediateInputStreamHandler`]: https://github.com/google/mediapipe/tree/master/mediapipe/framework/stream_handler/immediate_input_stream_handler.h
|
||||
[`SyncSetInputStreamHandler`]: https://github.com/google/mediapipe/tree/master/mediapipe/framework/stream_handler/sync_set_input_stream_handler.cc
|
||||
[`ImmediateInputStreamHandler`]: https://github.com/google/mediapipe/tree/master/mediapipe/framework/stream_handler/immediate_input_stream_handler.cc
|
||||
[`CalculatorGraphConfig::max_queue_size`]: https://github.com/google/mediapipe/tree/master/mediapipe/framework/calculator.proto
|
||||
[`FlowLimiterCalculator`]: https://github.com/google/mediapipe/tree/master/mediapipe/calculators/core/flow_limiter_calculator.cc
|
||||
|
||||
@@ -30,7 +30,7 @@ APIs (currently in alpha) that are now available in
|
||||
* Install MediaPipe following these [instructions](./install.md).
|
||||
* Setup Java Runtime.
|
||||
* Setup Android SDK release 30.0.0 and above.
|
||||
* Setup Android NDK version 18 and above.
|
||||
* Setup Android NDK version between 18 and 21.
|
||||
|
||||
MediaPipe recommends setting up Android SDK and NDK via Android Studio (and see
|
||||
below for Android Studio setup). However, if you prefer using MediaPipe without
|
||||
@@ -53,7 +53,7 @@ the following:
|
||||
|
||||
```bash
|
||||
$ echo "android_sdk_repository(name = \"androidsdk\")" >> WORKSPACE
|
||||
$ echo "android_ndk_repository(name = \"androidndk\")" >> WORKSPACE
|
||||
$ echo "android_ndk_repository(name = \"androidndk\", api_level=21)" >> WORKSPACE
|
||||
```
|
||||
|
||||
In order to use MediaPipe on earlier Android versions, MediaPipe needs to switch
|
||||
|
||||
@@ -48,6 +48,16 @@ each project.
|
||||
bazel build -c opt --strip=ALWAYS \
|
||||
--host_crosstool_top=@bazel_tools//tools/cpp:toolchain \
|
||||
--fat_apk_cpu=arm64-v8a,armeabi-v7a \
|
||||
--legacy_whole_archive=0 \
|
||||
--features=-legacy_whole_archive \
|
||||
--copt=-fvisibility=hidden \
|
||||
--copt=-ffunction-sections \
|
||||
--copt=-fdata-sections \
|
||||
--copt=-fstack-protector \
|
||||
--copt=-Oz \
|
||||
--copt=-fomit-frame-pointer \
|
||||
--copt=-DABSL_MIN_LOG_LEVEL=2 \
|
||||
--linkopt=-Wl,--gc-sections,--strip-all \
|
||||
//path/to/the/aar/build/file:aar_name.aar
|
||||
```
|
||||
|
||||
@@ -57,6 +67,16 @@ each project.
|
||||
bazel build -c opt --strip=ALWAYS \
|
||||
--host_crosstool_top=@bazel_tools//tools/cpp:toolchain \
|
||||
--fat_apk_cpu=arm64-v8a,armeabi-v7a \
|
||||
--legacy_whole_archive=0 \
|
||||
--features=-legacy_whole_archive \
|
||||
--copt=-fvisibility=hidden \
|
||||
--copt=-ffunction-sections \
|
||||
--copt=-fdata-sections \
|
||||
--copt=-fstack-protector \
|
||||
--copt=-Oz \
|
||||
--copt=-fomit-frame-pointer \
|
||||
--copt=-DABSL_MIN_LOG_LEVEL=2 \
|
||||
--linkopt=-Wl,--gc-sections,--strip-all \
|
||||
//mediapipe/examples/android/src/java/com/google/mediapipe/apps/aar_example:mediapipe_face_detection.aar
|
||||
|
||||
# It should print:
|
||||
|
||||
@@ -59,6 +59,21 @@ OpenGL ES profile shading language version string: OpenGL ES GLSL ES 3.20
|
||||
OpenGL ES profile extensions:
|
||||
```
|
||||
|
||||
If you have connected to your computer through SSH and find when you probe for
|
||||
GPU information you see the output:
|
||||
|
||||
```bash
|
||||
glxinfo | grep -i opengl
|
||||
Error: unable to open display
|
||||
```
|
||||
|
||||
Try re-establishing your SSH connection with the `-X` option and try again. For
|
||||
example:
|
||||
|
||||
```bash
|
||||
ssh -X <user>@<host>
|
||||
```
|
||||
|
||||
*Notice the ES 3.20 text above.*
|
||||
|
||||
You need to see ES 3.1 or greater printed in order to perform TFLite inference
|
||||
|
||||
@@ -131,7 +131,7 @@ Create a `BUILD` file in the `$APPLICATION_PATH` and add the following build
|
||||
rules:
|
||||
|
||||
```
|
||||
MIN_IOS_VERSION = "10.0"
|
||||
MIN_IOS_VERSION = "11.0"
|
||||
|
||||
load(
|
||||
"@build_bazel_rules_apple//apple:ios.bzl",
|
||||
|
||||
@@ -569,7 +569,7 @@ next section.
|
||||
|
||||
Option 1. Follow
|
||||
[the official Bazel documentation](https://docs.bazel.build/versions/master/install-windows.html)
|
||||
to install Bazel 4.2.1 or higher.
|
||||
to install Bazel 5.0.0 or higher.
|
||||
|
||||
Option 2. Follow the official
|
||||
[Bazel documentation](https://docs.bazel.build/versions/master/install-bazelisk.html)
|
||||
|
||||
@@ -32,9 +32,14 @@ example apps, start from, start from
|
||||
xcode-select --install
|
||||
```
|
||||
|
||||
3. Install [Bazel](https://bazel.build/).
|
||||
3. Install [Bazelisk](https://github.com/bazelbuild/bazelisk)
|
||||
.
|
||||
|
||||
We recommend using [Homebrew](https://brew.sh/) to get the latest version.
|
||||
We recommend using [Homebrew](https://brew.sh/) to get the latest versions.
|
||||
|
||||
```bash
|
||||
brew install bazelisk
|
||||
```
|
||||
|
||||
4. Set Python 3.7 as the default Python version and install the Python "six"
|
||||
library. This is needed for TensorFlow.
|
||||
@@ -187,6 +192,9 @@ Note: When you ask Xcode to run an app, by default it will use the Debug
|
||||
configuration. Some of our demos are computationally heavy; you may want to use
|
||||
the Release configuration for better performance.
|
||||
|
||||
Note: Due to an imcoptibility caused by one of our dependencies, MediaPipe
|
||||
cannot be used for apps running on the iPhone Simulator on Apple Silicon (M1).
|
||||
|
||||
Tip: To switch build configuration in Xcode, click on the target menu, choose
|
||||
"Edit Scheme...", select the Run action, and switch the Build Configuration from
|
||||
Debug to Release. Note that this is set independently for each target.
|
||||
|
||||
@@ -126,6 +126,7 @@ following steps:
|
||||
}
|
||||
return packet.Get<MyType>();
|
||||
});
|
||||
}
|
||||
} // namespace mediapipe
|
||||
```
|
||||
|
||||
|
||||
+2
-2
@@ -136,8 +136,8 @@ run code search using
|
||||
|
||||
## Community
|
||||
|
||||
* [Awesome MediaPipe](https://mediapipe.org) - A curated list of awesome
|
||||
MediaPipe related frameworks, libraries and software
|
||||
* [Awesome MediaPipe](https://mediapipe.page.link/awesome-mediapipe) - A
|
||||
curated list of awesome MediaPipe related frameworks, libraries and software
|
||||
* [Slack community](https://mediapipe.page.link/joinslack) for MediaPipe users
|
||||
* [Discuss](https://groups.google.com/forum/#!forum/mediapipe) - General
|
||||
community discussion around MediaPipe
|
||||
|
||||
@@ -26,7 +26,7 @@ MediaPipe Face Detection is an ultrafast face detection solution that comes with
|
||||
face detector tailored for mobile GPU inference. The detector's super-realtime
|
||||
performance enables it to be applied to any live viewfinder experience that
|
||||
requires an accurate facial region of interest as an input for other
|
||||
task-specific models, such as 3D facial keypoint or geometry estimation (e.g.,
|
||||
task-specific models, such as 3D facial keypoint estimation (e.g.,
|
||||
[MediaPipe Face Mesh](./face_mesh.md)), facial features or expression
|
||||
classification, and face region segmentation. BlazeFace uses a lightweight
|
||||
feature extraction network inspired by, but distinct from
|
||||
|
||||
+29
-29
@@ -20,34 +20,34 @@ nav_order: 2
|
||||
|
||||
## Overview
|
||||
|
||||
MediaPipe Face Mesh is a face geometry solution that estimates 468 3D face
|
||||
landmarks in real-time even on mobile devices. It employs machine learning (ML)
|
||||
to infer the 3D surface geometry, requiring only a single camera input without
|
||||
the need for a dedicated depth sensor. Utilizing lightweight model architectures
|
||||
together with GPU acceleration throughout the pipeline, the solution delivers
|
||||
real-time performance critical for live experiences.
|
||||
MediaPipe Face Mesh is a solution that estimates 468 3D face landmarks in
|
||||
real-time even on mobile devices. It employs machine learning (ML) to infer the
|
||||
3D facial surface, requiring only a single camera input without the need for a
|
||||
dedicated depth sensor. Utilizing lightweight model architectures together with
|
||||
GPU acceleration throughout the pipeline, the solution delivers real-time
|
||||
performance critical for live experiences.
|
||||
|
||||
Additionally, the solution is bundled with the Face Geometry module that bridges
|
||||
the gap between the face landmark estimation and useful real-time augmented
|
||||
reality (AR) applications. It establishes a metric 3D space and uses the face
|
||||
landmark screen positions to estimate face geometry within that space. The face
|
||||
geometry data consists of common 3D geometry primitives, including a face pose
|
||||
transformation matrix and a triangular face mesh. Under the hood, a lightweight
|
||||
statistical analysis method called
|
||||
Additionally, the solution is bundled with the Face Transform module that
|
||||
bridges the gap between the face landmark estimation and useful real-time
|
||||
augmented reality (AR) applications. It establishes a metric 3D space and uses
|
||||
the face landmark screen positions to estimate a face transform within that
|
||||
space. The face transform data consists of common 3D primitives, including a
|
||||
face pose transformation matrix and a triangular face mesh. Under the hood, a
|
||||
lightweight statistical analysis method called
|
||||
[Procrustes Analysis](https://en.wikipedia.org/wiki/Procrustes_analysis) is
|
||||
employed to drive a robust, performant and portable logic. The analysis runs on
|
||||
CPU and has a minimal speed/memory footprint on top of the ML model inference.
|
||||
|
||||
 |
|
||||
:-------------------------------------------------------------: |
|
||||
*Fig 1. AR effects utilizing facial surface geometry.* |
|
||||
*Fig 1. AR effects utilizing the 3D facial surface.* |
|
||||
|
||||
## ML Pipeline
|
||||
|
||||
Our ML pipeline consists of two real-time deep neural network models that work
|
||||
together: A detector that operates on the full image and computes face locations
|
||||
and a 3D face landmark model that operates on those locations and predicts the
|
||||
approximate surface geometry via regression. Having the face accurately cropped
|
||||
approximate 3D surface via regression. Having the face accurately cropped
|
||||
drastically reduces the need for common data augmentations like affine
|
||||
transformations consisting of rotations, translation and scale changes. Instead
|
||||
it allows the network to dedicate most of its capacity towards coordinate
|
||||
@@ -55,8 +55,8 @@ prediction accuracy. In addition, in our pipeline the crops can also be
|
||||
generated based on the face landmarks identified in the previous frame, and only
|
||||
when the landmark model could no longer identify face presence is the face
|
||||
detector invoked to relocalize the face. This strategy is similar to that
|
||||
employed in our [MediaPipe Hands](./hands.md) solution, which uses a palm detector
|
||||
together with a hand landmark model.
|
||||
employed in our [MediaPipe Hands](./hands.md) solution, which uses a palm
|
||||
detector together with a hand landmark model.
|
||||
|
||||
The pipeline is implemented as a MediaPipe
|
||||
[graph](https://github.com/google/mediapipe/tree/master/mediapipe/graphs/face_mesh/face_mesh_mobile.pbtxt)
|
||||
@@ -128,7 +128,7 @@ about the model in this [paper](https://arxiv.org/abs/2006.10962).
|
||||
:---------------------------------------------------------------------------: |
|
||||
*Fig 3. Attention Mesh: Overview of model architecture.* |
|
||||
|
||||
## Face Geometry Module
|
||||
## Face Transform Module
|
||||
|
||||
The [Face Landmark Model](#face-landmark-model) performs a single-camera face landmark
|
||||
detection in the screen coordinate space: the X- and Y- coordinates are
|
||||
@@ -140,7 +140,7 @@ enable the full spectrum of augmented reality (AR) features like aligning a
|
||||
virtual 3D object with a detected face.
|
||||
|
||||
The
|
||||
[Face Geometry module](https://github.com/google/mediapipe/tree/master/mediapipe/modules/face_geometry)
|
||||
[Face Transform module](https://github.com/google/mediapipe/tree/master/mediapipe/modules/face_geometry)
|
||||
moves away from the screen coordinate space towards a metric 3D space and
|
||||
provides necessary primitives to handle a detected face as a regular 3D object.
|
||||
By design, you'll be able to use a perspective camera to project the final 3D
|
||||
@@ -151,7 +151,7 @@ landmark positions are not changed.
|
||||
|
||||
#### Metric 3D Space
|
||||
|
||||
The **Metric 3D space** established within the Face Geometry module is a
|
||||
The **Metric 3D space** established within the Face Transform module is a
|
||||
right-handed orthonormal metric 3D coordinate space. Within the space, there is
|
||||
a **virtual perspective camera** located at the space origin and pointed in the
|
||||
negative direction of the Z-axis. In the current pipeline, it is assumed that
|
||||
@@ -184,11 +184,11 @@ functions:
|
||||
|
||||
### Components
|
||||
|
||||
#### Geometry Pipeline
|
||||
#### Transform Pipeline
|
||||
|
||||
The **Geometry Pipeline** is a key component, which is responsible for
|
||||
estimating face geometry objects within the Metric 3D space. On each frame, the
|
||||
following steps are executed in the given order:
|
||||
The **Transform Pipeline** is a key component, which is responsible for
|
||||
estimating the face transform objects within the Metric 3D space. On each frame,
|
||||
the following steps are executed in the given order:
|
||||
|
||||
- Face landmark screen coordinates are converted into the Metric 3D space
|
||||
coordinates;
|
||||
@@ -199,12 +199,12 @@ following steps are executed in the given order:
|
||||
positions (XYZ), while both the vertex texture coordinates (UV) and the
|
||||
triangular topology are inherited from the canonical face model.
|
||||
|
||||
The geometry pipeline is implemented as a MediaPipe
|
||||
The transform pipeline is implemented as a MediaPipe
|
||||
[calculator](https://github.com/google/mediapipe/tree/master/mediapipe/modules/face_geometry/geometry_pipeline_calculator.cc).
|
||||
For your convenience, the face geometry pipeline calculator is bundled together
|
||||
with corresponding metadata into a unified MediaPipe
|
||||
For your convenience, this calculator is bundled together with corresponding
|
||||
metadata into a unified MediaPipe
|
||||
[subgraph](https://github.com/google/mediapipe/tree/master/mediapipe/modules/face_geometry/face_geometry_from_landmarks.pbtxt).
|
||||
The face geometry format is defined as a Protocol Buffer
|
||||
The face transform format is defined as a Protocol Buffer
|
||||
[message](https://github.com/google/mediapipe/tree/master/mediapipe/modules/face_geometry/protos/face_geometry.proto).
|
||||
|
||||
#### Effect Renderer
|
||||
@@ -227,7 +227,7 @@ The effect renderer is implemented as a MediaPipe
|
||||
|
||||
|  |
|
||||
| :---------------------------------------------------------------------: |
|
||||
| *Fig 5. An example of face effects rendered by the Face Geometry Effect Renderer.* |
|
||||
| *Fig 5. An example of face effects rendered by the Face Transform Effect Renderer.* |
|
||||
|
||||
## Solution APIs
|
||||
|
||||
|
||||
@@ -116,7 +116,7 @@ on how to build MediaPipe examples.
|
||||
|
||||
Note: The following runs TensorFlow inference on CPU. If you would like to
|
||||
run inference on GPU (Linux only), please follow
|
||||
[TensorFlow CUDA Support and Setup on Linux Desktop](gpu.md#tensorflow-cuda-support-and-setup-on-linux-desktop)
|
||||
[TensorFlow CUDA Support and Setup on Linux Desktop](../getting_started/gpu_support.md#tensorflow-cuda-support-and-setup-on-linux-desktop)
|
||||
instead.
|
||||
|
||||
To build the TensorFlow CPU inference example on desktop, run:
|
||||
|
||||
@@ -384,7 +384,7 @@ Supported configuration options:
|
||||
<meta charset="utf-8">
|
||||
<script src="https://cdn.jsdelivr.net/npm/@mediapipe/camera_utils/camera_utils.js" crossorigin="anonymous"></script>
|
||||
<script src="https://cdn.jsdelivr.net/npm/@mediapipe/control_utils/control_utils.js" crossorigin="anonymous"></script>
|
||||
<script src="https://cdn.jsdelivr.net/npm/@mediapipe/drawing_utils/control_utils_3d.js" crossorigin="anonymous"></script>
|
||||
<script src="https://cdn.jsdelivr.net/npm/@mediapipe/control_utils_3d/control_utils_3d.js" crossorigin="anonymous"></script>
|
||||
<script src="https://cdn.jsdelivr.net/npm/@mediapipe/drawing_utils/drawing_utils.js" crossorigin="anonymous"></script>
|
||||
<script src="https://cdn.jsdelivr.net/npm/@mediapipe/objectron/objectron.js" crossorigin="anonymous"></script>
|
||||
</head>
|
||||
|
||||
@@ -359,7 +359,7 @@ Supported configuration options:
|
||||
<meta charset="utf-8">
|
||||
<script src="https://cdn.jsdelivr.net/npm/@mediapipe/camera_utils/camera_utils.js" crossorigin="anonymous"></script>
|
||||
<script src="https://cdn.jsdelivr.net/npm/@mediapipe/control_utils/control_utils.js" crossorigin="anonymous"></script>
|
||||
<script src="https://cdn.jsdelivr.net/npm/@mediapipe/drawing_utils/control_utils_3d.js" crossorigin="anonymous"></script>
|
||||
<script src="https://cdn.jsdelivr.net/npm/@mediapipe/control_utils_3d/control_utils_3d.js" crossorigin="anonymous"></script>
|
||||
<script src="https://cdn.jsdelivr.net/npm/@mediapipe/drawing_utils/drawing_utils.js" crossorigin="anonymous"></script>
|
||||
<script src="https://cdn.jsdelivr.net/npm/@mediapipe/pose/pose.js" crossorigin="anonymous"></script>
|
||||
</head>
|
||||
|
||||
@@ -258,13 +258,14 @@ Many of the following settings are advanced and not recommended for general
|
||||
usage. Consult [Enabling tracing and profiling](#enabling-tracing-and-profiling)
|
||||
for a friendlier introduction.
|
||||
|
||||
histogram_interval_size_usec :Specifies the size of the runtimes histogram
|
||||
intervals (in microseconds) to generate the histogram of the Process() time. The
|
||||
last interval extends to +inf. If not specified, the interval is 1000000 usec =
|
||||
1 sec.
|
||||
histogram_interval_size_usec
|
||||
: Specifies the size of the runtimes histogram intervals (in microseconds) to
|
||||
generate the histogram of the `Process()` time. The last interval extends to
|
||||
+inf. If not specified, the interval is 1000000 usec = 1 sec.
|
||||
|
||||
num_histogram_intervals :Specifies the number of intervals to generate the
|
||||
histogram of the `Process()` runtime. If not specified, one interval is used.
|
||||
num_histogram_intervals
|
||||
: Specifies the number of intervals to generate the histogram of the
|
||||
`Process()` runtime. If not specified, one interval is used.
|
||||
|
||||
enable_profiler
|
||||
: If true, the profiler starts profiling when graph is initialized.
|
||||
@@ -288,7 +289,7 @@ trace_event_types_disabled
|
||||
|
||||
trace_log_path
|
||||
: The output directory and base-name prefix for trace log files. Log files are
|
||||
written to: StrCat(trace_log_path, index, "`.binarypb`")
|
||||
written to: `StrCat(trace_log_path, index, ".binarypb")`
|
||||
|
||||
trace_log_count
|
||||
: The number of trace log files retained. The trace log files are named
|
||||
@@ -310,8 +311,8 @@ trace_log_instant_events
|
||||
|
||||
trace_log_interval_count
|
||||
: The number of trace log intervals per file. The total log duration is:
|
||||
`trace_log_interval_usec * trace_log_file_count * trace_log_interval_count`.
|
||||
The default value specifies 10 intervals per file.
|
||||
`trace_log_interval_usec * trace_log_count * trace_log_interval_count`. The
|
||||
default value specifies 10 intervals per file.
|
||||
|
||||
trace_log_disabled
|
||||
: An option to turn ON/OFF writing trace files to disk. Saving trace files to
|
||||
|
||||
@@ -75,6 +75,7 @@ alias(
|
||||
actual = select({
|
||||
":macos_i386": ":macos_i386",
|
||||
":macos_x86_64": ":macos_x86_64",
|
||||
":macos_arm64": ":macos_arm64",
|
||||
"//conditions:default": ":macos_i386", # Arbitrarily chosen from above.
|
||||
}),
|
||||
visibility = ["//visibility:public"],
|
||||
@@ -119,6 +120,15 @@ config_setting(
|
||||
visibility = ["//visibility:public"],
|
||||
)
|
||||
|
||||
config_setting(
|
||||
name = "macos_arm64",
|
||||
values = {
|
||||
"apple_platform_type": "macos",
|
||||
"cpu": "darwin_arm64",
|
||||
},
|
||||
visibility = ["//visibility:public"],
|
||||
)
|
||||
|
||||
[
|
||||
config_setting(
|
||||
name = arch,
|
||||
|
||||
@@ -244,6 +244,7 @@ cc_test(
|
||||
"//mediapipe/framework/formats:time_series_header_cc_proto",
|
||||
"//mediapipe/framework/port:gtest_main",
|
||||
"//mediapipe/framework/port:parse_text_proto",
|
||||
"//mediapipe/framework/tool:test_util",
|
||||
"@com_google_absl//absl/flags:flag",
|
||||
],
|
||||
)
|
||||
|
||||
@@ -20,8 +20,12 @@
|
||||
#include "mediapipe/framework/port/gtest.h"
|
||||
#include "mediapipe/framework/port/parse_text_proto.h"
|
||||
#include "mediapipe/framework/port/status_matchers.h"
|
||||
#include "mediapipe/framework/tool/test_util.h"
|
||||
|
||||
namespace mediapipe {
|
||||
namespace {
|
||||
|
||||
constexpr char kTestPackageRoot[] = "mediapipe/calculators/audio";
|
||||
|
||||
TEST(AudioDecoderCalculatorTest, TestWAV) {
|
||||
CalculatorGraphConfig::Node node_config =
|
||||
@@ -37,9 +41,8 @@ TEST(AudioDecoderCalculatorTest, TestWAV) {
|
||||
})pb");
|
||||
CalculatorRunner runner(node_config);
|
||||
runner.MutableSidePackets()->Tag("INPUT_FILE_PATH") = MakePacket<std::string>(
|
||||
file::JoinPath("./",
|
||||
"/mediapipe/calculators/audio/"
|
||||
"testdata/sine_wave_1k_44100_mono_2_sec_wav.audio"));
|
||||
file::JoinPath(GetTestDataDir(kTestPackageRoot),
|
||||
"sine_wave_1k_44100_mono_2_sec_wav.audio"));
|
||||
MP_ASSERT_OK(runner.Run());
|
||||
MP_EXPECT_OK(runner.Outputs()
|
||||
.Tag("AUDIO_HEADER")
|
||||
@@ -68,9 +71,8 @@ TEST(AudioDecoderCalculatorTest, Test48KWAV) {
|
||||
})pb");
|
||||
CalculatorRunner runner(node_config);
|
||||
runner.MutableSidePackets()->Tag("INPUT_FILE_PATH") = MakePacket<std::string>(
|
||||
file::JoinPath("./",
|
||||
"/mediapipe/calculators/audio/"
|
||||
"testdata/sine_wave_1k_48000_stereo_2_sec_wav.audio"));
|
||||
file::JoinPath(GetTestDataDir(kTestPackageRoot),
|
||||
"sine_wave_1k_48000_stereo_2_sec_wav.audio"));
|
||||
MP_ASSERT_OK(runner.Run());
|
||||
MP_EXPECT_OK(runner.Outputs()
|
||||
.Tag("AUDIO_HEADER")
|
||||
@@ -99,9 +101,8 @@ TEST(AudioDecoderCalculatorTest, TestMP3) {
|
||||
})pb");
|
||||
CalculatorRunner runner(node_config);
|
||||
runner.MutableSidePackets()->Tag("INPUT_FILE_PATH") = MakePacket<std::string>(
|
||||
file::JoinPath("./",
|
||||
"/mediapipe/calculators/audio/"
|
||||
"testdata/sine_wave_1k_44100_stereo_2_sec_mp3.audio"));
|
||||
file::JoinPath(GetTestDataDir(kTestPackageRoot),
|
||||
"sine_wave_1k_44100_stereo_2_sec_mp3.audio"));
|
||||
MP_ASSERT_OK(runner.Run());
|
||||
MP_EXPECT_OK(runner.Outputs()
|
||||
.Tag("AUDIO_HEADER")
|
||||
@@ -130,9 +131,8 @@ TEST(AudioDecoderCalculatorTest, TestAAC) {
|
||||
})pb");
|
||||
CalculatorRunner runner(node_config);
|
||||
runner.MutableSidePackets()->Tag("INPUT_FILE_PATH") = MakePacket<std::string>(
|
||||
file::JoinPath("./",
|
||||
"/mediapipe/calculators/audio/"
|
||||
"testdata/sine_wave_1k_44100_stereo_2_sec_aac.audio"));
|
||||
file::JoinPath(GetTestDataDir(kTestPackageRoot),
|
||||
"sine_wave_1k_44100_stereo_2_sec_aac.audio"));
|
||||
MP_ASSERT_OK(runner.Run());
|
||||
MP_EXPECT_OK(runner.Outputs()
|
||||
.Tag("AUDIO_HEADER")
|
||||
@@ -147,4 +147,5 @@ TEST(AudioDecoderCalculatorTest, TestAAC) {
|
||||
std::ceil(44100.0 * 2 / 1024));
|
||||
}
|
||||
|
||||
} // namespace
|
||||
} // namespace mediapipe
|
||||
|
||||
@@ -20,24 +20,22 @@
|
||||
#include <memory>
|
||||
#include <string>
|
||||
|
||||
#include "Eigen/Core"
|
||||
#include "absl/strings/string_view.h"
|
||||
#include "audio/dsp/spectrogram/spectrogram.h"
|
||||
#include "audio/dsp/window_functions.h"
|
||||
#include "mediapipe/calculators/audio/spectrogram_calculator.pb.h"
|
||||
#include "mediapipe/framework/calculator_framework.h"
|
||||
#include "mediapipe/framework/formats/matrix.h"
|
||||
#include "mediapipe/framework/formats/time_series_header.pb.h"
|
||||
#include "mediapipe/framework/port/core_proto_inc.h"
|
||||
#include "mediapipe/framework/port/integral_types.h"
|
||||
#include "mediapipe/framework/port/logging.h"
|
||||
#include "mediapipe/framework/port/ret_check.h"
|
||||
#include "mediapipe/framework/port/source_location.h"
|
||||
#include "mediapipe/framework/port/status_builder.h"
|
||||
#include "mediapipe/util/time_series_util.h"
|
||||
|
||||
namespace mediapipe {
|
||||
|
||||
namespace {
|
||||
constexpr char kFrameDurationTag[] = "FRAME_DURATION";
|
||||
constexpr char kFrameOverlapTag[] = "FRAME_OVERLAP";
|
||||
} // namespace
|
||||
// MediaPipe Calculator for computing the "spectrogram" (short-time Fourier
|
||||
// transform squared-magnitude, by default) of a multichannel input
|
||||
// time series, including optionally overlapping frames. Options are
|
||||
@@ -46,11 +44,14 @@ namespace mediapipe {
|
||||
//
|
||||
// Result is a MatrixData record (for single channel input and when the
|
||||
// allow_multichannel_input flag is false), or a vector of MatrixData records,
|
||||
// one for each channel (when the allow_multichannel_input flag is set). The
|
||||
// rows of each spectrogram matrix correspond to the n_fft/2+1 unique complex
|
||||
// values, or squared/linear/dB magnitudes, depending on the output_type option.
|
||||
// Each input packet will result in zero or one output packets, each containing
|
||||
// one Matrix for each channel of the input, where each Matrix has one or more
|
||||
// one for each channel (when the allow_multichannel_input flag is set). Each
|
||||
// waveform frame is converted to frequency by a fast Fourier transform whose
|
||||
// size, n_fft, is the smallest power of two large enough to enclose the frame
|
||||
// length of round(frame_duration_seconds * sample_rate).The rows of each
|
||||
// spectrogram matrix(result) correspond to the n_fft/2+1 unique complex values,
|
||||
// or squared/linear/dB magnitudes, depending on the output_type option. Each
|
||||
// input packet will result in zero or one output packets, each containing one
|
||||
// Matrix for each channel of the input, where each Matrix has one or more
|
||||
// columns of spectral values, one for each complete frame of input samples. If
|
||||
// the input packet contains too few samples to trigger a new output frame, no
|
||||
// output packet is generated (since zero-length packets are not legal since
|
||||
@@ -71,6 +72,22 @@ class SpectrogramCalculator : public CalculatorBase {
|
||||
// Input stream with TimeSeriesHeader.
|
||||
);
|
||||
|
||||
if (cc->InputSidePackets().HasTag(kFrameDurationTag)) {
|
||||
cc->InputSidePackets()
|
||||
.Tag(kFrameDurationTag)
|
||||
.Set<double>(
|
||||
// Optional side packet for frame_duration_seconds if provided.
|
||||
);
|
||||
}
|
||||
|
||||
if (cc->InputSidePackets().HasTag(kFrameOverlapTag)) {
|
||||
cc->InputSidePackets()
|
||||
.Tag(kFrameOverlapTag)
|
||||
.Set<double>(
|
||||
// Optional side packet for frame_overlap_seconds if provided.
|
||||
);
|
||||
}
|
||||
|
||||
SpectrogramCalculatorOptions spectrogram_options =
|
||||
cc->Options<SpectrogramCalculatorOptions>();
|
||||
if (!spectrogram_options.allow_multichannel_input()) {
|
||||
@@ -184,27 +201,47 @@ class SpectrogramCalculator : public CalculatorBase {
|
||||
// Fixed scale factor applied to output values (regardless of type).
|
||||
double output_scale_;
|
||||
|
||||
static const float kLnPowerToDb;
|
||||
static const float kLnSquaredMagnitudeToDb;
|
||||
};
|
||||
REGISTER_CALCULATOR(SpectrogramCalculator);
|
||||
|
||||
// Factor to convert ln(magnitude_squared) to deciBels = 10.0/ln(10.0).
|
||||
const float SpectrogramCalculator::kLnPowerToDb = 4.342944819032518;
|
||||
// DECIBELS = 20*log10(LINEAR_MAGNITUDE) = 10*Log10(SQUARED_MAGNITUDE)
|
||||
// =10/ln(10)*ln(SQUARED_MAGNITUDE).
|
||||
// Factor to convert ln(SQUARED_MAGNITUDE) to deciBels = 10.0/ln(10.0).
|
||||
const float SpectrogramCalculator::kLnSquaredMagnitudeToDb = 4.342944819032518;
|
||||
|
||||
absl::Status SpectrogramCalculator::Open(CalculatorContext* cc) {
|
||||
SpectrogramCalculatorOptions spectrogram_options =
|
||||
cc->Options<SpectrogramCalculatorOptions>();
|
||||
// Provide frame_duration_seconds and frame_overlap_seconds either from static
|
||||
// options, or dynamically from a side packet, the side packet one will
|
||||
// override the options one if provided.
|
||||
|
||||
double frame_duration_seconds = 0;
|
||||
double frame_overlap_seconds = 0;
|
||||
if (cc->InputSidePackets().HasTag(kFrameDurationTag)) {
|
||||
frame_duration_seconds =
|
||||
cc->InputSidePackets().Tag(kFrameDurationTag).Get<double>();
|
||||
} else {
|
||||
frame_duration_seconds = spectrogram_options.frame_duration_seconds();
|
||||
}
|
||||
|
||||
if (cc->InputSidePackets().HasTag(kFrameOverlapTag)) {
|
||||
frame_overlap_seconds =
|
||||
cc->InputSidePackets().Tag(kFrameOverlapTag).Get<double>();
|
||||
} else {
|
||||
frame_overlap_seconds = spectrogram_options.frame_overlap_seconds();
|
||||
}
|
||||
|
||||
use_local_timestamp_ = spectrogram_options.use_local_timestamp();
|
||||
|
||||
if (spectrogram_options.frame_duration_seconds() <= 0.0) {
|
||||
if (frame_duration_seconds <= 0.0) {
|
||||
// TODO: return an error.
|
||||
}
|
||||
if (spectrogram_options.frame_overlap_seconds() >=
|
||||
spectrogram_options.frame_duration_seconds()) {
|
||||
if (frame_overlap_seconds >= frame_duration_seconds) {
|
||||
// TODO: return an error.
|
||||
}
|
||||
if (spectrogram_options.frame_overlap_seconds() < 0.0) {
|
||||
if (frame_overlap_seconds < 0.0) {
|
||||
// TODO: return an error.
|
||||
}
|
||||
|
||||
@@ -220,10 +257,8 @@ absl::Status SpectrogramCalculator::Open(CalculatorContext* cc) {
|
||||
// TODO: return an error.
|
||||
}
|
||||
|
||||
frame_duration_samples_ =
|
||||
round(spectrogram_options.frame_duration_seconds() * input_sample_rate_);
|
||||
frame_overlap_samples_ =
|
||||
round(spectrogram_options.frame_overlap_seconds() * input_sample_rate_);
|
||||
frame_duration_samples_ = round(frame_duration_seconds * input_sample_rate_);
|
||||
frame_overlap_samples_ = round(frame_overlap_seconds * input_sample_rate_);
|
||||
|
||||
pad_final_packet_ = spectrogram_options.pad_final_packet();
|
||||
output_type_ = spectrogram_options.output_type();
|
||||
@@ -419,7 +454,7 @@ absl::Status SpectrogramCalculator::ProcessVector(const Matrix& input_stream,
|
||||
return ProcessVectorToOutput(
|
||||
input_stream,
|
||||
+[](const Matrix& col) -> const Matrix {
|
||||
return kLnPowerToDb * col.array().log().matrix();
|
||||
return kLnSquaredMagnitudeToDb * col.array().log().matrix();
|
||||
}, cc);
|
||||
}
|
||||
// clang-format on
|
||||
|
||||
@@ -32,7 +32,11 @@ message SpectrogramCalculatorOptions {
|
||||
|
||||
// Duration of overlap between adjacent windows.
|
||||
// Hence, frame_rate = 1/(frame_duration_seconds - frame_overlap_seconds).
|
||||
// Required that 0 <= frame_overlap_seconds < frame_duration_seconds.
|
||||
// Note the frame_rate here is not the MediaPipe packet rate, the frame here
|
||||
// means each Fourier transform analysis waveform frame, the output MediaPipe
|
||||
// packet rate will the the same as input, if frame rate is lower than input
|
||||
// packet rate, will result in intermittent empty output packets. Required
|
||||
// that 0 <= frame_overlap_seconds < frame_duration_seconds.
|
||||
optional double frame_overlap_seconds = 2 [default = 0.0];
|
||||
|
||||
// Whether to pad the final packet with zeros. If true, guarantees that
|
||||
@@ -42,6 +46,11 @@ message SpectrogramCalculatorOptions {
|
||||
|
||||
// Output value type can be squared-magnitude, linear-magnitude,
|
||||
// deciBels (dB, = 20*log10(linear_magnitude)), or std::complex.
|
||||
// Their relationship:
|
||||
// COMPLEX c = Re + Im*i;
|
||||
// SQUARED_MAGNITUDE = Re^2 + Im^2;
|
||||
// LINEAR_MAGNITUDE = sqrt(SQUARED_MAGNITUDE);
|
||||
// DECIBELS = 20*log10(LINEAR_MAGNITUDE) = 10*log10(SQUARED_MAGNITUDE);
|
||||
enum OutputType {
|
||||
SQUARED_MAGNITUDE = 0;
|
||||
LINEAR_MAGNITUDE = 1;
|
||||
|
||||
@@ -117,6 +117,7 @@ mediapipe_proto_library(
|
||||
"//mediapipe/framework:calculator_options_proto",
|
||||
"//mediapipe/framework:calculator_proto",
|
||||
"//mediapipe/framework/formats:classification_proto",
|
||||
"//mediapipe/framework/formats:landmark_proto",
|
||||
],
|
||||
)
|
||||
|
||||
@@ -213,6 +214,7 @@ cc_library(
|
||||
"//mediapipe/framework:collection_item_id",
|
||||
"//mediapipe/framework:packet",
|
||||
"//mediapipe/framework/formats:classification_cc_proto",
|
||||
"//mediapipe/framework/formats:detection_cc_proto",
|
||||
"//mediapipe/framework/formats:landmark_cc_proto",
|
||||
"//mediapipe/framework/formats:rect_cc_proto",
|
||||
"//mediapipe/framework/port:integral_types",
|
||||
@@ -309,8 +311,8 @@ cc_library(
|
||||
)
|
||||
|
||||
cc_library(
|
||||
name = "concatenate_normalized_landmark_list_calculator",
|
||||
srcs = ["concatenate_normalized_landmark_list_calculator.cc"],
|
||||
name = "concatenate_proto_list_calculator",
|
||||
srcs = ["concatenate_proto_list_calculator.cc"],
|
||||
visibility = ["//visibility:public"],
|
||||
deps = [
|
||||
":concatenate_vector_calculator_cc_proto",
|
||||
@@ -324,10 +326,10 @@ cc_library(
|
||||
)
|
||||
|
||||
cc_test(
|
||||
name = "concatenate_normalized_landmark_list_calculator_test",
|
||||
srcs = ["concatenate_normalized_landmark_list_calculator_test.cc"],
|
||||
name = "concatenate_proto_list_calculator_test",
|
||||
srcs = ["concatenate_proto_list_calculator_test.cc"],
|
||||
deps = [
|
||||
":concatenate_normalized_landmark_list_calculator",
|
||||
":concatenate_proto_list_calculator",
|
||||
":concatenate_vector_calculator_cc_proto",
|
||||
"//mediapipe/framework:calculator_framework",
|
||||
"//mediapipe/framework:calculator_runner",
|
||||
@@ -555,6 +557,22 @@ cc_library(
|
||||
alwayslink = 1,
|
||||
)
|
||||
|
||||
cc_test(
|
||||
name = "packet_cloner_calculator_test",
|
||||
srcs = ["packet_cloner_calculator_test.cc"],
|
||||
deps = [
|
||||
":packet_cloner_calculator",
|
||||
"//mediapipe/framework:calculator_framework",
|
||||
"//mediapipe/framework:timestamp",
|
||||
"//mediapipe/framework/port:gtest_main",
|
||||
"//mediapipe/framework/port:parse_text_proto",
|
||||
"//mediapipe/framework/port:status",
|
||||
"//mediapipe/framework/stream_handler:immediate_input_stream_handler",
|
||||
"//mediapipe/framework/tool:sink",
|
||||
"@com_google_absl//absl/strings",
|
||||
],
|
||||
)
|
||||
|
||||
cc_library(
|
||||
name = "packet_inner_join_calculator",
|
||||
srcs = ["packet_inner_join_calculator.cc"],
|
||||
@@ -964,8 +982,8 @@ cc_test(
|
||||
)
|
||||
|
||||
cc_library(
|
||||
name = "split_landmarks_calculator",
|
||||
srcs = ["split_landmarks_calculator.cc"],
|
||||
name = "split_proto_list_calculator",
|
||||
srcs = ["split_proto_list_calculator.cc"],
|
||||
visibility = ["//visibility:public"],
|
||||
deps = [
|
||||
":split_vector_calculator_cc_proto",
|
||||
@@ -979,10 +997,10 @@ cc_library(
|
||||
)
|
||||
|
||||
cc_test(
|
||||
name = "split_landmarks_calculator_test",
|
||||
srcs = ["split_landmarks_calculator_test.cc"],
|
||||
name = "split_proto_list_calculator_test",
|
||||
srcs = ["split_proto_list_calculator_test.cc"],
|
||||
deps = [
|
||||
":split_landmarks_calculator",
|
||||
":split_proto_list_calculator",
|
||||
":split_vector_calculator_cc_proto",
|
||||
"//mediapipe/framework:calculator_framework",
|
||||
"//mediapipe/framework:calculator_runner",
|
||||
@@ -1195,6 +1213,7 @@ cc_library(
|
||||
"//mediapipe/framework:calculator_framework",
|
||||
"//mediapipe/framework:collection_item_id",
|
||||
"//mediapipe/framework/formats:classification_cc_proto",
|
||||
"//mediapipe/framework/formats:landmark_cc_proto",
|
||||
"//mediapipe/framework/port:integral_types",
|
||||
"//mediapipe/framework/port:ret_check",
|
||||
"//mediapipe/framework/port:status",
|
||||
@@ -1255,3 +1274,36 @@ cc_test(
|
||||
"@com_google_absl//absl/time",
|
||||
],
|
||||
)
|
||||
|
||||
cc_library(
|
||||
name = "get_vector_item_calculator",
|
||||
srcs = ["get_vector_item_calculator.cc"],
|
||||
hdrs = ["get_vector_item_calculator.h"],
|
||||
visibility = ["//visibility:public"],
|
||||
deps = [
|
||||
"//mediapipe/framework:calculator_framework",
|
||||
"//mediapipe/framework:packet",
|
||||
"//mediapipe/framework/api2:node",
|
||||
"//mediapipe/framework/api2:port",
|
||||
"//mediapipe/framework/formats:classification_cc_proto",
|
||||
"//mediapipe/framework/formats:landmark_cc_proto",
|
||||
"//mediapipe/framework/port:ret_check",
|
||||
"//mediapipe/framework/port:status",
|
||||
],
|
||||
alwayslink = 1,
|
||||
)
|
||||
|
||||
cc_library(
|
||||
name = "vector_size_calculator",
|
||||
srcs = ["vector_size_calculator.cc"],
|
||||
hdrs = ["vector_size_calculator.h"],
|
||||
visibility = ["//visibility:public"],
|
||||
deps = [
|
||||
"//mediapipe/framework:calculator_framework",
|
||||
"//mediapipe/framework/api2:node",
|
||||
"//mediapipe/framework/formats:classification_cc_proto",
|
||||
"//mediapipe/framework/formats:landmark_cc_proto",
|
||||
"//mediapipe/framework/port:status",
|
||||
],
|
||||
alwayslink = 1,
|
||||
)
|
||||
|
||||
@@ -28,6 +28,10 @@ typedef BeginLoopCalculator<std::vector<::mediapipe::NormalizedLandmarkList>>
|
||||
BeginLoopNormalizedLandmarkListVectorCalculator;
|
||||
REGISTER_CALCULATOR(BeginLoopNormalizedLandmarkListVectorCalculator);
|
||||
|
||||
// A calculator to process std::vector<int>.
|
||||
typedef BeginLoopCalculator<std::vector<int>> BeginLoopIntCalculator;
|
||||
REGISTER_CALCULATOR(BeginLoopIntCalculator);
|
||||
|
||||
// A calculator to process std::vector<NormalizedRect>.
|
||||
typedef BeginLoopCalculator<std::vector<::mediapipe::NormalizedRect>>
|
||||
BeginLoopNormalizedRectCalculator;
|
||||
|
||||
@@ -1,79 +0,0 @@
|
||||
// Copyright 2019 The MediaPipe Authors.
|
||||
//
|
||||
// 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.
|
||||
|
||||
#ifndef MEDIAPIPE_CALCULATORS_CORE_CONCATENATE_NORMALIZED_LIST_CALCULATOR_H_ // NOLINT
|
||||
#define MEDIAPIPE_CALCULATORS_CORE_CONCATENATE_NORMALIZED_LIST_CALCULATOR_H_ // NOLINT
|
||||
|
||||
#include "mediapipe/calculators/core/concatenate_vector_calculator.pb.h"
|
||||
#include "mediapipe/framework/api2/node.h"
|
||||
#include "mediapipe/framework/calculator_framework.h"
|
||||
#include "mediapipe/framework/formats/landmark.pb.h"
|
||||
#include "mediapipe/framework/port/canonical_errors.h"
|
||||
#include "mediapipe/framework/port/ret_check.h"
|
||||
#include "mediapipe/framework/port/status.h"
|
||||
|
||||
namespace mediapipe {
|
||||
namespace api2 {
|
||||
|
||||
// Concatenates several NormalizedLandmarkList protos following stream index
|
||||
// order. This class assumes that every input stream contains a
|
||||
// NormalizedLandmarkList proto object.
|
||||
class ConcatenateNormalizedLandmarkListCalculator : public Node {
|
||||
public:
|
||||
static constexpr Input<NormalizedLandmarkList>::Multiple kIn{""};
|
||||
static constexpr Output<NormalizedLandmarkList> kOut{""};
|
||||
|
||||
MEDIAPIPE_NODE_CONTRACT(kIn, kOut);
|
||||
|
||||
static absl::Status UpdateContract(CalculatorContract* cc) {
|
||||
RET_CHECK_GE(kIn(cc).Count(), 1);
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
absl::Status Open(CalculatorContext* cc) override {
|
||||
only_emit_if_all_present_ =
|
||||
cc->Options<::mediapipe::ConcatenateVectorCalculatorOptions>()
|
||||
.only_emit_if_all_present();
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
absl::Status Process(CalculatorContext* cc) override {
|
||||
if (only_emit_if_all_present_) {
|
||||
for (const auto& input : kIn(cc)) {
|
||||
if (input.IsEmpty()) return absl::OkStatus();
|
||||
}
|
||||
}
|
||||
|
||||
NormalizedLandmarkList output;
|
||||
for (const auto& input : kIn(cc)) {
|
||||
if (input.IsEmpty()) continue;
|
||||
const NormalizedLandmarkList& list = *input;
|
||||
for (int j = 0; j < list.landmark_size(); ++j) {
|
||||
*output.add_landmark() = list.landmark(j);
|
||||
}
|
||||
}
|
||||
kOut(cc).Send(std::move(output));
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
private:
|
||||
bool only_emit_if_all_present_;
|
||||
};
|
||||
MEDIAPIPE_REGISTER_NODE(ConcatenateNormalizedLandmarkListCalculator);
|
||||
|
||||
} // namespace api2
|
||||
} // namespace mediapipe
|
||||
|
||||
// NOLINTNEXTLINE
|
||||
#endif // MEDIAPIPE_CALCULATORS_CORE_CONCATENATE_NORMALIZED_LIST_CALCULATOR_H_
|
||||
@@ -0,0 +1,118 @@
|
||||
// Copyright 2019 The MediaPipe Authors.
|
||||
//
|
||||
// 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.
|
||||
|
||||
#ifndef MEDIAPIPE_CALCULATORS_CORE_CONCATENATE_PROTO_LIST_CALCULATOR_H_ // NOLINT
|
||||
#define MEDIAPIPE_CALCULATORS_CORE_CONCATENATE_PROTO_LIST_CALCULATOR_H_ // NOLINT
|
||||
|
||||
#include "mediapipe/calculators/core/concatenate_vector_calculator.pb.h"
|
||||
#include "mediapipe/framework/api2/node.h"
|
||||
#include "mediapipe/framework/calculator_framework.h"
|
||||
#include "mediapipe/framework/formats/landmark.pb.h"
|
||||
#include "mediapipe/framework/port/canonical_errors.h"
|
||||
#include "mediapipe/framework/port/ret_check.h"
|
||||
#include "mediapipe/framework/port/status.h"
|
||||
|
||||
namespace mediapipe {
|
||||
namespace api2 {
|
||||
|
||||
// Concatenate several input packets of ListType with a repeated field of
|
||||
// ItemType into a single output packet of ListType following stream index
|
||||
// order.
|
||||
template <typename ItemType, typename ListType>
|
||||
class ConcatenateListsCalculator : public Node {
|
||||
public:
|
||||
static constexpr typename Input<ListType>::Multiple kIn{""};
|
||||
static constexpr Output<ListType> kOut{""};
|
||||
|
||||
MEDIAPIPE_NODE_CONTRACT(kIn, kOut);
|
||||
|
||||
static absl::Status UpdateContract(CalculatorContract* cc) {
|
||||
RET_CHECK_GE(kIn(cc).Count(), 1);
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
absl::Status Open(CalculatorContext* cc) override {
|
||||
only_emit_if_all_present_ =
|
||||
cc->Options<::mediapipe::ConcatenateVectorCalculatorOptions>()
|
||||
.only_emit_if_all_present();
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
absl::Status Process(CalculatorContext* cc) override {
|
||||
if (only_emit_if_all_present_) {
|
||||
for (const auto& input : kIn(cc)) {
|
||||
if (input.IsEmpty()) return absl::OkStatus();
|
||||
}
|
||||
}
|
||||
|
||||
ListType output;
|
||||
for (const auto& input : kIn(cc)) {
|
||||
if (input.IsEmpty()) continue;
|
||||
const ListType& list = *input;
|
||||
for (int j = 0; j < ListSize(list); ++j) {
|
||||
*AddItem(output) = GetItem(list, j);
|
||||
}
|
||||
}
|
||||
kOut(cc).Send(std::move(output));
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
protected:
|
||||
virtual int ListSize(const ListType& list) const = 0;
|
||||
virtual const ItemType GetItem(const ListType& list, int idx) const = 0;
|
||||
virtual ItemType* AddItem(ListType& list) const = 0;
|
||||
|
||||
private:
|
||||
bool only_emit_if_all_present_;
|
||||
};
|
||||
|
||||
// TODO: Move calculators to separate *.cc files
|
||||
|
||||
class ConcatenateNormalizedLandmarkListCalculator
|
||||
: public ConcatenateListsCalculator<NormalizedLandmark,
|
||||
NormalizedLandmarkList> {
|
||||
protected:
|
||||
int ListSize(const NormalizedLandmarkList& list) const override {
|
||||
return list.landmark_size();
|
||||
}
|
||||
const NormalizedLandmark GetItem(const NormalizedLandmarkList& list,
|
||||
int idx) const override {
|
||||
return list.landmark(idx);
|
||||
}
|
||||
NormalizedLandmark* AddItem(NormalizedLandmarkList& list) const override {
|
||||
return list.add_landmark();
|
||||
}
|
||||
};
|
||||
MEDIAPIPE_REGISTER_NODE(ConcatenateNormalizedLandmarkListCalculator);
|
||||
|
||||
class ConcatenateLandmarkListCalculator
|
||||
: public ConcatenateListsCalculator<Landmark, LandmarkList> {
|
||||
protected:
|
||||
int ListSize(const LandmarkList& list) const override {
|
||||
return list.landmark_size();
|
||||
}
|
||||
const Landmark GetItem(const LandmarkList& list, int idx) const override {
|
||||
return list.landmark(idx);
|
||||
}
|
||||
Landmark* AddItem(LandmarkList& list) const override {
|
||||
return list.add_landmark();
|
||||
}
|
||||
};
|
||||
MEDIAPIPE_REGISTER_NODE(ConcatenateLandmarkListCalculator);
|
||||
|
||||
} // namespace api2
|
||||
} // namespace mediapipe
|
||||
|
||||
// NOLINTNEXTLINE
|
||||
#endif // MEDIAPIPE_CALCULATORS_CORE_CONCATENATE_PROTO_LIST_CALCULATOR_H_
|
||||
@@ -73,8 +73,17 @@ typedef ConcatenateVectorCalculator<::mediapipe::NormalizedLandmark>
|
||||
ConcatenateLandmarkVectorCalculator;
|
||||
MEDIAPIPE_REGISTER_NODE(ConcatenateLandmarkVectorCalculator);
|
||||
|
||||
typedef ConcatenateVectorCalculator<::mediapipe::LandmarkList>
|
||||
ConcatenateLandmarkListVectorCalculator;
|
||||
MEDIAPIPE_REGISTER_NODE(ConcatenateLandmarkListVectorCalculator);
|
||||
|
||||
typedef ConcatenateVectorCalculator<::mediapipe::NormalizedLandmarkList>
|
||||
ConcatenateLandmarListVectorCalculator;
|
||||
ConcatenateNormalizedLandmarkListVectorCalculator;
|
||||
MEDIAPIPE_REGISTER_NODE(ConcatenateNormalizedLandmarkListVectorCalculator);
|
||||
|
||||
// For backwards compatibility, keep the version with the typo.
|
||||
using ConcatenateLandmarListVectorCalculator =
|
||||
ConcatenateNormalizedLandmarkListVectorCalculator;
|
||||
MEDIAPIPE_REGISTER_NODE(ConcatenateLandmarListVectorCalculator);
|
||||
|
||||
typedef ConcatenateVectorCalculator<mediapipe::ClassificationList>
|
||||
|
||||
@@ -18,6 +18,7 @@
|
||||
#include "mediapipe/framework/calculator_framework.h"
|
||||
#include "mediapipe/framework/collection_item_id.h"
|
||||
#include "mediapipe/framework/formats/classification.pb.h"
|
||||
#include "mediapipe/framework/formats/landmark.pb.h"
|
||||
#include "mediapipe/framework/port/canonical_errors.h"
|
||||
#include "mediapipe/framework/port/integral_types.h"
|
||||
#include "mediapipe/framework/port/ret_check.h"
|
||||
@@ -79,6 +80,8 @@ class ConstantSidePacketCalculator : public CalculatorBase {
|
||||
packet.Set<uint64>();
|
||||
} else if (packet_options.has_classification_list_value()) {
|
||||
packet.Set<ClassificationList>();
|
||||
} else if (packet_options.has_landmark_list_value()) {
|
||||
packet.Set<LandmarkList>();
|
||||
} else {
|
||||
return absl::InvalidArgumentError(
|
||||
"None of supported values were specified in options.");
|
||||
@@ -108,6 +111,9 @@ class ConstantSidePacketCalculator : public CalculatorBase {
|
||||
} else if (packet_options.has_classification_list_value()) {
|
||||
packet.Set(MakePacket<ClassificationList>(
|
||||
packet_options.classification_list_value()));
|
||||
} else if (packet_options.has_landmark_list_value()) {
|
||||
packet.Set(
|
||||
MakePacket<LandmarkList>(packet_options.landmark_list_value()));
|
||||
} else {
|
||||
return absl::InvalidArgumentError(
|
||||
"None of supported values were specified in options.");
|
||||
|
||||
@@ -18,6 +18,7 @@ package mediapipe;
|
||||
|
||||
import "mediapipe/framework/calculator.proto";
|
||||
import "mediapipe/framework/formats/classification.proto";
|
||||
import "mediapipe/framework/formats/landmark.proto";
|
||||
|
||||
option objc_class_prefix = "MediaPipe";
|
||||
|
||||
@@ -34,6 +35,7 @@ message ConstantSidePacketCalculatorOptions {
|
||||
string string_value = 4;
|
||||
uint64 uint64_value = 5;
|
||||
ClassificationList classification_list_value = 6;
|
||||
LandmarkList landmark_list_value = 7;
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -17,6 +17,7 @@
|
||||
#include <vector>
|
||||
|
||||
#include "mediapipe/framework/formats/classification.pb.h"
|
||||
#include "mediapipe/framework/formats/detection.pb.h"
|
||||
#include "mediapipe/framework/formats/landmark.pb.h"
|
||||
#include "mediapipe/framework/formats/rect.pb.h"
|
||||
#include "mediapipe/util/render_data.pb.h"
|
||||
@@ -50,4 +51,8 @@ REGISTER_CALCULATOR(EndLoopClassificationListCalculator);
|
||||
typedef EndLoopCalculator<std::vector<TfLiteTensor>> EndLoopTensorCalculator;
|
||||
REGISTER_CALCULATOR(EndLoopTensorCalculator);
|
||||
|
||||
typedef EndLoopCalculator<std::vector<::mediapipe::Detection>>
|
||||
EndLoopDetectionCalculator;
|
||||
REGISTER_CALCULATOR(EndLoopDetectionCalculator);
|
||||
|
||||
} // namespace mediapipe
|
||||
|
||||
@@ -32,8 +32,8 @@ constexpr char kOptionsTag[] = "OPTIONS";
|
||||
// FlowLimiterCalculator is used to limit the number of frames in flight
|
||||
// by dropping input frames when necessary.
|
||||
//
|
||||
// The input stream "FINISH" is used to signal the FlowLimiterCalculator
|
||||
// when a frame is finished processing. Either a non-empty "FINISH" packet
|
||||
// The input stream "FINISHED" is used to signal the FlowLimiterCalculator
|
||||
// when a frame is finished processing. Either a non-empty "FINISHED" packet
|
||||
// or a timestamp bound should be received for each processed frame.
|
||||
//
|
||||
// The combination of `max_in_flight: 1` and `max_in_queue: 1` generally gives
|
||||
|
||||
@@ -0,0 +1,32 @@
|
||||
// Copyright 2022 The MediaPipe Authors.
|
||||
//
|
||||
// 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.
|
||||
|
||||
#include "mediapipe/calculators/core/get_vector_item_calculator.h"
|
||||
|
||||
#include "mediapipe/framework/formats/classification.pb.h"
|
||||
#include "mediapipe/framework/formats/landmark.pb.h"
|
||||
|
||||
namespace mediapipe {
|
||||
namespace api2 {
|
||||
|
||||
using GetLandmarkListVectorItemCalculator =
|
||||
GetVectorItemCalculator<mediapipe::LandmarkList>;
|
||||
REGISTER_CALCULATOR(GetLandmarkListVectorItemCalculator);
|
||||
|
||||
using GetClassificationListVectorItemCalculator =
|
||||
GetVectorItemCalculator<mediapipe::ClassificationList>;
|
||||
REGISTER_CALCULATOR(GetClassificationListVectorItemCalculator);
|
||||
|
||||
} // namespace api2
|
||||
} // namespace mediapipe
|
||||
@@ -0,0 +1,77 @@
|
||||
// Copyright 2019 The MediaPipe Authors.
|
||||
//
|
||||
// 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.
|
||||
|
||||
#ifndef MEDIAPIPE_CALCULATORS_CORE_GET_VECTOR_ITEM_CALCULATOR_H_
|
||||
#define MEDIAPIPE_CALCULATORS_CORE_GET_VECTOR_ITEM_CALCULATOR_H_
|
||||
|
||||
#include <optional>
|
||||
|
||||
#include "mediapipe/framework/api2/node.h"
|
||||
#include "mediapipe/framework/api2/port.h"
|
||||
#include "mediapipe/framework/calculator_framework.h"
|
||||
#include "mediapipe/framework/packet.h"
|
||||
#include "mediapipe/framework/port/ret_check.h"
|
||||
#include "mediapipe/framework/port/status.h"
|
||||
|
||||
namespace mediapipe {
|
||||
namespace api2 {
|
||||
|
||||
// A calcutlator to return an item from the vector by its index.
|
||||
//
|
||||
// Inputs:
|
||||
// VECTOR - std::vector<T>
|
||||
// Vector to take an item from.
|
||||
// INDEX - int
|
||||
// Index of the item to return.
|
||||
//
|
||||
// Outputs:
|
||||
// ITEM - T
|
||||
// Item from the vector at given index.
|
||||
//
|
||||
// Example config:
|
||||
// node {
|
||||
// calculator: "Get{SpecificType}VectorItemCalculator"
|
||||
// input_stream: "VECTOR:vector"
|
||||
// input_stream: "INDEX:index"
|
||||
// input_stream: "ITEM:item"
|
||||
// }
|
||||
//
|
||||
template <typename T>
|
||||
class GetVectorItemCalculator : public Node {
|
||||
public:
|
||||
static constexpr Input<std::vector<T>> kIn{"VECTOR"};
|
||||
static constexpr Input<int> kIdx{"INDEX"};
|
||||
static constexpr Output<T> kOut{"ITEM"};
|
||||
|
||||
MEDIAPIPE_NODE_CONTRACT(kIn, kIdx, kOut);
|
||||
|
||||
absl::Status Process(CalculatorContext* cc) final {
|
||||
if (kIn(cc).IsEmpty() || kIdx(cc).IsEmpty()) {
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
const std::vector<T>& items = kIn(cc).Get();
|
||||
const int idx = kIdx(cc).Get();
|
||||
|
||||
RET_CHECK_LT(idx, items.size());
|
||||
kOut(cc).Send(items[idx]);
|
||||
|
||||
return absl::OkStatus();
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace api2
|
||||
} // namespace mediapipe
|
||||
|
||||
#endif // MEDIAPIPE_CALCULATORS_CORE_GET_VECTOR_ITEM_CALCULATOR_H_
|
||||
@@ -29,6 +29,11 @@ namespace api2 {
|
||||
// This calculator periodically copies the GraphProfile from
|
||||
// mediapipe::GraphProfiler::CaptureProfile to the "PROFILE" output stream.
|
||||
//
|
||||
// Similarly to the log files saved by GraphProfiler::WriteProfile when trace
|
||||
// logging is enabled, the first captured profile contains the full
|
||||
// canonicalized graph config and, if tracing is enabled, calculator names in
|
||||
// graph traces. Subsequent profiles omit this information.
|
||||
//
|
||||
// Example config:
|
||||
// node {
|
||||
// calculator: "GraphProfileCalculator"
|
||||
@@ -50,11 +55,14 @@ class GraphProfileCalculator : public Node {
|
||||
absl::Status Process(CalculatorContext* cc) final {
|
||||
auto options = cc->Options<::mediapipe::GraphProfileCalculatorOptions>();
|
||||
|
||||
if (prev_profile_ts_ == Timestamp::Unset() ||
|
||||
bool first_profile = prev_profile_ts_ == Timestamp::Unset();
|
||||
if (first_profile ||
|
||||
cc->InputTimestamp() - prev_profile_ts_ >= options.profile_interval()) {
|
||||
prev_profile_ts_ = cc->InputTimestamp();
|
||||
GraphProfile result;
|
||||
MP_RETURN_IF_ERROR(cc->GetProfilingContext()->CaptureProfile(&result));
|
||||
MP_RETURN_IF_ERROR(cc->GetProfilingContext()->CaptureProfile(
|
||||
&result, first_profile ? PopulateGraphConfig::kFull
|
||||
: PopulateGraphConfig::kNo));
|
||||
kProfileOut(cc).Send(result);
|
||||
}
|
||||
return absl::OkStatus();
|
||||
|
||||
@@ -202,6 +202,8 @@ TEST_F(GraphProfileCalculatorTest, GraphProfile) {
|
||||
}
|
||||
})pb");
|
||||
|
||||
ASSERT_EQ(output_packets.size(), 2);
|
||||
EXPECT_TRUE(output_packets[0].Get<GraphProfile>().has_config());
|
||||
EXPECT_THAT(output_packets[1].Get<GraphProfile>(),
|
||||
mediapipe::EqualsProto(expected_profile));
|
||||
}
|
||||
|
||||
@@ -16,9 +16,10 @@
|
||||
// For every packet that appears in B, outputs the most recent packet from each
|
||||
// of the A_i on a separate stream.
|
||||
|
||||
#include <string_view>
|
||||
#include <vector>
|
||||
|
||||
#include "absl/strings/str_cat.h"
|
||||
#include "absl/strings/string_view.h"
|
||||
#include "mediapipe/calculators/core/packet_cloner_calculator.pb.h"
|
||||
#include "mediapipe/framework/calculator_framework.h"
|
||||
|
||||
@@ -34,7 +35,18 @@ namespace mediapipe {
|
||||
// calculator: "PacketClonerCalculator"
|
||||
// input_stream: "first_base_signal"
|
||||
// input_stream: "second_base_signal"
|
||||
// input_stream: "tick_signal"
|
||||
// input_stream: "tick_signal" # or input_stream: "TICK:tick_signal"
|
||||
// output_stream: "cloned_first_base_signal"
|
||||
// output_stream: "cloned_second_base_signal"
|
||||
// }
|
||||
//
|
||||
// Or you can use "TICK" tag and put corresponding input stream at any location,
|
||||
// for example at the very beginning:
|
||||
// node {
|
||||
// calculator: "PacketClonerCalculator"
|
||||
// input_stream: "TICK:tick_signal"
|
||||
// input_stream: "first_base_signal"
|
||||
// input_stream: "second_base_signal"
|
||||
// output_stream: "cloned_first_base_signal"
|
||||
// output_stream: "cloned_second_base_signal"
|
||||
// }
|
||||
@@ -46,12 +58,13 @@ namespace mediapipe {
|
||||
class PacketClonerCalculator : public CalculatorBase {
|
||||
public:
|
||||
static absl::Status GetContract(CalculatorContract* cc) {
|
||||
const int tick_signal_index = cc->Inputs().NumEntries() - 1;
|
||||
for (int i = 0; i < tick_signal_index; ++i) {
|
||||
cc->Inputs().Index(i).SetAny();
|
||||
cc->Outputs().Index(i).SetSameAs(&cc->Inputs().Index(i));
|
||||
const Ids ids = GetIds(*cc);
|
||||
for (const auto& in_out : ids.inputs_outputs) {
|
||||
auto& input = cc->Inputs().Get(in_out.in);
|
||||
input.SetAny();
|
||||
cc->Outputs().Get(in_out.out).SetSameAs(&input);
|
||||
}
|
||||
cc->Inputs().Index(tick_signal_index).SetAny();
|
||||
cc->Inputs().Get(ids.tick_id).SetAny();
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
@@ -65,13 +78,15 @@ class PacketClonerCalculator : public CalculatorBase {
|
||||
output_empty_packets_before_all_inputs_received_ =
|
||||
calculator_options.output_packets_only_when_all_inputs_received();
|
||||
|
||||
// Parse input streams.
|
||||
tick_signal_index_ = cc->Inputs().NumEntries() - 1;
|
||||
current_.resize(tick_signal_index_);
|
||||
// Prepare input and output ids.
|
||||
ids_ = GetIds(*cc);
|
||||
current_.resize(ids_.inputs_outputs.size());
|
||||
|
||||
// Pass along the header for each stream if present.
|
||||
for (int i = 0; i < tick_signal_index_; ++i) {
|
||||
if (!cc->Inputs().Index(i).Header().IsEmpty()) {
|
||||
cc->Outputs().Index(i).SetHeader(cc->Inputs().Index(i).Header());
|
||||
for (const auto& in_out : ids_.inputs_outputs) {
|
||||
auto& input = cc->Inputs().Get(in_out.in);
|
||||
if (!input.Header().IsEmpty()) {
|
||||
cc->Outputs().Get(in_out.out).SetHeader(input.Header());
|
||||
}
|
||||
}
|
||||
return absl::OkStatus();
|
||||
@@ -79,17 +94,18 @@ class PacketClonerCalculator : public CalculatorBase {
|
||||
|
||||
absl::Status Process(CalculatorContext* cc) final {
|
||||
// Store input signals.
|
||||
for (int i = 0; i < tick_signal_index_; ++i) {
|
||||
if (!cc->Inputs().Index(i).Value().IsEmpty()) {
|
||||
current_[i] = cc->Inputs().Index(i).Value();
|
||||
for (int i = 0; i < ids_.inputs_outputs.size(); ++i) {
|
||||
const auto& input = cc->Inputs().Get(ids_.inputs_outputs[i].in);
|
||||
if (!input.IsEmpty()) {
|
||||
current_[i] = input.Value();
|
||||
}
|
||||
}
|
||||
|
||||
// Output according to the TICK signal.
|
||||
if (!cc->Inputs().Index(tick_signal_index_).Value().IsEmpty()) {
|
||||
if (!cc->Inputs().Get(ids_.tick_id).IsEmpty()) {
|
||||
if (output_only_when_all_inputs_received_) {
|
||||
// Return if one of the input is null.
|
||||
for (int i = 0; i < tick_signal_index_; ++i) {
|
||||
for (int i = 0; i < ids_.inputs_outputs.size(); ++i) {
|
||||
if (current_[i].IsEmpty()) {
|
||||
if (output_empty_packets_before_all_inputs_received_) {
|
||||
SetAllNextTimestampBounds(cc);
|
||||
@@ -99,12 +115,12 @@ class PacketClonerCalculator : public CalculatorBase {
|
||||
}
|
||||
}
|
||||
// Output each stream.
|
||||
for (int i = 0; i < tick_signal_index_; ++i) {
|
||||
for (int i = 0; i < ids_.inputs_outputs.size(); ++i) {
|
||||
auto& output = cc->Outputs().Get(ids_.inputs_outputs[i].out);
|
||||
if (!current_[i].IsEmpty()) {
|
||||
cc->Outputs().Index(i).AddPacket(
|
||||
current_[i].At(cc->InputTimestamp()));
|
||||
output.AddPacket(current_[i].At(cc->InputTimestamp()));
|
||||
} else {
|
||||
cc->Outputs().Index(i).SetNextTimestampBound(
|
||||
output.SetNextTimestampBound(
|
||||
cc->InputTimestamp().NextAllowedInStream());
|
||||
}
|
||||
}
|
||||
@@ -113,15 +129,44 @@ class PacketClonerCalculator : public CalculatorBase {
|
||||
}
|
||||
|
||||
private:
|
||||
struct Ids {
|
||||
struct InputOutput {
|
||||
CollectionItemId in;
|
||||
CollectionItemId out;
|
||||
};
|
||||
CollectionItemId tick_id;
|
||||
std::vector<InputOutput> inputs_outputs;
|
||||
};
|
||||
|
||||
template <typename CC>
|
||||
static Ids GetIds(CC& cc) {
|
||||
Ids ids;
|
||||
static constexpr absl::string_view kEmptyTag = "";
|
||||
int num_inputs_to_clone = cc.Inputs().NumEntries(kEmptyTag);
|
||||
static constexpr absl::string_view kTickTag = "TICK";
|
||||
if (cc.Inputs().HasTag(kTickTag)) {
|
||||
ids.tick_id = cc.Inputs().GetId(kTickTag, 0);
|
||||
} else {
|
||||
--num_inputs_to_clone;
|
||||
ids.tick_id = cc.Inputs().GetId(kEmptyTag, num_inputs_to_clone);
|
||||
}
|
||||
for (int i = 0; i < num_inputs_to_clone; ++i) {
|
||||
ids.inputs_outputs.push_back({.in = cc.Inputs().GetId(kEmptyTag, i),
|
||||
.out = cc.Outputs().GetId(kEmptyTag, i)});
|
||||
}
|
||||
return ids;
|
||||
}
|
||||
|
||||
void SetAllNextTimestampBounds(CalculatorContext* cc) {
|
||||
for (int j = 0; j < tick_signal_index_; ++j) {
|
||||
cc->Outputs().Index(j).SetNextTimestampBound(
|
||||
cc->InputTimestamp().NextAllowedInStream());
|
||||
for (const auto& in_out : ids_.inputs_outputs) {
|
||||
cc->Outputs()
|
||||
.Get(in_out.out)
|
||||
.SetNextTimestampBound(cc->InputTimestamp().NextAllowedInStream());
|
||||
}
|
||||
}
|
||||
|
||||
std::vector<Packet> current_;
|
||||
int tick_signal_index_;
|
||||
Ids ids_;
|
||||
bool output_only_when_all_inputs_received_;
|
||||
bool output_empty_packets_before_all_inputs_received_;
|
||||
};
|
||||
|
||||
@@ -0,0 +1,349 @@
|
||||
// Copyright 2022 The MediaPipe Authors.
|
||||
//
|
||||
// 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.
|
||||
|
||||
#include <algorithm>
|
||||
#include <functional>
|
||||
#include <memory>
|
||||
#include <string>
|
||||
#include <vector>
|
||||
|
||||
#include "absl/strings/str_cat.h"
|
||||
#include "mediapipe/framework/calculator_framework.h"
|
||||
#include "mediapipe/framework/port/gmock.h"
|
||||
#include "mediapipe/framework/port/gtest.h"
|
||||
#include "mediapipe/framework/port/parse_text_proto.h"
|
||||
#include "mediapipe/framework/port/status.h"
|
||||
#include "mediapipe/framework/port/status_matchers.h"
|
||||
#include "mediapipe/framework/timestamp.h"
|
||||
#include "mediapipe/framework/tool/sink.h"
|
||||
|
||||
namespace mediapipe {
|
||||
namespace {
|
||||
|
||||
using ::testing::ElementsAre;
|
||||
using ::testing::Eq;
|
||||
using ::testing::Value;
|
||||
|
||||
MATCHER_P2(IntPacket, value, ts, "") {
|
||||
return Value(arg.template Get<int>(), Eq(value)) &&
|
||||
Value(arg.Timestamp(), Eq(Timestamp(ts)));
|
||||
}
|
||||
|
||||
MATCHER_P2(FloatPacket, value, ts, "") {
|
||||
return Value(arg.template Get<float>(), Eq(value)) &&
|
||||
Value(arg.Timestamp(), Eq(Timestamp(ts)));
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
absl::Status SendPacket(const std::string& input_name, T value, int ts,
|
||||
CalculatorGraph& graph) {
|
||||
return graph.AddPacketToInputStream(input_name,
|
||||
MakePacket<T>(value).At(Timestamp(ts)));
|
||||
}
|
||||
|
||||
struct Params {
|
||||
bool use_tick_tag = false;
|
||||
};
|
||||
|
||||
class PacketClonerCalculatorTest : public testing::TestWithParam<Params> {};
|
||||
|
||||
TEST_P(PacketClonerCalculatorTest, ClonesSingleInputSameTimestamps) {
|
||||
CalculatorGraphConfig graph_config =
|
||||
ParseTextProtoOrDie<CalculatorGraphConfig>([&]() {
|
||||
if (GetParam().use_tick_tag) {
|
||||
return R"pb(
|
||||
input_stream: 'in1'
|
||||
input_stream: 'tick'
|
||||
node {
|
||||
calculator: 'PacketClonerCalculator'
|
||||
input_stream: 'in1'
|
||||
input_stream: 'TICK:tick'
|
||||
output_stream: 'out1'
|
||||
})pb";
|
||||
}
|
||||
return R"pb(
|
||||
input_stream: 'in1'
|
||||
input_stream: 'tick'
|
||||
node {
|
||||
calculator: 'PacketClonerCalculator'
|
||||
input_stream: 'in1'
|
||||
input_stream: 'tick'
|
||||
output_stream: 'out1'
|
||||
})pb";
|
||||
}());
|
||||
std::vector<Packet> out1;
|
||||
tool::AddVectorSink("out1", &graph_config, &out1);
|
||||
|
||||
CalculatorGraph graph;
|
||||
MP_ASSERT_OK(graph.Initialize(graph_config, {}));
|
||||
MP_ASSERT_OK(graph.StartRun({}));
|
||||
|
||||
MP_ASSERT_OK(SendPacket("in1", 1, /*ts=*/10000, graph));
|
||||
MP_ASSERT_OK(SendPacket("tick", 1000, /*ts=*/10000, graph));
|
||||
MP_ASSERT_OK(graph.WaitUntilIdle());
|
||||
|
||||
EXPECT_THAT(out1, ElementsAre(IntPacket(1, 10000)));
|
||||
}
|
||||
|
||||
TEST_P(PacketClonerCalculatorTest, ClonesSingleInputEarlierTimestamps) {
|
||||
CalculatorGraphConfig graph_config =
|
||||
ParseTextProtoOrDie<CalculatorGraphConfig>([&]() {
|
||||
if (GetParam().use_tick_tag) {
|
||||
return R"pb(
|
||||
input_stream: 'in1'
|
||||
input_stream: 'tick'
|
||||
node {
|
||||
calculator: 'PacketClonerCalculator'
|
||||
input_stream: 'in1'
|
||||
input_stream: 'TICK:tick'
|
||||
output_stream: 'out1'
|
||||
})pb";
|
||||
}
|
||||
return R"pb(
|
||||
input_stream: 'in1'
|
||||
input_stream: 'tick'
|
||||
node {
|
||||
calculator: 'PacketClonerCalculator'
|
||||
input_stream: 'in1'
|
||||
input_stream: 'tick'
|
||||
output_stream: 'out1'
|
||||
})pb";
|
||||
}());
|
||||
std::vector<Packet> out1;
|
||||
tool::AddVectorSink("out1", &graph_config, &out1);
|
||||
|
||||
CalculatorGraph graph;
|
||||
MP_ASSERT_OK(graph.Initialize(graph_config, {}));
|
||||
MP_ASSERT_OK(graph.StartRun({}));
|
||||
|
||||
// PacketClonerCalculator is non-ImmediateInputStreamHandler
|
||||
// PacketClonerCalculator waits for "in1" to arrive for ts=5000
|
||||
MP_ASSERT_OK(SendPacket("in1", 1, /*ts=*/5000, graph));
|
||||
// Newer tick at ts=10000, should NOT trigger output for ts=5000
|
||||
// PacketClonerCalculator waits for "in1" to arrive for ts=10000
|
||||
MP_ASSERT_OK(SendPacket("tick", 1000, /*ts=*/10000, graph));
|
||||
MP_ASSERT_OK(SendPacket("tick", 1001, /*ts=*/10001, graph));
|
||||
MP_ASSERT_OK(SendPacket("tick", 1002, /*ts=*/10002, graph));
|
||||
// Newer "in1" at ts=15000, should trigger output for ts=10000
|
||||
MP_ASSERT_OK(SendPacket("in1", 2, /*ts=*/15000, graph));
|
||||
MP_ASSERT_OK(graph.WaitUntilIdle());
|
||||
|
||||
EXPECT_THAT(out1, ElementsAre(IntPacket(1, 10000), IntPacket(1, 10001),
|
||||
IntPacket(1, 10002)));
|
||||
}
|
||||
|
||||
TEST_P(PacketClonerCalculatorTest, ClonesFiveInputs) {
|
||||
CalculatorGraphConfig graph_config =
|
||||
ParseTextProtoOrDie<CalculatorGraphConfig>([&]() {
|
||||
if (GetParam().use_tick_tag) {
|
||||
return R"pb(
|
||||
input_stream: 'in1'
|
||||
input_stream: 'in2'
|
||||
input_stream: 'in3'
|
||||
input_stream: 'in4'
|
||||
input_stream: 'in5'
|
||||
input_stream: 'tick'
|
||||
node {
|
||||
calculator: 'PacketClonerCalculator'
|
||||
input_stream: 'in1'
|
||||
input_stream: 'in2'
|
||||
input_stream: 'in3'
|
||||
input_stream: 'in4'
|
||||
input_stream: 'in5'
|
||||
output_stream: 'out1'
|
||||
output_stream: 'out2'
|
||||
output_stream: 'out3'
|
||||
input_stream: 'TICK:tick' # arbitrary location
|
||||
output_stream: 'out4'
|
||||
output_stream: 'out5'
|
||||
}
|
||||
)pb";
|
||||
}
|
||||
return R"pb(
|
||||
input_stream: 'in1'
|
||||
input_stream: 'in2'
|
||||
input_stream: 'in3'
|
||||
input_stream: 'in4'
|
||||
input_stream: 'in5'
|
||||
input_stream: 'tick'
|
||||
node {
|
||||
calculator: 'PacketClonerCalculator'
|
||||
input_stream: 'in1'
|
||||
input_stream: 'in2'
|
||||
input_stream: 'in3'
|
||||
input_stream: 'in4'
|
||||
input_stream: 'in5'
|
||||
input_stream: 'tick'
|
||||
output_stream: 'out1'
|
||||
output_stream: 'out2'
|
||||
output_stream: 'out3'
|
||||
output_stream: 'out4'
|
||||
output_stream: 'out5'
|
||||
}
|
||||
)pb";
|
||||
}());
|
||||
constexpr int kNumToClone = 5;
|
||||
std::array<std::vector<Packet>, kNumToClone> outs;
|
||||
for (int i = 0; i < kNumToClone; ++i) {
|
||||
tool::AddVectorSink(absl::StrCat("out", i + 1), &graph_config, &outs[i]);
|
||||
}
|
||||
|
||||
CalculatorGraph graph;
|
||||
MP_ASSERT_OK(graph.Initialize(graph_config, {}));
|
||||
MP_ASSERT_OK(graph.StartRun({}));
|
||||
|
||||
MP_ASSERT_OK(SendPacket("in1", 10, /*ts=*/10000, graph));
|
||||
MP_ASSERT_OK(SendPacket("in2", 20.0f, /*ts=*/10000, graph));
|
||||
MP_ASSERT_OK(SendPacket("in3", 30, /*ts=*/10000, graph));
|
||||
MP_ASSERT_OK(SendPacket("in4", 40.0f, /*ts=*/10000, graph));
|
||||
MP_ASSERT_OK(SendPacket("in5", 50, /*ts=*/10000, graph));
|
||||
MP_ASSERT_OK(SendPacket("tick", 1000, /*ts=*/10000, graph));
|
||||
// Below "tick" packets won't trigger output, until newer inputs are sent,
|
||||
// because inputs are missing and ImmediateInputStreamHandler is not
|
||||
// configured.
|
||||
MP_ASSERT_OK(SendPacket("tick", 1001, /*ts=*/10001, graph));
|
||||
MP_ASSERT_OK(SendPacket("tick", 1002, /*ts=*/10002, graph));
|
||||
MP_ASSERT_OK(graph.WaitUntilIdle());
|
||||
|
||||
EXPECT_THAT(outs, ElementsAre(ElementsAre(IntPacket(10, 10000)),
|
||||
ElementsAre(FloatPacket(20.0f, 10000)),
|
||||
ElementsAre(IntPacket(30, 10000)),
|
||||
ElementsAre(FloatPacket(40.0f, 10000)),
|
||||
ElementsAre(IntPacket(50, 10000))));
|
||||
|
||||
MP_ASSERT_OK(SendPacket("in1", 100, /*ts=*/20000, graph));
|
||||
MP_ASSERT_OK(SendPacket("in2", 200.0f, /*ts=*/20000, graph));
|
||||
MP_ASSERT_OK(SendPacket("in3", 300, /*ts=*/20000, graph));
|
||||
MP_ASSERT_OK(SendPacket("in4", 400.0f, /*ts=*/20000, graph));
|
||||
MP_ASSERT_OK(SendPacket("in5", 500, /*ts=*/20000, graph));
|
||||
MP_ASSERT_OK(SendPacket("tick", 2000, /*ts=*/20000, graph));
|
||||
// Below "tick" packets won't trigger output, because inputs are missing and
|
||||
// ImmediateInputStreamHandler is not configured.
|
||||
MP_ASSERT_OK(SendPacket("tick", 2001, /*ts=*/20001, graph));
|
||||
MP_ASSERT_OK(SendPacket("tick", 2002, /*ts=*/20002, graph));
|
||||
MP_ASSERT_OK(graph.WaitUntilIdle());
|
||||
|
||||
EXPECT_THAT(
|
||||
outs,
|
||||
ElementsAre(
|
||||
ElementsAre(IntPacket(10, 10000), IntPacket(10, 10001),
|
||||
IntPacket(10, 10002), IntPacket(100, 20000)),
|
||||
ElementsAre(FloatPacket(20.0f, 10000), FloatPacket(20.0f, 10001),
|
||||
FloatPacket(20.0f, 10002), FloatPacket(200.0f, 20000)),
|
||||
ElementsAre(IntPacket(30, 10000), IntPacket(30, 10001),
|
||||
IntPacket(30, 10002), IntPacket(300, 20000)),
|
||||
ElementsAre(FloatPacket(40.0f, 10000), FloatPacket(40.0f, 10001),
|
||||
FloatPacket(40.0f, 10002), FloatPacket(400.0f, 20000)),
|
||||
ElementsAre(IntPacket(50, 10000), IntPacket(50, 10001),
|
||||
IntPacket(50, 10002), IntPacket(500, 20000))));
|
||||
}
|
||||
|
||||
TEST_P(PacketClonerCalculatorTest,
|
||||
ClonesTwoInputsWithImmediateInputStreamHandler) {
|
||||
CalculatorGraphConfig graph_config =
|
||||
ParseTextProtoOrDie<CalculatorGraphConfig>([&]() {
|
||||
if (GetParam().use_tick_tag) {
|
||||
return R"pb(
|
||||
input_stream: 'in1'
|
||||
input_stream: 'in2'
|
||||
input_stream: 'tick'
|
||||
node {
|
||||
calculator: 'PacketClonerCalculator'
|
||||
input_stream: 'TICK:tick'
|
||||
input_stream: 'in1'
|
||||
input_stream: 'in2'
|
||||
output_stream: 'out1'
|
||||
output_stream: 'out2'
|
||||
input_stream_handler {
|
||||
input_stream_handler: "ImmediateInputStreamHandler"
|
||||
}
|
||||
})pb";
|
||||
}
|
||||
return R"pb(
|
||||
input_stream: 'in1'
|
||||
input_stream: 'in2'
|
||||
input_stream: 'tick'
|
||||
node {
|
||||
calculator: 'PacketClonerCalculator'
|
||||
input_stream: 'in1'
|
||||
input_stream: 'in2'
|
||||
input_stream: 'tick'
|
||||
output_stream: 'out1'
|
||||
output_stream: 'out2'
|
||||
input_stream_handler {
|
||||
input_stream_handler: "ImmediateInputStreamHandler"
|
||||
}
|
||||
})pb";
|
||||
}());
|
||||
constexpr int kNumToClone = 2;
|
||||
std::array<std::vector<Packet>, kNumToClone> outs;
|
||||
for (int i = 0; i < kNumToClone; ++i) {
|
||||
tool::AddVectorSink(absl::StrCat("out", i + 1), &graph_config, &outs[i]);
|
||||
}
|
||||
|
||||
CalculatorGraph graph;
|
||||
MP_ASSERT_OK(graph.Initialize(graph_config, {}));
|
||||
MP_ASSERT_OK(graph.StartRun({}));
|
||||
|
||||
// No packets to clone.
|
||||
MP_ASSERT_OK(SendPacket("tick", 0, /*ts=*/0, graph));
|
||||
MP_ASSERT_OK(graph.WaitUntilIdle());
|
||||
|
||||
// Cloning current packets.
|
||||
MP_ASSERT_OK(SendPacket("in1", 1, /*ts=*/10000, graph));
|
||||
MP_ASSERT_OK(SendPacket("in2", 10.0f, /*ts=*/10000, graph));
|
||||
MP_ASSERT_OK(SendPacket("tick", 1000, /*ts=*/10000, graph));
|
||||
MP_ASSERT_OK(graph.WaitUntilIdle());
|
||||
|
||||
// Cloning past packets.
|
||||
MP_ASSERT_OK(SendPacket("tick", 1500, /*ts=*/15000, graph));
|
||||
MP_ASSERT_OK(graph.WaitUntilIdle());
|
||||
|
||||
// Cloning past packets.
|
||||
MP_ASSERT_OK(SendPacket("in1", 2, /*ts=*/10001, graph));
|
||||
MP_ASSERT_OK(SendPacket("in2", 20.0f, /*ts=*/10001, graph));
|
||||
MP_ASSERT_OK(SendPacket("tick", 2000, /*ts=*/20000, graph));
|
||||
MP_ASSERT_OK(graph.WaitUntilIdle());
|
||||
|
||||
// Cloning future packets.
|
||||
MP_ASSERT_OK(SendPacket("in1", 3, /*ts=*/30000, graph));
|
||||
MP_ASSERT_OK(SendPacket("in2", 30.0f, /*ts=*/30000, graph));
|
||||
// Waiting to ensure newer packets (ts=30000) to clone would get into the
|
||||
// cloner before tick (ts=25000) does.
|
||||
MP_ASSERT_OK(graph.WaitUntilIdle());
|
||||
MP_ASSERT_OK(SendPacket("tick", 3000, /*ts=*/25000, graph));
|
||||
MP_ASSERT_OK(graph.WaitUntilIdle());
|
||||
|
||||
// Cloning packets having different timestamps.
|
||||
MP_ASSERT_OK(SendPacket("in1", 4, /*ts=*/38000, graph));
|
||||
MP_ASSERT_OK(SendPacket("in2", 40.0f, /*ts=*/39000, graph));
|
||||
MP_ASSERT_OK(SendPacket("tick", 4000, /*ts=*/40000, graph));
|
||||
MP_ASSERT_OK(graph.WaitUntilIdle());
|
||||
|
||||
EXPECT_THAT(
|
||||
outs,
|
||||
ElementsAre(
|
||||
ElementsAre(IntPacket(1, 10000), IntPacket(1, 15000),
|
||||
IntPacket(2, 20000), IntPacket(3, 25000),
|
||||
IntPacket(4, 40000)),
|
||||
ElementsAre(FloatPacket(10.0f, 10000), FloatPacket(10.0f, 15000),
|
||||
FloatPacket(20.0f, 20000), FloatPacket(30.0f, 25000),
|
||||
FloatPacket(40.0f, 40000))));
|
||||
}
|
||||
|
||||
INSTANTIATE_TEST_SUITE_P(PacketClonerCalculator, PacketClonerCalculatorTest,
|
||||
testing::ValuesIn({Params{.use_tick_tag = false},
|
||||
Params{.use_tick_tag = true}}));
|
||||
} // anonymous namespace
|
||||
} // namespace mediapipe
|
||||
@@ -157,9 +157,7 @@ absl::Status PacketResamplerCalculator::Process(CalculatorContext* cc) {
|
||||
}
|
||||
}
|
||||
|
||||
if (absl::Status status = strategy_->Process(cc); !status.ok()) {
|
||||
return status; // Avoid MP_RETURN_IF_ERROR macro for external release.
|
||||
}
|
||||
MP_RETURN_IF_ERROR(strategy_->Process(cc));
|
||||
|
||||
last_packet_ = cc->Inputs().Get(input_data_id_).Value();
|
||||
|
||||
|
||||
@@ -23,8 +23,8 @@
|
||||
#include "mediapipe/framework/port/canonical_errors.h"
|
||||
#include "mediapipe/framework/port/status.h"
|
||||
|
||||
// Quantizes a vector of floats to a std::string so that each float becomes a
|
||||
// byte in the [0, 255] range. Any value above max_quantized_value or below
|
||||
// Quantizes a vector of floats to a string so that each float becomes a byte
|
||||
// in the [0, 255] range. Any value above max_quantized_value or below
|
||||
// min_quantized_value will be saturated to '/xFF' or '/0'.
|
||||
//
|
||||
// Example config:
|
||||
|
||||
+62
-33
@@ -12,8 +12,8 @@
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
#ifndef MEDIAPIPE_CALCULATORS_CORE_SPLIT_LANDMARKS_CALCULATOR_H_ // NOLINT
|
||||
#define MEDIAPIPE_CALCULATORS_CORE_SPLIT_LANDMARKS_CALCULATOR_H_ // NOLINT
|
||||
#ifndef MEDIAPIPE_CALCULATORS_CORE_SPLIT_PROTO_LIST_CALCULATOR_H_ // NOLINT
|
||||
#define MEDIAPIPE_CALCULATORS_CORE_SPLIT_PROTO_LIST_CALCULATOR_H_ // NOLINT
|
||||
|
||||
#include "mediapipe/calculators/core/split_vector_calculator.pb.h"
|
||||
#include "mediapipe/framework/calculator_framework.h"
|
||||
@@ -24,30 +24,30 @@
|
||||
|
||||
namespace mediapipe {
|
||||
|
||||
// Splits an input packet with LandmarkListType into
|
||||
// multiple LandmarkListType output packets using the [begin, end) ranges
|
||||
// Splits an input packet of ListType with a repeated field of ItemType
|
||||
// into multiple ListType output packets using the [begin, end) ranges
|
||||
// specified in SplitVectorCalculatorOptions. If the option "element_only" is
|
||||
// set to true, all ranges should be of size 1 and all outputs will be elements
|
||||
// of type LandmarkType. If "element_only" is false, ranges can be
|
||||
// non-zero in size and all outputs will be of type LandmarkListType.
|
||||
// of type ItemType. If "element_only" is false, ranges can be
|
||||
// non-zero in size and all outputs will be of type ListType.
|
||||
// If the option "combine_outputs" is set to true, only one output stream can be
|
||||
// specified and all ranges of elements will be combined into one
|
||||
// LandmarkListType.
|
||||
template <typename LandmarkType, typename LandmarkListType>
|
||||
class SplitLandmarksCalculator : public CalculatorBase {
|
||||
// ListType.
|
||||
template <typename ItemType, typename ListType>
|
||||
class SplitListsCalculator : public CalculatorBase {
|
||||
public:
|
||||
static absl::Status GetContract(CalculatorContract* cc) {
|
||||
RET_CHECK(cc->Inputs().NumEntries() == 1);
|
||||
RET_CHECK(cc->Outputs().NumEntries() != 0);
|
||||
|
||||
cc->Inputs().Index(0).Set<LandmarkListType>();
|
||||
cc->Inputs().Index(0).Set<ListType>();
|
||||
|
||||
const auto& options =
|
||||
cc->Options<::mediapipe::SplitVectorCalculatorOptions>();
|
||||
|
||||
if (options.combine_outputs()) {
|
||||
RET_CHECK_EQ(cc->Outputs().NumEntries(), 1);
|
||||
cc->Outputs().Index(0).Set<LandmarkListType>();
|
||||
cc->Outputs().Index(0).Set<ListType>();
|
||||
for (int i = 0; i < options.ranges_size() - 1; ++i) {
|
||||
for (int j = i + 1; j < options.ranges_size(); ++j) {
|
||||
const auto& range_0 = options.ranges(i);
|
||||
@@ -82,9 +82,9 @@ class SplitLandmarksCalculator : public CalculatorBase {
|
||||
return absl::InvalidArgumentError(
|
||||
"Since element_only is true, all ranges should be of size 1.");
|
||||
}
|
||||
cc->Outputs().Index(i).Set<LandmarkType>();
|
||||
cc->Outputs().Index(i).Set<ItemType>();
|
||||
} else {
|
||||
cc->Outputs().Index(i).Set<LandmarkListType>();
|
||||
cc->Outputs().Index(i).Set<ListType>();
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -111,39 +111,38 @@ class SplitLandmarksCalculator : public CalculatorBase {
|
||||
}
|
||||
|
||||
absl::Status Process(CalculatorContext* cc) override {
|
||||
const LandmarkListType& input =
|
||||
cc->Inputs().Index(0).Get<LandmarkListType>();
|
||||
RET_CHECK_GE(input.landmark_size(), max_range_end_)
|
||||
<< "Max range end " << max_range_end_ << " exceeds landmarks size "
|
||||
<< input.landmark_size();
|
||||
const ListType& input = cc->Inputs().Index(0).Get<ListType>();
|
||||
RET_CHECK_GE(ListSize(input), max_range_end_)
|
||||
<< "Max range end " << max_range_end_ << " exceeds list size "
|
||||
<< ListSize(input);
|
||||
|
||||
if (combine_outputs_) {
|
||||
LandmarkListType output;
|
||||
ListType output;
|
||||
for (int i = 0; i < ranges_.size(); ++i) {
|
||||
for (int j = ranges_[i].first; j < ranges_[i].second; ++j) {
|
||||
const LandmarkType& input_landmark = input.landmark(j);
|
||||
*output.add_landmark() = input_landmark;
|
||||
const ItemType& input_item = GetItem(input, j);
|
||||
*AddItem(output) = input_item;
|
||||
}
|
||||
}
|
||||
RET_CHECK_EQ(output.landmark_size(), total_elements_);
|
||||
RET_CHECK_EQ(ListSize(output), total_elements_);
|
||||
cc->Outputs().Index(0).AddPacket(
|
||||
MakePacket<LandmarkListType>(output).At(cc->InputTimestamp()));
|
||||
MakePacket<ListType>(output).At(cc->InputTimestamp()));
|
||||
} else {
|
||||
if (element_only_) {
|
||||
for (int i = 0; i < ranges_.size(); ++i) {
|
||||
cc->Outputs().Index(i).AddPacket(
|
||||
MakePacket<LandmarkType>(input.landmark(ranges_[i].first))
|
||||
MakePacket<ItemType>(GetItem(input, ranges_[i].first))
|
||||
.At(cc->InputTimestamp()));
|
||||
}
|
||||
} else {
|
||||
for (int i = 0; i < ranges_.size(); ++i) {
|
||||
LandmarkListType output;
|
||||
ListType output;
|
||||
for (int j = ranges_[i].first; j < ranges_[i].second; ++j) {
|
||||
const LandmarkType& input_landmark = input.landmark(j);
|
||||
*output.add_landmark() = input_landmark;
|
||||
const ItemType& input_item = GetItem(input, j);
|
||||
*AddItem(output) = input_item;
|
||||
}
|
||||
cc->Outputs().Index(i).AddPacket(
|
||||
MakePacket<LandmarkListType>(output).At(cc->InputTimestamp()));
|
||||
MakePacket<ListType>(output).At(cc->InputTimestamp()));
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -151,6 +150,11 @@ class SplitLandmarksCalculator : public CalculatorBase {
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
protected:
|
||||
virtual int ListSize(const ListType& list) const = 0;
|
||||
virtual const ItemType GetItem(const ListType& list, int idx) const = 0;
|
||||
virtual ItemType* AddItem(ListType& list) const = 0;
|
||||
|
||||
private:
|
||||
std::vector<std::pair<int32, int32>> ranges_;
|
||||
int32 max_range_end_ = -1;
|
||||
@@ -159,15 +163,40 @@ class SplitLandmarksCalculator : public CalculatorBase {
|
||||
bool combine_outputs_ = false;
|
||||
};
|
||||
|
||||
typedef SplitLandmarksCalculator<NormalizedLandmark, NormalizedLandmarkList>
|
||||
SplitNormalizedLandmarkListCalculator;
|
||||
// TODO: Move calculators to separate *.cc files
|
||||
|
||||
class SplitNormalizedLandmarkListCalculator
|
||||
: public SplitListsCalculator<NormalizedLandmark, NormalizedLandmarkList> {
|
||||
protected:
|
||||
int ListSize(const NormalizedLandmarkList& list) const override {
|
||||
return list.landmark_size();
|
||||
}
|
||||
const NormalizedLandmark GetItem(const NormalizedLandmarkList& list,
|
||||
int idx) const override {
|
||||
return list.landmark(idx);
|
||||
}
|
||||
NormalizedLandmark* AddItem(NormalizedLandmarkList& list) const override {
|
||||
return list.add_landmark();
|
||||
}
|
||||
};
|
||||
REGISTER_CALCULATOR(SplitNormalizedLandmarkListCalculator);
|
||||
|
||||
typedef SplitLandmarksCalculator<Landmark, LandmarkList>
|
||||
SplitLandmarkListCalculator;
|
||||
class SplitLandmarkListCalculator
|
||||
: public SplitListsCalculator<Landmark, LandmarkList> {
|
||||
protected:
|
||||
int ListSize(const LandmarkList& list) const override {
|
||||
return list.landmark_size();
|
||||
}
|
||||
const Landmark GetItem(const LandmarkList& list, int idx) const override {
|
||||
return list.landmark(idx);
|
||||
}
|
||||
Landmark* AddItem(LandmarkList& list) const override {
|
||||
return list.add_landmark();
|
||||
}
|
||||
};
|
||||
REGISTER_CALCULATOR(SplitLandmarkListCalculator);
|
||||
|
||||
} // namespace mediapipe
|
||||
|
||||
// NOLINTNEXTLINE
|
||||
#endif // MEDIAPIPE_CALCULATORS_CORE_SPLIT_LANDMARKS_CALCULATOR_H_
|
||||
#endif // MEDIAPIPE_CALCULATORS_CORE_SPLIT_PROTO_LIST_CALCULATOR_H_
|
||||
@@ -83,4 +83,7 @@ REGISTER_CALCULATOR(SplitClassificationListVectorCalculator);
|
||||
typedef SplitVectorCalculator<uint64_t, false> SplitUint64tVectorCalculator;
|
||||
REGISTER_CALCULATOR(SplitUint64tVectorCalculator);
|
||||
|
||||
typedef SplitVectorCalculator<float, false> SplitFloatVectorCalculator;
|
||||
REGISTER_CALCULATOR(SplitFloatVectorCalculator);
|
||||
|
||||
} // namespace mediapipe
|
||||
|
||||
@@ -24,7 +24,7 @@
|
||||
|
||||
namespace mediapipe {
|
||||
|
||||
// Calculator that converts a std::string into an integer type, or fails if the
|
||||
// Calculator that converts a string into an integer type, or fails if the
|
||||
// conversion is not possible.
|
||||
//
|
||||
// Example config:
|
||||
@@ -47,7 +47,7 @@ class StringToIntCalculatorTemplate : public CalculatorBase {
|
||||
if (!absl::SimpleAtoi(cc->InputSidePackets().Index(0).Get<std::string>(),
|
||||
&number)) {
|
||||
return absl::InvalidArgumentError(
|
||||
"The std::string could not be parsed as an integer.");
|
||||
"The string could not be parsed as an integer.");
|
||||
}
|
||||
cc->OutputSidePackets().Index(0).Set(MakePacket<IntType>(number));
|
||||
return absl::OkStatus();
|
||||
|
||||
@@ -0,0 +1,32 @@
|
||||
// Copyright 2022 The MediaPipe Authors.
|
||||
//
|
||||
// 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.
|
||||
|
||||
#include "mediapipe/calculators/core/vector_size_calculator.h"
|
||||
|
||||
#include "mediapipe/framework/formats/classification.pb.h"
|
||||
#include "mediapipe/framework/formats/landmark.pb.h"
|
||||
|
||||
namespace mediapipe {
|
||||
namespace api2 {
|
||||
|
||||
using LandmarkListVectorSizeCalculator =
|
||||
VectorSizeCalculator<mediapipe::LandmarkList>;
|
||||
REGISTER_CALCULATOR(LandmarkListVectorSizeCalculator);
|
||||
|
||||
using ClassificationListVectorSizeCalculator =
|
||||
VectorSizeCalculator<mediapipe::ClassificationList>;
|
||||
REGISTER_CALCULATOR(ClassificationListVectorSizeCalculator);
|
||||
|
||||
} // namespace api2
|
||||
} // namespace mediapipe
|
||||
@@ -0,0 +1,64 @@
|
||||
// Copyright 2022 The MediaPipe Authors.
|
||||
//
|
||||
// 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.
|
||||
|
||||
#ifndef MEDIAPIPE_CALCULATORS_CORE_VECTOR_SIZE_CALCULATOR_H_
|
||||
#define MEDIAPIPE_CALCULATORS_CORE_VECTOR_SIZE_CALCULATOR_H_
|
||||
|
||||
#include <optional>
|
||||
|
||||
#include "mediapipe/framework/api2/node.h"
|
||||
#include "mediapipe/framework/calculator_framework.h"
|
||||
#include "mediapipe/framework/port/status.h"
|
||||
|
||||
namespace mediapipe {
|
||||
namespace api2 {
|
||||
|
||||
// A calcutlator to return vector size.
|
||||
//
|
||||
// Inputs:
|
||||
// VECTOR - std::vector<T>
|
||||
// Vector which size to return.
|
||||
//
|
||||
// Outputs:
|
||||
// SIZE - int
|
||||
// Size of the input vector.
|
||||
//
|
||||
// Example config:
|
||||
// node {
|
||||
// calculator: "{SpecificType}VectorSizeCalculator"
|
||||
// input_stream: "VECTOR:vector"
|
||||
// output_stream: "SIZE:vector_size"
|
||||
// }
|
||||
//
|
||||
template <typename T>
|
||||
class VectorSizeCalculator : public Node {
|
||||
public:
|
||||
static constexpr Input<std::vector<T>> kIn{"VECTOR"};
|
||||
static constexpr Output<int> kOut{"SIZE"};
|
||||
|
||||
MEDIAPIPE_NODE_CONTRACT(kIn, kOut);
|
||||
|
||||
absl::Status Process(CalculatorContext* cc) final {
|
||||
if (kIn(cc).IsEmpty()) {
|
||||
return absl::OkStatus();
|
||||
}
|
||||
kOut(cc).Send(kIn(cc).Get().size());
|
||||
return absl::OkStatus();
|
||||
}
|
||||
};
|
||||
|
||||
} // namespace api2
|
||||
} // namespace mediapipe
|
||||
|
||||
#endif // MEDIAPIPE_CALCULATORS_CORE_VECTOR_SIZE_CALCULATOR_H_
|
||||
@@ -239,10 +239,13 @@ cc_library(
|
||||
visibility = ["//visibility:public"],
|
||||
deps = [
|
||||
":image_transformation_calculator_cc_proto",
|
||||
"//mediapipe/framework:packet",
|
||||
"//mediapipe/framework:timestamp",
|
||||
"//mediapipe/gpu:scale_mode_cc_proto",
|
||||
"//mediapipe/framework:calculator_framework",
|
||||
"//mediapipe/framework/formats:image_frame",
|
||||
"//mediapipe/framework/formats:image_frame_opencv",
|
||||
"//mediapipe/framework/formats:video_stream_header",
|
||||
"//mediapipe/framework/port:opencv_core",
|
||||
"//mediapipe/framework/port:opencv_imgproc",
|
||||
"//mediapipe/framework/port:ret_check",
|
||||
@@ -623,11 +626,8 @@ cc_library(
|
||||
"//mediapipe/framework/formats:image_format_cc_proto",
|
||||
"//mediapipe/framework:calculator_framework",
|
||||
"//mediapipe/framework/formats:image_frame",
|
||||
"//mediapipe/framework/formats:image_frame_opencv",
|
||||
"//mediapipe/framework/formats:image",
|
||||
"//mediapipe/framework/formats:image_opencv",
|
||||
"//mediapipe/framework/port:logging",
|
||||
"//mediapipe/framework/port:opencv_core",
|
||||
"//mediapipe/framework/port:status",
|
||||
"//mediapipe/framework/port:vector",
|
||||
] + select({
|
||||
@@ -638,6 +638,13 @@ cc_library(
|
||||
"//mediapipe/gpu:gl_quad_renderer",
|
||||
"//mediapipe/gpu:shader_util",
|
||||
],
|
||||
}) + select({
|
||||
"//mediapipe/framework/port:disable_opencv": [],
|
||||
"//conditions:default": [
|
||||
"//mediapipe/framework/formats:image_frame_opencv",
|
||||
"//mediapipe/framework/formats:image_opencv",
|
||||
"//mediapipe/framework/port:opencv_core",
|
||||
],
|
||||
}),
|
||||
alwayslink = 1,
|
||||
)
|
||||
@@ -724,7 +731,6 @@ cc_library(
|
||||
visibility = ["//visibility:public"],
|
||||
deps = [
|
||||
":affine_transformation",
|
||||
":affine_transformation_runner_opencv",
|
||||
":warp_affine_calculator_cc_proto",
|
||||
"@com_google_absl//absl/status",
|
||||
"@com_google_absl//absl/status:statusor",
|
||||
@@ -742,6 +748,9 @@ cc_library(
|
||||
"//mediapipe/gpu:gpu_buffer",
|
||||
":affine_transformation_runner_gl",
|
||||
],
|
||||
}) + select({
|
||||
"//mediapipe/framework/port:disable_opencv": [],
|
||||
"//conditions:default": [":affine_transformation_runner_opencv"],
|
||||
}),
|
||||
alwayslink = 1,
|
||||
)
|
||||
@@ -796,3 +805,21 @@ cc_test(
|
||||
"@com_google_absl//absl/strings",
|
||||
],
|
||||
)
|
||||
|
||||
cc_library(
|
||||
name = "yuv_to_image_calculator",
|
||||
srcs = ["yuv_to_image_calculator.cc"],
|
||||
visibility = ["//visibility:public"],
|
||||
deps = [
|
||||
"//mediapipe/framework:calculator_context",
|
||||
"//mediapipe/framework:calculator_framework",
|
||||
"//mediapipe/framework/api2:node",
|
||||
"//mediapipe/framework/formats:image",
|
||||
"//mediapipe/framework/formats:image_frame",
|
||||
"//mediapipe/framework/formats:yuv_image",
|
||||
"//third_party/libyuv",
|
||||
"@com_google_absl//absl/status",
|
||||
"@com_google_absl//absl/strings:str_format",
|
||||
],
|
||||
alwayslink = 1,
|
||||
)
|
||||
|
||||
@@ -105,7 +105,7 @@ absl::StatusOr<ImageFileProperties> GetImageFileProperites(
|
||||
} // namespace
|
||||
|
||||
// Calculator to extract EXIF information from an image file. The input is
|
||||
// a std::string containing raw byte data from a file, and the output is an
|
||||
// a string containing raw byte data from a file, and the output is an
|
||||
// ImageFileProperties proto object with the relevant fields filled in.
|
||||
// The calculator accepts the input as a stream or a side packet, and can output
|
||||
// the result as a stream or a side packet. The calculator checks that if an
|
||||
|
||||
@@ -16,10 +16,13 @@
|
||||
#include "mediapipe/framework/calculator_framework.h"
|
||||
#include "mediapipe/framework/formats/image_frame.h"
|
||||
#include "mediapipe/framework/formats/image_frame_opencv.h"
|
||||
#include "mediapipe/framework/formats/video_stream_header.h"
|
||||
#include "mediapipe/framework/packet.h"
|
||||
#include "mediapipe/framework/port/opencv_core_inc.h"
|
||||
#include "mediapipe/framework/port/opencv_imgproc_inc.h"
|
||||
#include "mediapipe/framework/port/ret_check.h"
|
||||
#include "mediapipe/framework/port/status.h"
|
||||
#include "mediapipe/framework/timestamp.h"
|
||||
#include "mediapipe/gpu/scale_mode.pb.h"
|
||||
|
||||
#if !MEDIAPIPE_DISABLE_GPU
|
||||
@@ -52,6 +55,7 @@ namespace mediapipe {
|
||||
namespace {
|
||||
constexpr char kImageFrameTag[] = "IMAGE";
|
||||
constexpr char kGpuBufferTag[] = "IMAGE_GPU";
|
||||
constexpr char kVideoPrestreamTag[] = "VIDEO_PRESTREAM";
|
||||
|
||||
int RotationModeToDegrees(mediapipe::RotationMode_Mode rotation) {
|
||||
switch (rotation) {
|
||||
@@ -122,6 +126,12 @@ mediapipe::ScaleMode_Mode ParseScaleMode(
|
||||
// provided, it overrides the FLIP_VERTICALLY input side packet and/or
|
||||
// corresponding field in the calculator options.
|
||||
//
|
||||
// VIDEO_PRESTREAM (optional): VideoHeader for the input ImageFrames, if
|
||||
// rotating or scaling the frames, the header width and height will be updated
|
||||
// appropriately. Note the header is updated only based on dimensions and
|
||||
// rotations specified as side packets or options, input_stream
|
||||
// transformations will not update the header.
|
||||
//
|
||||
// Output:
|
||||
// One of the following tags:
|
||||
// IMAGE - ImageFrame representing the output image.
|
||||
@@ -242,6 +252,21 @@ absl::Status ImageTransformationCalculator::GetContract(
|
||||
cc->Inputs().Tag("FLIP_VERTICALLY").Set<bool>();
|
||||
}
|
||||
|
||||
RET_CHECK(cc->Inputs().HasTag(kVideoPrestreamTag) ==
|
||||
cc->Outputs().HasTag(kVideoPrestreamTag))
|
||||
<< "If VIDEO_PRESTREAM is provided, it must be provided both as an "
|
||||
"inputs and output stream.";
|
||||
if (cc->Inputs().HasTag(kVideoPrestreamTag)) {
|
||||
RET_CHECK(!(cc->Inputs().HasTag("OUTPUT_DIMENSIONS") ||
|
||||
cc->Inputs().HasTag("ROTATION_DEGREES")))
|
||||
<< "If specifying VIDEO_PRESTREAM, the transformations that affect the "
|
||||
"dimensions of the frames (OUTPUT_DIMENSIONS and ROTATION_DEGREES) "
|
||||
"need to be constant for every frame, meaning they can only be "
|
||||
"provided in the calculator options or side packets.";
|
||||
cc->Inputs().Tag(kVideoPrestreamTag).Set<mediapipe::VideoHeader>();
|
||||
cc->Outputs().Tag(kVideoPrestreamTag).Set<mediapipe::VideoHeader>();
|
||||
}
|
||||
|
||||
if (cc->InputSidePackets().HasTag("OUTPUT_DIMENSIONS")) {
|
||||
cc->InputSidePackets().Tag("OUTPUT_DIMENSIONS").Set<DimensionsPacketType>();
|
||||
}
|
||||
@@ -326,6 +351,24 @@ absl::Status ImageTransformationCalculator::Open(CalculatorContext* cc) {
|
||||
}
|
||||
|
||||
absl::Status ImageTransformationCalculator::Process(CalculatorContext* cc) {
|
||||
// First update the video header if it is given, based on the rotation and
|
||||
// dimensions specified as side packets or options. This will only be done
|
||||
// once, so streaming transformation changes will not be reflected in
|
||||
// the header.
|
||||
if (cc->Inputs().HasTag(kVideoPrestreamTag) &&
|
||||
!cc->Inputs().Tag(kVideoPrestreamTag).IsEmpty() &&
|
||||
cc->Outputs().HasTag(kVideoPrestreamTag)) {
|
||||
mediapipe::VideoHeader header =
|
||||
cc->Inputs().Tag(kVideoPrestreamTag).Get<mediapipe::VideoHeader>();
|
||||
// Update the header's width and height if needed.
|
||||
ComputeOutputDimensions(header.width, header.height, &header.width,
|
||||
&header.height);
|
||||
cc->Outputs()
|
||||
.Tag(kVideoPrestreamTag)
|
||||
.AddPacket(mediapipe::MakePacket<mediapipe::VideoHeader>(header).At(
|
||||
mediapipe::Timestamp::PreStream()));
|
||||
}
|
||||
|
||||
// Override values if specified so.
|
||||
if (cc->Inputs().HasTag("ROTATION_DEGREES") &&
|
||||
!cc->Inputs().Tag("ROTATION_DEGREES").IsEmpty()) {
|
||||
|
||||
@@ -22,9 +22,9 @@
|
||||
|
||||
namespace mediapipe {
|
||||
|
||||
// Takes in an encoded image std::string, decodes it by OpenCV, and converts to
|
||||
// an ImageFrame. Note that this calculator only supports grayscale and RGB
|
||||
// images for now.
|
||||
// Takes in an encoded image string, decodes it by OpenCV, and converts to an
|
||||
// ImageFrame. Note that this calculator only supports grayscale and RGB images
|
||||
// for now.
|
||||
//
|
||||
// Example config:
|
||||
// node {
|
||||
|
||||
@@ -20,8 +20,8 @@
|
||||
|
||||
namespace mediapipe {
|
||||
|
||||
// Takes in a std::string, draws the text std::string by cv::putText(), and
|
||||
// outputs an ImageFrame.
|
||||
// Takes in a string, draws the text string by cv::putText(), and outputs an
|
||||
// ImageFrame.
|
||||
//
|
||||
// Example config:
|
||||
// node {
|
||||
|
||||
@@ -421,6 +421,10 @@ absl::Status ScaleImageCalculator::InitializeFromOptions() {
|
||||
alignment_boundary_ = options_.alignment_boundary();
|
||||
}
|
||||
|
||||
if (options_.has_output_format()) {
|
||||
output_format_ = options_.output_format();
|
||||
}
|
||||
|
||||
downscaler_.reset(new ImageResizer(options_.post_sharpening_coefficient()));
|
||||
|
||||
return absl::OkStatus();
|
||||
@@ -433,13 +437,17 @@ absl::Status ScaleImageCalculator::ValidateImageFormats() const {
|
||||
<< "The output image format was set to UNKNOWN.";
|
||||
// TODO Remove these conditions.
|
||||
RET_CHECK(output_format_ == ImageFormat::SRGB ||
|
||||
output_format_ == ImageFormat::SRGBA ||
|
||||
(input_format_ == output_format_ &&
|
||||
output_format_ == ImageFormat::YCBCR420P))
|
||||
<< "Outputting YCbCr420P images from SRGB input is not yet supported";
|
||||
RET_CHECK(input_format_ == output_format_ ||
|
||||
input_format_ == ImageFormat::YCBCR420P)
|
||||
(input_format_ == ImageFormat::YCBCR420P &&
|
||||
output_format_ == ImageFormat::SRGB) ||
|
||||
(input_format_ == ImageFormat::SRGB &&
|
||||
output_format_ == ImageFormat::SRGBA))
|
||||
<< "Conversion of the color space (except from "
|
||||
"YCbCr420P to SRGB) is not yet supported.";
|
||||
"YCbCr420P to SRGB or SRGB to SRBGA) is not yet supported.";
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
@@ -553,7 +561,6 @@ absl::Status ScaleImageCalculator::Process(CalculatorContext* cc) {
|
||||
}
|
||||
}
|
||||
|
||||
cc->GetCounter("Inputs")->Increment();
|
||||
const ImageFrame* image_frame;
|
||||
ImageFrame converted_image_frame;
|
||||
if (input_format_ == ImageFormat::YCBCR420P) {
|
||||
@@ -605,6 +612,15 @@ absl::Status ScaleImageCalculator::Process(CalculatorContext* cc) {
|
||||
.Add(output_image.release(), cc->InputTimestamp());
|
||||
return absl::OkStatus();
|
||||
}
|
||||
} else if (input_format_ == ImageFormat::SRGB &&
|
||||
output_format_ == ImageFormat::SRGBA) {
|
||||
image_frame = &cc->Inputs().Get(input_data_id_).Get<ImageFrame>();
|
||||
cv::Mat input_mat = ::mediapipe::formats::MatView(image_frame);
|
||||
converted_image_frame.Reset(ImageFormat::SRGBA, image_frame->Width(),
|
||||
image_frame->Height(), alignment_boundary_);
|
||||
cv::Mat output_mat = ::mediapipe::formats::MatView(&converted_image_frame);
|
||||
cv::cvtColor(input_mat, output_mat, cv::COLOR_RGB2RGBA, 4);
|
||||
image_frame = &converted_image_frame;
|
||||
} else {
|
||||
image_frame = &cc->Inputs().Get(input_data_id_).Get<ImageFrame>();
|
||||
MP_RETURN_IF_ERROR(ValidateImageFrame(cc, *image_frame));
|
||||
|
||||
@@ -21,10 +21,7 @@
|
||||
#include "mediapipe/framework/formats/image.h"
|
||||
#include "mediapipe/framework/formats/image_format.pb.h"
|
||||
#include "mediapipe/framework/formats/image_frame.h"
|
||||
#include "mediapipe/framework/formats/image_frame_opencv.h"
|
||||
#include "mediapipe/framework/formats/image_opencv.h"
|
||||
#include "mediapipe/framework/port/logging.h"
|
||||
#include "mediapipe/framework/port/opencv_core_inc.h"
|
||||
#include "mediapipe/framework/port/status.h"
|
||||
#include "mediapipe/framework/port/vector.h"
|
||||
|
||||
@@ -34,6 +31,12 @@
|
||||
#include "mediapipe/gpu/shader_util.h"
|
||||
#endif // !MEDIAPIPE_DISABLE_GPU
|
||||
|
||||
#if !MEDIAPIPE_DISABLE_OPENCV
|
||||
#include "mediapipe/framework/formats/image_frame_opencv.h"
|
||||
#include "mediapipe/framework/formats/image_opencv.h"
|
||||
#include "mediapipe/framework/port/opencv_core_inc.h"
|
||||
#endif // !MEDIAPIPE_DISABLE_OPENCV
|
||||
|
||||
namespace mediapipe {
|
||||
|
||||
namespace {
|
||||
@@ -163,7 +166,11 @@ absl::Status SegmentationSmoothingCalculator::Process(CalculatorContext* cc) {
|
||||
return absl::InternalError("GPU processing is disabled.");
|
||||
#endif // !MEDIAPIPE_DISABLE_GPU
|
||||
} else {
|
||||
#if !MEDIAPIPE_DISABLE_OPENCV
|
||||
MP_RETURN_IF_ERROR(RenderCpu(cc));
|
||||
#else
|
||||
return absl::InternalError("OpenCV processing is disabled.");
|
||||
#endif // !MEDIAPIPE_DISABLE_OPENCV
|
||||
}
|
||||
|
||||
return absl::OkStatus();
|
||||
@@ -181,24 +188,25 @@ absl::Status SegmentationSmoothingCalculator::Close(CalculatorContext* cc) {
|
||||
}
|
||||
|
||||
absl::Status SegmentationSmoothingCalculator::RenderCpu(CalculatorContext* cc) {
|
||||
#if !MEDIAPIPE_DISABLE_OPENCV
|
||||
// Setup source images.
|
||||
const auto& current_frame = cc->Inputs().Tag(kCurrentMaskTag).Get<Image>();
|
||||
const cv::Mat current_mat = mediapipe::formats::MatView(¤t_frame);
|
||||
RET_CHECK_EQ(current_mat.type(), CV_32FC1)
|
||||
auto current_mat = mediapipe::formats::MatView(¤t_frame);
|
||||
RET_CHECK_EQ(current_mat->type(), CV_32FC1)
|
||||
<< "Only 1-channel float input image is supported.";
|
||||
|
||||
const auto& previous_frame = cc->Inputs().Tag(kPreviousMaskTag).Get<Image>();
|
||||
const cv::Mat previous_mat = mediapipe::formats::MatView(&previous_frame);
|
||||
RET_CHECK_EQ(previous_mat.type(), current_mat.type())
|
||||
<< "Warning: mixing input format types: " << previous_mat.type()
|
||||
<< " != " << previous_mat.type();
|
||||
auto previous_mat = mediapipe::formats::MatView(&previous_frame);
|
||||
RET_CHECK_EQ(previous_mat->type(), current_mat->type())
|
||||
<< "Warning: mixing input format types: " << previous_mat->type()
|
||||
<< " != " << previous_mat->type();
|
||||
|
||||
RET_CHECK_EQ(current_mat.rows, previous_mat.rows);
|
||||
RET_CHECK_EQ(current_mat.cols, previous_mat.cols);
|
||||
RET_CHECK_EQ(current_mat->rows, previous_mat->rows);
|
||||
RET_CHECK_EQ(current_mat->cols, previous_mat->cols);
|
||||
|
||||
// Setup destination image.
|
||||
auto output_frame = std::make_shared<ImageFrame>(
|
||||
current_frame.image_format(), current_mat.cols, current_mat.rows);
|
||||
current_frame.image_format(), current_mat->cols, current_mat->rows);
|
||||
cv::Mat output_mat = mediapipe::formats::MatView(output_frame.get());
|
||||
output_mat.setTo(cv::Scalar(0));
|
||||
|
||||
@@ -233,8 +241,8 @@ absl::Status SegmentationSmoothingCalculator::RenderCpu(CalculatorContext* cc) {
|
||||
// Write directly to the first channel of output.
|
||||
for (int i = 0; i < output_mat.rows; ++i) {
|
||||
float* out_ptr = output_mat.ptr<float>(i);
|
||||
const float* curr_ptr = current_mat.ptr<float>(i);
|
||||
const float* prev_ptr = previous_mat.ptr<float>(i);
|
||||
const float* curr_ptr = current_mat->ptr<float>(i);
|
||||
const float* prev_ptr = previous_mat->ptr<float>(i);
|
||||
for (int j = 0; j < output_mat.cols; ++j) {
|
||||
const float new_mask_value = curr_ptr[j];
|
||||
const float prev_mask_value = prev_ptr[j];
|
||||
@@ -245,6 +253,7 @@ absl::Status SegmentationSmoothingCalculator::RenderCpu(CalculatorContext* cc) {
|
||||
cc->Outputs()
|
||||
.Tag(kOutputMaskTag)
|
||||
.AddPacket(MakePacket<Image>(output_frame).At(cc->InputTimestamp()));
|
||||
#endif // !MEDIAPIPE_DISABLE_OPENCV
|
||||
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
@@ -116,8 +116,8 @@ void RunGraph(Packet curr_packet, Packet prev_packet, bool use_gpu, float ratio,
|
||||
ASSERT_EQ(1, output_packets.size());
|
||||
|
||||
Image result_image = output_packets[0].Get<Image>();
|
||||
cv::Mat result_mat = formats::MatView(&result_image);
|
||||
result_mat.copyTo(*result);
|
||||
auto result_mat = formats::MatView(&result_image);
|
||||
result_mat->copyTo(*result);
|
||||
|
||||
// Fully close graph at end, otherwise calculator+Images are destroyed
|
||||
// after calling WaitUntilDone().
|
||||
@@ -135,10 +135,10 @@ void RunTest(bool use_gpu, float mix_ratio, cv::Mat& test_result) {
|
||||
|
||||
Packet curr_packet = MakePacket<Image>(std::make_unique<ImageFrame>(
|
||||
ImageFormat::VEC32F1, curr_mat.size().width, curr_mat.size().height));
|
||||
curr_mat.copyTo(formats::MatView(&(curr_packet.Get<Image>())));
|
||||
curr_mat.copyTo(*formats::MatView(&(curr_packet.Get<Image>())));
|
||||
Packet prev_packet = MakePacket<Image>(std::make_unique<ImageFrame>(
|
||||
ImageFormat::VEC32F1, prev_mat.size().width, prev_mat.size().height));
|
||||
prev_mat.copyTo(formats::MatView(&(prev_packet.Get<Image>())));
|
||||
prev_mat.copyTo(*formats::MatView(&(prev_packet.Get<Image>())));
|
||||
|
||||
cv::Mat result;
|
||||
RunGraph(curr_packet, prev_packet, use_gpu, mix_ratio, &result);
|
||||
|
||||
@@ -24,7 +24,9 @@
|
||||
#endif // !MEDIAPIPE_DISABLE_GPU
|
||||
#include "absl/status/status.h"
|
||||
#include "absl/status/statusor.h"
|
||||
#if !MEDIAPIPE_DISABLE_OPENCV
|
||||
#include "mediapipe/calculators/image/affine_transformation_runner_opencv.h"
|
||||
#endif // !MEDIAPIPE_DISABLE_OPENCV
|
||||
#include "mediapipe/calculators/image/warp_affine_calculator.pb.h"
|
||||
#include "mediapipe/framework/api2/node.h"
|
||||
#include "mediapipe/framework/calculator_framework.h"
|
||||
@@ -54,6 +56,7 @@ AffineTransformation::BorderMode GetBorderMode(
|
||||
template <typename ImageT>
|
||||
class WarpAffineRunnerHolder {};
|
||||
|
||||
#if !MEDIAPIPE_DISABLE_OPENCV
|
||||
template <>
|
||||
class WarpAffineRunnerHolder<ImageFrame> {
|
||||
public:
|
||||
@@ -69,6 +72,7 @@ class WarpAffineRunnerHolder<ImageFrame> {
|
||||
private:
|
||||
std::unique_ptr<RunnerType> runner_;
|
||||
};
|
||||
#endif // !MEDIAPIPE_DISABLE_OPENCV
|
||||
|
||||
#if !MEDIAPIPE_DISABLE_GPU
|
||||
template <>
|
||||
@@ -113,7 +117,9 @@ class WarpAffineRunnerHolder<mediapipe::Image> {
|
||||
mediapipe::Image> {
|
||||
public:
|
||||
absl::Status Open(CalculatorContext* cc) {
|
||||
#if !MEDIAPIPE_DISABLE_OPENCV
|
||||
MP_RETURN_IF_ERROR(cpu_holder_.Open(cc));
|
||||
#endif // !MEDIAPIPE_DISABLE_OPENCV
|
||||
#if !MEDIAPIPE_DISABLE_GPU
|
||||
MP_RETURN_IF_ERROR(gpu_holder_.Open(cc));
|
||||
#endif // !MEDIAPIPE_DISABLE_GPU
|
||||
@@ -133,20 +139,26 @@ class WarpAffineRunnerHolder<mediapipe::Image> {
|
||||
return absl::UnavailableError("GPU support is disabled");
|
||||
#endif // !MEDIAPIPE_DISABLE_GPU
|
||||
}
|
||||
#if !MEDIAPIPE_DISABLE_OPENCV
|
||||
ASSIGN_OR_RETURN(auto* runner, cpu_holder_.GetRunner());
|
||||
const auto& frame_ptr = input.GetImageFrameSharedPtr();
|
||||
// Wrap image into image frame.
|
||||
const ImageFrame image_frame(frame_ptr->Format(), frame_ptr->Width(),
|
||||
frame_ptr->Height(), frame_ptr->WidthStep(),
|
||||
const_cast<uint8_t*>(frame_ptr->PixelData()),
|
||||
[](uint8* data) {});
|
||||
[](uint8* data){});
|
||||
ASSIGN_OR_RETURN(auto result,
|
||||
runner->Run(image_frame, matrix, size, border_mode));
|
||||
return mediapipe::Image(std::make_shared<ImageFrame>(std::move(result)));
|
||||
#else
|
||||
return absl::UnavailableError("OpenCV support is disabled");
|
||||
#endif // !MEDIAPIPE_DISABLE_OPENCV
|
||||
}
|
||||
|
||||
private:
|
||||
#if !MEDIAPIPE_DISABLE_OPENCV
|
||||
WarpAffineRunnerHolder<ImageFrame> cpu_holder_;
|
||||
#endif // !MEDIAPIPE_DISABLE_OPENCV
|
||||
#if !MEDIAPIPE_DISABLE_GPU
|
||||
WarpAffineRunnerHolder<mediapipe::GpuBuffer> gpu_holder_;
|
||||
#endif // !MEDIAPIPE_DISABLE_GPU
|
||||
@@ -200,8 +212,10 @@ class WarpAffineCalculatorImpl : public mediapipe::api2::NodeImpl<InterfaceT> {
|
||||
|
||||
} // namespace
|
||||
|
||||
#if !MEDIAPIPE_DISABLE_OPENCV
|
||||
MEDIAPIPE_NODE_IMPLEMENTATION(
|
||||
WarpAffineCalculatorImpl<WarpAffineCalculatorCpu>);
|
||||
#endif // !MEDIAPIPE_DISABLE_OPENCV
|
||||
#if !MEDIAPIPE_DISABLE_GPU
|
||||
MEDIAPIPE_NODE_IMPLEMENTATION(
|
||||
WarpAffineCalculatorImpl<WarpAffineCalculatorGpu>);
|
||||
|
||||
@@ -70,11 +70,13 @@ class WarpAffineCalculatorIntf : public mediapipe::api2::NodeIntf {
|
||||
static constexpr mediapipe::api2::Output<ImageT> kOutImage{"IMAGE"};
|
||||
};
|
||||
|
||||
#if !MEDIAPIPE_DISABLE_OPENCV
|
||||
class WarpAffineCalculatorCpu : public WarpAffineCalculatorIntf<ImageFrame> {
|
||||
public:
|
||||
MEDIAPIPE_NODE_INTERFACE(WarpAffineCalculatorCpu, kInImage, kMatrix,
|
||||
kOutputSize, kOutImage);
|
||||
};
|
||||
#endif // !MEDIAPIPE_DISABLE_OPENCV
|
||||
#if !MEDIAPIPE_DISABLE_GPU
|
||||
class WarpAffineCalculatorGpu
|
||||
: public WarpAffineCalculatorIntf<mediapipe::GpuBuffer> {
|
||||
|
||||
@@ -0,0 +1,123 @@
|
||||
// Copyright 2022 The MediaPipe Authors.
|
||||
//
|
||||
// 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.
|
||||
|
||||
#include <memory>
|
||||
#include <string>
|
||||
#include <utility>
|
||||
|
||||
#include "absl/status/status.h"
|
||||
#include "absl/strings/str_format.h"
|
||||
#include "libyuv/convert_argb.h"
|
||||
#include "libyuv/video_common.h"
|
||||
#include "mediapipe/framework/api2/node.h"
|
||||
#include "mediapipe/framework/calculator_context.h"
|
||||
#include "mediapipe/framework/calculator_framework.h"
|
||||
#include "mediapipe/framework/formats/image.h"
|
||||
#include "mediapipe/framework/formats/image_frame.h"
|
||||
#include "mediapipe/framework/formats/yuv_image.h"
|
||||
|
||||
namespace mediapipe {
|
||||
namespace api2 {
|
||||
|
||||
namespace {
|
||||
|
||||
// Utility function to convert FourCC enum to string, for error messages.
|
||||
std::string FourCCToString(libyuv::FourCC fourcc) {
|
||||
char buf[5];
|
||||
buf[0] = (fourcc >> 24) & 0xff;
|
||||
buf[1] = (fourcc >> 16) & 0xff;
|
||||
buf[2] = (fourcc >> 8) & 0xff;
|
||||
buf[3] = (fourcc)&0xff;
|
||||
buf[4] = 0;
|
||||
return std::string(buf);
|
||||
}
|
||||
} // namespace
|
||||
|
||||
// Converts a `YUVImage` into an RGB `Image` using libyuv.
|
||||
//
|
||||
// The input `YUVImage` is expected to be in the NV12, NV21, YV12 or I420 (aka
|
||||
// YV21) format (as per the `fourcc()` property). This covers the most commonly
|
||||
// used YUV image formats used on mobile devices. Other formats are not
|
||||
// supported and wil result in an `InvalidArgumentError`.
|
||||
class YUVToImageCalculator : public Node {
|
||||
public:
|
||||
static constexpr Input<YUVImage> kInput{"YUV_IMAGE"};
|
||||
static constexpr Output<Image> kOutput{"IMAGE"};
|
||||
|
||||
MEDIAPIPE_NODE_CONTRACT(kInput, kOutput);
|
||||
|
||||
absl::Status Process(CalculatorContext* cc) override {
|
||||
const auto& yuv_image = *kInput(cc);
|
||||
// Check that the format is supported.
|
||||
auto format = yuv_image.fourcc();
|
||||
if (format != libyuv::FOURCC_NV12 && format != libyuv::FOURCC_NV21 &&
|
||||
format != libyuv::FOURCC_YV12 && format != libyuv::FOURCC_I420) {
|
||||
return absl::InvalidArgumentError(
|
||||
absl::StrFormat("Unsupported YUVImage format: %s. Only NV12, NV21, "
|
||||
"YV12 and I420 (aka YV21) are supported.",
|
||||
FourCCToString(format)));
|
||||
}
|
||||
// Build a transient ImageFrameSharedPtr with default alignment to host
|
||||
// conversion results.
|
||||
ImageFrameSharedPtr image_frame = std::make_shared<ImageFrame>(
|
||||
ImageFormat::SRGB, yuv_image.width(), yuv_image.height());
|
||||
// Perform actual conversion.
|
||||
switch (format) {
|
||||
case libyuv::FOURCC_NV12:
|
||||
// 8-bit Y plane followed by an interleaved 8-bit U/V plane with 2×2
|
||||
// subsampling.
|
||||
libyuv::NV12ToRAW(
|
||||
yuv_image.data(0), yuv_image.stride(0), yuv_image.data(1),
|
||||
yuv_image.stride(1), image_frame->MutablePixelData(),
|
||||
image_frame->WidthStep(), yuv_image.width(), yuv_image.height());
|
||||
break;
|
||||
case libyuv::FOURCC_NV21:
|
||||
// 8-bit Y plane followed by an interleaved 8-bit V/U plane with 2×2
|
||||
// subsampling.
|
||||
libyuv::NV21ToRAW(
|
||||
yuv_image.data(0), yuv_image.stride(0), yuv_image.data(1),
|
||||
yuv_image.stride(1), image_frame->MutablePixelData(),
|
||||
image_frame->WidthStep(), yuv_image.width(), yuv_image.height());
|
||||
break;
|
||||
case libyuv::FOURCC_I420:
|
||||
// Also known as YV21.
|
||||
// 8-bit Y plane followed by 8-bit 2×2 subsampled U and V planes.
|
||||
libyuv::I420ToRAW(
|
||||
yuv_image.data(0), yuv_image.stride(0), yuv_image.data(1),
|
||||
yuv_image.stride(1), yuv_image.data(2), yuv_image.stride(2),
|
||||
image_frame->MutablePixelData(), image_frame->WidthStep(),
|
||||
yuv_image.width(), yuv_image.height());
|
||||
break;
|
||||
case libyuv::FOURCC_YV12:
|
||||
// 8-bit Y plane followed by 8-bit 2×2 subsampled V and U planes.
|
||||
libyuv::I420ToRAW(
|
||||
yuv_image.data(0), yuv_image.stride(0), yuv_image.data(2),
|
||||
yuv_image.stride(2), yuv_image.data(1), yuv_image.stride(1),
|
||||
image_frame->MutablePixelData(), image_frame->WidthStep(),
|
||||
yuv_image.width(), yuv_image.height());
|
||||
break;
|
||||
default:
|
||||
// This should never happen (caught by checks above).
|
||||
return absl::InternalError("Unsupported YUVImage format.");
|
||||
}
|
||||
// Finally, build and send an Image object that takes ownership of the
|
||||
// transient ImageFrameSharedPtr object.
|
||||
kOutput(cc).Send(std::make_unique<Image>(std::move(image_frame)));
|
||||
return absl::OkStatus();
|
||||
}
|
||||
};
|
||||
MEDIAPIPE_REGISTER_NODE(YUVToImageCalculator);
|
||||
|
||||
} // namespace api2
|
||||
} // namespace mediapipe
|
||||
@@ -28,7 +28,9 @@ package(default_visibility = ["//visibility:private"])
|
||||
|
||||
exports_files(
|
||||
glob(["testdata/image_to_tensor/*"]),
|
||||
visibility = ["//mediapipe/calculators/image:__subpackages__"],
|
||||
visibility = [
|
||||
"//mediapipe/calculators/image:__subpackages__",
|
||||
],
|
||||
)
|
||||
|
||||
selects.config_setting_group(
|
||||
@@ -38,6 +40,63 @@ selects.config_setting_group(
|
||||
],
|
||||
)
|
||||
|
||||
mediapipe_proto_library(
|
||||
name = "audio_to_tensor_calculator_proto",
|
||||
srcs = ["audio_to_tensor_calculator.proto"],
|
||||
visibility = [
|
||||
"//mediapipe/framework:mediapipe_internal",
|
||||
],
|
||||
deps = [
|
||||
"//mediapipe/framework:calculator_options_proto",
|
||||
"//mediapipe/framework:calculator_proto",
|
||||
],
|
||||
)
|
||||
|
||||
cc_library(
|
||||
name = "audio_to_tensor_calculator",
|
||||
srcs = ["audio_to_tensor_calculator.cc"],
|
||||
visibility = [
|
||||
"//mediapipe/framework:mediapipe_internal",
|
||||
],
|
||||
deps = [
|
||||
":audio_to_tensor_calculator_cc_proto",
|
||||
"//mediapipe/framework:calculator_framework",
|
||||
"//mediapipe/framework/api2:node",
|
||||
"//mediapipe/framework/api2:packet",
|
||||
"//mediapipe/framework/api2:port",
|
||||
"//mediapipe/framework/formats:matrix",
|
||||
"//mediapipe/framework/formats:tensor",
|
||||
"//mediapipe/framework/formats:time_series_header_cc_proto",
|
||||
"//mediapipe/util:time_series_util",
|
||||
"@com_google_absl//absl/memory",
|
||||
"@com_google_absl//absl/status",
|
||||
"@com_google_absl//absl/status:statusor",
|
||||
"@com_google_absl//absl/strings:str_format",
|
||||
"@com_google_audio_tools//audio/dsp:resampler_q",
|
||||
"@org_tensorflow//tensorflow/lite/c:common",
|
||||
],
|
||||
alwayslink = 1,
|
||||
)
|
||||
|
||||
cc_test(
|
||||
name = "audio_to_tensor_calculator_test",
|
||||
srcs = ["audio_to_tensor_calculator_test.cc"],
|
||||
deps = [
|
||||
":audio_to_tensor_calculator",
|
||||
"//mediapipe/framework:calculator_cc_proto",
|
||||
"//mediapipe/framework:calculator_framework",
|
||||
"//mediapipe/framework:timestamp",
|
||||
"//mediapipe/framework/api2:packet",
|
||||
"//mediapipe/framework/formats:matrix",
|
||||
"//mediapipe/framework/formats:tensor",
|
||||
"//mediapipe/framework/port:gtest_main",
|
||||
"//mediapipe/framework/port:parse_text_proto",
|
||||
"@com_google_absl//absl/strings",
|
||||
"@com_google_audio_tools//audio/dsp:resampler_q",
|
||||
"@org_tensorflow//tensorflow/lite/c:common",
|
||||
],
|
||||
)
|
||||
|
||||
mediapipe_proto_library(
|
||||
name = "inference_calculator_proto",
|
||||
srcs = ["inference_calculator.proto"],
|
||||
@@ -48,6 +107,14 @@ mediapipe_proto_library(
|
||||
],
|
||||
)
|
||||
|
||||
# This target defines the "InferenceCalculator" component, which looks for the available concrete
|
||||
# implementations linked into the current binary and picks the one to use.
|
||||
# You can depend on :inference_calculator instead if you want to automatically include a default
|
||||
# set of implementations tailored for the current build configuration.
|
||||
# If you want to have precise control of which implementations to include (e.g. for strict binary
|
||||
# size concerns), depend on those implementations directly, and do not depend on
|
||||
# :inference_calculator.
|
||||
# In all cases, use "InferenceCalulator" in your graphs.
|
||||
cc_library(
|
||||
name = "inference_calculator_interface",
|
||||
srcs = ["inference_calculator.cc"],
|
||||
@@ -60,19 +127,21 @@ cc_library(
|
||||
],
|
||||
"//conditions:default": [],
|
||||
}),
|
||||
visibility = ["//visibility:public"],
|
||||
deps = [
|
||||
":inference_calculator_cc_proto",
|
||||
":inference_calculator_options_lib",
|
||||
"//mediapipe/framework:calculator_framework",
|
||||
"//mediapipe/framework/api2:node",
|
||||
"//mediapipe/framework/api2:packet",
|
||||
"//mediapipe/framework/formats:tensor",
|
||||
"//mediapipe/framework/port:ret_check",
|
||||
"//mediapipe/framework/stream_handler:fixed_size_input_stream_handler",
|
||||
"//mediapipe/framework/tool:subgraph_expansion",
|
||||
"//mediapipe/util/tflite:config",
|
||||
"//mediapipe/util/tflite:tflite_model_loader",
|
||||
"@com_google_absl//absl/memory",
|
||||
"@com_google_absl//absl/strings",
|
||||
"@org_tensorflow//tensorflow/lite:framework",
|
||||
"@org_tensorflow//tensorflow/lite/core/api:op_resolver",
|
||||
"@org_tensorflow//tensorflow/lite/kernels:builtin_ops",
|
||||
],
|
||||
alwayslink = 1,
|
||||
@@ -82,16 +151,31 @@ cc_library(
|
||||
name = "inference_calculator_gl",
|
||||
srcs = ["inference_calculator_gl.cc"],
|
||||
tags = ["nomac"], # config problem with cpuinfo via TF
|
||||
visibility = ["//visibility:public"],
|
||||
deps = [
|
||||
"inference_calculator_interface",
|
||||
":inference_calculator_interface",
|
||||
"//mediapipe/gpu:gl_calculator_helper",
|
||||
"//mediapipe/gpu:gpu_buffer",
|
||||
"//mediapipe/util/tflite:tflite_gpu_runner",
|
||||
"@com_google_absl//absl/memory",
|
||||
"@com_google_absl//absl/status",
|
||||
"@org_tensorflow//tensorflow/lite:framework_stable",
|
||||
"@org_tensorflow//tensorflow/lite/delegates/gpu:gl_delegate",
|
||||
"@org_tensorflow//tensorflow/lite/delegates/gpu/common:shape",
|
||||
"@org_tensorflow//tensorflow/lite/delegates/gpu/gl:gl_buffer",
|
||||
"@org_tensorflow//tensorflow/lite/delegates/gpu/gl:gl_program",
|
||||
"@org_tensorflow//tensorflow/lite/delegates/gpu/gl:gl_shader",
|
||||
],
|
||||
alwayslink = 1,
|
||||
)
|
||||
|
||||
cc_library(
|
||||
name = "inference_calculator_gl_advanced",
|
||||
srcs = ["inference_calculator_gl_advanced.cc"],
|
||||
tags = ["nomac"],
|
||||
visibility = ["//visibility:public"],
|
||||
deps = [
|
||||
":inference_calculator_interface",
|
||||
"//mediapipe/framework/deps:file_path",
|
||||
"//mediapipe/gpu:gl_calculator_helper",
|
||||
"//mediapipe/util/tflite:tflite_gpu_runner",
|
||||
"@com_google_absl//absl/memory",
|
||||
"@com_google_absl//absl/status",
|
||||
"@org_tensorflow//tensorflow/lite:framework_stable",
|
||||
],
|
||||
alwayslink = 1,
|
||||
)
|
||||
@@ -108,6 +192,7 @@ cc_library(
|
||||
"-framework MetalKit",
|
||||
],
|
||||
tags = ["ios"],
|
||||
visibility = ["//visibility:public"],
|
||||
deps = [
|
||||
"inference_calculator_interface",
|
||||
"//mediapipe/gpu:MPPMetalHelper",
|
||||
@@ -137,10 +222,13 @@ cc_library(
|
||||
],
|
||||
"//conditions:default": [],
|
||||
}),
|
||||
visibility = ["//visibility:public"],
|
||||
deps = [
|
||||
":inference_calculator_interface",
|
||||
"@com_google_absl//absl/memory",
|
||||
"@org_tensorflow//tensorflow/lite/delegates/xnnpack:xnnpack_delegate",
|
||||
"@org_tensorflow//tensorflow/lite:framework_stable",
|
||||
"@org_tensorflow//tensorflow/lite/c:c_api_types",
|
||||
] + select({
|
||||
"//conditions:default": [
|
||||
"//mediapipe/util:cpu_util",
|
||||
@@ -154,9 +242,13 @@ cc_library(
|
||||
|
||||
cc_library(
|
||||
name = "inference_calculator_gl_if_compute_shader_available",
|
||||
deps = select({
|
||||
visibility = ["//visibility:public"],
|
||||
deps = selects.with_or({
|
||||
":compute_shader_unavailable": [],
|
||||
"//conditions:default": [":inference_calculator_gl"],
|
||||
"//conditions:default": [
|
||||
":inference_calculator_gl",
|
||||
":inference_calculator_gl_advanced",
|
||||
],
|
||||
}),
|
||||
)
|
||||
|
||||
@@ -303,7 +395,7 @@ cc_library(
|
||||
"//mediapipe/framework/formats:tensor",
|
||||
"//mediapipe/framework/formats/object_detection:anchor_cc_proto",
|
||||
"//mediapipe/framework/port:ret_check",
|
||||
] + select({
|
||||
] + selects.with_or({
|
||||
":compute_shader_unavailable": [],
|
||||
"//conditions:default": [":tensors_to_detections_calculator_gpu_deps"],
|
||||
}),
|
||||
@@ -477,6 +569,7 @@ cc_library(
|
||||
"//mediapipe/framework/formats:location",
|
||||
"//mediapipe/framework/port:ret_check",
|
||||
"//mediapipe/framework/formats:tensor",
|
||||
"//mediapipe/util:label_map_cc_proto",
|
||||
"//mediapipe/util:resource_util",
|
||||
] + select({
|
||||
"//mediapipe:android": [
|
||||
@@ -499,6 +592,7 @@ mediapipe_proto_library(
|
||||
deps = [
|
||||
"//mediapipe/framework:calculator_options_proto",
|
||||
"//mediapipe/framework:calculator_proto",
|
||||
"//mediapipe/util:label_map_proto",
|
||||
],
|
||||
)
|
||||
|
||||
@@ -560,7 +654,7 @@ cc_library(
|
||||
|
||||
cc_library(
|
||||
name = "image_to_tensor_calculator_gpu_deps",
|
||||
deps = select({
|
||||
deps = selects.with_or({
|
||||
"//mediapipe:android": [
|
||||
":image_to_tensor_converter_gl_buffer",
|
||||
"//mediapipe/gpu:gl_calculator_helper",
|
||||
@@ -665,6 +759,7 @@ cc_library(
|
||||
],
|
||||
"//conditions:default": [],
|
||||
}),
|
||||
visibility = ["//visibility:public"],
|
||||
deps = [
|
||||
":image_to_tensor_converter",
|
||||
":image_to_tensor_utils",
|
||||
@@ -684,7 +779,7 @@ cc_library(
|
||||
name = "image_to_tensor_converter_gl_buffer",
|
||||
srcs = ["image_to_tensor_converter_gl_buffer.cc"],
|
||||
hdrs = ["image_to_tensor_converter_gl_buffer.h"],
|
||||
deps = ["//mediapipe/framework:port"] + select({
|
||||
deps = ["//mediapipe/framework:port"] + selects.with_or({
|
||||
"//mediapipe:apple": [],
|
||||
"//conditions:default": [
|
||||
":image_to_tensor_converter",
|
||||
@@ -851,9 +946,7 @@ cc_library(
|
||||
"@com_google_absl//absl/types:span",
|
||||
"//mediapipe/framework/formats:image",
|
||||
"//mediapipe/framework/formats:image_frame",
|
||||
"//mediapipe/framework/formats:image_opencv",
|
||||
"//mediapipe/framework/formats:tensor",
|
||||
"//mediapipe/framework/port:opencv_imgproc",
|
||||
"//mediapipe/framework/port:ret_check",
|
||||
"//mediapipe/framework:calculator_context",
|
||||
"//mediapipe/framework:calculator_framework",
|
||||
@@ -883,6 +976,12 @@ cc_library(
|
||||
"@org_tensorflow//tensorflow/lite/delegates/gpu/gl:gl_texture",
|
||||
"@org_tensorflow//tensorflow/lite/delegates/gpu/gl/converters:util",
|
||||
],
|
||||
}) + select({
|
||||
"//mediapipe/framework/port:disable_opencv": [],
|
||||
"//conditions:default": [
|
||||
"//mediapipe/framework/formats:image_opencv",
|
||||
"//mediapipe/framework/port:opencv_imgproc",
|
||||
],
|
||||
}),
|
||||
alwayslink = 1,
|
||||
)
|
||||
|
||||
@@ -0,0 +1,401 @@
|
||||
// Copyright 2022 The MediaPipe Authors.
|
||||
//
|
||||
// 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.
|
||||
|
||||
#include <math.h>
|
||||
|
||||
#include <algorithm>
|
||||
#include <cstring>
|
||||
#include <memory>
|
||||
#include <string>
|
||||
#include <utility>
|
||||
#include <vector>
|
||||
|
||||
#include "absl/memory/memory.h"
|
||||
#include "absl/status/status.h"
|
||||
#include "absl/status/statusor.h"
|
||||
#include "absl/strings/str_format.h"
|
||||
#include "audio/dsp/resampler_q.h"
|
||||
#include "mediapipe/calculators/tensor/audio_to_tensor_calculator.pb.h"
|
||||
#include "mediapipe/framework/api2/node.h"
|
||||
#include "mediapipe/framework/api2/packet.h"
|
||||
#include "mediapipe/framework/api2/port.h"
|
||||
#include "mediapipe/framework/calculator_framework.h"
|
||||
#include "mediapipe/framework/formats/matrix.h"
|
||||
#include "mediapipe/framework/formats/tensor.h"
|
||||
#include "mediapipe/framework/formats/time_series_header.pb.h"
|
||||
#include "mediapipe/util/time_series_util.h"
|
||||
|
||||
namespace mediapipe {
|
||||
namespace api2 {
|
||||
|
||||
// Converts audio buffers into tensors, possibly with resampling, buffering
|
||||
// and framing, according to specified inputs and options. All input audio
|
||||
// buffers will be first resampled from the input sample rate to the target
|
||||
// sample rate if they are not equal. The resampled audio data (with the
|
||||
// buffered samples from the previous runs in the streaming mode) will be broken
|
||||
// into fixed-sized, possibly overlapping frames. Finally, all frames will be
|
||||
// converted to and outputted as MediaPipe Tensors. The last output tensor will
|
||||
// be zero-padding if the remaining samples are insufficient.
|
||||
//
|
||||
// This calculator assumes that the input timestamps refer to the first
|
||||
// sample in each Matrix. The output timestamps follow this same convention.
|
||||
// One Process() call may output multiple tensors packets. The timestamps of
|
||||
// the output packets are determined by the timestamp of the previous output
|
||||
// packet, the target sample rate, and the number of samples advanced after the
|
||||
// previous output.
|
||||
//
|
||||
// The calculator has two running modes:
|
||||
// Streaming mode: when "streaming_mode" is set to true in the calculator
|
||||
// options, the calculator treats the input audio stream as a continuous
|
||||
// stream. Thus, any samples that are not consumed in the previous runs will
|
||||
// be cached in a global sample buffer. The audio data resampled from the
|
||||
// current raw audio input will be appended to the global sample buffer.
|
||||
// The calculator will process the global sample buffer and output as many
|
||||
// tensors as possible.
|
||||
// Non-streaming mode: when "streaming_mode" is set to false in the calculator
|
||||
// options, the calculators treats the packets in the input audio stream as
|
||||
// a batch of unrelated audio buffers. In each Process() call, the input
|
||||
// buffer will be frist resampled, and framed as fixed-sized, possibly
|
||||
// overlapping tensors. The last tensor produced by a Process() invocation
|
||||
// will be zero-padding if the remaining samples are insufficient. As the
|
||||
// calculator treats the input packets as unrelated, all samples will be
|
||||
// processed immediately and no samples will be cached in the global sample
|
||||
// buffer.
|
||||
//
|
||||
// Inputs:
|
||||
// AUDIO - mediapipe::Matrix
|
||||
// The audio data represented as mediapipe::Matrix.
|
||||
// SAMPLE_RATE - double @Optional
|
||||
// The sample rate of the corresponding audio data in the "AUDIO" stream.
|
||||
// If a sample rate packet is provided at Timestamp::PreStream(), the sample
|
||||
// rate will be used as the sample rate of every audio packets in the
|
||||
// "AUDIO" stream. Note that one and only one of the "AUDIO" stream's time
|
||||
// series header or the "SAMPLE_RATE" stream can exist.
|
||||
//
|
||||
// Outputs:
|
||||
// TENSORS - std::vector<Tensor>
|
||||
// Vector containing a single Tensor that represents a fix-sized audio
|
||||
// frame.
|
||||
// TIMESTAMPS - std::vector<Timestamp> @Optional
|
||||
// Vector containing the output timestamps emitted by the current Process()
|
||||
// invocation. In the non-streaming mode, the vector contains all of the
|
||||
// output timestamps for an input audio buffer.
|
||||
//
|
||||
// Example:
|
||||
// node {
|
||||
// calculator: "AudioToTensorCalculator"
|
||||
// input_stream: "AUDIO:audio"
|
||||
// output_stream: "TENSORS:tensors"
|
||||
// output_stream: "TIMESTAMPS:timestamps"
|
||||
// options {
|
||||
// [mediapipe.AudioToTensorCalculatorOptions.ext] {
|
||||
// num_channels: 2
|
||||
// num_samples: 512
|
||||
// num_overlapping_samples: 64
|
||||
// target_sample_rate: 16000
|
||||
// streaming_mode: true # or false
|
||||
// }
|
||||
// }
|
||||
// }
|
||||
class AudioToTensorCalculator : public Node {
|
||||
public:
|
||||
static constexpr Input<Matrix> kAudioIn{"AUDIO"};
|
||||
// TODO: Removes this optional input stream when the "AUDIO" stream
|
||||
// uses the new mediapipe audio data containers that carry audio metatdata,
|
||||
// such as sample rate.
|
||||
static constexpr Input<double>::Optional kAudioSampleRateIn{"SAMPLE_RATE"};
|
||||
static constexpr Output<std::vector<Tensor>> kTensorsOut{"TENSORS"};
|
||||
// A vector of the output timestamps emitted by the current Process()
|
||||
// invocation. The packet timestamp is the last emitted timestamp.
|
||||
static constexpr Output<std::vector<Timestamp>>::Optional kTimestampsOut{
|
||||
"TIMESTAMPS"};
|
||||
MEDIAPIPE_NODE_CONTRACT(kAudioIn, kAudioSampleRateIn, kTensorsOut,
|
||||
kTimestampsOut);
|
||||
|
||||
static absl::Status UpdateContract(CalculatorContract* cc);
|
||||
absl::Status Open(CalculatorContext* cc);
|
||||
absl::Status Process(CalculatorContext* cc);
|
||||
absl::Status Close(CalculatorContext* cc);
|
||||
|
||||
private:
|
||||
// The target number of channels.
|
||||
int num_channels_;
|
||||
// The target number of samples per channel.
|
||||
int num_samples_;
|
||||
// The number of samples per channel to advance after the current frame is
|
||||
// processed.
|
||||
int frame_step_;
|
||||
bool streaming_mode_;
|
||||
bool check_inconsistent_timestamps_;
|
||||
Timestamp initial_timestamp_ = Timestamp::Unstarted();
|
||||
int64 cumulative_input_samples_ = 0;
|
||||
Timestamp next_output_timestamp_ = Timestamp::Unstarted();
|
||||
|
||||
double source_sample_rate_ = -1;
|
||||
double target_sample_rate_ = -1;
|
||||
// TODO: Configures QResamplerParams through calculator options.
|
||||
audio_dsp::QResamplerParams params_;
|
||||
// A QResampler instance to resample an audio stream.
|
||||
std::unique_ptr<audio_dsp::QResampler<float>> resampler_;
|
||||
Matrix sample_buffer_;
|
||||
int processed_buffer_cols_ = 0;
|
||||
|
||||
absl::Status ProcessStreamingData(CalculatorContext* cc);
|
||||
absl::Status ProcessNonStreamingData(CalculatorContext* cc);
|
||||
|
||||
absl::Status SetupStreamingResampler(double input_sample_rate_);
|
||||
void AppendToSampleBuffer(Matrix buffer_to_append);
|
||||
|
||||
absl::StatusOr<std::vector<Tensor>> ConvertToTensor(
|
||||
const Matrix& frame_to_convert);
|
||||
absl::Status OutputTensors(const Matrix& buffer, bool should_flush,
|
||||
CalculatorContext* cc);
|
||||
};
|
||||
|
||||
absl::Status AudioToTensorCalculator::UpdateContract(CalculatorContract* cc) {
|
||||
const auto& options =
|
||||
cc->Options<mediapipe::AudioToTensorCalculatorOptions>();
|
||||
if (!options.has_num_channels() || !options.has_num_samples() ||
|
||||
!options.has_target_sample_rate()) {
|
||||
return absl::InvalidArgumentError(
|
||||
"AudioToTensorCalculatorOptions must specifiy "
|
||||
"`num_channels`, `num_samples`, and `target_sample_rate`.");
|
||||
}
|
||||
if (options.streaming_mode()) {
|
||||
// Explicitly disables tiemstamp offset to disallow the timestamp bound
|
||||
// from the input streams to be propagated to the output streams.
|
||||
// In the streaming mode, the output timestamp bound is based on
|
||||
// next_output_timestamp_, which can be smaller than the current input
|
||||
// timestamps.
|
||||
cc->SetTimestampOffset(TimestampDiff::Unset());
|
||||
}
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
absl::Status AudioToTensorCalculator::Open(CalculatorContext* cc) {
|
||||
const auto& options =
|
||||
cc->Options<mediapipe::AudioToTensorCalculatorOptions>();
|
||||
num_channels_ = options.num_channels();
|
||||
num_samples_ = options.num_samples();
|
||||
if (options.has_num_overlapping_samples()) {
|
||||
RET_CHECK_GE(options.num_overlapping_samples(), 0);
|
||||
RET_CHECK_LT(options.num_overlapping_samples(), num_samples_);
|
||||
frame_step_ = num_samples_ - options.num_overlapping_samples();
|
||||
} else {
|
||||
frame_step_ = num_samples_;
|
||||
}
|
||||
target_sample_rate_ = options.target_sample_rate();
|
||||
streaming_mode_ = options.streaming_mode();
|
||||
if (streaming_mode_) {
|
||||
check_inconsistent_timestamps_ = options.check_inconsistent_timestamps();
|
||||
sample_buffer_.resize(num_channels_, Eigen::NoChange);
|
||||
}
|
||||
|
||||
RET_CHECK(kAudioSampleRateIn(cc).IsConnected() ^
|
||||
!kAudioIn(cc).Header().IsEmpty())
|
||||
<< "Must either specify the time series header of the \"AUDIO\" stream "
|
||||
"or have the \"SAMPLE_RATE\" stream connected.";
|
||||
if (!kAudioIn(cc).Header().IsEmpty()) {
|
||||
mediapipe::TimeSeriesHeader input_header;
|
||||
MP_RETURN_IF_ERROR(mediapipe::time_series_util::FillTimeSeriesHeaderIfValid(
|
||||
kAudioIn(cc).Header(), &input_header));
|
||||
if (streaming_mode_) {
|
||||
MP_RETURN_IF_ERROR(SetupStreamingResampler(input_header.sample_rate()));
|
||||
} else {
|
||||
source_sample_rate_ = input_header.sample_rate();
|
||||
}
|
||||
}
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
absl::Status AudioToTensorCalculator::Process(CalculatorContext* cc) {
|
||||
if (cc->InputTimestamp() == Timestamp::PreStream()) {
|
||||
double current_source_sample_rate = kAudioSampleRateIn(cc).Get();
|
||||
if (cc->Options<mediapipe::AudioToTensorCalculatorOptions>()
|
||||
.streaming_mode()) {
|
||||
return SetupStreamingResampler(current_source_sample_rate);
|
||||
} else {
|
||||
source_sample_rate_ = current_source_sample_rate;
|
||||
return absl::OkStatus();
|
||||
}
|
||||
}
|
||||
// Sanity checks.
|
||||
const auto& input_frame = kAudioIn(cc).Get();
|
||||
if (input_frame.rows() != num_channels_) {
|
||||
return absl::InvalidArgumentError(absl::StrFormat(
|
||||
"Audio input has %d channel(s) but the model requires %d channel(s).",
|
||||
input_frame.rows(), num_channels_));
|
||||
}
|
||||
if (num_channels_ > 1 && input_frame.IsRowMajor) {
|
||||
return absl::InvalidArgumentError(
|
||||
"The audio data should be stored in column-major.");
|
||||
}
|
||||
return streaming_mode_ ? ProcessStreamingData(cc)
|
||||
: ProcessNonStreamingData(cc);
|
||||
}
|
||||
|
||||
absl::Status AudioToTensorCalculator::Close(CalculatorContext* cc) {
|
||||
if (!streaming_mode_) {
|
||||
return absl::OkStatus();
|
||||
}
|
||||
if (resampler_) {
|
||||
Matrix resampled_buffer(num_channels_, 0);
|
||||
resampler_->Flush(&resampled_buffer);
|
||||
AppendToSampleBuffer(std::move(resampled_buffer));
|
||||
}
|
||||
return OutputTensors(sample_buffer_, /*should_flush=*/true, cc);
|
||||
}
|
||||
|
||||
absl::Status AudioToTensorCalculator::ProcessStreamingData(
|
||||
CalculatorContext* cc) {
|
||||
const auto& input_buffer = kAudioIn(cc).Get();
|
||||
if (initial_timestamp_ == Timestamp::Unstarted()) {
|
||||
initial_timestamp_ = cc->InputTimestamp();
|
||||
next_output_timestamp_ = initial_timestamp_;
|
||||
}
|
||||
if (source_sample_rate_ != -1 && check_inconsistent_timestamps_) {
|
||||
mediapipe::time_series_util::LogWarningIfTimestampIsInconsistent(
|
||||
cc->InputTimestamp(), initial_timestamp_, cumulative_input_samples_,
|
||||
source_sample_rate_);
|
||||
cumulative_input_samples_ += input_buffer.cols();
|
||||
}
|
||||
if (!kAudioSampleRateIn(cc).IsEmpty()) {
|
||||
double current_source_sample_rate = kAudioSampleRateIn(cc).Get();
|
||||
if (resampler_) {
|
||||
RET_CHECK_EQ(current_source_sample_rate, source_sample_rate_);
|
||||
} else {
|
||||
MP_RETURN_IF_ERROR(SetupStreamingResampler(current_source_sample_rate));
|
||||
}
|
||||
}
|
||||
|
||||
if (resampler_) {
|
||||
Matrix resampled_buffer(num_channels_, 0);
|
||||
resampler_->ProcessSamples(input_buffer, &resampled_buffer);
|
||||
AppendToSampleBuffer(std::move(resampled_buffer));
|
||||
} else {
|
||||
// Tries to consume the input matrix first to avoid extra data copy.
|
||||
auto status_or_matrix = kAudioIn(cc).packet().Consume<Matrix>();
|
||||
if (status_or_matrix.ok()) {
|
||||
Matrix local_matrix(num_channels_, 0);
|
||||
local_matrix.swap(*status_or_matrix.value());
|
||||
AppendToSampleBuffer(std::move(local_matrix));
|
||||
} else {
|
||||
AppendToSampleBuffer(input_buffer);
|
||||
}
|
||||
}
|
||||
|
||||
MP_RETURN_IF_ERROR(OutputTensors(sample_buffer_, /*should_flush=*/false, cc));
|
||||
// Removes the processed samples from the global sample buffer.
|
||||
sample_buffer_ = Matrix(sample_buffer_.rightCols(sample_buffer_.cols() -
|
||||
processed_buffer_cols_ - 1));
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
absl::Status AudioToTensorCalculator::ProcessNonStreamingData(
|
||||
CalculatorContext* cc) {
|
||||
initial_timestamp_ = cc->InputTimestamp();
|
||||
next_output_timestamp_ = initial_timestamp_;
|
||||
const auto& input_frame = kAudioIn(cc).Get();
|
||||
double source_sample_rate = kAudioSampleRateIn(cc).GetOr(source_sample_rate_);
|
||||
|
||||
if (source_sample_rate != -1 && source_sample_rate != target_sample_rate_) {
|
||||
std::vector<float> resampled = audio_dsp::QResampleSignal<float>(
|
||||
source_sample_rate, target_sample_rate_, num_channels_, params_,
|
||||
input_frame);
|
||||
Eigen::Map<const Matrix> matrix_mapping(resampled.data(), num_channels_,
|
||||
resampled.size() / num_channels_);
|
||||
return OutputTensors(matrix_mapping, /*should_flush=*/true, cc);
|
||||
}
|
||||
return OutputTensors(input_frame, /*should_flush=*/true, cc);
|
||||
}
|
||||
|
||||
absl::Status AudioToTensorCalculator::SetupStreamingResampler(
|
||||
double input_sample_rate) {
|
||||
if (input_sample_rate == source_sample_rate_) {
|
||||
return absl::OkStatus();
|
||||
}
|
||||
source_sample_rate_ = input_sample_rate;
|
||||
if (source_sample_rate_ != target_sample_rate_) {
|
||||
resampler_ = absl::make_unique<audio_dsp::QResampler<float>>(
|
||||
source_sample_rate_, target_sample_rate_, num_channels_, params_);
|
||||
if (!resampler_) {
|
||||
return absl::InternalError("Failed to initialize resampler.");
|
||||
}
|
||||
}
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
void AudioToTensorCalculator::AppendToSampleBuffer(Matrix buffer_to_append) {
|
||||
sample_buffer_.conservativeResize(
|
||||
Eigen::NoChange, sample_buffer_.cols() + buffer_to_append.cols());
|
||||
sample_buffer_.rightCols(buffer_to_append.cols()).swap(buffer_to_append);
|
||||
}
|
||||
|
||||
absl::StatusOr<std::vector<Tensor>> AudioToTensorCalculator::ConvertToTensor(
|
||||
const Matrix& frame_to_convert) {
|
||||
Tensor tensor(Tensor::ElementType::kFloat32,
|
||||
Tensor::Shape({num_channels_, num_samples_}));
|
||||
auto buffer_view = tensor.GetCpuWriteView();
|
||||
if (frame_to_convert.size() < num_channels_ * num_samples_) {
|
||||
std::memset(buffer_view.buffer<float>(), 0, tensor.bytes());
|
||||
}
|
||||
std::memcpy(buffer_view.buffer<float>(), frame_to_convert.data(),
|
||||
frame_to_convert.size() * sizeof(float));
|
||||
std::vector<Tensor> tensor_vector;
|
||||
tensor_vector.push_back(std::move(tensor));
|
||||
return tensor_vector;
|
||||
}
|
||||
|
||||
absl::Status AudioToTensorCalculator::OutputTensors(const Matrix& buffer,
|
||||
bool should_flush,
|
||||
CalculatorContext* cc) {
|
||||
int next_frame_first_col = 0;
|
||||
std::vector<Timestamp> timestamps;
|
||||
while ((!streaming_mode_ || !should_flush) &&
|
||||
next_frame_first_col + num_samples_ <= buffer.cols()) {
|
||||
ASSIGN_OR_RETURN(auto output_tensor, ConvertToTensor(buffer.block(
|
||||
0, next_frame_first_col,
|
||||
num_channels_, num_samples_)));
|
||||
kTensorsOut(cc).Send(std::move(output_tensor), next_output_timestamp_);
|
||||
timestamps.push_back(next_output_timestamp_);
|
||||
next_output_timestamp_ += round(frame_step_ / target_sample_rate_ *
|
||||
Timestamp::kTimestampUnitsPerSecond);
|
||||
next_frame_first_col += frame_step_;
|
||||
}
|
||||
if (should_flush && next_frame_first_col < buffer.cols()) {
|
||||
ASSIGN_OR_RETURN(auto output_tensor,
|
||||
ConvertToTensor(buffer.block(
|
||||
0, next_frame_first_col, num_channels_,
|
||||
std::min(num_samples_,
|
||||
(int)buffer.cols() - next_frame_first_col))));
|
||||
// In the streaming mode, the flush happens in Close() and a packet at
|
||||
// Timestamp::Max() will be emitted. In the non-streaming mode, each
|
||||
// Process() invocation will process the entire buffer completely.
|
||||
Timestamp timestamp =
|
||||
streaming_mode_ ? Timestamp::Max() : next_output_timestamp_;
|
||||
timestamps.push_back(timestamp);
|
||||
kTensorsOut(cc).Send(std::move(output_tensor), timestamp);
|
||||
}
|
||||
if (kTimestampsOut(cc).IsConnected()) {
|
||||
Timestamp timestamp = timestamps.back();
|
||||
kTimestampsOut(cc).Send(std::move(timestamps), timestamp);
|
||||
}
|
||||
processed_buffer_cols_ = next_frame_first_col - 1;
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
MEDIAPIPE_REGISTER_NODE(AudioToTensorCalculator);
|
||||
|
||||
} // namespace api2
|
||||
} // namespace mediapipe
|
||||
@@ -0,0 +1,46 @@
|
||||
// Copyright 2022 The MediaPipe Authors.
|
||||
//
|
||||
// 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.
|
||||
|
||||
syntax = "proto2";
|
||||
|
||||
package mediapipe;
|
||||
|
||||
import "mediapipe/framework/calculator.proto";
|
||||
|
||||
message AudioToTensorCalculatorOptions {
|
||||
extend mediapipe.CalculatorOptions {
|
||||
optional AudioToTensorCalculatorOptions ext = 448635064;
|
||||
}
|
||||
|
||||
// The required number of channels the output audio tensor has.
|
||||
optional int64 num_channels = 1;
|
||||
|
||||
// The required number of samples per channel the output audio tensor has.
|
||||
optional int64 num_samples = 2;
|
||||
|
||||
// The number of overlapping samples per channel the output audio tensor has.
|
||||
optional int64 num_overlapping_samples = 3 [default = 0];
|
||||
|
||||
// The target number of samples per second (hertz) of the audio buffers that
|
||||
// will be converted into tensors.
|
||||
optional double target_sample_rate = 4;
|
||||
|
||||
// Whether to treat the input audio stream as a continous stream or a batch
|
||||
// of unrelated audio buffers.
|
||||
optional bool streaming_mode = 5 [default = true];
|
||||
|
||||
// Set to false to disable checks for jitter in timestamp values. Useful with
|
||||
// live audio input.
|
||||
optional bool check_inconsistent_timestamps = 6 [default = true];
|
||||
}
|
||||
@@ -0,0 +1,483 @@
|
||||
// Copyright 2022 The MediaPipe Authors.
|
||||
//
|
||||
// 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.
|
||||
|
||||
#include <cmath>
|
||||
#include <memory>
|
||||
#include <string>
|
||||
#include <vector>
|
||||
|
||||
#include "absl/strings/substitute.h"
|
||||
#include "audio/dsp/resampler_q.h"
|
||||
#include "mediapipe/framework/api2/packet.h"
|
||||
#include "mediapipe/framework/calculator.pb.h"
|
||||
#include "mediapipe/framework/calculator_framework.h"
|
||||
#include "mediapipe/framework/formats/matrix.h"
|
||||
#include "mediapipe/framework/formats/tensor.h"
|
||||
#include "mediapipe/framework/port/gtest.h"
|
||||
#include "mediapipe/framework/port/parse_text_proto.h"
|
||||
#include "mediapipe/framework/port/status_matchers.h"
|
||||
#include "mediapipe/framework/timestamp.h"
|
||||
|
||||
namespace mediapipe {
|
||||
namespace {
|
||||
|
||||
std::unique_ptr<Matrix> CreateTestMatrix(int num_channels, int num_samples,
|
||||
int timestamp) {
|
||||
auto matrix = std::make_unique<Matrix>(num_channels, num_samples);
|
||||
for (int c = 0; c < num_channels; ++c) {
|
||||
for (int i = 0; i < num_samples; ++i) {
|
||||
// A float value with the sample, channel, and timestamp separated by a
|
||||
// few orders of magnitude, for easy parsing by humans.
|
||||
(*matrix)(c, i) = timestamp / 10000 + i + c / 100.0;
|
||||
}
|
||||
}
|
||||
return matrix;
|
||||
}
|
||||
|
||||
std::unique_ptr<Matrix> ResampleBuffer(const Matrix& input_matrix,
|
||||
double resampling_factor) {
|
||||
audio_dsp::QResamplerParams params;
|
||||
std::vector<float> resampled;
|
||||
int num_channels = input_matrix.rows();
|
||||
std::vector<float> input_data(input_matrix.data(),
|
||||
input_matrix.data() + input_matrix.size());
|
||||
resampled = audio_dsp::QResampleSignal<float>(
|
||||
1, resampling_factor, num_channels, params, input_data);
|
||||
Matrix res = Eigen::Map<Matrix>(resampled.data(), num_channels,
|
||||
resampled.size() / num_channels);
|
||||
return std::make_unique<Matrix>(std::move(res));
|
||||
}
|
||||
|
||||
class AudioToTensorCalculatorNonStreamingModeTest : public ::testing::Test {
|
||||
protected:
|
||||
void SetUp() override {}
|
||||
void Run(int num_samples, int num_overlapping_samples,
|
||||
double resampling_factor, const Matrix& input_matrix) {
|
||||
double input_sample_rate = 10000;
|
||||
double target_sample_rate = input_sample_rate * resampling_factor;
|
||||
auto graph_config = ParseTextProtoOrDie<CalculatorGraphConfig>(
|
||||
absl::Substitute(R"(
|
||||
input_stream: "audio"
|
||||
input_stream: "sample_rate"
|
||||
output_stream: "tensors"
|
||||
output_stream: "timestamps"
|
||||
node {
|
||||
calculator: "AudioToTensorCalculator"
|
||||
input_stream: "AUDIO:audio"
|
||||
input_stream: "SAMPLE_RATE:sample_rate"
|
||||
output_stream: "TENSORS:tensors"
|
||||
output_stream: "TIMESTAMPS:timestamps"
|
||||
options {
|
||||
[mediapipe.AudioToTensorCalculatorOptions.ext] {
|
||||
num_channels: $0
|
||||
num_samples: $1
|
||||
num_overlapping_samples: $2
|
||||
target_sample_rate: $3
|
||||
streaming_mode: false
|
||||
}
|
||||
}
|
||||
}
|
||||
)",
|
||||
/*$0=*/input_matrix.rows(),
|
||||
/*$1=*/num_samples, /*$2=*/num_overlapping_samples,
|
||||
/*$3=*/target_sample_rate));
|
||||
tool::AddVectorSink("tensors", &graph_config, &tensors_packets_);
|
||||
tool::AddVectorSink("timestamps", &graph_config, ×tamps_packets_);
|
||||
|
||||
// Run the graph.
|
||||
MP_ASSERT_OK(graph_.Initialize(graph_config));
|
||||
MP_ASSERT_OK(graph_.StartRun({}));
|
||||
// Run with the input matrix multiple times.
|
||||
for (int i = 0; i < num_iterations_; ++i) {
|
||||
MP_ASSERT_OK(graph_.AddPacketToInputStream(
|
||||
"audio",
|
||||
MakePacket<Matrix>(input_matrix)
|
||||
.At(Timestamp(i * Timestamp::kTimestampUnitsPerSecond))));
|
||||
MP_ASSERT_OK(graph_.AddPacketToInputStream(
|
||||
"sample_rate",
|
||||
MakePacket<double>(input_sample_rate)
|
||||
.At(Timestamp(i * Timestamp::kTimestampUnitsPerSecond))));
|
||||
}
|
||||
MP_ASSERT_OK(graph_.CloseAllInputStreams());
|
||||
MP_ASSERT_OK(graph_.WaitUntilIdle());
|
||||
}
|
||||
|
||||
void CheckTensorsOutputPackets(const Matrix& expected_matrix,
|
||||
int sample_offset, int num_tensors_per_input) {
|
||||
ASSERT_EQ(num_iterations_ * num_tensors_per_input, tensors_packets_.size());
|
||||
for (int i = 0; i < num_iterations_; ++i) {
|
||||
for (int j = 0; j < num_tensors_per_input; ++j) {
|
||||
CheckTensorsOutputPacket(
|
||||
expected_matrix, tensors_packets_[i * num_tensors_per_input + j],
|
||||
/*sample_offset*/ sample_offset * j, /*index=*/j);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void CheckTensorsOutputPacket(const Matrix& expected_matrix,
|
||||
const Packet& packet, int sample_offset,
|
||||
int index) {
|
||||
MP_ASSERT_OK(packet.ValidateAsType<std::vector<Tensor>>());
|
||||
ASSERT_EQ(1, packet.Get<std::vector<Tensor>>().size());
|
||||
const Tensor& output_tensor = packet.Get<std::vector<Tensor>>()[0];
|
||||
auto* buffer = output_tensor.GetCpuReadView().buffer<float>();
|
||||
int num_values = output_tensor.shape().num_elements();
|
||||
const std::vector<float> output_floats(buffer, buffer + num_values);
|
||||
for (int i = 0; i < num_values; ++i) {
|
||||
if (i + sample_offset >= expected_matrix.size()) {
|
||||
EXPECT_FLOAT_EQ(output_floats[i], 0);
|
||||
} else {
|
||||
EXPECT_FLOAT_EQ(output_floats[i],
|
||||
expected_matrix.coeff((i + sample_offset) % 2,
|
||||
(i + sample_offset) / 2))
|
||||
<< "i=" << i << ", sample_offset=" << sample_offset;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void CheckTimestampsOutputPackets(
|
||||
std::vector<int64> expected_timestamp_values) {
|
||||
ASSERT_EQ(num_iterations_, timestamps_packets_.size());
|
||||
for (int i = 0; i < timestamps_packets_.size(); ++i) {
|
||||
const auto& p = timestamps_packets_[i];
|
||||
MP_ASSERT_OK(p.ValidateAsType<std::vector<Timestamp>>());
|
||||
auto output_timestamps = p.Get<std::vector<Timestamp>>();
|
||||
int64 base_timestamp = i * Timestamp::kTimestampUnitsPerSecond;
|
||||
std::vector<Timestamp> expected_timestamps;
|
||||
expected_timestamps.resize(expected_timestamp_values.size());
|
||||
std::transform(
|
||||
expected_timestamp_values.begin(), expected_timestamp_values.end(),
|
||||
expected_timestamps.begin(), [base_timestamp](int64 v) -> Timestamp {
|
||||
return Timestamp(v + base_timestamp);
|
||||
});
|
||||
EXPECT_EQ(expected_timestamps, output_timestamps);
|
||||
EXPECT_EQ(p.Timestamp(), expected_timestamps.back());
|
||||
}
|
||||
}
|
||||
|
||||
// Fully close graph at end, otherwise calculator+tensors are destroyed
|
||||
// after calling WaitUntilDone().
|
||||
void CloseGraph() { MP_EXPECT_OK(graph_.WaitUntilDone()); }
|
||||
|
||||
private:
|
||||
CalculatorGraph graph_;
|
||||
int num_iterations_ = 10;
|
||||
std::vector<Packet> tensors_packets_;
|
||||
std::vector<Packet> timestamps_packets_;
|
||||
};
|
||||
|
||||
TEST_F(AudioToTensorCalculatorNonStreamingModeTest,
|
||||
ConvertToNoOverlappingFp32Tensors) {
|
||||
auto input_matrix = CreateTestMatrix(2, 8, 0);
|
||||
Run(/*num_samples=*/4, /*num_overlapping_samples=*/0,
|
||||
/*resampling_factor=*/1.0f, *input_matrix);
|
||||
CheckTensorsOutputPackets(*input_matrix, /*sample_offset=*/8,
|
||||
/*num_tensors_per_input=*/2);
|
||||
CheckTimestampsOutputPackets({0, 400});
|
||||
CloseGraph();
|
||||
}
|
||||
|
||||
TEST_F(AudioToTensorCalculatorNonStreamingModeTest,
|
||||
ConvertToOverlappingFp32Tensors) {
|
||||
auto input_matrix = CreateTestMatrix(2, 8, 0);
|
||||
Run(/*num_samples=*/4, /*num_overlapping_samples=*/2,
|
||||
/*resampling_factor=*/1.0f, *input_matrix);
|
||||
CheckTensorsOutputPackets(*input_matrix, /*sample_offset=*/4,
|
||||
/*num_tensors_per_input=*/4);
|
||||
CheckTimestampsOutputPackets({0, 200, 400, 600});
|
||||
CloseGraph();
|
||||
}
|
||||
|
||||
TEST_F(AudioToTensorCalculatorNonStreamingModeTest, TensorsWithZeroPadding) {
|
||||
auto input_matrix = CreateTestMatrix(2, 7, 0);
|
||||
Run(/*num_samples=*/4, /*num_overlapping_samples=*/2,
|
||||
/*resampling_factor=*/1.0f, *input_matrix);
|
||||
CheckTensorsOutputPackets(*input_matrix, /*sample_offset=*/4,
|
||||
/*num_tensors_per_input=*/3);
|
||||
CheckTimestampsOutputPackets({0, 200, 400});
|
||||
CloseGraph();
|
||||
}
|
||||
|
||||
TEST_F(AudioToTensorCalculatorNonStreamingModeTest, Downsampling) {
|
||||
auto input_matrix = CreateTestMatrix(2, 1024, 0);
|
||||
Run(/*num_samples=*/256, /*num_overlapping_samples=*/0,
|
||||
/*resampling_factor=*/0.5f, *input_matrix);
|
||||
auto expected_matrix =
|
||||
ResampleBuffer(*input_matrix, /*resampling_factor=*/0.5f);
|
||||
CheckTensorsOutputPackets(*expected_matrix, /*sample_offset=*/512,
|
||||
/*num_tensors_per_input=*/3);
|
||||
CheckTimestampsOutputPackets({0, 51200, 102400});
|
||||
CloseGraph();
|
||||
}
|
||||
|
||||
TEST_F(AudioToTensorCalculatorNonStreamingModeTest,
|
||||
DownsamplingWithOverlapping) {
|
||||
auto input_matrix = CreateTestMatrix(2, 1024, 0);
|
||||
Run(/*num_samples=*/256, /*num_overlapping_samples=*/64,
|
||||
/*resampling_factor=*/0.5f, *input_matrix);
|
||||
auto expected_matrix =
|
||||
ResampleBuffer(*input_matrix, /*resampling_factor=*/0.5f);
|
||||
CheckTensorsOutputPackets(*expected_matrix, /*sample_offset=*/384,
|
||||
/*num_tensors_per_input=*/3);
|
||||
CheckTimestampsOutputPackets({0, 38400, 76800});
|
||||
CloseGraph();
|
||||
}
|
||||
|
||||
TEST_F(AudioToTensorCalculatorNonStreamingModeTest, Upsampling) {
|
||||
auto input_matrix = CreateTestMatrix(2, 1024, 0);
|
||||
Run(/*num_samples=*/256, /*num_overlapping_samples=*/0,
|
||||
/*resampling_factor=*/2.0f, *input_matrix);
|
||||
auto expected_matrix =
|
||||
ResampleBuffer(*input_matrix, /*resampling_factor=*/2.0f);
|
||||
CheckTensorsOutputPackets(*expected_matrix,
|
||||
/*sample_offset=*/512,
|
||||
/*num_tensors_per_input=*/9);
|
||||
CheckTimestampsOutputPackets(
|
||||
{0, 12800, 25600, 38400, 51200, 64000, 76800, 89600, 102400});
|
||||
CloseGraph();
|
||||
}
|
||||
|
||||
TEST_F(AudioToTensorCalculatorNonStreamingModeTest, UpsamplingWithOverlapping) {
|
||||
auto input_matrix = CreateTestMatrix(2, 256, 0);
|
||||
Run(/*num_samples=*/256, /*num_overlapping_samples=*/64,
|
||||
/*resampling_factor=*/2.0f, *input_matrix);
|
||||
auto expected_matrix =
|
||||
ResampleBuffer(*input_matrix, /*resampling_factor=*/2.0f);
|
||||
CheckTensorsOutputPackets(*expected_matrix,
|
||||
/*sample_offset=*/384,
|
||||
/*num_tensors_per_input=*/3);
|
||||
CheckTimestampsOutputPackets({0, 9600, 19200});
|
||||
CloseGraph();
|
||||
}
|
||||
|
||||
class AudioToTensorCalculatorStreamingModeTest : public ::testing::Test {
|
||||
protected:
|
||||
void SetUp() override { sample_buffer_ = std::make_unique<Matrix>(2, 0); }
|
||||
|
||||
void SetInputBufferNumSamplesPerChannel(int num_samples) {
|
||||
input_buffer_num_samples_ = num_samples;
|
||||
}
|
||||
|
||||
void SetNumIterations(int num_iterations) {
|
||||
num_iterations_ = num_iterations;
|
||||
}
|
||||
|
||||
int GetExpectedNumOfSamples() {
|
||||
Matrix* expected_matrix =
|
||||
resampled_buffer_ ? resampled_buffer_.get() : sample_buffer_.get();
|
||||
return expected_matrix->cols();
|
||||
}
|
||||
|
||||
void Run(int num_samples, int num_overlapping_samples,
|
||||
double resampling_factor) {
|
||||
double input_sample_rate = 10000;
|
||||
double target_sample_rate = input_sample_rate * resampling_factor;
|
||||
auto graph_config = ParseTextProtoOrDie<CalculatorGraphConfig>(
|
||||
absl::Substitute(R"(
|
||||
input_stream: "audio"
|
||||
input_stream: "sample_rate"
|
||||
output_stream: "tensors"
|
||||
node {
|
||||
calculator: "AudioToTensorCalculator"
|
||||
input_stream: "AUDIO:audio"
|
||||
input_stream: "SAMPLE_RATE:sample_rate"
|
||||
output_stream: "TENSORS:tensors"
|
||||
options {
|
||||
[mediapipe.AudioToTensorCalculatorOptions.ext] {
|
||||
num_channels: 2
|
||||
num_samples: $0
|
||||
num_overlapping_samples: $1
|
||||
target_sample_rate: $2
|
||||
streaming_mode:true
|
||||
}
|
||||
}
|
||||
}
|
||||
)",
|
||||
/*$0=*/num_samples, /*$1=*/num_overlapping_samples,
|
||||
/*$2=*/target_sample_rate));
|
||||
tool::AddVectorSink("tensors", &graph_config, &tensors_packets_);
|
||||
|
||||
// Run the graph.
|
||||
MP_ASSERT_OK(graph_.Initialize(graph_config));
|
||||
MP_ASSERT_OK(graph_.StartRun({}));
|
||||
for (int i = 0; i < num_iterations_; ++i) {
|
||||
Timestamp input_timestamp(Timestamp::kTimestampUnitsPerSecond * i);
|
||||
auto new_data = CreateTestMatrix(2, input_buffer_num_samples_,
|
||||
input_timestamp.Value());
|
||||
MP_ASSERT_OK(graph_.AddPacketToInputStream(
|
||||
"audio", MakePacket<Matrix>(*new_data).At(input_timestamp)));
|
||||
MP_ASSERT_OK(graph_.AddPacketToInputStream(
|
||||
"sample_rate",
|
||||
MakePacket<double>(input_sample_rate).At(input_timestamp)));
|
||||
sample_buffer_->conservativeResize(
|
||||
Eigen::NoChange, sample_buffer_->cols() + new_data->cols());
|
||||
sample_buffer_->rightCols(new_data->cols()).swap(*new_data);
|
||||
}
|
||||
MP_ASSERT_OK(graph_.CloseAllInputStreams());
|
||||
MP_ASSERT_OK(graph_.WaitUntilIdle());
|
||||
if (resampling_factor != 1) {
|
||||
resampled_buffer_ = ResampleBuffer(*sample_buffer_, resampling_factor);
|
||||
}
|
||||
}
|
||||
|
||||
void CheckTensorsOutputPackets(int sample_offset, int num_packets,
|
||||
int64 timestamp_interval,
|
||||
bool output_last_at_close) {
|
||||
ASSERT_EQ(num_packets, tensors_packets_.size());
|
||||
for (int i = 0; i < num_packets; ++i) {
|
||||
if (i == num_packets - 1 && output_last_at_close) {
|
||||
CheckTensorsOutputPacket(sample_offset * i, i, Timestamp::Max());
|
||||
} else {
|
||||
CheckTensorsOutputPacket(sample_offset * i, i,
|
||||
Timestamp(timestamp_interval * i));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
void CheckTensorsOutputPacket(int sample_offset, int index,
|
||||
Timestamp expected_timestamp) {
|
||||
const Packet& p = tensors_packets_[index];
|
||||
MP_ASSERT_OK(p.ValidateAsType<std::vector<Tensor>>());
|
||||
const Tensor& output_tensor = p.Get<std::vector<Tensor>>()[0];
|
||||
auto buffer = output_tensor.GetCpuReadView().buffer<float>();
|
||||
int num_values = output_tensor.shape().num_elements();
|
||||
std::vector<float> output_floats(buffer, buffer + num_values);
|
||||
Matrix* expected_matrix =
|
||||
resampled_buffer_ ? resampled_buffer_.get() : sample_buffer_.get();
|
||||
for (int i = 0; i < num_values; ++i) {
|
||||
if (i + sample_offset >= expected_matrix->size()) {
|
||||
EXPECT_FLOAT_EQ(output_floats[i], 0);
|
||||
} else {
|
||||
EXPECT_NEAR(output_floats[i],
|
||||
expected_matrix->coeff((i + sample_offset) % 2,
|
||||
(i + sample_offset) / 2),
|
||||
0.001)
|
||||
<< "i=" << i << ", sample_offset=" << sample_offset
|
||||
<< ", packet index=" << index;
|
||||
}
|
||||
}
|
||||
EXPECT_EQ(p.Timestamp(), expected_timestamp);
|
||||
}
|
||||
|
||||
// Fully close graph at end, otherwise calculator+tensors are destroyed
|
||||
// after calling WaitUntilDone().
|
||||
void CloseGraph() { MP_EXPECT_OK(graph_.WaitUntilDone()); }
|
||||
|
||||
private:
|
||||
int input_buffer_num_samples_ = 10;
|
||||
int num_iterations_ = 10;
|
||||
CalculatorGraph graph_;
|
||||
std::vector<Packet> tensors_packets_;
|
||||
std::unique_ptr<Matrix> sample_buffer_;
|
||||
std::unique_ptr<Matrix> resampled_buffer_;
|
||||
};
|
||||
|
||||
TEST_F(AudioToTensorCalculatorStreamingModeTest,
|
||||
OutputNoOverlappingFp32Tensors) {
|
||||
Run(/*num_samples=*/5, /*num_overlapping_samples=*/0,
|
||||
/*resampling_factor=*/1.0f);
|
||||
CheckTensorsOutputPackets(
|
||||
/*sample_offset=*/10,
|
||||
/*num_packets=*/std::ceil((float)GetExpectedNumOfSamples() / 5),
|
||||
/*timestamp_interval=*/500,
|
||||
/*output_last_at_close=*/false);
|
||||
CloseGraph();
|
||||
}
|
||||
|
||||
TEST_F(AudioToTensorCalculatorStreamingModeTest, OutputRemainingInCloseMethod) {
|
||||
Run(/*num_samples=*/6, /*num_overlapping_samples=*/0,
|
||||
/*resampling_factor=*/1.0f);
|
||||
CheckTensorsOutputPackets(
|
||||
/*sample_offset=*/12,
|
||||
/*num_packets=*/std::ceil((float)GetExpectedNumOfSamples() / 6),
|
||||
/*timestamp_interval=*/600,
|
||||
/*output_last_at_close=*/true);
|
||||
CloseGraph();
|
||||
}
|
||||
|
||||
TEST_F(AudioToTensorCalculatorStreamingModeTest, OutputOverlappingFp32Tensors) {
|
||||
SetInputBufferNumSamplesPerChannel(12);
|
||||
Run(/*num_samples=*/10, /*num_overlapping_samples=*/2,
|
||||
/*resampling_factor=*/1.0f);
|
||||
CheckTensorsOutputPackets(
|
||||
/*sample_offset=*/16,
|
||||
/*num_packets=*/std::ceil((float)GetExpectedNumOfSamples() / 8),
|
||||
/*timestamp_interval=*/800,
|
||||
/*output_last_at_close=*/true);
|
||||
CloseGraph();
|
||||
}
|
||||
|
||||
TEST_F(AudioToTensorCalculatorStreamingModeTest, Downsampling) {
|
||||
SetInputBufferNumSamplesPerChannel(1000);
|
||||
Run(/*num_samples=*/256, /*num_overlapping_samples=*/0,
|
||||
/*resampling_factor=*/0.5f);
|
||||
CheckTensorsOutputPackets(
|
||||
/*sample_offset=*/512,
|
||||
/*num_packets=*/std::ceil((float)GetExpectedNumOfSamples() / 256),
|
||||
/*timestamp_interval=*/51200,
|
||||
/*output_last_at_close=*/true);
|
||||
CloseGraph();
|
||||
}
|
||||
|
||||
TEST_F(AudioToTensorCalculatorStreamingModeTest, DownsamplingWithOverlapping) {
|
||||
SetInputBufferNumSamplesPerChannel(1024);
|
||||
Run(/*num_samples=*/256, /*num_overlapping_samples=*/64,
|
||||
/*resampling_factor=*/0.5f);
|
||||
CheckTensorsOutputPackets(
|
||||
/*sample_offset=*/384,
|
||||
/*num_packets=*/std::ceil((float)GetExpectedNumOfSamples() / 192),
|
||||
/*timestamp_interval=*/38400,
|
||||
/*output_last_at_close=*/true);
|
||||
CloseGraph();
|
||||
}
|
||||
|
||||
TEST_F(AudioToTensorCalculatorStreamingModeTest, Upsampling) {
|
||||
SetInputBufferNumSamplesPerChannel(1000);
|
||||
Run(/*num_samples=*/256, /*num_overlapping_samples=*/0,
|
||||
/*resampling_factor=*/2.0f);
|
||||
CheckTensorsOutputPackets(
|
||||
/*sample_offset=*/512,
|
||||
/*num_packets=*/std::ceil((float)GetExpectedNumOfSamples() / 256),
|
||||
/*timestamp_interval=*/12800,
|
||||
/*output_last_at_close=*/true);
|
||||
CloseGraph();
|
||||
}
|
||||
|
||||
TEST_F(AudioToTensorCalculatorStreamingModeTest, UpsamplingWithOverlapping) {
|
||||
SetInputBufferNumSamplesPerChannel(1024);
|
||||
Run(/*num_samples=*/256, /*num_overlapping_samples=*/64,
|
||||
/*resampling_factor=*/2.0f);
|
||||
CheckTensorsOutputPackets(
|
||||
/*sample_offset=*/384,
|
||||
/*num_packets=*/std::ceil((float)GetExpectedNumOfSamples() / 192),
|
||||
/*timestamp_interval=*/9600,
|
||||
/*output_last_at_close=*/true);
|
||||
CloseGraph();
|
||||
}
|
||||
|
||||
TEST_F(AudioToTensorCalculatorStreamingModeTest,
|
||||
OnlyOutputInCloseIfNoSufficientSamples) {
|
||||
SetNumIterations(1);
|
||||
Run(/*num_samples=*/8, /*num_overlapping_samples=*/0,
|
||||
/*resampling_factor=*/0.5f);
|
||||
CheckTensorsOutputPackets(
|
||||
/*sample_offset=*/0,
|
||||
/*num_packets=*/1,
|
||||
/*timestamp_interval=*/0,
|
||||
/*output_last_at_close=*/true);
|
||||
CloseGraph();
|
||||
}
|
||||
|
||||
} // namespace
|
||||
} // namespace mediapipe
|
||||
@@ -49,7 +49,6 @@
|
||||
#include "mediapipe/calculators/tensor/image_to_tensor_converter_gl_texture.h"
|
||||
#include "mediapipe/gpu/gl_calculator_helper.h"
|
||||
#endif // MEDIAPIPE_METAL_ENABLED
|
||||
|
||||
#endif // !MEDIAPIPE_DISABLE_GPU
|
||||
|
||||
namespace mediapipe {
|
||||
@@ -142,11 +141,37 @@ class ImageToTensorCalculator : public Node {
|
||||
const auto& options =
|
||||
cc->Options<mediapipe::ImageToTensorCalculatorOptions>();
|
||||
|
||||
RET_CHECK(options.has_output_tensor_float_range())
|
||||
RET_CHECK(options.has_output_tensor_float_range() ||
|
||||
options.has_output_tensor_int_range() ||
|
||||
options.has_output_tensor_uint_range())
|
||||
<< "Output tensor range is required.";
|
||||
RET_CHECK_LT(options.output_tensor_float_range().min(),
|
||||
options.output_tensor_float_range().max())
|
||||
<< "Valid output tensor range is required.";
|
||||
if (options.has_output_tensor_float_range()) {
|
||||
RET_CHECK_LT(options.output_tensor_float_range().min(),
|
||||
options.output_tensor_float_range().max())
|
||||
<< "Valid output float tensor range is required.";
|
||||
}
|
||||
if (options.has_output_tensor_uint_range()) {
|
||||
RET_CHECK_LT(options.output_tensor_uint_range().min(),
|
||||
options.output_tensor_uint_range().max())
|
||||
<< "Valid output uint tensor range is required.";
|
||||
RET_CHECK_GE(options.output_tensor_uint_range().min(), 0)
|
||||
<< "The minimum of the output uint tensor range must be "
|
||||
"non-negative.";
|
||||
RET_CHECK_LE(options.output_tensor_uint_range().max(), 255)
|
||||
<< "The maximum of the output uint tensor range must be less than or "
|
||||
"equal to 255.";
|
||||
}
|
||||
if (options.has_output_tensor_int_range()) {
|
||||
RET_CHECK_LT(options.output_tensor_int_range().min(),
|
||||
options.output_tensor_int_range().max())
|
||||
<< "Valid output int tensor range is required.";
|
||||
RET_CHECK_GE(options.output_tensor_int_range().min(), -128)
|
||||
<< "The minimum of the output int tensor range must be greater than "
|
||||
"or equal to -128.";
|
||||
RET_CHECK_LE(options.output_tensor_int_range().max(), 127)
|
||||
<< "The maximum of the output int tensor range must be less than or "
|
||||
"equal to 127.";
|
||||
}
|
||||
RET_CHECK_GT(options.output_tensor_width(), 0)
|
||||
<< "Valid output tensor width is required.";
|
||||
RET_CHECK_GT(options.output_tensor_height(), 0)
|
||||
@@ -175,9 +200,19 @@ class ImageToTensorCalculator : public Node {
|
||||
options_ = cc->Options<mediapipe::ImageToTensorCalculatorOptions>();
|
||||
output_width_ = options_.output_tensor_width();
|
||||
output_height_ = options_.output_tensor_height();
|
||||
range_min_ = options_.output_tensor_float_range().min();
|
||||
range_max_ = options_.output_tensor_float_range().max();
|
||||
|
||||
is_float_output_ = options_.has_output_tensor_float_range();
|
||||
if (options_.has_output_tensor_uint_range()) {
|
||||
range_min_ =
|
||||
static_cast<float>(options_.output_tensor_uint_range().min());
|
||||
range_max_ =
|
||||
static_cast<float>(options_.output_tensor_uint_range().max());
|
||||
} else if (options_.has_output_tensor_int_range()) {
|
||||
range_min_ = static_cast<float>(options_.output_tensor_int_range().min());
|
||||
range_max_ = static_cast<float>(options_.output_tensor_int_range().max());
|
||||
} else {
|
||||
range_min_ = options_.output_tensor_float_range().min();
|
||||
range_max_ = options_.output_tensor_float_range().max();
|
||||
}
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
@@ -225,7 +260,7 @@ class ImageToTensorCalculator : public Node {
|
||||
}
|
||||
|
||||
// Lazy initialization of the GPU or CPU converter.
|
||||
MP_RETURN_IF_ERROR(InitConverterIfNecessary(cc, image->UsesGpu()));
|
||||
MP_RETURN_IF_ERROR(InitConverterIfNecessary(cc, *image.get()));
|
||||
|
||||
ASSIGN_OR_RETURN(Tensor tensor,
|
||||
(image->UsesGpu() ? gpu_converter_ : cpu_converter_)
|
||||
@@ -257,6 +292,17 @@ class ImageToTensorCalculator : public Node {
|
||||
}
|
||||
}
|
||||
|
||||
Tensor::ElementType GetOutputTensorType() {
|
||||
if (is_float_output_) {
|
||||
return Tensor::ElementType::kFloat32;
|
||||
}
|
||||
if (range_min_ < 0) {
|
||||
return Tensor::ElementType::kInt8;
|
||||
} else {
|
||||
return Tensor::ElementType::kUInt8;
|
||||
}
|
||||
}
|
||||
|
||||
absl::StatusOr<std::shared_ptr<const mediapipe::Image>> GetInputImage(
|
||||
CalculatorContext* cc) {
|
||||
if (kIn(cc).IsConnected()) {
|
||||
@@ -283,9 +329,15 @@ class ImageToTensorCalculator : public Node {
|
||||
}
|
||||
}
|
||||
|
||||
absl::Status InitConverterIfNecessary(CalculatorContext* cc, bool use_gpu) {
|
||||
absl::Status InitConverterIfNecessary(CalculatorContext* cc,
|
||||
const Image& image) {
|
||||
// Lazy initialization of the GPU or CPU converter.
|
||||
if (use_gpu) {
|
||||
if (image.UsesGpu()) {
|
||||
if (!is_float_output_) {
|
||||
return absl::UnimplementedError(
|
||||
"ImageToTensorConverter for the input GPU image currently doesn't "
|
||||
"support quantization.");
|
||||
}
|
||||
if (!gpu_converter_) {
|
||||
#if !MEDIAPIPE_DISABLE_GPU
|
||||
#if MEDIAPIPE_METAL_ENABLED
|
||||
@@ -296,17 +348,26 @@ class ImageToTensorCalculator : public Node {
|
||||
CreateImageToGlBufferTensorConverter(
|
||||
cc, DoesGpuInputStartAtBottom(), GetBorderMode()));
|
||||
#else
|
||||
ASSIGN_OR_RETURN(gpu_converter_,
|
||||
CreateImageToGlTextureTensorConverter(
|
||||
cc, DoesGpuInputStartAtBottom(), GetBorderMode()));
|
||||
// Check whether the underlying storage object is a GL texture.
|
||||
if (image.GetGpuBuffer()
|
||||
.internal_storage<mediapipe::GlTextureBuffer>()) {
|
||||
ASSIGN_OR_RETURN(
|
||||
gpu_converter_,
|
||||
CreateImageToGlTextureTensorConverter(
|
||||
cc, DoesGpuInputStartAtBottom(), GetBorderMode()));
|
||||
} else {
|
||||
return absl::UnimplementedError(
|
||||
"ImageToTensorConverter for the input GPU image is unavailable.");
|
||||
}
|
||||
#endif // MEDIAPIPE_METAL_ENABLED
|
||||
#endif // !MEDIAPIPE_DISABLE_GPU
|
||||
}
|
||||
} else {
|
||||
if (!cpu_converter_) {
|
||||
#if !MEDIAPIPE_DISABLE_OPENCV
|
||||
ASSIGN_OR_RETURN(cpu_converter_,
|
||||
CreateOpenCvConverter(cc, GetBorderMode()));
|
||||
ASSIGN_OR_RETURN(
|
||||
cpu_converter_,
|
||||
CreateOpenCvConverter(cc, GetBorderMode(), GetOutputTensorType()));
|
||||
#else
|
||||
LOG(FATAL) << "Cannot create image to tensor opencv converter since "
|
||||
"MEDIAPIPE_DISABLE_OPENCV is defined.";
|
||||
@@ -321,6 +382,7 @@ class ImageToTensorCalculator : public Node {
|
||||
mediapipe::ImageToTensorCalculatorOptions options_;
|
||||
int output_width_ = 0;
|
||||
int output_height_ = 0;
|
||||
bool is_float_output_ = false;
|
||||
float range_min_ = 0.0f;
|
||||
float range_max_ = 1.0f;
|
||||
};
|
||||
|
||||
@@ -31,6 +31,22 @@ message ImageToTensorCalculatorOptions {
|
||||
optional float max = 2;
|
||||
}
|
||||
|
||||
// Range of int values [min, max].
|
||||
// min, must be strictly less than max.
|
||||
// Please note that IntRange is supported for CPU tensors only.
|
||||
message IntRange {
|
||||
optional int64 min = 1;
|
||||
optional int64 max = 2;
|
||||
}
|
||||
|
||||
// Range of uint values [min, max].
|
||||
// min, must be strictly less than max.
|
||||
// Please note that UIntRange is supported for CPU tensors only.
|
||||
message UIntRange {
|
||||
optional uint64 min = 1;
|
||||
optional uint64 max = 2;
|
||||
}
|
||||
|
||||
// Pixel extrapolation methods. See @border_mode.
|
||||
enum BorderMode {
|
||||
BORDER_UNSPECIFIED = 0;
|
||||
@@ -49,6 +65,8 @@ message ImageToTensorCalculatorOptions {
|
||||
// Output tensor element range/type image pixels are converted to.
|
||||
oneof range {
|
||||
FloatRange output_tensor_float_range = 4;
|
||||
IntRange output_tensor_int_range = 7;
|
||||
UIntRange output_tensor_uint_range = 8;
|
||||
}
|
||||
|
||||
// For CONVENTIONAL mode for OpenGL, input image starts at bottom and needs
|
||||
|
||||
@@ -61,7 +61,8 @@ void RunTestWithInputImagePacket(const Packet& input_image_packet,
|
||||
float range_max, int tensor_width,
|
||||
int tensor_height, bool keep_aspect,
|
||||
absl::optional<BorderMode> border_mode,
|
||||
const mediapipe::NormalizedRect& roi) {
|
||||
const mediapipe::NormalizedRect& roi,
|
||||
bool output_int_tensor) {
|
||||
std::string border_mode_str;
|
||||
if (border_mode) {
|
||||
switch (*border_mode) {
|
||||
@@ -73,6 +74,30 @@ void RunTestWithInputImagePacket(const Packet& input_image_packet,
|
||||
break;
|
||||
}
|
||||
}
|
||||
std::string output_tensor_range;
|
||||
if (output_int_tensor) {
|
||||
if (range_min < 0) {
|
||||
output_tensor_range = absl::Substitute(R"(output_tensor_int_range {
|
||||
min: $0
|
||||
max: $1
|
||||
})",
|
||||
static_cast<int>(range_min),
|
||||
static_cast<int>(range_max));
|
||||
} else {
|
||||
output_tensor_range = absl::Substitute(R"(output_tensor_uint_range {
|
||||
min: $0
|
||||
max: $1
|
||||
})",
|
||||
static_cast<uint>(range_min),
|
||||
static_cast<uint>(range_max));
|
||||
}
|
||||
} else {
|
||||
output_tensor_range = absl::Substitute(R"(output_tensor_float_range {
|
||||
min: $0
|
||||
max: $1
|
||||
})",
|
||||
range_min, range_max);
|
||||
}
|
||||
auto graph_config = mediapipe::ParseTextProtoOrDie<CalculatorGraphConfig>(
|
||||
absl::Substitute(R"(
|
||||
input_stream: "input_image"
|
||||
@@ -86,22 +111,18 @@ void RunTestWithInputImagePacket(const Packet& input_image_packet,
|
||||
[mediapipe.ImageToTensorCalculatorOptions.ext] {
|
||||
output_tensor_width: $0
|
||||
output_tensor_height: $1
|
||||
keep_aspect_ratio: $4
|
||||
output_tensor_float_range {
|
||||
min: $2
|
||||
max: $3
|
||||
}
|
||||
$5 # border mode
|
||||
keep_aspect_ratio: $2
|
||||
$3 # output range
|
||||
$4 # border mode
|
||||
}
|
||||
}
|
||||
}
|
||||
)",
|
||||
/*$0=*/tensor_width,
|
||||
/*$1=*/tensor_height,
|
||||
/*$2=*/range_min,
|
||||
/*$3=*/range_max,
|
||||
/*$4=*/keep_aspect ? "true" : "false",
|
||||
/*$5=*/border_mode_str));
|
||||
/*$2=*/keep_aspect ? "true" : "false",
|
||||
/*$3=*/output_tensor_range,
|
||||
/*$4=*/border_mode_str));
|
||||
|
||||
std::vector<Packet> output_packets;
|
||||
tool::AddVectorSink("tensor", &graph_config, &output_packets);
|
||||
@@ -126,11 +147,24 @@ void RunTestWithInputImagePacket(const Packet& input_image_packet,
|
||||
ASSERT_THAT(tensor_vec, testing::SizeIs(1));
|
||||
|
||||
const Tensor& tensor = tensor_vec[0];
|
||||
EXPECT_EQ(tensor.element_type(), Tensor::ElementType::kFloat32);
|
||||
|
||||
auto view = tensor.GetCpuReadView();
|
||||
cv::Mat tensor_mat(tensor_height, tensor_width, CV_32FC3,
|
||||
const_cast<float*>(view.buffer<float>()));
|
||||
cv::Mat tensor_mat;
|
||||
if (output_int_tensor) {
|
||||
if (range_min < 0) {
|
||||
EXPECT_EQ(tensor.element_type(), Tensor::ElementType::kInt8);
|
||||
tensor_mat = cv::Mat(tensor_height, tensor_width, CV_8SC3,
|
||||
const_cast<int8*>(view.buffer<int8>()));
|
||||
} else {
|
||||
EXPECT_EQ(tensor.element_type(), Tensor::ElementType::kUInt8);
|
||||
tensor_mat = cv::Mat(tensor_height, tensor_width, CV_8UC3,
|
||||
const_cast<uint8*>(view.buffer<uint8>()));
|
||||
}
|
||||
} else {
|
||||
EXPECT_EQ(tensor.element_type(), Tensor::ElementType::kFloat32);
|
||||
tensor_mat = cv::Mat(tensor_height, tensor_width, CV_32FC3,
|
||||
const_cast<float*>(view.buffer<float>()));
|
||||
}
|
||||
|
||||
cv::Mat result_rgb;
|
||||
auto transformation =
|
||||
GetValueRangeTransformation(range_min, range_max, 0.0f, 255.0f).value();
|
||||
@@ -170,16 +204,29 @@ enum class InputType { kImageFrame, kImage };
|
||||
const std::vector<InputType> kInputTypesToTest = {InputType::kImageFrame,
|
||||
InputType::kImage};
|
||||
|
||||
void RunTest(cv::Mat input, cv::Mat expected_result, float range_min,
|
||||
float range_max, int tensor_width, int tensor_height,
|
||||
bool keep_aspect, absl::optional<BorderMode> border_mode,
|
||||
void RunTest(cv::Mat input, cv::Mat expected_result,
|
||||
std::vector<std::pair<float, float>> float_ranges,
|
||||
std::vector<std::pair<int, int>> int_ranges, int tensor_width,
|
||||
int tensor_height, bool keep_aspect,
|
||||
absl::optional<BorderMode> border_mode,
|
||||
const mediapipe::NormalizedRect& roi) {
|
||||
for (auto input_type : kInputTypesToTest) {
|
||||
RunTestWithInputImagePacket(
|
||||
input_type == InputType::kImageFrame ? MakeImageFramePacket(input)
|
||||
: MakeImagePacket(input),
|
||||
expected_result, range_min, range_max, tensor_width, tensor_height,
|
||||
keep_aspect, border_mode, roi);
|
||||
for (auto float_range : float_ranges) {
|
||||
RunTestWithInputImagePacket(
|
||||
input_type == InputType::kImageFrame ? MakeImageFramePacket(input)
|
||||
: MakeImagePacket(input),
|
||||
expected_result, float_range.first, float_range.second, tensor_width,
|
||||
tensor_height, keep_aspect, border_mode, roi,
|
||||
/*output_int_tensor=*/false);
|
||||
}
|
||||
for (auto int_range : int_ranges) {
|
||||
RunTestWithInputImagePacket(
|
||||
input_type == InputType::kImageFrame ? MakeImageFramePacket(input)
|
||||
: MakeImagePacket(input),
|
||||
expected_result, int_range.first, int_range.second, tensor_width,
|
||||
tensor_height, keep_aspect, border_mode, roi,
|
||||
/*output_int_tensor=*/true);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -195,8 +242,8 @@ TEST(ImageToTensorCalculatorTest, MediumSubRectKeepAspect) {
|
||||
"tensor/testdata/image_to_tensor/input.jpg"),
|
||||
GetRgb("/mediapipe/calculators/"
|
||||
"tensor/testdata/image_to_tensor/medium_sub_rect_keep_aspect.png"),
|
||||
/*range_min=*/0.0f,
|
||||
/*range_max=*/1.0f,
|
||||
/*float_ranges=*/{{0.0f, 1.0f}},
|
||||
/*int_ranges=*/{{0, 255}, {-128, 127}},
|
||||
/*tensor_width=*/256, /*tensor_height=*/256, /*keep_aspect=*/true,
|
||||
/*border mode*/ {}, roi);
|
||||
}
|
||||
@@ -213,8 +260,8 @@ TEST(ImageToTensorCalculatorTest, MediumSubRectKeepAspectBorderZero) {
|
||||
GetRgb("/mediapipe/calculators/"
|
||||
"tensor/testdata/image_to_tensor/"
|
||||
"medium_sub_rect_keep_aspect_border_zero.png"),
|
||||
/*range_min=*/0.0f,
|
||||
/*range_max=*/1.0f,
|
||||
/*float_ranges=*/{{0.0f, 1.0f}},
|
||||
/*int_ranges=*/{{0, 255}, {-128, 127}},
|
||||
/*tensor_width=*/256, /*tensor_height=*/256, /*keep_aspect=*/true,
|
||||
BorderMode::kZero, roi);
|
||||
}
|
||||
@@ -231,7 +278,8 @@ TEST(ImageToTensorCalculatorTest, MediumSubRectKeepAspectWithRotation) {
|
||||
GetRgb("/mediapipe/calculators/"
|
||||
"tensor/testdata/image_to_tensor/"
|
||||
"medium_sub_rect_keep_aspect_with_rotation.png"),
|
||||
/*range_min=*/0.0f, /*range_max=*/1.0f,
|
||||
/*float_ranges=*/{{0.0f, 1.0f}},
|
||||
/*int_ranges=*/{{0, 255}},
|
||||
/*tensor_width=*/256, /*tensor_height=*/256, /*keep_aspect=*/true,
|
||||
BorderMode::kReplicate, roi);
|
||||
}
|
||||
@@ -249,7 +297,8 @@ TEST(ImageToTensorCalculatorTest,
|
||||
GetRgb("/mediapipe/calculators/"
|
||||
"tensor/testdata/image_to_tensor/"
|
||||
"medium_sub_rect_keep_aspect_with_rotation_border_zero.png"),
|
||||
/*range_min=*/0.0f, /*range_max=*/1.0f,
|
||||
/*float_ranges=*/{{0.0f, 1.0f}},
|
||||
/*int_ranges=*/{{0, 255}, {-128, 127}},
|
||||
/*tensor_width=*/256, /*tensor_height=*/256, /*keep_aspect=*/true,
|
||||
BorderMode::kZero, roi);
|
||||
}
|
||||
@@ -267,8 +316,8 @@ TEST(ImageToTensorCalculatorTest, MediumSubRectWithRotation) {
|
||||
GetRgb(
|
||||
"/mediapipe/calculators/"
|
||||
"tensor/testdata/image_to_tensor/medium_sub_rect_with_rotation.png"),
|
||||
/*range_min=*/-1.0f,
|
||||
/*range_max=*/1.0f,
|
||||
/*float_ranges=*/{{-1.0f, 1.0f}},
|
||||
/*int_ranges=*/{{0, 255}, {-128, 127}},
|
||||
/*tensor_width=*/256, /*tensor_height=*/256, /*keep_aspect=*/false,
|
||||
BorderMode::kReplicate, roi);
|
||||
}
|
||||
@@ -285,8 +334,8 @@ TEST(ImageToTensorCalculatorTest, MediumSubRectWithRotationBorderZero) {
|
||||
GetRgb("/mediapipe/calculators/"
|
||||
"tensor/testdata/image_to_tensor/"
|
||||
"medium_sub_rect_with_rotation_border_zero.png"),
|
||||
/*range_min=*/-1.0f,
|
||||
/*range_max=*/1.0f,
|
||||
/*float_ranges=*/{{-1.0f, 1.0f}},
|
||||
/*int_ranges=*/{{0, 255}, {-128, 127}},
|
||||
/*tensor_width=*/256, /*tensor_height=*/256, /*keep_aspect=*/false,
|
||||
BorderMode::kZero, roi);
|
||||
}
|
||||
@@ -302,8 +351,8 @@ TEST(ImageToTensorCalculatorTest, LargeSubRect) {
|
||||
"tensor/testdata/image_to_tensor/input.jpg"),
|
||||
GetRgb("/mediapipe/calculators/"
|
||||
"tensor/testdata/image_to_tensor/large_sub_rect.png"),
|
||||
/*range_min=*/0.0f,
|
||||
/*range_max=*/1.0f,
|
||||
/*float_ranges=*/{{0.0f, 1.0f}},
|
||||
/*int_ranges=*/{{0, 255}},
|
||||
/*tensor_width=*/128, /*tensor_height=*/128, /*keep_aspect=*/false,
|
||||
BorderMode::kReplicate, roi);
|
||||
}
|
||||
@@ -320,8 +369,8 @@ TEST(ImageToTensorCalculatorTest, LargeSubRectBorderZero) {
|
||||
"tensor/testdata/image_to_tensor/input.jpg"),
|
||||
GetRgb("/mediapipe/calculators/"
|
||||
"tensor/testdata/image_to_tensor/large_sub_rect_border_zero.png"),
|
||||
/*range_min=*/0.0f,
|
||||
/*range_max=*/1.0f,
|
||||
/*float_ranges=*/{{0.0f, 1.0f}},
|
||||
/*int_ranges=*/{{0, 255}, {-128, 127}},
|
||||
/*tensor_width=*/128, /*tensor_height=*/128, /*keep_aspect=*/false,
|
||||
BorderMode::kZero, roi);
|
||||
}
|
||||
@@ -338,8 +387,8 @@ TEST(ImageToTensorCalculatorTest, LargeSubRectKeepAspect) {
|
||||
"tensor/testdata/image_to_tensor/input.jpg"),
|
||||
GetRgb("/mediapipe/calculators/"
|
||||
"tensor/testdata/image_to_tensor/large_sub_rect_keep_aspect.png"),
|
||||
/*range_min=*/0.0f,
|
||||
/*range_max=*/1.0f,
|
||||
/*float_ranges=*/{{0.0f, 1.0f}},
|
||||
/*int_ranges=*/{{0, 255}, {-128, 127}},
|
||||
/*tensor_width=*/128, /*tensor_height=*/128, /*keep_aspect=*/true,
|
||||
BorderMode::kReplicate, roi);
|
||||
}
|
||||
@@ -356,8 +405,8 @@ TEST(ImageToTensorCalculatorTest, LargeSubRectKeepAspectBorderZero) {
|
||||
GetRgb("/mediapipe/calculators/"
|
||||
"tensor/testdata/image_to_tensor/"
|
||||
"large_sub_rect_keep_aspect_border_zero.png"),
|
||||
/*range_min=*/0.0f,
|
||||
/*range_max=*/1.0f,
|
||||
/*float_ranges=*/{{0.0f, 1.0f}},
|
||||
/*int_ranges=*/{{0, 255}, {-128, 127}},
|
||||
/*tensor_width=*/128, /*tensor_height=*/128, /*keep_aspect=*/true,
|
||||
BorderMode::kZero, roi);
|
||||
}
|
||||
@@ -374,8 +423,8 @@ TEST(ImageToTensorCalculatorTest, LargeSubRectKeepAspectWithRotation) {
|
||||
GetRgb("/mediapipe/calculators/"
|
||||
"tensor/testdata/image_to_tensor/"
|
||||
"large_sub_rect_keep_aspect_with_rotation.png"),
|
||||
/*range_min=*/0.0f,
|
||||
/*range_max=*/1.0f,
|
||||
/*float_ranges=*/{{0.0f, 1.0f}},
|
||||
/*int_ranges=*/{{0, 255}, {-128, 127}},
|
||||
/*tensor_width=*/128, /*tensor_height=*/128, /*keep_aspect=*/true,
|
||||
/*border_mode=*/{}, roi);
|
||||
}
|
||||
@@ -393,8 +442,8 @@ TEST(ImageToTensorCalculatorTest,
|
||||
GetRgb("/mediapipe/calculators/"
|
||||
"tensor/testdata/image_to_tensor/"
|
||||
"large_sub_rect_keep_aspect_with_rotation_border_zero.png"),
|
||||
/*range_min=*/0.0f,
|
||||
/*range_max=*/1.0f,
|
||||
/*float_ranges=*/{{0.0f, 1.0f}},
|
||||
/*int_ranges=*/{{0, 255}},
|
||||
/*tensor_width=*/128, /*tensor_height=*/128, /*keep_aspect=*/true,
|
||||
/*border_mode=*/BorderMode::kZero, roi);
|
||||
}
|
||||
@@ -410,8 +459,8 @@ TEST(ImageToTensorCalculatorTest, NoOpExceptRange) {
|
||||
"tensor/testdata/image_to_tensor/input.jpg"),
|
||||
GetRgb("/mediapipe/calculators/"
|
||||
"tensor/testdata/image_to_tensor/noop_except_range.png"),
|
||||
/*range_min=*/0.0f,
|
||||
/*range_max=*/1.0f,
|
||||
/*float_ranges=*/{{0.0f, 1.0f}},
|
||||
/*int_ranges=*/{{0, 255}, {-128, 127}},
|
||||
/*tensor_width=*/64, /*tensor_height=*/128, /*keep_aspect=*/true,
|
||||
BorderMode::kReplicate, roi);
|
||||
}
|
||||
@@ -427,8 +476,8 @@ TEST(ImageToTensorCalculatorTest, NoOpExceptRangeBorderZero) {
|
||||
"tensor/testdata/image_to_tensor/input.jpg"),
|
||||
GetRgb("/mediapipe/calculators/"
|
||||
"tensor/testdata/image_to_tensor/noop_except_range.png"),
|
||||
/*range_min=*/0.0f,
|
||||
/*range_max=*/1.0f,
|
||||
/*float_ranges=*/{{0.0f, 1.0f}},
|
||||
/*int_ranges=*/{{0, 255}, {-128, 127}},
|
||||
/*tensor_width=*/64, /*tensor_height=*/128, /*keep_aspect=*/true,
|
||||
BorderMode::kZero, roi);
|
||||
}
|
||||
|
||||
@@ -268,10 +268,12 @@ class GlProcessor : public ImageToTensorConverter {
|
||||
const RotatedRect& roi,
|
||||
const Size& output_dims, float range_min,
|
||||
float range_max) override {
|
||||
if (input.format() != mediapipe::GpuBufferFormat::kBGRA32) {
|
||||
return InvalidArgumentError(
|
||||
absl::StrCat("Only BGRA/RGBA textures are supported, passed format: ",
|
||||
static_cast<uint32_t>(input.format())));
|
||||
if (input.format() != mediapipe::GpuBufferFormat::kBGRA32 &&
|
||||
input.format() != mediapipe::GpuBufferFormat::kRGBAHalf64 &&
|
||||
input.format() != mediapipe::GpuBufferFormat::kRGBAFloat128) {
|
||||
return InvalidArgumentError(absl::StrCat(
|
||||
"Only 4-channel texture input formats are supported, passed format: ",
|
||||
static_cast<uint32_t>(input.format())));
|
||||
}
|
||||
|
||||
constexpr int kNumChannels = 3;
|
||||
|
||||
@@ -16,7 +16,7 @@
|
||||
|
||||
#include "mediapipe/framework/port.h"
|
||||
|
||||
#if MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_20
|
||||
#if MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_30
|
||||
|
||||
#include <array>
|
||||
#include <memory>
|
||||
@@ -172,10 +172,12 @@ class GlProcessor : public ImageToTensorConverter {
|
||||
const RotatedRect& roi,
|
||||
const Size& output_dims, float range_min,
|
||||
float range_max) override {
|
||||
if (input.format() != mediapipe::GpuBufferFormat::kBGRA32) {
|
||||
return InvalidArgumentError(
|
||||
absl::StrCat("Only BGRA/RGBA textures are supported, passed format: ",
|
||||
static_cast<uint32_t>(input.format())));
|
||||
if (input.format() != mediapipe::GpuBufferFormat::kBGRA32 &&
|
||||
input.format() != mediapipe::GpuBufferFormat::kRGBAHalf64 &&
|
||||
input.format() != mediapipe::GpuBufferFormat::kRGBAFloat128) {
|
||||
return InvalidArgumentError(absl::StrCat(
|
||||
"Only 4-channel texture input formats are supported, passed format: ",
|
||||
static_cast<uint32_t>(input.format())));
|
||||
}
|
||||
|
||||
constexpr int kNumChannels = 3;
|
||||
@@ -339,4 +341,4 @@ CreateImageToGlTextureTensorConverter(CalculatorContext* cc,
|
||||
|
||||
} // namespace mediapipe
|
||||
|
||||
#endif // MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_20
|
||||
#endif // MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_30
|
||||
|
||||
@@ -17,7 +17,7 @@
|
||||
|
||||
#include "mediapipe/framework/port.h"
|
||||
|
||||
#if MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_20
|
||||
#if MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_30
|
||||
|
||||
#include <memory>
|
||||
|
||||
@@ -37,6 +37,6 @@ CreateImageToGlTextureTensorConverter(CalculatorContext* cc,
|
||||
|
||||
} // namespace mediapipe
|
||||
|
||||
#endif // MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_20
|
||||
#endif // MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_30
|
||||
|
||||
#endif // MEDIAPIPE_CALCULATORS_TENSOR_IMAGE_TO_TENSOR_CONVERTER_GL_TEXTURE_H_
|
||||
|
||||
@@ -2,7 +2,7 @@
|
||||
|
||||
#include "mediapipe/framework/port.h"
|
||||
|
||||
#if MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_20
|
||||
#if MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_30
|
||||
|
||||
#include <array>
|
||||
#include <memory>
|
||||
@@ -85,4 +85,4 @@ bool IsGlClampToBorderSupported(const mediapipe::GlContext& gl_context) {
|
||||
|
||||
} // namespace mediapipe
|
||||
|
||||
#endif // MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_20
|
||||
#endif // MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_30
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
|
||||
#include "mediapipe/framework/port.h"
|
||||
|
||||
#if MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_20
|
||||
#if MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_30
|
||||
|
||||
#include <array>
|
||||
#include <memory>
|
||||
@@ -40,6 +40,6 @@ bool IsGlClampToBorderSupported(const mediapipe::GlContext& gl_context);
|
||||
|
||||
} // namespace mediapipe
|
||||
|
||||
#endif // MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_20
|
||||
#endif // MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_30
|
||||
|
||||
#endif // MEDIAPIPE_CALCULATORS_TENSOR_IMAGE_TO_TENSOR_CONVERTER_GL_UTILS_H_
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
#include "mediapipe/framework/port.h"
|
||||
|
||||
#if MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_20
|
||||
#if MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_30
|
||||
|
||||
#include "mediapipe/calculators/tensor/image_to_tensor_converter_gl_utils.h"
|
||||
#include "mediapipe/framework/port/gtest.h"
|
||||
@@ -46,4 +46,4 @@ TEST(ImageToTensorConverterGlUtilsTest, GlTexParameteriOverrider) {
|
||||
} // namespace
|
||||
} // namespace mediapipe
|
||||
|
||||
#endif // MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_20
|
||||
#endif // MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_30
|
||||
|
||||
@@ -352,11 +352,12 @@ class MetalProcessor : public ImageToTensorConverter {
|
||||
const RotatedRect& roi,
|
||||
const Size& output_dims, float range_min,
|
||||
float range_max) override {
|
||||
if (input.format() != mediapipe::GpuBufferFormat::kBGRA32) {
|
||||
return InvalidArgumentError(
|
||||
absl::StrCat("Only BGRA/RGBA textures are supported, passed "
|
||||
"format: ",
|
||||
static_cast<uint32_t>(input.format())));
|
||||
if (input.format() != mediapipe::GpuBufferFormat::kBGRA32 &&
|
||||
input.format() != mediapipe::GpuBufferFormat::kRGBAHalf64 &&
|
||||
input.format() != mediapipe::GpuBufferFormat::kRGBAFloat128) {
|
||||
return InvalidArgumentError(absl::StrCat(
|
||||
"Only 4-channel texture input formats are supported, passed format: ",
|
||||
static_cast<uint32_t>(input.format())));
|
||||
}
|
||||
|
||||
@autoreleasepool {
|
||||
|
||||
@@ -35,7 +35,8 @@ namespace {
|
||||
|
||||
class OpenCvProcessor : public ImageToTensorConverter {
|
||||
public:
|
||||
OpenCvProcessor(BorderMode border_mode) {
|
||||
OpenCvProcessor(BorderMode border_mode, Tensor::ElementType tensor_type)
|
||||
: tensor_type_(tensor_type) {
|
||||
switch (border_mode) {
|
||||
case BorderMode::kReplicate:
|
||||
border_mode_ = cv::BORDER_REPLICATE;
|
||||
@@ -44,6 +45,19 @@ class OpenCvProcessor : public ImageToTensorConverter {
|
||||
border_mode_ = cv::BORDER_CONSTANT;
|
||||
break;
|
||||
}
|
||||
switch (tensor_type_) {
|
||||
case Tensor::ElementType::kInt8:
|
||||
mat_type_ = CV_8SC3;
|
||||
break;
|
||||
case Tensor::ElementType::kFloat32:
|
||||
mat_type_ = CV_32FC3;
|
||||
break;
|
||||
case Tensor::ElementType::kUInt8:
|
||||
mat_type_ = CV_8UC3;
|
||||
break;
|
||||
default:
|
||||
mat_type_ = -1;
|
||||
}
|
||||
}
|
||||
|
||||
absl::StatusOr<Tensor> Convert(const mediapipe::Image& input,
|
||||
@@ -56,15 +70,30 @@ class OpenCvProcessor : public ImageToTensorConverter {
|
||||
absl::StrCat("Only RGBA/RGB formats are supported, passed format: ",
|
||||
static_cast<uint32_t>(input.image_format())));
|
||||
}
|
||||
cv::Mat src = mediapipe::formats::MatView(&input);
|
||||
auto src = mediapipe::formats::MatView(&input);
|
||||
|
||||
constexpr int kNumChannels = 3;
|
||||
Tensor tensor(
|
||||
Tensor::ElementType::kFloat32,
|
||||
Tensor::Shape{1, output_dims.height, output_dims.width, kNumChannels});
|
||||
Tensor tensor(tensor_type_, Tensor::Shape{1, output_dims.height,
|
||||
output_dims.width, kNumChannels});
|
||||
auto buffer_view = tensor.GetCpuWriteView();
|
||||
cv::Mat dst(output_dims.height, output_dims.width, CV_32FC3,
|
||||
buffer_view.buffer<float>());
|
||||
cv::Mat dst;
|
||||
switch (tensor_type_) {
|
||||
case Tensor::ElementType::kInt8:
|
||||
dst = cv::Mat(output_dims.height, output_dims.width, mat_type_,
|
||||
buffer_view.buffer<int8>());
|
||||
break;
|
||||
case Tensor::ElementType::kFloat32:
|
||||
dst = cv::Mat(output_dims.height, output_dims.width, mat_type_,
|
||||
buffer_view.buffer<float>());
|
||||
break;
|
||||
case Tensor::ElementType::kUInt8:
|
||||
dst = cv::Mat(output_dims.height, output_dims.width, mat_type_,
|
||||
buffer_view.buffer<uint8>());
|
||||
break;
|
||||
default:
|
||||
return InvalidArgumentError(
|
||||
absl::StrCat("Unsupported tensor type: ", tensor_type_));
|
||||
}
|
||||
|
||||
const cv::RotatedRect rotated_rect(cv::Point2f(roi.center_x, roi.center_y),
|
||||
cv::Size2f(roi.width, roi.height),
|
||||
@@ -85,7 +114,7 @@ class OpenCvProcessor : public ImageToTensorConverter {
|
||||
cv::Mat projection_matrix =
|
||||
cv::getPerspectiveTransform(src_points, dst_points);
|
||||
cv::Mat transformed;
|
||||
cv::warpPerspective(src, transformed, projection_matrix,
|
||||
cv::warpPerspective(*src, transformed, projection_matrix,
|
||||
cv::Size(dst_width, dst_height),
|
||||
/*flags=*/cv::INTER_LINEAR,
|
||||
/*borderMode=*/border_mode_);
|
||||
@@ -102,19 +131,29 @@ class OpenCvProcessor : public ImageToTensorConverter {
|
||||
auto transform,
|
||||
GetValueRangeTransformation(kInputImageRangeMin, kInputImageRangeMax,
|
||||
range_min, range_max));
|
||||
transformed.convertTo(dst, CV_32FC3, transform.scale, transform.offset);
|
||||
transformed.convertTo(dst, mat_type_, transform.scale, transform.offset);
|
||||
return tensor;
|
||||
}
|
||||
|
||||
private:
|
||||
enum cv::BorderTypes border_mode_;
|
||||
Tensor::ElementType tensor_type_;
|
||||
int mat_type_;
|
||||
};
|
||||
|
||||
} // namespace
|
||||
|
||||
absl::StatusOr<std::unique_ptr<ImageToTensorConverter>> CreateOpenCvConverter(
|
||||
CalculatorContext* cc, BorderMode border_mode) {
|
||||
return absl::make_unique<OpenCvProcessor>(border_mode);
|
||||
CalculatorContext* cc, BorderMode border_mode,
|
||||
Tensor::ElementType tensor_type) {
|
||||
if (tensor_type != Tensor::ElementType::kInt8 &&
|
||||
tensor_type != Tensor::ElementType::kFloat32 &&
|
||||
tensor_type != Tensor::ElementType::kUInt8) {
|
||||
return absl::InvalidArgumentError(absl::StrCat(
|
||||
"Tensor type is currently not supported by OpenCvProcessor, type: ",
|
||||
tensor_type));
|
||||
}
|
||||
return absl::make_unique<OpenCvProcessor>(border_mode, tensor_type);
|
||||
}
|
||||
|
||||
} // namespace mediapipe
|
||||
|
||||
@@ -25,7 +25,8 @@ namespace mediapipe {
|
||||
|
||||
// Creates OpenCV image-to-tensor converter.
|
||||
absl::StatusOr<std::unique_ptr<ImageToTensorConverter>> CreateOpenCvConverter(
|
||||
CalculatorContext* cc, BorderMode border_mode);
|
||||
CalculatorContext* cc, BorderMode border_mode,
|
||||
Tensor::ElementType tensor_type);
|
||||
|
||||
} // namespace mediapipe
|
||||
|
||||
|
||||
@@ -19,9 +19,11 @@
|
||||
#include <string>
|
||||
#include <vector>
|
||||
|
||||
#include "absl/memory/memory.h"
|
||||
#include "absl/strings/string_view.h"
|
||||
#include "mediapipe/calculators/tensor/inference_calculator.pb.h"
|
||||
#include "mediapipe/framework/api2/packet.h"
|
||||
#include "mediapipe/framework/tool/subgraph_expansion.h"
|
||||
#include "tensorflow/lite/core/api/op_resolver.h"
|
||||
|
||||
namespace mediapipe {
|
||||
namespace api2 {
|
||||
@@ -36,12 +38,24 @@ class InferenceCalculatorSelectorImpl
|
||||
Subgraph::GetOptions<mediapipe::InferenceCalculatorOptions>(
|
||||
subgraph_node);
|
||||
std::vector<absl::string_view> impls;
|
||||
|
||||
const bool should_use_gpu =
|
||||
!options.has_delegate() || // Use GPU delegate if not specified
|
||||
(options.has_delegate() && options.delegate().has_gpu());
|
||||
if (should_use_gpu) {
|
||||
const auto& api = options.delegate().gpu().api();
|
||||
using Gpu = ::mediapipe::InferenceCalculatorOptions::Delegate::Gpu;
|
||||
impls.emplace_back("Metal");
|
||||
impls.emplace_back("Gl");
|
||||
const bool prefer_gl_advanced =
|
||||
options.delegate().gpu().use_advanced_gpu_api() &&
|
||||
(api == Gpu::ANY || api == Gpu::OPENGL || api == Gpu::OPENCL);
|
||||
if (prefer_gl_advanced) {
|
||||
impls.emplace_back("GlAdvanced");
|
||||
impls.emplace_back("Gl");
|
||||
} else {
|
||||
impls.emplace_back("Gl");
|
||||
impls.emplace_back("GlAdvanced");
|
||||
}
|
||||
}
|
||||
impls.emplace_back("Cpu");
|
||||
for (const auto& suffix : impls) {
|
||||
@@ -66,5 +80,17 @@ absl::StatusOr<Packet<TfLiteModelPtr>> InferenceCalculator::GetModelAsPacket(
|
||||
"Must specify TFLite model as path or loaded model.");
|
||||
}
|
||||
|
||||
absl::StatusOr<Packet<tflite::OpResolver>>
|
||||
InferenceCalculator::GetOpResolverAsPacket(CalculatorContext* cc) {
|
||||
if (kSideInOpResolver(cc).IsConnected()) {
|
||||
return kSideInOpResolver(cc).As<tflite::OpResolver>();
|
||||
} else if (kSideInCustomOpResolver(cc).IsConnected()) {
|
||||
return kSideInCustomOpResolver(cc).As<tflite::OpResolver>();
|
||||
}
|
||||
return PacketAdopting<tflite::OpResolver>(
|
||||
std::make_unique<
|
||||
tflite::ops::builtin::BuiltinOpResolverWithoutDefaultDelegates>());
|
||||
}
|
||||
|
||||
} // namespace api2
|
||||
} // namespace mediapipe
|
||||
|
||||
@@ -27,6 +27,7 @@
|
||||
#include "mediapipe/framework/formats/tensor.h"
|
||||
#include "mediapipe/framework/port/ret_check.h"
|
||||
#include "mediapipe/util/tflite/tflite_model_loader.h"
|
||||
#include "tensorflow/lite/core/api/op_resolver.h"
|
||||
#include "tensorflow/lite/error_reporter.h"
|
||||
#include "tensorflow/lite/interpreter.h"
|
||||
#include "tensorflow/lite/kernels/register.h"
|
||||
@@ -55,8 +56,11 @@ namespace api2 {
|
||||
// TENSORS - Vector of Tensors
|
||||
//
|
||||
// Input side packet:
|
||||
// DEPRECATED: Prefer to use the "OP_RESOLVER" input side packet instead.
|
||||
// CUSTOM_OP_RESOLVER (optional) - Use a custom op resolver,
|
||||
// instead of the builtin one.
|
||||
// OP_RESOLVER (optional) - Use to provide tflite op resolver
|
||||
// (tflite::OpResolver)
|
||||
// MODEL (optional) - Use to specify TfLite model
|
||||
// (std::unique_ptr<tflite::FlatBufferModel,
|
||||
// std::function<void(tflite::FlatBufferModel*)>>)
|
||||
@@ -95,15 +99,21 @@ namespace api2 {
|
||||
class InferenceCalculator : public NodeIntf {
|
||||
public:
|
||||
static constexpr Input<std::vector<Tensor>> kInTensors{"TENSORS"};
|
||||
// Deprecated. Prefers to use "OP_RESOLVER" input side packet instead.
|
||||
// TODO: Removes the "CUSTOM_OP_RESOLVER" side input after the
|
||||
// migration.
|
||||
static constexpr SideInput<tflite::ops::builtin::BuiltinOpResolver>::Optional
|
||||
kSideInCustomOpResolver{"CUSTOM_OP_RESOLVER"};
|
||||
static constexpr SideInput<tflite::OpResolver>::Optional kSideInOpResolver{
|
||||
"OP_RESOLVER"};
|
||||
static constexpr SideInput<TfLiteModelPtr>::Optional kSideInModel{"MODEL"};
|
||||
static constexpr Output<std::vector<Tensor>> kOutTensors{"TENSORS"};
|
||||
static constexpr SideInput<
|
||||
mediapipe::InferenceCalculatorOptions::Delegate>::Optional kDelegate{
|
||||
"DELEGATE"};
|
||||
MEDIAPIPE_NODE_CONTRACT(kInTensors, kSideInCustomOpResolver, kSideInModel,
|
||||
kOutTensors, kDelegate);
|
||||
MEDIAPIPE_NODE_CONTRACT(kInTensors, kSideInCustomOpResolver,
|
||||
kSideInOpResolver, kSideInModel, kOutTensors,
|
||||
kDelegate);
|
||||
|
||||
protected:
|
||||
using TfLiteDelegatePtr =
|
||||
@@ -111,6 +121,9 @@ class InferenceCalculator : public NodeIntf {
|
||||
|
||||
absl::StatusOr<Packet<TfLiteModelPtr>> GetModelAsPacket(
|
||||
CalculatorContext* cc);
|
||||
|
||||
absl::StatusOr<Packet<tflite::OpResolver>> GetOpResolverAsPacket(
|
||||
CalculatorContext* cc);
|
||||
};
|
||||
|
||||
struct InferenceCalculatorSelector : public InferenceCalculator {
|
||||
@@ -121,6 +134,10 @@ struct InferenceCalculatorGl : public InferenceCalculator {
|
||||
static constexpr char kCalculatorName[] = "InferenceCalculatorGl";
|
||||
};
|
||||
|
||||
struct InferenceCalculatorGlAdvanced : public InferenceCalculator {
|
||||
static constexpr char kCalculatorName[] = "InferenceCalculatorGlAdvanced";
|
||||
};
|
||||
|
||||
struct InferenceCalculatorMetal : public InferenceCalculator {
|
||||
static constexpr char kCalculatorName[] = "InferenceCalculatorMetal";
|
||||
};
|
||||
|
||||
@@ -116,6 +116,9 @@ message InferenceCalculatorOptions {
|
||||
// to ensure there is no clash of the tokens. If unspecified, NNAPI will
|
||||
// not try caching the compilation.
|
||||
optional string model_token = 2;
|
||||
// The name of an accelerator to be used for NNAPI delegate, e.g.
|
||||
// "google-edgetpu". When not specified, it will be selected by NNAPI.
|
||||
optional string accelerator_name = 3;
|
||||
}
|
||||
message Xnnpack {
|
||||
// Number of threads for XNNPACK delegate. (By default, calculator tries
|
||||
|
||||
@@ -19,7 +19,7 @@
|
||||
|
||||
#include "absl/memory/memory.h"
|
||||
#include "mediapipe/calculators/tensor/inference_calculator.h"
|
||||
|
||||
#include "tensorflow/lite/interpreter_builder.h"
|
||||
#if defined(MEDIAPIPE_ANDROID)
|
||||
#include "tensorflow/lite/delegates/nnapi/nnapi_delegate.h"
|
||||
#endif // ANDROID
|
||||
@@ -28,6 +28,7 @@
|
||||
#include "mediapipe/util/cpu_util.h"
|
||||
#endif // !__EMSCRIPTEN__ || __EMSCRIPTEN_PTHREADS__
|
||||
|
||||
#include "tensorflow/lite/c/c_api_types.h"
|
||||
#include "tensorflow/lite/delegates/xnnpack/xnnpack_delegate.h"
|
||||
|
||||
namespace mediapipe {
|
||||
@@ -61,6 +62,17 @@ int GetXnnpackNumThreads(
|
||||
return GetXnnpackDefaultNumThreads();
|
||||
}
|
||||
|
||||
template <typename T>
|
||||
void CopyTensorBuffer(const Tensor& input_tensor,
|
||||
tflite::Interpreter* interpreter,
|
||||
int input_tensor_index) {
|
||||
auto input_tensor_view = input_tensor.GetCpuReadView();
|
||||
auto input_tensor_buffer = input_tensor_view.buffer<T>();
|
||||
T* local_tensor_buffer =
|
||||
interpreter->typed_input_tensor<T>(input_tensor_index);
|
||||
std::memcpy(local_tensor_buffer, input_tensor_buffer, input_tensor.bytes());
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
class InferenceCalculatorCpuImpl
|
||||
@@ -73,14 +85,16 @@ class InferenceCalculatorCpuImpl
|
||||
absl::Status Close(CalculatorContext* cc) override;
|
||||
|
||||
private:
|
||||
absl::Status LoadModel(CalculatorContext* cc);
|
||||
absl::Status LoadDelegate(CalculatorContext* cc);
|
||||
absl::Status LoadDelegateAndAllocateTensors(CalculatorContext* cc);
|
||||
absl::Status InitInterpreter(CalculatorContext* cc);
|
||||
absl::Status LoadDelegate(CalculatorContext* cc,
|
||||
tflite::InterpreterBuilder* interpreter_builder);
|
||||
absl::Status AllocateTensors();
|
||||
|
||||
// TfLite requires us to keep the model alive as long as the interpreter is.
|
||||
Packet<TfLiteModelPtr> model_packet_;
|
||||
std::unique_ptr<tflite::Interpreter> interpreter_;
|
||||
TfLiteDelegatePtr delegate_;
|
||||
TfLiteType input_tensor_type_ = TfLiteType::kTfLiteNoType;
|
||||
};
|
||||
|
||||
absl::Status InferenceCalculatorCpuImpl::UpdateContract(
|
||||
@@ -93,8 +107,7 @@ absl::Status InferenceCalculatorCpuImpl::UpdateContract(
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorCpuImpl::Open(CalculatorContext* cc) {
|
||||
MP_RETURN_IF_ERROR(LoadModel(cc));
|
||||
return LoadDelegateAndAllocateTensors(cc);
|
||||
return InitInterpreter(cc);
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorCpuImpl::Process(CalculatorContext* cc) {
|
||||
@@ -107,12 +120,24 @@ absl::Status InferenceCalculatorCpuImpl::Process(CalculatorContext* cc) {
|
||||
|
||||
// Read CPU input into tensors.
|
||||
for (int i = 0; i < input_tensors.size(); ++i) {
|
||||
const Tensor* input_tensor = &input_tensors[i];
|
||||
auto input_tensor_view = input_tensor->GetCpuReadView();
|
||||
auto input_tensor_buffer = input_tensor_view.buffer<float>();
|
||||
float* local_tensor_buffer = interpreter_->typed_input_tensor<float>(i);
|
||||
std::memcpy(local_tensor_buffer, input_tensor_buffer,
|
||||
input_tensor->bytes());
|
||||
switch (input_tensor_type_) {
|
||||
case TfLiteType::kTfLiteFloat16:
|
||||
case TfLiteType::kTfLiteFloat32: {
|
||||
CopyTensorBuffer<float>(input_tensors[i], interpreter_.get(), i);
|
||||
break;
|
||||
}
|
||||
case TfLiteType::kTfLiteUInt8: {
|
||||
CopyTensorBuffer<uint8>(input_tensors[i], interpreter_.get(), i);
|
||||
break;
|
||||
}
|
||||
case TfLiteType::kTfLiteInt8: {
|
||||
CopyTensorBuffer<int8>(input_tensors[i], interpreter_.get(), i);
|
||||
break;
|
||||
}
|
||||
default:
|
||||
return absl::InvalidArgumentError(
|
||||
absl::StrCat("Unsupported input tensor type:", input_tensor_type_));
|
||||
}
|
||||
}
|
||||
|
||||
// Run inference.
|
||||
@@ -141,40 +166,34 @@ absl::Status InferenceCalculatorCpuImpl::Close(CalculatorContext* cc) {
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorCpuImpl::LoadModel(CalculatorContext* cc) {
|
||||
absl::Status InferenceCalculatorCpuImpl::InitInterpreter(
|
||||
CalculatorContext* cc) {
|
||||
ASSIGN_OR_RETURN(model_packet_, GetModelAsPacket(cc));
|
||||
const auto& model = *model_packet_.Get();
|
||||
tflite::ops::builtin::BuiltinOpResolver op_resolver =
|
||||
kSideInCustomOpResolver(cc).GetOr(
|
||||
tflite::ops::builtin::BuiltinOpResolverWithoutDefaultDelegates());
|
||||
|
||||
tflite::InterpreterBuilder(model, op_resolver)(&interpreter_);
|
||||
RET_CHECK(interpreter_);
|
||||
|
||||
ASSIGN_OR_RETURN(auto op_resolver_packet, GetOpResolverAsPacket(cc));
|
||||
const auto& op_resolver = op_resolver_packet.Get();
|
||||
tflite::InterpreterBuilder interpreter_builder(model, op_resolver);
|
||||
MP_RETURN_IF_ERROR(LoadDelegate(cc, &interpreter_builder));
|
||||
#if defined(__EMSCRIPTEN__)
|
||||
interpreter_->SetNumThreads(1);
|
||||
interpreter_builder.SetNumThreads(1);
|
||||
#else
|
||||
interpreter_->SetNumThreads(
|
||||
interpreter_builder.SetNumThreads(
|
||||
cc->Options<mediapipe::InferenceCalculatorOptions>().cpu_num_thread());
|
||||
#endif // __EMSCRIPTEN__
|
||||
|
||||
return absl::OkStatus();
|
||||
RET_CHECK_EQ(interpreter_builder(&interpreter_), kTfLiteOk);
|
||||
RET_CHECK(interpreter_);
|
||||
return AllocateTensors();
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorCpuImpl::LoadDelegateAndAllocateTensors(
|
||||
CalculatorContext* cc) {
|
||||
MP_RETURN_IF_ERROR(LoadDelegate(cc));
|
||||
|
||||
// AllocateTensors() can be called only after ModifyGraphWithDelegate.
|
||||
absl::Status InferenceCalculatorCpuImpl::AllocateTensors() {
|
||||
RET_CHECK_EQ(interpreter_->AllocateTensors(), kTfLiteOk);
|
||||
// TODO: Support quantized tensors.
|
||||
RET_CHECK_NE(
|
||||
interpreter_->tensor(interpreter_->inputs()[0])->quantization.type,
|
||||
kTfLiteAffineQuantization);
|
||||
input_tensor_type_ = interpreter_->tensor(interpreter_->inputs()[0])->type;
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorCpuImpl::LoadDelegate(CalculatorContext* cc) {
|
||||
absl::Status InferenceCalculatorCpuImpl::LoadDelegate(
|
||||
CalculatorContext* cc, tflite::InterpreterBuilder* interpreter_builder) {
|
||||
const auto& calculator_opts =
|
||||
cc->Options<mediapipe::InferenceCalculatorOptions>();
|
||||
auto opts_delegate = calculator_opts.delegate();
|
||||
@@ -203,18 +222,20 @@ absl::Status InferenceCalculatorCpuImpl::LoadDelegate(CalculatorContext* cc) {
|
||||
if (nnapi_requested) {
|
||||
// Attempt to use NNAPI.
|
||||
// If not supported, the default CPU delegate will be created and used.
|
||||
interpreter_->SetAllowFp16PrecisionForFp32(1);
|
||||
tflite::StatefulNnApiDelegate::Options options;
|
||||
const auto& nnapi = opts_delegate.nnapi();
|
||||
options.allow_fp16 = true;
|
||||
// Set up cache_dir and model_token for NNAPI compilation cache.
|
||||
options.cache_dir =
|
||||
nnapi.has_cache_dir() ? nnapi.cache_dir().c_str() : nullptr;
|
||||
options.model_token =
|
||||
nnapi.has_model_token() ? nnapi.model_token().c_str() : nullptr;
|
||||
options.accelerator_name = nnapi.has_accelerator_name()
|
||||
? nnapi.accelerator_name().c_str()
|
||||
: nullptr;
|
||||
delegate_ = TfLiteDelegatePtr(new tflite::StatefulNnApiDelegate(options),
|
||||
[](TfLiteDelegate*) {});
|
||||
RET_CHECK_EQ(interpreter_->ModifyGraphWithDelegate(delegate_.get()),
|
||||
kTfLiteOk);
|
||||
interpreter_builder->AddDelegate(delegate_.get());
|
||||
return absl::OkStatus();
|
||||
}
|
||||
#endif // MEDIAPIPE_ANDROID
|
||||
@@ -226,13 +247,12 @@ absl::Status InferenceCalculatorCpuImpl::LoadDelegate(CalculatorContext* cc) {
|
||||
#endif // defined(__EMSCRIPTEN__)
|
||||
|
||||
if (use_xnnpack) {
|
||||
TfLiteXNNPackDelegateOptions xnnpack_opts{};
|
||||
auto xnnpack_opts = TfLiteXNNPackDelegateOptionsDefault();
|
||||
xnnpack_opts.num_threads =
|
||||
GetXnnpackNumThreads(opts_has_delegate, opts_delegate);
|
||||
delegate_ = TfLiteDelegatePtr(TfLiteXNNPackDelegateCreate(&xnnpack_opts),
|
||||
&TfLiteXNNPackDelegateDelete);
|
||||
RET_CHECK_EQ(interpreter_->ModifyGraphWithDelegate(delegate_.get()),
|
||||
kTfLiteOk);
|
||||
interpreter_builder->AddDelegate(delegate_.get());
|
||||
}
|
||||
|
||||
return absl::OkStatus();
|
||||
|
||||
@@ -75,9 +75,10 @@ const std::vector<Param>& GetParams() {
|
||||
class InferenceCalculatorTest : public testing::TestWithParam<Param> {
|
||||
protected:
|
||||
void SetDelegateForParam(mediapipe::CalculatorGraphConfig_Node* node) {
|
||||
*node->mutable_options()
|
||||
->MutableExtension(mediapipe::InferenceCalculatorOptions::ext)
|
||||
->mutable_delegate() = GetParam().delegate;
|
||||
auto options_map = tool::MutableOptionsMap().Initialize(*node);
|
||||
auto options = options_map.Get<mediapipe::InferenceCalculatorOptions>();
|
||||
*options.mutable_delegate() = GetParam().delegate;
|
||||
options_map.Set(options);
|
||||
}
|
||||
};
|
||||
|
||||
@@ -154,8 +155,9 @@ TEST_P(InferenceCalculatorTest, TestFaceDetection) {
|
||||
detection_packets[0].Get<std::vector<Detection>>();
|
||||
#if !defined(MEDIAPIPE_PROTO_LITE)
|
||||
// Approximately is not available with lite protos (b/178137094).
|
||||
EXPECT_THAT(dets,
|
||||
ElementsAre(Approximately(EqualsProto(expected_detection))));
|
||||
constexpr float kEpison = 0.001;
|
||||
EXPECT_THAT(dets, ElementsAre(Approximately(EqualsProto(expected_detection),
|
||||
kEpison)));
|
||||
#endif
|
||||
}
|
||||
|
||||
|
||||
@@ -20,22 +20,8 @@
|
||||
#include "absl/memory/memory.h"
|
||||
#include "absl/status/status.h"
|
||||
#include "mediapipe/calculators/tensor/inference_calculator.h"
|
||||
#include "mediapipe/framework/deps/file_path.h"
|
||||
#include "mediapipe/util/tflite/config.h"
|
||||
|
||||
#if MEDIAPIPE_TFLITE_GL_INFERENCE
|
||||
#include "mediapipe/gpu/gl_calculator_helper.h"
|
||||
#include "mediapipe/gpu/gpu_buffer.h"
|
||||
#include "mediapipe/util/tflite/tflite_gpu_runner.h"
|
||||
#include "tensorflow/lite/delegates/gpu/common/shape.h"
|
||||
#include "tensorflow/lite/delegates/gpu/gl_delegate.h"
|
||||
#endif // MEDIAPIPE_TFLITE_GL_INFERENCE
|
||||
|
||||
#if defined(MEDIAPIPE_ANDROID)
|
||||
#include "mediapipe/util/android/file/base/file.h"
|
||||
#include "mediapipe/util/android/file/base/filesystem.h"
|
||||
#include "mediapipe/util/android/file/base/helpers.h"
|
||||
#endif // ANDROID
|
||||
|
||||
namespace mediapipe {
|
||||
namespace api2 {
|
||||
@@ -50,41 +36,22 @@ class InferenceCalculatorGlImpl
|
||||
absl::Status Close(CalculatorContext* cc) override;
|
||||
|
||||
private:
|
||||
absl::Status ReadGpuCaches();
|
||||
absl::Status SaveGpuCaches();
|
||||
absl::Status LoadModel(CalculatorContext* cc);
|
||||
absl::Status LoadDelegate(CalculatorContext* cc);
|
||||
absl::Status LoadDelegateAndAllocateTensors(CalculatorContext* cc);
|
||||
absl::Status InitTFLiteGPURunner(CalculatorContext* cc);
|
||||
|
||||
// TfLite requires us to keep the model alive as long as the interpreter is.
|
||||
Packet<TfLiteModelPtr> model_packet_;
|
||||
std::unique_ptr<tflite::Interpreter> interpreter_;
|
||||
TfLiteDelegatePtr delegate_;
|
||||
|
||||
#if MEDIAPIPE_TFLITE_GL_INFERENCE
|
||||
mediapipe::GlCalculatorHelper gpu_helper_;
|
||||
std::unique_ptr<tflite::gpu::TFLiteGPURunner> tflite_gpu_runner_;
|
||||
bool allow_precision_loss_ = false;
|
||||
mediapipe::InferenceCalculatorOptions::Delegate::Gpu::Api
|
||||
tflite_gpu_runner_api_;
|
||||
mediapipe::InferenceCalculatorOptions::Delegate::Gpu::InferenceUsage
|
||||
tflite_gpu_runner_usage_;
|
||||
#endif // MEDIAPIPE_TFLITE_GL_INFERENCE
|
||||
|
||||
#if MEDIAPIPE_TFLITE_GPU_SUPPORTED
|
||||
TfLiteDelegatePtr delegate_;
|
||||
std::unique_ptr<tflite::Interpreter> interpreter_;
|
||||
|
||||
std::vector<Tensor::Shape> output_shapes_;
|
||||
std::vector<std::unique_ptr<Tensor>> gpu_buffers_in_;
|
||||
std::vector<std::unique_ptr<Tensor>> gpu_buffers_out_;
|
||||
#endif // MEDIAPIPE_TFLITE_GPU_SUPPORTED
|
||||
|
||||
bool use_advanced_gpu_api_ = false;
|
||||
bool use_gpu_delegate_ = false;
|
||||
|
||||
bool use_kernel_caching_ = false;
|
||||
std::string cached_kernel_filename_;
|
||||
bool use_serialized_model_ = false;
|
||||
std::string serialized_model_path_;
|
||||
};
|
||||
|
||||
absl::Status InferenceCalculatorGlImpl::UpdateContract(CalculatorContract* cc) {
|
||||
@@ -92,8 +59,7 @@ absl::Status InferenceCalculatorGlImpl::UpdateContract(CalculatorContract* cc) {
|
||||
RET_CHECK(!options.model_path().empty() ^ kSideInModel(cc).IsConnected())
|
||||
<< "Either model as side packet or model path in options is required.";
|
||||
|
||||
MP_RETURN_IF_ERROR(mediapipe::GlCalculatorHelper::UpdateContract(cc));
|
||||
return absl::OkStatus();
|
||||
return mediapipe::GlCalculatorHelper::UpdateContract(cc);
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorGlImpl::Open(CalculatorContext* cc) {
|
||||
@@ -109,46 +75,12 @@ absl::Status InferenceCalculatorGlImpl::Open(CalculatorContext* cc) {
|
||||
<< "for Gpu";
|
||||
delegate.MergeFrom(input_side_packet_delegate);
|
||||
}
|
||||
const bool has_delegate = options.has_delegate() || !kDelegate(cc).IsEmpty();
|
||||
use_advanced_gpu_api_ = has_delegate && delegate.has_gpu() &&
|
||||
delegate.gpu().use_advanced_gpu_api();
|
||||
allow_precision_loss_ = delegate.gpu().allow_precision_loss();
|
||||
tflite_gpu_runner_api_ = delegate.gpu().api();
|
||||
tflite_gpu_runner_usage_ = delegate.gpu().usage();
|
||||
use_kernel_caching_ =
|
||||
use_advanced_gpu_api_ && delegate.gpu().has_cached_kernel_path();
|
||||
use_serialized_model_ = use_advanced_gpu_api_ &&
|
||||
delegate.gpu().has_serialized_model_dir() &&
|
||||
delegate.gpu().has_model_token();
|
||||
use_gpu_delegate_ = !use_advanced_gpu_api_;
|
||||
|
||||
if (use_kernel_caching_) {
|
||||
#ifdef MEDIAPIPE_ANDROID
|
||||
cached_kernel_filename_ = delegate.gpu().cached_kernel_path() +
|
||||
mediapipe::File::Basename(options.model_path()) +
|
||||
".ker";
|
||||
#endif // MEDIAPIPE_ANDROID
|
||||
}
|
||||
if (use_serialized_model_) {
|
||||
#ifdef MEDIAPIPE_ANDROID
|
||||
serialized_model_path_ = mediapipe::file::JoinPath(
|
||||
delegate.gpu().serialized_model_dir(), delegate.gpu().model_token());
|
||||
#endif // MEDIAPIPE_ANDROID
|
||||
}
|
||||
|
||||
// When use_advanced_gpu_api_, model loading is handled in InitTFLiteGPURunner
|
||||
// for everything.
|
||||
if (!use_advanced_gpu_api_) {
|
||||
MP_RETURN_IF_ERROR(LoadModel(cc));
|
||||
}
|
||||
|
||||
MP_RETURN_IF_ERROR(LoadModel(cc));
|
||||
MP_RETURN_IF_ERROR(gpu_helper_.Open(cc));
|
||||
MP_RETURN_IF_ERROR(
|
||||
gpu_helper_.RunInGlContext([this, &cc]() -> ::mediapipe::Status {
|
||||
return use_advanced_gpu_api_ ? InitTFLiteGPURunner(cc)
|
||||
: LoadDelegateAndAllocateTensors(cc);
|
||||
}));
|
||||
return absl::OkStatus();
|
||||
return gpu_helper_.RunInGlContext([this, &cc]() -> ::mediapipe::Status {
|
||||
return LoadDelegateAndAllocateTensors(cc);
|
||||
});
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorGlImpl::Process(CalculatorContext* cc) {
|
||||
@@ -159,212 +91,71 @@ absl::Status InferenceCalculatorGlImpl::Process(CalculatorContext* cc) {
|
||||
RET_CHECK(!input_tensors.empty());
|
||||
auto output_tensors = absl::make_unique<std::vector<Tensor>>();
|
||||
|
||||
if (use_advanced_gpu_api_) {
|
||||
MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext(
|
||||
[this, &input_tensors, &output_tensors]() -> ::mediapipe::Status {
|
||||
for (int i = 0; i < input_tensors.size(); ++i) {
|
||||
MP_RETURN_IF_ERROR(tflite_gpu_runner_->BindSSBOToInputTensor(
|
||||
input_tensors[i].GetOpenGlBufferReadView().name(), i));
|
||||
}
|
||||
output_tensors->reserve(output_shapes_.size());
|
||||
for (int i = 0; i < output_shapes_.size(); ++i) {
|
||||
output_tensors->emplace_back(Tensor::ElementType::kFloat32,
|
||||
output_shapes_[i]);
|
||||
MP_RETURN_IF_ERROR(tflite_gpu_runner_->BindSSBOToOutputTensor(
|
||||
output_tensors->back().GetOpenGlBufferWriteView().name(), i));
|
||||
}
|
||||
return absl::OkStatus();
|
||||
}));
|
||||
} else {
|
||||
MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext(
|
||||
[this, &input_tensors]() -> ::mediapipe::Status {
|
||||
// Explicitly copy input.
|
||||
for (int i = 0; i < input_tensors.size(); ++i) {
|
||||
glBindBuffer(GL_COPY_READ_BUFFER,
|
||||
input_tensors[i].GetOpenGlBufferReadView().name());
|
||||
glBindBuffer(GL_COPY_WRITE_BUFFER,
|
||||
gpu_buffers_in_[i]->GetOpenGlBufferWriteView().name());
|
||||
glCopyBufferSubData(GL_COPY_READ_BUFFER, GL_COPY_WRITE_BUFFER, 0, 0,
|
||||
input_tensors[i].bytes());
|
||||
}
|
||||
return absl::OkStatus();
|
||||
}));
|
||||
}
|
||||
MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext(
|
||||
[this, &input_tensors]() -> ::mediapipe::Status {
|
||||
// Explicitly copy input.
|
||||
for (int i = 0; i < input_tensors.size(); ++i) {
|
||||
glBindBuffer(GL_COPY_READ_BUFFER,
|
||||
input_tensors[i].GetOpenGlBufferReadView().name());
|
||||
glBindBuffer(GL_COPY_WRITE_BUFFER,
|
||||
gpu_buffers_in_[i]->GetOpenGlBufferWriteView().name());
|
||||
glCopyBufferSubData(GL_COPY_READ_BUFFER, GL_COPY_WRITE_BUFFER, 0, 0,
|
||||
input_tensors[i].bytes());
|
||||
}
|
||||
return absl::OkStatus();
|
||||
}));
|
||||
|
||||
// Run inference.
|
||||
if (use_advanced_gpu_api_) {
|
||||
RET_CHECK(tflite_gpu_runner_->Invoke().ok());
|
||||
} else {
|
||||
RET_CHECK_EQ(interpreter_->Invoke(), kTfLiteOk);
|
||||
}
|
||||
RET_CHECK_EQ(interpreter_->Invoke(), kTfLiteOk);
|
||||
|
||||
if (use_gpu_delegate_) {
|
||||
MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext(
|
||||
[this, &output_tensors]() -> ::mediapipe::Status {
|
||||
output_tensors->reserve(output_shapes_.size());
|
||||
for (int i = 0; i < output_shapes_.size(); ++i) {
|
||||
const auto& t = gpu_buffers_out_[i];
|
||||
output_tensors->emplace_back(Tensor::ElementType::kFloat32,
|
||||
gpu_buffers_out_[i]->shape());
|
||||
auto read_view = t->GetOpenGlBufferReadView();
|
||||
glBindBuffer(GL_COPY_READ_BUFFER, read_view.name());
|
||||
auto write_view = output_tensors->back().GetOpenGlBufferWriteView();
|
||||
glBindBuffer(GL_COPY_WRITE_BUFFER, write_view.name());
|
||||
glCopyBufferSubData(GL_COPY_READ_BUFFER, GL_COPY_WRITE_BUFFER, 0, 0,
|
||||
t->bytes());
|
||||
}
|
||||
return absl::OkStatus();
|
||||
}));
|
||||
}
|
||||
// Output tensors are already bound if use_advanced_gpu_api_ is true.
|
||||
MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext(
|
||||
[this, &output_tensors]() -> ::mediapipe::Status {
|
||||
output_tensors->reserve(output_shapes_.size());
|
||||
for (int i = 0; i < output_shapes_.size(); ++i) {
|
||||
const auto& t = gpu_buffers_out_[i];
|
||||
output_tensors->emplace_back(Tensor::ElementType::kFloat32,
|
||||
gpu_buffers_out_[i]->shape());
|
||||
auto read_view = t->GetOpenGlBufferReadView();
|
||||
glBindBuffer(GL_COPY_READ_BUFFER, read_view.name());
|
||||
auto write_view = output_tensors->back().GetOpenGlBufferWriteView();
|
||||
glBindBuffer(GL_COPY_WRITE_BUFFER, write_view.name());
|
||||
glCopyBufferSubData(GL_COPY_READ_BUFFER, GL_COPY_WRITE_BUFFER, 0, 0,
|
||||
t->bytes());
|
||||
}
|
||||
return absl::OkStatus();
|
||||
}));
|
||||
|
||||
kOutTensors(cc).Send(std::move(output_tensors));
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorGlImpl::SaveGpuCaches() {
|
||||
#ifdef MEDIAPIPE_ANDROID
|
||||
if (use_kernel_caching_) {
|
||||
// Save kernel file.
|
||||
auto kernel_cache = absl::make_unique<std::vector<uint8_t>>(
|
||||
tflite_gpu_runner_->GetSerializedBinaryCache());
|
||||
std::string cache_str(kernel_cache->begin(), kernel_cache->end());
|
||||
MP_RETURN_IF_ERROR(
|
||||
mediapipe::file::SetContents(cached_kernel_filename_, cache_str));
|
||||
}
|
||||
if (use_serialized_model_) {
|
||||
// Save serialized model file.
|
||||
ASSIGN_OR_RETURN(std::vector<uint8_t> serialized_model_vec,
|
||||
tflite_gpu_runner_->GetSerializedModel());
|
||||
absl::string_view serialized_model(
|
||||
reinterpret_cast<char*>(serialized_model_vec.data()),
|
||||
serialized_model_vec.size());
|
||||
MP_RETURN_IF_ERROR(
|
||||
mediapipe::file::SetContents(serialized_model_path_, serialized_model));
|
||||
}
|
||||
#endif // MEDIAPIPE_ANDROID
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorGlImpl::Close(CalculatorContext* cc) {
|
||||
MP_RETURN_IF_ERROR(SaveGpuCaches());
|
||||
if (use_gpu_delegate_) {
|
||||
MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext([this]() -> Status {
|
||||
gpu_buffers_in_.clear();
|
||||
gpu_buffers_out_.clear();
|
||||
return absl::OkStatus();
|
||||
}));
|
||||
}
|
||||
|
||||
interpreter_ = nullptr;
|
||||
delegate_ = nullptr;
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorGlImpl::ReadGpuCaches() {
|
||||
#ifdef MEDIAPIPE_ANDROID
|
||||
if (use_kernel_caching_ && File::Exists(cached_kernel_filename_)) {
|
||||
// Load pre-compiled kernel file.
|
||||
std::string cache_str;
|
||||
MP_RETURN_IF_ERROR(
|
||||
mediapipe::file::GetContents(cached_kernel_filename_, &cache_str));
|
||||
std::vector<uint8_t> cache_vec(cache_str.begin(), cache_str.end());
|
||||
tflite_gpu_runner_->SetSerializedBinaryCache(std::move(cache_vec));
|
||||
}
|
||||
if (use_serialized_model_ && File::Exists(serialized_model_path_)) {
|
||||
// Load serialized model file.
|
||||
std::string serialized_model_str;
|
||||
MP_RETURN_IF_ERROR(
|
||||
file::GetContents(serialized_model_path_, &serialized_model_str));
|
||||
std::vector<uint8_t> serialized_model_vec(serialized_model_str.begin(),
|
||||
serialized_model_str.end());
|
||||
tflite_gpu_runner_->SetSerializedModel(std::move(serialized_model_vec));
|
||||
}
|
||||
#endif // MEDIAPIPE_ANDROID
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorGlImpl::InitTFLiteGPURunner(
|
||||
CalculatorContext* cc) {
|
||||
ASSIGN_OR_RETURN(model_packet_, GetModelAsPacket(cc));
|
||||
const auto& model = *model_packet_.Get();
|
||||
tflite::ops::builtin::BuiltinOpResolver op_resolver =
|
||||
kSideInCustomOpResolver(cc).GetOr(
|
||||
tflite::ops::builtin::BuiltinOpResolverWithoutDefaultDelegates());
|
||||
|
||||
// Create runner
|
||||
tflite::gpu::InferenceOptions options;
|
||||
options.priority1 = allow_precision_loss_
|
||||
? tflite::gpu::InferencePriority::MIN_LATENCY
|
||||
: tflite::gpu::InferencePriority::MAX_PRECISION;
|
||||
options.priority2 = tflite::gpu::InferencePriority::AUTO;
|
||||
options.priority3 = tflite::gpu::InferencePriority::AUTO;
|
||||
switch (tflite_gpu_runner_usage_) {
|
||||
case mediapipe::InferenceCalculatorOptions::Delegate::Gpu::
|
||||
FAST_SINGLE_ANSWER: {
|
||||
options.usage = tflite::gpu::InferenceUsage::FAST_SINGLE_ANSWER;
|
||||
break;
|
||||
}
|
||||
case mediapipe::InferenceCalculatorOptions::Delegate::Gpu::
|
||||
SUSTAINED_SPEED: {
|
||||
options.usage = tflite::gpu::InferenceUsage::SUSTAINED_SPEED;
|
||||
break;
|
||||
}
|
||||
case mediapipe::InferenceCalculatorOptions::Delegate::Gpu::UNSPECIFIED: {
|
||||
return absl::InternalError("inference usage need to be specified.");
|
||||
}
|
||||
}
|
||||
tflite_gpu_runner_ = std::make_unique<tflite::gpu::TFLiteGPURunner>(options);
|
||||
switch (tflite_gpu_runner_api_) {
|
||||
case mediapipe::InferenceCalculatorOptions::Delegate::Gpu::ANY: {
|
||||
// Do not need to force any specific API.
|
||||
break;
|
||||
}
|
||||
case mediapipe::InferenceCalculatorOptions::Delegate::Gpu::OPENGL: {
|
||||
tflite_gpu_runner_->ForceOpenGL();
|
||||
break;
|
||||
}
|
||||
case mediapipe::InferenceCalculatorOptions::Delegate::Gpu::OPENCL: {
|
||||
tflite_gpu_runner_->ForceOpenCL();
|
||||
break;
|
||||
}
|
||||
}
|
||||
MP_RETURN_IF_ERROR(tflite_gpu_runner_->InitializeWithModel(
|
||||
model, op_resolver, /*allow_quant_ops=*/true));
|
||||
|
||||
// Create and bind OpenGL buffers for outputs.
|
||||
// The buffers are created once and their ids are passed to calculator outputs
|
||||
output_shapes_.resize(tflite_gpu_runner_->outputs_size());
|
||||
for (int i = 0; i < tflite_gpu_runner_->outputs_size(); ++i) {
|
||||
output_shapes_[i] = {tflite_gpu_runner_->GetOutputShapes()[i].b,
|
||||
tflite_gpu_runner_->GetOutputShapes()[i].h,
|
||||
tflite_gpu_runner_->GetOutputShapes()[i].w,
|
||||
tflite_gpu_runner_->GetOutputShapes()[i].c};
|
||||
}
|
||||
|
||||
MP_RETURN_IF_ERROR(ReadGpuCaches());
|
||||
|
||||
MP_RETURN_IF_ERROR(tflite_gpu_runner_->Build());
|
||||
|
||||
return absl::OkStatus();
|
||||
return gpu_helper_.RunInGlContext([this]() -> absl::Status {
|
||||
gpu_buffers_in_.clear();
|
||||
gpu_buffers_out_.clear();
|
||||
// Delegate must outlive the interpreter, hence the order is important.
|
||||
interpreter_ = nullptr;
|
||||
delegate_ = nullptr;
|
||||
return absl::OkStatus();
|
||||
});
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorGlImpl::LoadModel(CalculatorContext* cc) {
|
||||
ASSIGN_OR_RETURN(model_packet_, GetModelAsPacket(cc));
|
||||
const auto& model = *model_packet_.Get();
|
||||
tflite::ops::builtin::BuiltinOpResolver op_resolver =
|
||||
kSideInCustomOpResolver(cc).GetOr(
|
||||
tflite::ops::builtin::BuiltinOpResolverWithoutDefaultDelegates());
|
||||
|
||||
tflite::InterpreterBuilder(model, op_resolver)(&interpreter_);
|
||||
if (kSideInOpResolver(cc).IsConnected()) {
|
||||
const tflite::OpResolver& op_resolver = kSideInOpResolver(cc).Get();
|
||||
tflite::InterpreterBuilder(model, op_resolver)(&interpreter_);
|
||||
} else {
|
||||
tflite::ops::builtin::BuiltinOpResolver op_resolver =
|
||||
kSideInCustomOpResolver(cc).GetOr(
|
||||
tflite::ops::builtin::BuiltinOpResolverWithoutDefaultDelegates());
|
||||
tflite::InterpreterBuilder(model, op_resolver)(&interpreter_);
|
||||
}
|
||||
RET_CHECK(interpreter_);
|
||||
|
||||
#if defined(__EMSCRIPTEN__)
|
||||
interpreter_->SetNumThreads(1);
|
||||
#else
|
||||
interpreter_->SetNumThreads(
|
||||
cc->Options<mediapipe::InferenceCalculatorOptions>().cpu_num_thread());
|
||||
#endif // __EMSCRIPTEN__
|
||||
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
@@ -0,0 +1,285 @@
|
||||
// Copyright 2022 The MediaPipe Authors.
|
||||
//
|
||||
// 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.
|
||||
|
||||
#include <cstring>
|
||||
#include <memory>
|
||||
#include <string>
|
||||
#include <vector>
|
||||
|
||||
#include "absl/memory/memory.h"
|
||||
#include "absl/status/status.h"
|
||||
#include "mediapipe/calculators/tensor/inference_calculator.h"
|
||||
#include "mediapipe/gpu/gl_calculator_helper.h"
|
||||
#include "mediapipe/util/tflite/tflite_gpu_runner.h"
|
||||
|
||||
#if defined(MEDIAPIPE_ANDROID)
|
||||
#include "mediapipe/framework/deps/file_path.h"
|
||||
#include "mediapipe/util/android/file/base/file.h"
|
||||
#include "mediapipe/util/android/file/base/filesystem.h"
|
||||
#include "mediapipe/util/android/file/base/helpers.h"
|
||||
#endif // ANDROID
|
||||
|
||||
namespace mediapipe {
|
||||
namespace api2 {
|
||||
|
||||
// Runs TFLite GPU delegate API2 directly, bypassing interpreter usage, and
|
||||
// allows choosing specific API.
|
||||
//
|
||||
// To trigger this code path:
|
||||
// [mediapipe.InferenceCalculatorOptions.ext] {
|
||||
// delegate {
|
||||
// gpu {
|
||||
// use_advanced_gpu_api: true
|
||||
// api: OPENCL # or OPENGL or ANY
|
||||
// }
|
||||
// }
|
||||
// }
|
||||
class InferenceCalculatorGlAdvancedImpl
|
||||
: public NodeImpl<InferenceCalculatorGlAdvanced,
|
||||
InferenceCalculatorGlAdvancedImpl> {
|
||||
public:
|
||||
static absl::Status UpdateContract(CalculatorContract* cc);
|
||||
|
||||
absl::Status Open(CalculatorContext* cc) override;
|
||||
absl::Status Process(CalculatorContext* cc) override;
|
||||
absl::Status Close(CalculatorContext* cc) override;
|
||||
|
||||
private:
|
||||
absl::Status ReadGpuCaches();
|
||||
absl::Status SaveGpuCaches();
|
||||
absl::Status InitTFLiteGPURunner(CalculatorContext* cc);
|
||||
|
||||
// TfLite requires us to keep the model alive as long as the interpreter is.
|
||||
Packet<TfLiteModelPtr> model_packet_;
|
||||
|
||||
mediapipe::GlCalculatorHelper gpu_helper_;
|
||||
std::unique_ptr<tflite::gpu::TFLiteGPURunner> tflite_gpu_runner_;
|
||||
bool allow_precision_loss_ = false;
|
||||
mediapipe::InferenceCalculatorOptions::Delegate::Gpu::Api
|
||||
tflite_gpu_runner_api_;
|
||||
mediapipe::InferenceCalculatorOptions::Delegate::Gpu::InferenceUsage
|
||||
tflite_gpu_runner_usage_;
|
||||
|
||||
std::vector<Tensor::Shape> output_shapes_;
|
||||
|
||||
bool use_kernel_caching_ = false;
|
||||
std::string cached_kernel_filename_;
|
||||
bool use_serialized_model_ = false;
|
||||
std::string serialized_model_path_;
|
||||
};
|
||||
|
||||
absl::Status InferenceCalculatorGlAdvancedImpl::UpdateContract(
|
||||
CalculatorContract* cc) {
|
||||
const auto& options = cc->Options<::mediapipe::InferenceCalculatorOptions>();
|
||||
RET_CHECK(!options.model_path().empty() ^ kSideInModel(cc).IsConnected())
|
||||
<< "Either model as side packet or model path in options is required.";
|
||||
|
||||
MP_RETURN_IF_ERROR(mediapipe::GlCalculatorHelper::UpdateContract(cc));
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorGlAdvancedImpl::Open(CalculatorContext* cc) {
|
||||
const auto& options = cc->Options<::mediapipe::InferenceCalculatorOptions>();
|
||||
mediapipe::InferenceCalculatorOptions::Delegate delegate = options.delegate();
|
||||
if (!kDelegate(cc).IsEmpty()) {
|
||||
mediapipe::InferenceCalculatorOptions::Delegate input_side_packet_delegate =
|
||||
kDelegate(cc).Get();
|
||||
CHECK(input_side_packet_delegate.has_gpu() ||
|
||||
input_side_packet_delegate.delegate_case() ==
|
||||
mediapipe::InferenceCalculatorOptions::Delegate::DELEGATE_NOT_SET)
|
||||
<< "inference_calculator_gl_advanced only supports delegate input side "
|
||||
"packet for Gpu";
|
||||
delegate.MergeFrom(input_side_packet_delegate);
|
||||
}
|
||||
allow_precision_loss_ = delegate.gpu().allow_precision_loss();
|
||||
tflite_gpu_runner_api_ = delegate.gpu().api();
|
||||
tflite_gpu_runner_usage_ = delegate.gpu().usage();
|
||||
use_kernel_caching_ = delegate.gpu().has_cached_kernel_path();
|
||||
use_serialized_model_ = delegate.gpu().has_serialized_model_dir() &&
|
||||
delegate.gpu().has_model_token();
|
||||
|
||||
if (use_kernel_caching_) {
|
||||
#ifdef MEDIAPIPE_ANDROID
|
||||
cached_kernel_filename_ = delegate.gpu().cached_kernel_path() +
|
||||
mediapipe::File::Basename(options.model_path()) +
|
||||
".ker";
|
||||
#endif // MEDIAPIPE_ANDROID
|
||||
}
|
||||
if (use_serialized_model_) {
|
||||
#ifdef MEDIAPIPE_ANDROID
|
||||
serialized_model_path_ = mediapipe::file::JoinPath(
|
||||
delegate.gpu().serialized_model_dir(), delegate.gpu().model_token());
|
||||
#endif // MEDIAPIPE_ANDROID
|
||||
}
|
||||
|
||||
MP_RETURN_IF_ERROR(gpu_helper_.Open(cc));
|
||||
return gpu_helper_.RunInGlContext(
|
||||
[this, &cc]() -> absl::Status { return InitTFLiteGPURunner(cc); });
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorGlAdvancedImpl::Process(CalculatorContext* cc) {
|
||||
if (kInTensors(cc).IsEmpty()) {
|
||||
return absl::OkStatus();
|
||||
}
|
||||
const auto& input_tensors = *kInTensors(cc);
|
||||
RET_CHECK(!input_tensors.empty());
|
||||
auto output_tensors = absl::make_unique<std::vector<Tensor>>();
|
||||
|
||||
MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext(
|
||||
[this, &input_tensors, &output_tensors]() -> absl::Status {
|
||||
for (int i = 0; i < input_tensors.size(); ++i) {
|
||||
MP_RETURN_IF_ERROR(tflite_gpu_runner_->BindSSBOToInputTensor(
|
||||
input_tensors[i].GetOpenGlBufferReadView().name(), i));
|
||||
}
|
||||
output_tensors->reserve(output_shapes_.size());
|
||||
for (int i = 0; i < output_shapes_.size(); ++i) {
|
||||
output_tensors->emplace_back(Tensor::ElementType::kFloat32,
|
||||
output_shapes_[i]);
|
||||
MP_RETURN_IF_ERROR(tflite_gpu_runner_->BindSSBOToOutputTensor(
|
||||
output_tensors->back().GetOpenGlBufferWriteView().name(), i));
|
||||
}
|
||||
return absl::OkStatus();
|
||||
}));
|
||||
|
||||
// Run inference.
|
||||
MP_RETURN_IF_ERROR(tflite_gpu_runner_->Invoke());
|
||||
kOutTensors(cc).Send(std::move(output_tensors));
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorGlAdvancedImpl::SaveGpuCaches() {
|
||||
#ifdef MEDIAPIPE_ANDROID
|
||||
if (use_kernel_caching_) {
|
||||
// Save kernel file.
|
||||
auto kernel_cache = absl::make_unique<std::vector<uint8_t>>(
|
||||
tflite_gpu_runner_->GetSerializedBinaryCache());
|
||||
std::string cache_str(kernel_cache->begin(), kernel_cache->end());
|
||||
MP_RETURN_IF_ERROR(
|
||||
mediapipe::file::SetContents(cached_kernel_filename_, cache_str));
|
||||
}
|
||||
if (use_serialized_model_) {
|
||||
// Save serialized model file.
|
||||
ASSIGN_OR_RETURN(std::vector<uint8_t> serialized_model_vec,
|
||||
tflite_gpu_runner_->GetSerializedModel());
|
||||
absl::string_view serialized_model(
|
||||
reinterpret_cast<char*>(serialized_model_vec.data()),
|
||||
serialized_model_vec.size());
|
||||
MP_RETURN_IF_ERROR(
|
||||
mediapipe::file::SetContents(serialized_model_path_, serialized_model));
|
||||
}
|
||||
#endif // MEDIAPIPE_ANDROID
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorGlAdvancedImpl::Close(CalculatorContext* cc) {
|
||||
MP_RETURN_IF_ERROR(SaveGpuCaches());
|
||||
return gpu_helper_.RunInGlContext([this]() -> absl::Status {
|
||||
tflite_gpu_runner_.reset();
|
||||
return absl::OkStatus();
|
||||
});
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorGlAdvancedImpl::ReadGpuCaches() {
|
||||
#ifdef MEDIAPIPE_ANDROID
|
||||
if (use_kernel_caching_ && File::Exists(cached_kernel_filename_)) {
|
||||
// Load pre-compiled kernel file.
|
||||
std::string cache_str;
|
||||
MP_RETURN_IF_ERROR(
|
||||
mediapipe::file::GetContents(cached_kernel_filename_, &cache_str));
|
||||
std::vector<uint8_t> cache_vec(cache_str.begin(), cache_str.end());
|
||||
tflite_gpu_runner_->SetSerializedBinaryCache(std::move(cache_vec));
|
||||
}
|
||||
if (use_serialized_model_ && File::Exists(serialized_model_path_)) {
|
||||
// Load serialized model file.
|
||||
std::string serialized_model_str;
|
||||
MP_RETURN_IF_ERROR(
|
||||
file::GetContents(serialized_model_path_, &serialized_model_str));
|
||||
std::vector<uint8_t> serialized_model_vec(serialized_model_str.begin(),
|
||||
serialized_model_str.end());
|
||||
tflite_gpu_runner_->SetSerializedModel(std::move(serialized_model_vec));
|
||||
}
|
||||
#endif // MEDIAPIPE_ANDROID
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorGlAdvancedImpl::InitTFLiteGPURunner(
|
||||
CalculatorContext* cc) {
|
||||
ASSIGN_OR_RETURN(model_packet_, GetModelAsPacket(cc));
|
||||
const auto& model = *model_packet_.Get();
|
||||
|
||||
// Create runner
|
||||
tflite::gpu::InferenceOptions options;
|
||||
options.priority1 = allow_precision_loss_
|
||||
? tflite::gpu::InferencePriority::MIN_LATENCY
|
||||
: tflite::gpu::InferencePriority::MAX_PRECISION;
|
||||
options.priority2 = tflite::gpu::InferencePriority::AUTO;
|
||||
options.priority3 = tflite::gpu::InferencePriority::AUTO;
|
||||
switch (tflite_gpu_runner_usage_) {
|
||||
case mediapipe::InferenceCalculatorOptions::Delegate::Gpu::
|
||||
FAST_SINGLE_ANSWER: {
|
||||
options.usage = tflite::gpu::InferenceUsage::FAST_SINGLE_ANSWER;
|
||||
break;
|
||||
}
|
||||
case mediapipe::InferenceCalculatorOptions::Delegate::Gpu::
|
||||
SUSTAINED_SPEED: {
|
||||
options.usage = tflite::gpu::InferenceUsage::SUSTAINED_SPEED;
|
||||
break;
|
||||
}
|
||||
case mediapipe::InferenceCalculatorOptions::Delegate::Gpu::UNSPECIFIED: {
|
||||
return absl::InternalError("inference usage need to be specified.");
|
||||
}
|
||||
}
|
||||
tflite_gpu_runner_ = std::make_unique<tflite::gpu::TFLiteGPURunner>(options);
|
||||
switch (tflite_gpu_runner_api_) {
|
||||
case mediapipe::InferenceCalculatorOptions::Delegate::Gpu::ANY: {
|
||||
// Do not need to force any specific API.
|
||||
break;
|
||||
}
|
||||
case mediapipe::InferenceCalculatorOptions::Delegate::Gpu::OPENGL: {
|
||||
tflite_gpu_runner_->ForceOpenGL();
|
||||
break;
|
||||
}
|
||||
case mediapipe::InferenceCalculatorOptions::Delegate::Gpu::OPENCL: {
|
||||
tflite_gpu_runner_->ForceOpenCL();
|
||||
break;
|
||||
}
|
||||
}
|
||||
if (kSideInOpResolver(cc).IsConnected()) {
|
||||
const tflite::OpResolver& op_resolver = kSideInOpResolver(cc).Get();
|
||||
MP_RETURN_IF_ERROR(tflite_gpu_runner_->InitializeWithModel(
|
||||
model, op_resolver, /*allow_quant_ops=*/true));
|
||||
} else {
|
||||
tflite::ops::builtin::BuiltinOpResolver op_resolver =
|
||||
kSideInCustomOpResolver(cc).GetOr(
|
||||
tflite::ops::builtin::BuiltinOpResolverWithoutDefaultDelegates());
|
||||
MP_RETURN_IF_ERROR(tflite_gpu_runner_->InitializeWithModel(
|
||||
model, op_resolver, /*allow_quant_ops=*/true));
|
||||
}
|
||||
|
||||
// Create and bind OpenGL buffers for outputs.
|
||||
// The buffers are created once and their ids are passed to calculator outputs
|
||||
output_shapes_.resize(tflite_gpu_runner_->outputs_size());
|
||||
for (int i = 0; i < tflite_gpu_runner_->outputs_size(); ++i) {
|
||||
output_shapes_[i] = {tflite_gpu_runner_->GetOutputShapes()[i].b,
|
||||
tflite_gpu_runner_->GetOutputShapes()[i].h,
|
||||
tflite_gpu_runner_->GetOutputShapes()[i].w,
|
||||
tflite_gpu_runner_->GetOutputShapes()[i].c};
|
||||
}
|
||||
|
||||
MP_RETURN_IF_ERROR(ReadGpuCaches());
|
||||
return tflite_gpu_runner_->Build();
|
||||
}
|
||||
|
||||
} // namespace api2
|
||||
} // namespace mediapipe
|
||||
@@ -90,9 +90,10 @@ class InferenceCalculatorMetalImpl
|
||||
absl::Status Close(CalculatorContext* cc) override;
|
||||
|
||||
private:
|
||||
absl::Status LoadModel(CalculatorContext* cc);
|
||||
absl::Status LoadDelegate(CalculatorContext* cc);
|
||||
absl::Status LoadDelegateAndAllocateTensors(CalculatorContext* cc);
|
||||
absl::Status InitInterpreter(CalculatorContext* cc);
|
||||
void AddDelegate(CalculatorContext* cc,
|
||||
tflite::InterpreterBuilder* interpreter_builder);
|
||||
absl::Status CreateConverters(CalculatorContext* cc);
|
||||
|
||||
// TfLite requires us to keep the model alive as long as the interpreter is.
|
||||
Packet<TfLiteModelPtr> model_packet_;
|
||||
@@ -127,11 +128,9 @@ absl::Status InferenceCalculatorMetalImpl::Open(CalculatorContext* cc) {
|
||||
const auto& options = cc->Options<::mediapipe::InferenceCalculatorOptions>();
|
||||
allow_precision_loss_ = options.delegate().gpu().allow_precision_loss();
|
||||
|
||||
MP_RETURN_IF_ERROR(LoadModel(cc));
|
||||
|
||||
gpu_helper_ = [[MPPMetalHelper alloc] initWithCalculatorContext:cc];
|
||||
RET_CHECK(gpu_helper_);
|
||||
return LoadDelegateAndAllocateTensors(cc);
|
||||
return InitInterpreter(cc);
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorMetalImpl::Process(CalculatorContext* cc) {
|
||||
@@ -199,27 +198,20 @@ absl::Status InferenceCalculatorMetalImpl::Close(CalculatorContext* cc) {
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorMetalImpl::LoadModel(CalculatorContext* cc) {
|
||||
absl::Status InferenceCalculatorMetalImpl::InitInterpreter(
|
||||
CalculatorContext* cc) {
|
||||
ASSIGN_OR_RETURN(model_packet_, GetModelAsPacket(cc));
|
||||
const auto& model = *model_packet_.Get();
|
||||
tflite::ops::builtin::BuiltinOpResolver op_resolver =
|
||||
kSideInCustomOpResolver(cc).GetOr(
|
||||
tflite::ops::builtin::BuiltinOpResolverWithoutDefaultDelegates());
|
||||
|
||||
tflite::InterpreterBuilder(model, op_resolver)(&interpreter_);
|
||||
ASSIGN_OR_RETURN(auto op_resolver_packet, GetOpResolverAsPacket(cc));
|
||||
const auto& op_resolver = op_resolver_packet.Get();
|
||||
tflite::InterpreterBuilder interpreter_builder(model, op_resolver);
|
||||
AddDelegate(cc, &interpreter_builder);
|
||||
interpreter_builder.SetNumThreads(
|
||||
cc->Options<mediapipe::InferenceCalculatorOptions>().cpu_num_thread());
|
||||
RET_CHECK_EQ(interpreter_builder(&interpreter_), kTfLiteOk);
|
||||
RET_CHECK(interpreter_);
|
||||
|
||||
interpreter_->SetNumThreads(
|
||||
cc->Options<mediapipe::InferenceCalculatorOptions>().cpu_num_thread());
|
||||
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorMetalImpl::LoadDelegateAndAllocateTensors(
|
||||
CalculatorContext* cc) {
|
||||
MP_RETURN_IF_ERROR(LoadDelegate(cc));
|
||||
|
||||
// AllocateTensors() can be called only after ModifyGraphWithDelegate.
|
||||
MP_RETURN_IF_ERROR(CreateConverters(cc));
|
||||
RET_CHECK_EQ(interpreter_->AllocateTensors(), kTfLiteOk);
|
||||
// TODO: Support quantized tensors.
|
||||
RET_CHECK_NE(
|
||||
@@ -228,7 +220,8 @@ absl::Status InferenceCalculatorMetalImpl::LoadDelegateAndAllocateTensors(
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorMetalImpl::LoadDelegate(CalculatorContext* cc) {
|
||||
void InferenceCalculatorMetalImpl::AddDelegate(
|
||||
CalculatorContext* cc, tflite::InterpreterBuilder* interpreter_builder) {
|
||||
const auto& calculator_opts =
|
||||
cc->Options<mediapipe::InferenceCalculatorOptions>();
|
||||
|
||||
@@ -242,9 +235,11 @@ absl::Status InferenceCalculatorMetalImpl::LoadDelegate(CalculatorContext* cc) {
|
||||
options.wait_type = TFLGpuDelegateWaitType::TFLGpuDelegateWaitTypeDoNotWait;
|
||||
delegate_ =
|
||||
TfLiteDelegatePtr(TFLGpuDelegateCreate(&options), &TFLGpuDelegateDelete);
|
||||
RET_CHECK_EQ(interpreter_->ModifyGraphWithDelegate(delegate_.get()),
|
||||
kTfLiteOk);
|
||||
interpreter_builder->AddDelegate(delegate_.get());
|
||||
}
|
||||
|
||||
absl::Status InferenceCalculatorMetalImpl::CreateConverters(
|
||||
CalculatorContext* cc) {
|
||||
id<MTLDevice> device = gpu_helper_.mtlDevice;
|
||||
|
||||
// Get input image sizes.
|
||||
|
||||
@@ -38,61 +38,13 @@
|
||||
#endif // defined(__APPLE__)
|
||||
|
||||
namespace mediapipe {
|
||||
namespace {
|
||||
|
||||
void DoSmokeTest(const std::string& graph_proto) {
|
||||
const int width = 8;
|
||||
const int height = 8;
|
||||
const int channels = 3;
|
||||
// Prepare input tensor.
|
||||
auto input_vec = absl::make_unique<std::vector<Tensor>>();
|
||||
input_vec->emplace_back(Tensor::ElementType::kFloat32,
|
||||
Tensor::Shape{1, height, width, channels});
|
||||
{
|
||||
auto view1 = input_vec->back().GetCpuWriteView();
|
||||
auto tensor_buffer = view1.buffer<float>();
|
||||
ASSERT_NE(tensor_buffer, nullptr);
|
||||
for (int i = 0; i < width * height * channels - 1; i++) {
|
||||
tensor_buffer[i] = 1;
|
||||
}
|
||||
}
|
||||
constexpr int kTensorWidth = 8;
|
||||
constexpr int kTensorHeight = 8;
|
||||
constexpr int kTensorChannels = 3;
|
||||
|
||||
// Prepare single calculator graph to and wait for packets.
|
||||
CalculatorGraphConfig graph_config =
|
||||
ParseTextProtoOrDie<CalculatorGraphConfig>(graph_proto);
|
||||
std::vector<Packet> output_packets;
|
||||
tool::AddVectorSink("tensor_out", &graph_config, &output_packets);
|
||||
CalculatorGraph graph(graph_config);
|
||||
MP_ASSERT_OK(graph.StartRun({}));
|
||||
|
||||
// Push the tensor into the graph.
|
||||
MP_ASSERT_OK(graph.AddPacketToInputStream(
|
||||
"tensor_in", Adopt(input_vec.release()).At(Timestamp(0))));
|
||||
// Wait until the calculator done processing.
|
||||
MP_ASSERT_OK(graph.WaitUntilIdle());
|
||||
ASSERT_EQ(1, output_packets.size());
|
||||
|
||||
// Get and process results.
|
||||
const std::vector<Tensor>& result_vec =
|
||||
output_packets[0].Get<std::vector<Tensor>>();
|
||||
ASSERT_EQ(1, result_vec.size());
|
||||
|
||||
const Tensor& result = result_vec[0];
|
||||
auto view = result.GetCpuReadView();
|
||||
auto result_buffer = view.buffer<float>();
|
||||
ASSERT_NE(result_buffer, nullptr);
|
||||
for (int i = 0; i < width * height * channels - 1; i++) {
|
||||
ASSERT_EQ(3, result_buffer[i]);
|
||||
}
|
||||
|
||||
// Fully close graph at end, otherwise calculator+tensors are destroyed
|
||||
// after calling WaitUntilDone().
|
||||
MP_ASSERT_OK(graph.CloseInputStream("tensor_in"));
|
||||
MP_ASSERT_OK(graph.WaitUntilDone());
|
||||
}
|
||||
|
||||
// Tests a simple add model that adds an input tensor to itself.
|
||||
TEST(InferenceCalculatorTest, SmokeTest) {
|
||||
std::string graph_proto = R"(
|
||||
constexpr char kGraphWithModelPathInOption[] = R"(
|
||||
input_stream: "tensor_in"
|
||||
node {
|
||||
calculator: "InferenceCalculator"
|
||||
@@ -106,18 +58,7 @@ TEST(InferenceCalculatorTest, SmokeTest) {
|
||||
}
|
||||
}
|
||||
)";
|
||||
// Test CPU inference only.
|
||||
DoSmokeTest(/*graph_proto=*/absl::StrReplaceAll(
|
||||
graph_proto, {{"$delegate", "delegate { tflite {} }"}}));
|
||||
DoSmokeTest(absl::StrReplaceAll(graph_proto,
|
||||
{{"$delegate", "delegate { xnnpack {} }"}}));
|
||||
DoSmokeTest(absl::StrReplaceAll(
|
||||
graph_proto,
|
||||
{{"$delegate", "delegate { xnnpack { num_threads: 10 } }"}}));
|
||||
}
|
||||
|
||||
TEST(InferenceCalculatorTest, SmokeTest_ModelAsInputSidePacket) {
|
||||
std::string graph_proto = R"(
|
||||
constexpr char kGraphWithModelAsInputSidePacket[] = R"(
|
||||
input_stream: "tensor_in"
|
||||
|
||||
node {
|
||||
@@ -154,7 +95,84 @@ TEST(InferenceCalculatorTest, SmokeTest_ModelAsInputSidePacket) {
|
||||
}
|
||||
}
|
||||
)";
|
||||
DoSmokeTest(graph_proto);
|
||||
|
||||
std::vector<Tensor> CreateInputs() {
|
||||
std::vector<Tensor> input_vec;
|
||||
// Prepare input tensor.
|
||||
input_vec.emplace_back(
|
||||
Tensor::ElementType::kFloat32,
|
||||
Tensor::Shape{1, kTensorHeight, kTensorWidth, kTensorChannels});
|
||||
{
|
||||
auto view = input_vec.back().GetCpuWriteView();
|
||||
auto num_elements = input_vec.back().shape().num_elements();
|
||||
auto tensor_buffer = view.buffer<float>();
|
||||
for (int i = 0; i < num_elements; i++) {
|
||||
tensor_buffer[i] = 1;
|
||||
}
|
||||
}
|
||||
|
||||
return input_vec;
|
||||
}
|
||||
|
||||
void RunGraphThenClose(CalculatorGraph& graph, std::vector<Tensor> input_vec) {
|
||||
MP_ASSERT_OK(graph.StartRun({}));
|
||||
|
||||
// Push the tensor into the graph.
|
||||
MP_ASSERT_OK(graph.AddPacketToInputStream(
|
||||
"tensor_in",
|
||||
MakePacket<std::vector<Tensor>>(std::move(input_vec)).At(Timestamp(0))));
|
||||
// Wait until the calculator done processing.
|
||||
MP_ASSERT_OK(graph.WaitUntilIdle());
|
||||
|
||||
// Fully close graph at end, otherwise calculator+tensors are destroyed
|
||||
// after calling WaitUntilDone().
|
||||
MP_ASSERT_OK(graph.CloseInputStream("tensor_in"));
|
||||
MP_ASSERT_OK(graph.WaitUntilDone());
|
||||
}
|
||||
|
||||
void DoSmokeTest(const std::string& graph_proto) {
|
||||
auto input_vec = CreateInputs();
|
||||
|
||||
// Prepare single calculator graph to and wait for packets.
|
||||
CalculatorGraphConfig graph_config =
|
||||
ParseTextProtoOrDie<CalculatorGraphConfig>(graph_proto);
|
||||
std::vector<Packet> output_packets;
|
||||
tool::AddVectorSink("tensor_out", &graph_config, &output_packets);
|
||||
CalculatorGraph graph(graph_config);
|
||||
|
||||
RunGraphThenClose(graph, std::move(input_vec));
|
||||
|
||||
ASSERT_EQ(1, output_packets.size());
|
||||
|
||||
// Get and process results.
|
||||
const std::vector<Tensor>& result_vec =
|
||||
output_packets[0].Get<std::vector<Tensor>>();
|
||||
ASSERT_EQ(1, result_vec.size());
|
||||
|
||||
const Tensor& result = result_vec[0];
|
||||
auto view = result.GetCpuReadView();
|
||||
auto result_buffer = view.buffer<float>();
|
||||
ASSERT_NE(result_buffer, nullptr);
|
||||
for (int i = 0; i < result.shape().num_elements(); i++) {
|
||||
ASSERT_EQ(3, result_buffer[i]);
|
||||
}
|
||||
}
|
||||
|
||||
// Tests a simple add model that adds an input tensor to itself.
|
||||
TEST(InferenceCalculatorTest, SmokeTest) {
|
||||
// Test CPU inference only.
|
||||
DoSmokeTest(/*graph_proto=*/absl::StrReplaceAll(
|
||||
kGraphWithModelPathInOption, {{"$delegate", "delegate { tflite {} }"}}));
|
||||
DoSmokeTest(absl::StrReplaceAll(kGraphWithModelPathInOption,
|
||||
{{"$delegate", "delegate { xnnpack {} }"}}));
|
||||
DoSmokeTest(absl::StrReplaceAll(
|
||||
kGraphWithModelPathInOption,
|
||||
{{"$delegate", "delegate { xnnpack { num_threads: 10 } }"}}));
|
||||
}
|
||||
|
||||
TEST(InferenceCalculatorTest, ModelAsInputSidePacketSmokeTest) {
|
||||
DoSmokeTest(kGraphWithModelAsInputSidePacket);
|
||||
}
|
||||
|
||||
} // namespace
|
||||
} // namespace mediapipe
|
||||
|
||||
@@ -16,7 +16,6 @@
|
||||
#include <unordered_map>
|
||||
#include <vector>
|
||||
|
||||
#include "absl/container/node_hash_map.h"
|
||||
#include "absl/strings/str_format.h"
|
||||
#include "absl/types/span.h"
|
||||
#include "mediapipe/calculators/tensor/tensors_to_classification_calculator.pb.h"
|
||||
@@ -25,6 +24,7 @@
|
||||
#include "mediapipe/framework/formats/classification.pb.h"
|
||||
#include "mediapipe/framework/formats/tensor.h"
|
||||
#include "mediapipe/framework/port/ret_check.h"
|
||||
#include "mediapipe/util/label_map.pb.h"
|
||||
#include "mediapipe/util/resource_util.h"
|
||||
#if defined(MEDIAPIPE_MOBILE)
|
||||
#include "mediapipe/util/android/file/base/file.h"
|
||||
@@ -35,6 +35,17 @@
|
||||
|
||||
namespace mediapipe {
|
||||
namespace api2 {
|
||||
namespace {
|
||||
|
||||
void SetClassificationLabel(const LabelMapItem label_map_item,
|
||||
Classification* classification) {
|
||||
classification->set_label(label_map_item.name());
|
||||
if (label_map_item.has_display_name()) {
|
||||
classification->set_display_name(label_map_item.display_name());
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
// Convert result tensors from classification models into MediaPipe
|
||||
// classifications.
|
||||
@@ -54,7 +65,6 @@ namespace api2 {
|
||||
// output_stream: "CLASSIFICATIONS:classifications"
|
||||
// options: {
|
||||
// [mediapipe.TensorsToClassificationCalculatorOptions.ext] {
|
||||
// num_classes: 1024
|
||||
// min_score_threshold: 0.1
|
||||
// label_map_path: "labelmap.txt"
|
||||
// }
|
||||
@@ -72,22 +82,35 @@ class TensorsToClassificationCalculator : public Node {
|
||||
absl::Status Close(CalculatorContext* cc) override;
|
||||
|
||||
private:
|
||||
::mediapipe::TensorsToClassificationCalculatorOptions options_;
|
||||
int top_k_ = 0;
|
||||
absl::node_hash_map<int, std::string> label_map_;
|
||||
bool sort_by_descending_score_ = false;
|
||||
proto_ns::Map<int64, LabelMapItem> local_label_map_;
|
||||
bool label_map_loaded_ = false;
|
||||
bool is_binary_classification_ = false;
|
||||
float min_score_threshold_ = std::numeric_limits<float>::lowest();
|
||||
|
||||
// Set of allowed or ignored class indices.
|
||||
struct ClassIndexSet {
|
||||
absl::flat_hash_set<int> values;
|
||||
bool is_allowlist;
|
||||
};
|
||||
// Allowed or ignored class indices based on provided options.
|
||||
// These are used to filter out the output classification results.
|
||||
ClassIndexSet class_index_set_;
|
||||
bool IsClassIndexAllowed(int class_index);
|
||||
const proto_ns::Map<int64, LabelMapItem>& GetLabelMap(CalculatorContext* cc);
|
||||
};
|
||||
MEDIAPIPE_REGISTER_NODE(TensorsToClassificationCalculator);
|
||||
|
||||
absl::Status TensorsToClassificationCalculator::Open(CalculatorContext* cc) {
|
||||
options_ =
|
||||
cc->Options<::mediapipe::TensorsToClassificationCalculatorOptions>();
|
||||
const auto& options = cc->Options<TensorsToClassificationCalculatorOptions>();
|
||||
|
||||
top_k_ = options_.top_k();
|
||||
if (options_.has_label_map_path()) {
|
||||
top_k_ = options.top_k();
|
||||
sort_by_descending_score_ = options.sort_by_descending_score();
|
||||
if (options.has_label_map_path()) {
|
||||
std::string string_path;
|
||||
ASSIGN_OR_RETURN(string_path,
|
||||
PathToResourceAsFile(options_.label_map_path()));
|
||||
PathToResourceAsFile(options.label_map_path()));
|
||||
std::string label_map_string;
|
||||
MP_RETURN_IF_ERROR(
|
||||
mediapipe::GetResourceContents(string_path, &label_map_string));
|
||||
@@ -96,18 +119,45 @@ absl::Status TensorsToClassificationCalculator::Open(CalculatorContext* cc) {
|
||||
std::string line;
|
||||
int i = 0;
|
||||
while (std::getline(stream, line)) {
|
||||
label_map_[i++] = line;
|
||||
LabelMapItem item;
|
||||
item.set_name(line);
|
||||
local_label_map_[i++] = item;
|
||||
}
|
||||
label_map_loaded_ = true;
|
||||
} else if (options_.has_label_map()) {
|
||||
for (int i = 0; i < options_.label_map().entries_size(); ++i) {
|
||||
const auto& entry = options_.label_map().entries(i);
|
||||
RET_CHECK(!label_map_.contains(entry.id()))
|
||||
} else if (!options.label_items().empty()) {
|
||||
label_map_loaded_ = true;
|
||||
} else if (options.has_label_map()) {
|
||||
for (int i = 0; i < options.label_map().entries_size(); ++i) {
|
||||
const auto& entry = options.label_map().entries(i);
|
||||
RET_CHECK(!local_label_map_.contains(entry.id()))
|
||||
<< "Duplicate id found: " << entry.id();
|
||||
label_map_[entry.id()] = entry.label();
|
||||
LabelMapItem item;
|
||||
item.set_name(entry.label());
|
||||
local_label_map_[entry.id()] = item;
|
||||
}
|
||||
label_map_loaded_ = true;
|
||||
}
|
||||
if (options.has_min_score_threshold()) {
|
||||
min_score_threshold_ = options.min_score_threshold();
|
||||
}
|
||||
is_binary_classification_ = options.binary_classification();
|
||||
|
||||
if (is_binary_classification_) {
|
||||
RET_CHECK(options.allow_classes().empty() &&
|
||||
options.ignore_classes().empty());
|
||||
}
|
||||
if (!options.allow_classes().empty()) {
|
||||
RET_CHECK(options.ignore_classes().empty());
|
||||
class_index_set_.is_allowlist = true;
|
||||
for (int i = 0; i < options.allow_classes_size(); ++i) {
|
||||
class_index_set_.values.insert(options.allow_classes(i));
|
||||
}
|
||||
} else {
|
||||
class_index_set_.is_allowlist = false;
|
||||
for (int i = 0; i < options.ignore_classes_size(); ++i) {
|
||||
class_index_set_.values.insert(options.ignore_classes(i));
|
||||
}
|
||||
}
|
||||
|
||||
return absl::OkStatus();
|
||||
}
|
||||
@@ -118,19 +168,19 @@ absl::Status TensorsToClassificationCalculator::Process(CalculatorContext* cc) {
|
||||
|
||||
int num_classes = input_tensors[0].shape().num_elements();
|
||||
|
||||
if (options_.binary_classification()) {
|
||||
if (is_binary_classification_) {
|
||||
RET_CHECK_EQ(num_classes, 1);
|
||||
// Number of classes for binary classification.
|
||||
num_classes = 2;
|
||||
}
|
||||
if (label_map_loaded_) {
|
||||
RET_CHECK_EQ(num_classes, label_map_.size());
|
||||
RET_CHECK_EQ(num_classes, GetLabelMap(cc).size());
|
||||
}
|
||||
auto view = input_tensors[0].GetCpuReadView();
|
||||
auto raw_scores = view.buffer<float>();
|
||||
|
||||
auto classification_list = absl::make_unique<ClassificationList>();
|
||||
if (options_.binary_classification()) {
|
||||
if (is_binary_classification_) {
|
||||
Classification* class_first = classification_list->add_classification();
|
||||
Classification* class_second = classification_list->add_classification();
|
||||
class_first->set_index(0);
|
||||
@@ -139,41 +189,48 @@ absl::Status TensorsToClassificationCalculator::Process(CalculatorContext* cc) {
|
||||
class_second->set_score(1. - raw_scores[0]);
|
||||
|
||||
if (label_map_loaded_) {
|
||||
class_first->set_label(label_map_[0]);
|
||||
class_second->set_label(label_map_[1]);
|
||||
SetClassificationLabel(GetLabelMap(cc).at(0), class_first);
|
||||
SetClassificationLabel(GetLabelMap(cc).at(1), class_second);
|
||||
}
|
||||
} else {
|
||||
for (int i = 0; i < num_classes; ++i) {
|
||||
if (options_.has_min_score_threshold() &&
|
||||
raw_scores[i] < options_.min_score_threshold()) {
|
||||
if (!IsClassIndexAllowed(i)) {
|
||||
continue;
|
||||
}
|
||||
if (raw_scores[i] < min_score_threshold_) {
|
||||
continue;
|
||||
}
|
||||
Classification* classification =
|
||||
classification_list->add_classification();
|
||||
classification->set_index(i);
|
||||
classification->set_score(raw_scores[i]);
|
||||
|
||||
if (label_map_loaded_) {
|
||||
classification->set_label(label_map_[i]);
|
||||
SetClassificationLabel(GetLabelMap(cc).at(i), classification);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Note that partial_sort will raise error when top_k_ >
|
||||
// classification_list->classification_size().
|
||||
CHECK_GE(classification_list->classification_size(), top_k_);
|
||||
auto raw_classification_list = classification_list->mutable_classification();
|
||||
if (top_k_ > 0 && classification_list->classification_size() >= top_k_) {
|
||||
if (top_k_ > 0) {
|
||||
int desired_size =
|
||||
std::min(classification_list->classification_size(), top_k_);
|
||||
std::partial_sort(raw_classification_list->begin(),
|
||||
raw_classification_list->begin() + top_k_,
|
||||
raw_classification_list->begin() + desired_size,
|
||||
raw_classification_list->end(),
|
||||
[](const Classification a, const Classification b) {
|
||||
return a.score() > b.score();
|
||||
});
|
||||
|
||||
// Resizes the underlying list to have only top_k_ classifications.
|
||||
raw_classification_list->DeleteSubrange(
|
||||
top_k_, raw_classification_list->size() - top_k_);
|
||||
if (desired_size >= top_k_) {
|
||||
// Resizes the underlying list to have only top_k_ classifications.
|
||||
raw_classification_list->DeleteSubrange(
|
||||
top_k_, raw_classification_list->size() - top_k_);
|
||||
}
|
||||
} else if (sort_by_descending_score_) {
|
||||
std::sort(raw_classification_list->begin(), raw_classification_list->end(),
|
||||
[](const Classification a, const Classification b) {
|
||||
return a.score() > b.score();
|
||||
});
|
||||
}
|
||||
kOutClassificationList(cc).Send(std::move(classification_list));
|
||||
return absl::OkStatus();
|
||||
@@ -183,5 +240,24 @@ absl::Status TensorsToClassificationCalculator::Close(CalculatorContext* cc) {
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
bool TensorsToClassificationCalculator::IsClassIndexAllowed(int class_index) {
|
||||
if (class_index_set_.values.empty()) {
|
||||
return true;
|
||||
}
|
||||
if (class_index_set_.is_allowlist) {
|
||||
return class_index_set_.values.contains(class_index);
|
||||
} else {
|
||||
return !class_index_set_.values.contains(class_index);
|
||||
}
|
||||
}
|
||||
|
||||
const proto_ns::Map<int64, LabelMapItem>&
|
||||
TensorsToClassificationCalculator::GetLabelMap(CalculatorContext* cc) {
|
||||
return !local_label_map_.empty()
|
||||
? local_label_map_
|
||||
: cc->Options<TensorsToClassificationCalculatorOptions>()
|
||||
.label_items();
|
||||
}
|
||||
|
||||
} // namespace api2
|
||||
} // namespace mediapipe
|
||||
|
||||
@@ -19,6 +19,7 @@ syntax = "proto2";
|
||||
package mediapipe;
|
||||
|
||||
import "mediapipe/framework/calculator.proto";
|
||||
import "mediapipe/util/label_map.proto";
|
||||
|
||||
message TensorsToClassificationCalculatorOptions {
|
||||
extend .mediapipe.CalculatorOptions {
|
||||
@@ -38,16 +39,37 @@ message TensorsToClassificationCalculatorOptions {
|
||||
// Number of highest scoring labels to output. If top_k is not positive then
|
||||
// all labels are used.
|
||||
optional int32 top_k = 2;
|
||||
// Whether results should be sorted by descending score. By default, results
|
||||
// may or may not be sorted: setting this to true guarantees that the returned
|
||||
// results will be sorted by descending score.
|
||||
optional bool sort_by_descending_score = 9;
|
||||
// Path to a label map file for getting the actual name of class ids.
|
||||
optional string label_map_path = 3;
|
||||
// Label map. (Can be used instead of label_map_path.)
|
||||
// NOTE: "label_map_path", if specified, takes precedence over "label_map".
|
||||
// NOTE: either "label_map_path" or "label_items", if specified, takes
|
||||
// precedence over "label_map".
|
||||
// Deprecated: please use `label_items` instead.
|
||||
optional LabelMap label_map = 5;
|
||||
|
||||
// Label items. (Can be used instead of label_map_path.)
|
||||
// NOTE: "label_map_path", if specified, takes precedence over "label_items".
|
||||
map<int64, LabelMapItem> label_items = 6;
|
||||
|
||||
// Whether the input is a single float for binary classification.
|
||||
// When true, only a single float is expected in the input tensor and the
|
||||
// label map, if provided, is expected to have exactly two labels.
|
||||
// The single score(float) represent the probability of first label, and
|
||||
// 1 - score is the probabilility of the second label.
|
||||
optional bool binary_classification = 4;
|
||||
|
||||
// The ids of classes that should be ignored during decoding the score for
|
||||
// each classification. If `ignore_classes` is specified, all the other
|
||||
// classes that are not in the `ignore_class` field will be considered during
|
||||
// decoding. `ignore_classes` and `allow_classes` are mutually exclusive.
|
||||
repeated int32 ignore_classes = 7 [packed = true];
|
||||
// The ids of classes that will be allowed during decoding the score for
|
||||
// each classification. If `allow_classes` is specified, all the other classes
|
||||
// that are not in the `allow_classes` field will be completely ignored.
|
||||
// `ignore_classes` and `allow_classes` are mutually exclusive.
|
||||
repeated int32 allow_classes = 8 [packed = true];
|
||||
}
|
||||
|
||||
@@ -12,6 +12,8 @@
|
||||
// See the License for the specific language governing permissions and
|
||||
// limitations under the License.
|
||||
|
||||
#include <algorithm>
|
||||
#include <limits>
|
||||
#include <vector>
|
||||
|
||||
#include "absl/memory/memory.h"
|
||||
@@ -206,4 +208,119 @@ TEST_F(TensorsToClassificationCalculatorTest, CorrectOutputWithTopK) {
|
||||
}
|
||||
}
|
||||
|
||||
TEST_F(TensorsToClassificationCalculatorTest,
|
||||
CorrectOutputWithSortByDescendingScore) {
|
||||
mediapipe::CalculatorRunner runner(ParseTextProtoOrDie<Node>(R"pb(
|
||||
calculator: "TensorsToClassificationCalculator"
|
||||
input_stream: "TENSORS:tensors"
|
||||
output_stream: "CLASSIFICATIONS:classifications"
|
||||
options {
|
||||
[mediapipe.TensorsToClassificationCalculatorOptions.ext] {
|
||||
sort_by_descending_score: true
|
||||
}
|
||||
}
|
||||
)pb"));
|
||||
|
||||
BuildGraph(&runner, {0, 0.5, 1});
|
||||
MP_ASSERT_OK(runner.Run());
|
||||
|
||||
const auto& output_packets_ = runner.Outputs().Tag("CLASSIFICATIONS").packets;
|
||||
|
||||
EXPECT_EQ(1, output_packets_.size());
|
||||
|
||||
const auto& classification_list =
|
||||
output_packets_[0].Get<ClassificationList>();
|
||||
|
||||
// Verify results are sorted by descending score.
|
||||
EXPECT_EQ(3, classification_list.classification_size());
|
||||
float score = std::numeric_limits<float>::max();
|
||||
for (int i = 0; i < classification_list.classification_size(); ++i) {
|
||||
EXPECT_LE(classification_list.classification(i).score(), score);
|
||||
score = classification_list.classification(i).score();
|
||||
}
|
||||
}
|
||||
|
||||
TEST_F(TensorsToClassificationCalculatorTest,
|
||||
ClassNameAllowlistWithLabelItems) {
|
||||
mediapipe::CalculatorRunner runner(ParseTextProtoOrDie<Node>(R"pb(
|
||||
calculator: "TensorsToClassificationCalculator"
|
||||
input_stream: "TENSORS:tensors"
|
||||
output_stream: "CLASSIFICATIONS:classifications"
|
||||
options {
|
||||
[mediapipe.TensorsToClassificationCalculatorOptions.ext] {
|
||||
label_items {
|
||||
key: 0
|
||||
value { name: "ClassA" }
|
||||
}
|
||||
label_items {
|
||||
key: 1
|
||||
value { name: "ClassB" }
|
||||
}
|
||||
label_items {
|
||||
key: 2
|
||||
value { name: "ClassC" }
|
||||
}
|
||||
allow_classes: 1
|
||||
}
|
||||
}
|
||||
)pb"));
|
||||
|
||||
BuildGraph(&runner, {0, 0.5, 1});
|
||||
MP_ASSERT_OK(runner.Run());
|
||||
|
||||
const auto& output_packets_ = runner.Outputs().Tag("CLASSIFICATIONS").packets;
|
||||
|
||||
EXPECT_EQ(1, output_packets_.size());
|
||||
|
||||
const auto& classification_list =
|
||||
output_packets_[0].Get<ClassificationList>();
|
||||
EXPECT_EQ(1, classification_list.classification_size());
|
||||
EXPECT_EQ(1, classification_list.classification(0).index());
|
||||
EXPECT_EQ(0.5, classification_list.classification(0).score());
|
||||
ASSERT_TRUE(classification_list.classification(0).has_label());
|
||||
}
|
||||
|
||||
TEST_F(TensorsToClassificationCalculatorTest,
|
||||
ClassNameIgnorelistWithLabelItems) {
|
||||
mediapipe::CalculatorRunner runner(ParseTextProtoOrDie<Node>(R"pb(
|
||||
calculator: "TensorsToClassificationCalculator"
|
||||
input_stream: "TENSORS:tensors"
|
||||
output_stream: "CLASSIFICATIONS:classifications"
|
||||
options {
|
||||
[mediapipe.TensorsToClassificationCalculatorOptions.ext] {
|
||||
label_items {
|
||||
key: 0
|
||||
value { name: "ClassA" }
|
||||
}
|
||||
label_items {
|
||||
key: 1
|
||||
value { name: "ClassB" }
|
||||
}
|
||||
label_items {
|
||||
key: 2
|
||||
value { name: "ClassC" }
|
||||
}
|
||||
ignore_classes: 1
|
||||
}
|
||||
}
|
||||
)pb"));
|
||||
|
||||
BuildGraph(&runner, {0, 0.5, 1});
|
||||
MP_ASSERT_OK(runner.Run());
|
||||
|
||||
const auto& output_packets_ = runner.Outputs().Tag("CLASSIFICATIONS").packets;
|
||||
|
||||
EXPECT_EQ(1, output_packets_.size());
|
||||
|
||||
const auto& classification_list =
|
||||
output_packets_[0].Get<ClassificationList>();
|
||||
EXPECT_EQ(2, classification_list.classification_size());
|
||||
EXPECT_EQ(0, classification_list.classification(0).index());
|
||||
EXPECT_EQ(0, classification_list.classification(0).score());
|
||||
ASSERT_TRUE(classification_list.classification(0).has_label());
|
||||
EXPECT_EQ(2, classification_list.classification(1).index());
|
||||
EXPECT_EQ(1, classification_list.classification(1).score());
|
||||
ASSERT_TRUE(classification_list.classification(1).has_label());
|
||||
}
|
||||
|
||||
} // namespace mediapipe
|
||||
|
||||
@@ -91,6 +91,40 @@ void ConvertAnchorsToRawValues(const std::vector<Anchor>& anchors,
|
||||
}
|
||||
}
|
||||
|
||||
absl::Status CheckCustomTensorMapping(
|
||||
const TensorsToDetectionsCalculatorOptions::TensorMapping& tensor_mapping) {
|
||||
RET_CHECK(tensor_mapping.has_detections_tensor_index() &&
|
||||
tensor_mapping.has_scores_tensor_index());
|
||||
int bitmap = 0;
|
||||
bitmap |= 1 << tensor_mapping.detections_tensor_index();
|
||||
bitmap |= 1 << tensor_mapping.scores_tensor_index();
|
||||
if (!tensor_mapping.has_num_detections_tensor_index() &&
|
||||
!tensor_mapping.has_classes_tensor_index() &&
|
||||
!tensor_mapping.has_anchors_tensor_index()) {
|
||||
// Only allows the output tensor index 0 and 1 to be occupied.
|
||||
RET_CHECK_EQ(3, bitmap) << "The custom output tensor indices should only "
|
||||
"cover index 0 and 1.";
|
||||
} else if (tensor_mapping.has_anchors_tensor_index()) {
|
||||
RET_CHECK(!tensor_mapping.has_classes_tensor_index() &&
|
||||
!tensor_mapping.has_num_detections_tensor_index());
|
||||
bitmap |= 1 << tensor_mapping.anchors_tensor_index();
|
||||
// If the"anchors" tensor will be available, only allows the output tensor
|
||||
// index 0, 1, 2 to be occupied.
|
||||
RET_CHECK_EQ(7, bitmap) << "The custom output tensor indices should only "
|
||||
"cover index 0, 1 and 2.";
|
||||
} else {
|
||||
RET_CHECK(tensor_mapping.has_classes_tensor_index() &&
|
||||
tensor_mapping.has_num_detections_tensor_index());
|
||||
// If the "classes" and the "number of detections" tensors will be
|
||||
// available, only allows the output tensor index 0, 1, 2, 3 to be occupied.
|
||||
bitmap |= 1 << tensor_mapping.classes_tensor_index();
|
||||
bitmap |= 1 << tensor_mapping.num_detections_tensor_index();
|
||||
RET_CHECK_EQ(15, bitmap) << "The custom output tensor indices should only "
|
||||
"cover index 0, 1, 2 and 3.";
|
||||
}
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
} // namespace
|
||||
|
||||
// Convert result Tensors from object detection models into MediaPipe
|
||||
@@ -170,13 +204,27 @@ class TensorsToDetectionsCalculator : public Node {
|
||||
Detection ConvertToDetection(float box_ymin, float box_xmin, float box_ymax,
|
||||
float box_xmax, float score, int class_id,
|
||||
bool flip_vertically);
|
||||
bool IsClassIndexAllowed(int class_index);
|
||||
|
||||
int num_classes_ = 0;
|
||||
int num_boxes_ = 0;
|
||||
int num_coords_ = 0;
|
||||
std::set<int> ignore_classes_;
|
||||
int max_results_ = -1;
|
||||
|
||||
::mediapipe::TensorsToDetectionsCalculatorOptions options_;
|
||||
// Set of allowed or ignored class indices.
|
||||
struct ClassIndexSet {
|
||||
absl::flat_hash_set<int> values;
|
||||
bool is_allowlist;
|
||||
};
|
||||
// Allowed or ignored class indices based on provided options or side packet.
|
||||
// These are used to filter out the output detection results.
|
||||
ClassIndexSet class_index_set_;
|
||||
|
||||
TensorsToDetectionsCalculatorOptions options_;
|
||||
bool scores_tensor_index_is_set_ = false;
|
||||
TensorsToDetectionsCalculatorOptions::TensorMapping tensor_mapping_;
|
||||
std::vector<int> box_indices_ = {0, 1, 2, 3};
|
||||
bool has_custom_box_indices_ = false;
|
||||
std::vector<Anchor> anchors_;
|
||||
|
||||
#ifndef MEDIAPIPE_DISABLE_GL_COMPUTE
|
||||
@@ -239,6 +287,21 @@ absl::Status TensorsToDetectionsCalculator::Process(CalculatorContext* cc) {
|
||||
}
|
||||
}
|
||||
}
|
||||
const int num_input_tensors = kInTensors(cc)->size();
|
||||
if (!scores_tensor_index_is_set_) {
|
||||
if (num_input_tensors == 2 ||
|
||||
num_input_tensors == kNumInputTensorsWithAnchors) {
|
||||
tensor_mapping_.set_scores_tensor_index(1);
|
||||
} else {
|
||||
tensor_mapping_.set_scores_tensor_index(2);
|
||||
}
|
||||
scores_tensor_index_is_set_ = true;
|
||||
}
|
||||
if (gpu_processing || num_input_tensors != 4) {
|
||||
// Allows custom bounding box indices when receiving 4 cpu tensors.
|
||||
// Uses the default bbox indices in other cases.
|
||||
RET_CHECK(!has_custom_box_indices_);
|
||||
}
|
||||
|
||||
if (gpu_processing) {
|
||||
if (!gpu_inited_) {
|
||||
@@ -263,12 +326,15 @@ absl::Status TensorsToDetectionsCalculator::ProcessCPU(
|
||||
// Postprocessing on CPU for model without postprocessing op. E.g. output
|
||||
// raw score tensor and box tensor. Anchor decoding will be handled below.
|
||||
// TODO: Add flexible input tensor size handling.
|
||||
auto raw_box_tensor = &input_tensors[0];
|
||||
auto raw_box_tensor =
|
||||
&input_tensors[tensor_mapping_.detections_tensor_index()];
|
||||
RET_CHECK_EQ(raw_box_tensor->shape().dims.size(), 3);
|
||||
RET_CHECK_EQ(raw_box_tensor->shape().dims[0], 1);
|
||||
RET_CHECK_GT(num_boxes_, 0) << "Please set num_boxes in calculator options";
|
||||
RET_CHECK_EQ(raw_box_tensor->shape().dims[1], num_boxes_);
|
||||
RET_CHECK_EQ(raw_box_tensor->shape().dims[2], num_coords_);
|
||||
auto raw_score_tensor = &input_tensors[1];
|
||||
auto raw_score_tensor =
|
||||
&input_tensors[tensor_mapping_.scores_tensor_index()];
|
||||
RET_CHECK_EQ(raw_score_tensor->shape().dims.size(), 3);
|
||||
RET_CHECK_EQ(raw_score_tensor->shape().dims[0], 1);
|
||||
RET_CHECK_EQ(raw_score_tensor->shape().dims[1], num_boxes_);
|
||||
@@ -281,7 +347,8 @@ absl::Status TensorsToDetectionsCalculator::ProcessCPU(
|
||||
// TODO: Support other options to load anchors.
|
||||
if (!anchors_init_) {
|
||||
if (input_tensors.size() == kNumInputTensorsWithAnchors) {
|
||||
auto anchor_tensor = &input_tensors[2];
|
||||
auto anchor_tensor =
|
||||
&input_tensors[tensor_mapping_.anchors_tensor_index()];
|
||||
RET_CHECK_EQ(anchor_tensor->shape().dims.size(), 2);
|
||||
RET_CHECK_EQ(anchor_tensor->shape().dims[0], num_boxes_);
|
||||
RET_CHECK_EQ(anchor_tensor->shape().dims[1], kNumCoordsPerBox);
|
||||
@@ -307,7 +374,7 @@ absl::Status TensorsToDetectionsCalculator::ProcessCPU(
|
||||
float max_score = -std::numeric_limits<float>::max();
|
||||
// Find the top score for box i.
|
||||
for (int score_idx = 0; score_idx < num_classes_; ++score_idx) {
|
||||
if (ignore_classes_.find(score_idx) == ignore_classes_.end()) {
|
||||
if (IsClassIndexAllowed(score_idx)) {
|
||||
auto score = raw_scores[i * num_classes_ + score_idx];
|
||||
if (options_.sigmoid_score()) {
|
||||
if (options_.has_score_clipping_thresh()) {
|
||||
@@ -337,23 +404,26 @@ absl::Status TensorsToDetectionsCalculator::ProcessCPU(
|
||||
// Postprocessing on CPU with postprocessing op (e.g. anchor decoding and
|
||||
// non-maximum suppression) within the model.
|
||||
RET_CHECK_EQ(input_tensors.size(), 4);
|
||||
|
||||
auto num_boxes_tensor = &input_tensors[3];
|
||||
auto num_boxes_tensor =
|
||||
&input_tensors[tensor_mapping_.num_detections_tensor_index()];
|
||||
RET_CHECK_EQ(num_boxes_tensor->shape().dims.size(), 1);
|
||||
RET_CHECK_EQ(num_boxes_tensor->shape().dims[0], 1);
|
||||
|
||||
auto detection_boxes_tensor = &input_tensors[0];
|
||||
auto detection_boxes_tensor =
|
||||
&input_tensors[tensor_mapping_.detections_tensor_index()];
|
||||
RET_CHECK_EQ(detection_boxes_tensor->shape().dims.size(), 3);
|
||||
RET_CHECK_EQ(detection_boxes_tensor->shape().dims[0], 1);
|
||||
const int max_detections = detection_boxes_tensor->shape().dims[1];
|
||||
RET_CHECK_EQ(detection_boxes_tensor->shape().dims[2], num_coords_);
|
||||
|
||||
auto detection_classes_tensor = &input_tensors[1];
|
||||
auto detection_classes_tensor =
|
||||
&input_tensors[tensor_mapping_.classes_tensor_index()];
|
||||
RET_CHECK_EQ(detection_classes_tensor->shape().dims.size(), 2);
|
||||
RET_CHECK_EQ(detection_classes_tensor->shape().dims[0], 1);
|
||||
RET_CHECK_EQ(detection_classes_tensor->shape().dims[1], max_detections);
|
||||
|
||||
auto detection_scores_tensor = &input_tensors[2];
|
||||
auto detection_scores_tensor =
|
||||
&input_tensors[tensor_mapping_.scores_tensor_index()];
|
||||
RET_CHECK_EQ(detection_scores_tensor->shape().dims.size(), 2);
|
||||
RET_CHECK_EQ(detection_scores_tensor->shape().dims[0], 1);
|
||||
RET_CHECK_EQ(detection_scores_tensor->shape().dims[1], max_detections);
|
||||
@@ -385,6 +455,7 @@ absl::Status TensorsToDetectionsCalculator::ProcessGPU(
|
||||
CalculatorContext* cc, std::vector<Detection>* output_detections) {
|
||||
const auto& input_tensors = *kInTensors(cc);
|
||||
RET_CHECK_GE(input_tensors.size(), 2);
|
||||
RET_CHECK_GT(num_boxes_, 0) << "Please set num_boxes in calculator options";
|
||||
#ifndef MEDIAPIPE_DISABLE_GL_COMPUTE
|
||||
|
||||
MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext([this, &input_tensors, &cc,
|
||||
@@ -392,12 +463,14 @@ absl::Status TensorsToDetectionsCalculator::ProcessGPU(
|
||||
-> absl::Status {
|
||||
if (!anchors_init_) {
|
||||
if (input_tensors.size() == kNumInputTensorsWithAnchors) {
|
||||
auto read_view = input_tensors[2].GetOpenGlBufferReadView();
|
||||
auto read_view = input_tensors[tensor_mapping_.anchors_tensor_index()]
|
||||
.GetOpenGlBufferReadView();
|
||||
glBindBuffer(GL_COPY_READ_BUFFER, read_view.name());
|
||||
auto write_view = raw_anchors_buffer_->GetOpenGlBufferWriteView();
|
||||
glBindBuffer(GL_COPY_WRITE_BUFFER, write_view.name());
|
||||
glCopyBufferSubData(GL_COPY_READ_BUFFER, GL_COPY_WRITE_BUFFER, 0, 0,
|
||||
input_tensors[2].bytes());
|
||||
glCopyBufferSubData(
|
||||
GL_COPY_READ_BUFFER, GL_COPY_WRITE_BUFFER, 0, 0,
|
||||
input_tensors[tensor_mapping_.anchors_tensor_index()].bytes());
|
||||
} else if (!kInAnchors(cc).IsEmpty()) {
|
||||
const auto& anchors = *kInAnchors(cc);
|
||||
auto anchors_view = raw_anchors_buffer_->GetCpuWriteView();
|
||||
@@ -416,7 +489,9 @@ absl::Status TensorsToDetectionsCalculator::ProcessGPU(
|
||||
auto decoded_boxes_view =
|
||||
decoded_boxes_buffer_->GetOpenGlBufferWriteView();
|
||||
glBindBufferBase(GL_SHADER_STORAGE_BUFFER, 0, decoded_boxes_view.name());
|
||||
auto input0_view = input_tensors[0].GetOpenGlBufferReadView();
|
||||
auto input0_view =
|
||||
input_tensors[tensor_mapping_.detections_tensor_index()]
|
||||
.GetOpenGlBufferReadView();
|
||||
glBindBufferBase(GL_SHADER_STORAGE_BUFFER, 1, input0_view.name());
|
||||
auto raw_anchors_view = raw_anchors_buffer_->GetOpenGlBufferReadView();
|
||||
glBindBufferBase(GL_SHADER_STORAGE_BUFFER, 2, raw_anchors_view.name());
|
||||
@@ -425,7 +500,8 @@ absl::Status TensorsToDetectionsCalculator::ProcessGPU(
|
||||
|
||||
// Score boxes.
|
||||
glBindBufferBase(GL_SHADER_STORAGE_BUFFER, 0, scored_boxes_view.name());
|
||||
auto input1_view = input_tensors[1].GetOpenGlBufferReadView();
|
||||
auto input1_view = input_tensors[tensor_mapping_.scores_tensor_index()]
|
||||
.GetOpenGlBufferReadView();
|
||||
glBindBufferBase(GL_SHADER_STORAGE_BUFFER, 1, input1_view.name());
|
||||
glUseProgram(score_program_);
|
||||
glDispatchCompute(num_boxes_, 1, 1);
|
||||
@@ -457,7 +533,8 @@ absl::Status TensorsToDetectionsCalculator::ProcessGPU(
|
||||
if (input_tensors.size() == kNumInputTensorsWithAnchors) {
|
||||
RET_CHECK_EQ(input_tensors.size(), kNumInputTensorsWithAnchors);
|
||||
auto command_buffer = [gpu_helper_ commandBuffer];
|
||||
auto src_buffer = input_tensors[2].GetMtlBufferReadView(command_buffer);
|
||||
auto src_buffer = input_tensors[tensor_mapping_.anchors_tensor_index()]
|
||||
.GetMtlBufferReadView(command_buffer);
|
||||
auto dest_buffer =
|
||||
raw_anchors_buffer_->GetMtlBufferWriteView(command_buffer);
|
||||
id<MTLBlitCommandEncoder> blit_command =
|
||||
@@ -466,7 +543,9 @@ absl::Status TensorsToDetectionsCalculator::ProcessGPU(
|
||||
sourceOffset:0
|
||||
toBuffer:dest_buffer.buffer()
|
||||
destinationOffset:0
|
||||
size:input_tensors[2].bytes()];
|
||||
size:input_tensors[tensor_mapping_
|
||||
.anchors_tensor_index()]
|
||||
.bytes()];
|
||||
[blit_command endEncoding];
|
||||
[command_buffer commit];
|
||||
} else if (!kInAnchors(cc).IsEmpty()) {
|
||||
@@ -493,7 +572,8 @@ absl::Status TensorsToDetectionsCalculator::ProcessGPU(
|
||||
auto decoded_boxes_view =
|
||||
decoded_boxes_buffer_->GetMtlBufferWriteView(command_buffer);
|
||||
[command_encoder setBuffer:decoded_boxes_view.buffer() offset:0 atIndex:0];
|
||||
auto input0_view = input_tensors[0].GetMtlBufferReadView(command_buffer);
|
||||
auto input0_view = input_tensors[tensor_mapping_.detections_tensor_index()]
|
||||
.GetMtlBufferReadView(command_buffer);
|
||||
[command_encoder setBuffer:input0_view.buffer() offset:0 atIndex:1];
|
||||
auto raw_anchors_view =
|
||||
raw_anchors_buffer_->GetMtlBufferReadView(command_buffer);
|
||||
@@ -505,7 +585,8 @@ absl::Status TensorsToDetectionsCalculator::ProcessGPU(
|
||||
|
||||
[command_encoder setComputePipelineState:score_program_];
|
||||
[command_encoder setBuffer:scored_boxes_view.buffer() offset:0 atIndex:0];
|
||||
auto input1_view = input_tensors[1].GetMtlBufferReadView(command_buffer);
|
||||
auto input1_view = input_tensors[tensor_mapping_.scores_tensor_index()]
|
||||
.GetMtlBufferReadView(command_buffer);
|
||||
[command_encoder setBuffer:input1_view.buffer() offset:0 atIndex:1];
|
||||
MTLSize score_threads_per_group = MTLSizeMake(1, num_classes_, 1);
|
||||
MTLSize score_threadgroups = MTLSizeMake(num_boxes_, 1, 1);
|
||||
@@ -563,12 +644,15 @@ absl::Status TensorsToDetectionsCalculator::LoadOptions(CalculatorContext* cc) {
|
||||
// Get calculator options specified in the graph.
|
||||
options_ = cc->Options<::mediapipe::TensorsToDetectionsCalculatorOptions>();
|
||||
RET_CHECK(options_.has_num_classes());
|
||||
RET_CHECK(options_.has_num_boxes());
|
||||
RET_CHECK(options_.has_num_coords());
|
||||
|
||||
num_classes_ = options_.num_classes();
|
||||
num_boxes_ = options_.num_boxes();
|
||||
num_coords_ = options_.num_coords();
|
||||
CHECK_NE(options_.max_results(), 0)
|
||||
<< "The maximum number of the top-scored detection results must be "
|
||||
"non-zero.";
|
||||
max_results_ = options_.max_results();
|
||||
|
||||
// Currently only support 2D when num_values_per_keypoint equals to 2.
|
||||
CHECK_EQ(options_.num_values_per_keypoint(), 2);
|
||||
@@ -580,15 +664,55 @@ absl::Status TensorsToDetectionsCalculator::LoadOptions(CalculatorContext* cc) {
|
||||
|
||||
if (kSideInIgnoreClasses(cc).IsConnected()) {
|
||||
RET_CHECK(!kSideInIgnoreClasses(cc).IsEmpty());
|
||||
RET_CHECK(options_.allow_classes().empty());
|
||||
class_index_set_.is_allowlist = false;
|
||||
for (int ignore_class : *kSideInIgnoreClasses(cc)) {
|
||||
ignore_classes_.insert(ignore_class);
|
||||
class_index_set_.values.insert(ignore_class);
|
||||
}
|
||||
} else if (!options_.allow_classes().empty()) {
|
||||
RET_CHECK(options_.ignore_classes().empty());
|
||||
class_index_set_.is_allowlist = true;
|
||||
for (int i = 0; i < options_.allow_classes_size(); ++i) {
|
||||
class_index_set_.values.insert(options_.allow_classes(i));
|
||||
}
|
||||
} else {
|
||||
class_index_set_.is_allowlist = false;
|
||||
for (int i = 0; i < options_.ignore_classes_size(); ++i) {
|
||||
ignore_classes_.insert(options_.ignore_classes(i));
|
||||
class_index_set_.values.insert(options_.ignore_classes(i));
|
||||
}
|
||||
}
|
||||
|
||||
if (options_.has_tensor_mapping()) {
|
||||
RET_CHECK_OK(CheckCustomTensorMapping(options_.tensor_mapping()));
|
||||
tensor_mapping_ = options_.tensor_mapping();
|
||||
scores_tensor_index_is_set_ = true;
|
||||
} else {
|
||||
// Assigns the default tensor indices.
|
||||
tensor_mapping_.set_detections_tensor_index(0);
|
||||
tensor_mapping_.set_classes_tensor_index(1);
|
||||
tensor_mapping_.set_anchors_tensor_index(2);
|
||||
tensor_mapping_.set_num_detections_tensor_index(3);
|
||||
// The scores tensor index needs to be determined based on the number of
|
||||
// model's output tensors, which will be available in the first invocation
|
||||
// of the Process() method.
|
||||
tensor_mapping_.set_scores_tensor_index(-1);
|
||||
scores_tensor_index_is_set_ = false;
|
||||
}
|
||||
|
||||
if (options_.has_box_boundaries_indices()) {
|
||||
box_indices_ = {options_.box_boundaries_indices().ymin(),
|
||||
options_.box_boundaries_indices().xmin(),
|
||||
options_.box_boundaries_indices().ymax(),
|
||||
options_.box_boundaries_indices().xmax()};
|
||||
int bitmap = 0;
|
||||
for (int i : box_indices_) {
|
||||
bitmap |= 1 << i;
|
||||
}
|
||||
RET_CHECK_EQ(bitmap, 15) << "The custom box boundaries indices should only "
|
||||
"cover index 0, 1, 2, and 3.";
|
||||
has_custom_box_indices_ = true;
|
||||
}
|
||||
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
@@ -660,14 +784,22 @@ absl::Status TensorsToDetectionsCalculator::ConvertToDetections(
|
||||
const float* detection_boxes, const float* detection_scores,
|
||||
const int* detection_classes, std::vector<Detection>* output_detections) {
|
||||
for (int i = 0; i < num_boxes_; ++i) {
|
||||
if (max_results_ > 0 && output_detections->size() == max_results_) {
|
||||
break;
|
||||
}
|
||||
if (options_.has_min_score_thresh() &&
|
||||
detection_scores[i] < options_.min_score_thresh()) {
|
||||
continue;
|
||||
}
|
||||
if (!IsClassIndexAllowed(detection_classes[i])) {
|
||||
continue;
|
||||
}
|
||||
const int box_offset = i * num_coords_;
|
||||
Detection detection = ConvertToDetection(
|
||||
detection_boxes[box_offset + 0], detection_boxes[box_offset + 1],
|
||||
detection_boxes[box_offset + 2], detection_boxes[box_offset + 3],
|
||||
/*box_ymin=*/detection_boxes[box_offset + box_indices_[0]],
|
||||
/*box_xmin=*/detection_boxes[box_offset + box_indices_[1]],
|
||||
/*box_ymax=*/detection_boxes[box_offset + box_indices_[2]],
|
||||
/*box_xmax=*/detection_boxes[box_offset + box_indices_[3]],
|
||||
detection_scores[i], detection_classes[i], options_.flip_vertically());
|
||||
const auto& bbox = detection.location_data().relative_bounding_box();
|
||||
if (bbox.width() < 0 || bbox.height() < 0 || std::isnan(bbox.width()) ||
|
||||
@@ -909,7 +1041,7 @@ void main() {
|
||||
options_.has_score_clipping_thresh() ? 1 : 0,
|
||||
options_.has_score_clipping_thresh() ? options_.score_clipping_thresh()
|
||||
: 0,
|
||||
!ignore_classes_.empty() ? 1 : 0);
|
||||
!IsClassIndexAllowed(0));
|
||||
|
||||
// # filter classes supported is hardware dependent.
|
||||
int max_wg_size; // typically <= 1024
|
||||
@@ -918,7 +1050,14 @@ void main() {
|
||||
CHECK_LT(num_classes_, max_wg_size)
|
||||
<< "# classes must be < " << max_wg_size;
|
||||
// TODO support better filtering.
|
||||
CHECK_LE(ignore_classes_.size(), 1) << "Only ignore class 0 is allowed";
|
||||
if (class_index_set_.is_allowlist) {
|
||||
CHECK_EQ(class_index_set_.values.size(),
|
||||
IsClassIndexAllowed(0) ? num_classes_ : num_classes_ - 1)
|
||||
<< "Only all classes >= class 0 or >= class 1";
|
||||
} else {
|
||||
CHECK_EQ(class_index_set_.values.size(), IsClassIndexAllowed(0) ? 0 : 1)
|
||||
<< "Only ignore class 0 is allowed";
|
||||
}
|
||||
|
||||
// Shader program
|
||||
{
|
||||
@@ -1125,10 +1264,17 @@ kernel void scoreKernel(
|
||||
options_.has_score_clipping_thresh() ? 1 : 0,
|
||||
options_.has_score_clipping_thresh() ? options_.score_clipping_thresh()
|
||||
: 0,
|
||||
ignore_classes_.size() ? 1 : 0);
|
||||
!IsClassIndexAllowed(0));
|
||||
|
||||
// TODO support better filtering.
|
||||
CHECK_LE(ignore_classes_.size(), 1) << "Only ignore class 0 is allowed";
|
||||
if (class_index_set_.is_allowlist) {
|
||||
CHECK_EQ(class_index_set_.values.size(),
|
||||
IsClassIndexAllowed(0) ? num_classes_ : num_classes_ - 1)
|
||||
<< "Only all classes >= class 0 or >= class 1";
|
||||
} else {
|
||||
CHECK_EQ(class_index_set_.values.size(), IsClassIndexAllowed(0) ? 0 : 1)
|
||||
<< "Only ignore class 0 is allowed";
|
||||
}
|
||||
|
||||
{
|
||||
// Shader program
|
||||
@@ -1160,5 +1306,16 @@ kernel void scoreKernel(
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
bool TensorsToDetectionsCalculator::IsClassIndexAllowed(int class_index) {
|
||||
if (class_index_set_.values.empty()) {
|
||||
return true;
|
||||
}
|
||||
if (class_index_set_.is_allowlist) {
|
||||
return class_index_set_.values.contains(class_index);
|
||||
} else {
|
||||
return !class_index_set_.values.contains(class_index);
|
||||
}
|
||||
}
|
||||
|
||||
} // namespace api2
|
||||
} // namespace mediapipe
|
||||
|
||||
@@ -57,7 +57,12 @@ message TensorsToDetectionsCalculatorOptions {
|
||||
optional bool reverse_output_order = 14 [default = false];
|
||||
// The ids of classes that should be ignored during decoding the score for
|
||||
// each predicted box. Can be overridden with IGNORE_CLASSES side packet.
|
||||
// `ignore_classes` and `allow_classes` are mutually exclusive.
|
||||
repeated int32 ignore_classes = 8;
|
||||
// The ids of classes that should be allowed during decoding the score for
|
||||
// each predicted box. `ignore_classes` and `allow_classes` are mutually
|
||||
// exclusive.
|
||||
repeated int32 allow_classes = 21 [packed = true];
|
||||
|
||||
optional bool sigmoid_score = 15 [default = false];
|
||||
optional float score_clipping_thresh = 16;
|
||||
@@ -71,4 +76,40 @@ message TensorsToDetectionsCalculatorOptions {
|
||||
|
||||
// Score threshold for perserving decoded detections.
|
||||
optional float min_score_thresh = 19;
|
||||
|
||||
// The maximum number of the detection results to return. If < 0, all
|
||||
// available results will be returned.
|
||||
// For the detection models that have built-in non max suppression op, the
|
||||
// output detections are the top-scored results. Otherwise, the output
|
||||
// detections are the first N results that have higher scores than
|
||||
// `min_score_thresh`.
|
||||
optional int32 max_results = 20 [default = -1];
|
||||
|
||||
// The custom model output tensor mapping.
|
||||
// The indices of the "detections" tensor and the "scores" tensor are always
|
||||
// required. If the model outputs an "anchors" tensor, `anchors_tensor_index`
|
||||
// must be specified. If the model outputs both "classes" tensor and "number
|
||||
// of detections" tensors, `classes_tensor_index` and
|
||||
// `num_detections_tensor_index` must be set.
|
||||
message TensorMapping {
|
||||
optional int32 detections_tensor_index = 1;
|
||||
optional int32 classes_tensor_index = 2;
|
||||
optional int32 scores_tensor_index = 3;
|
||||
optional int32 num_detections_tensor_index = 4;
|
||||
optional int32 anchors_tensor_index = 5;
|
||||
}
|
||||
optional TensorMapping tensor_mapping = 22;
|
||||
|
||||
// Represents the bounding box by using the combination of boundaries,
|
||||
// {ymin, xmin, ymax, xmax}.
|
||||
// The default order is {ymin, xmin, ymax, xmax}.
|
||||
message BoxBoundariesIndices {
|
||||
optional int32 ymin = 1 [default = 0];
|
||||
optional int32 xmin = 2 [default = 1];
|
||||
optional int32 ymax = 3 [default = 2];
|
||||
optional int32 xmax = 4 [default = 3];
|
||||
}
|
||||
oneof box_indices {
|
||||
BoxBoundariesIndices box_boundaries_indices = 23;
|
||||
}
|
||||
}
|
||||
|
||||
@@ -20,10 +20,8 @@
|
||||
#include "mediapipe/framework/calculator_context.h"
|
||||
#include "mediapipe/framework/calculator_framework.h"
|
||||
#include "mediapipe/framework/formats/image.h"
|
||||
#include "mediapipe/framework/formats/image_opencv.h"
|
||||
#include "mediapipe/framework/formats/tensor.h"
|
||||
#include "mediapipe/framework/port.h"
|
||||
#include "mediapipe/framework/port/opencv_imgproc_inc.h"
|
||||
#include "mediapipe/framework/port/ret_check.h"
|
||||
#include "mediapipe/framework/port/statusor.h"
|
||||
#include "mediapipe/gpu/gpu_origin.pb.h"
|
||||
@@ -37,6 +35,11 @@
|
||||
#include "mediapipe/gpu/shader_util.h"
|
||||
#endif // !MEDIAPIPE_DISABLE_GPU
|
||||
|
||||
#if !MEDIAPIPE_DISABLE_OPENCV
|
||||
#include "mediapipe/framework/formats/image_opencv.h"
|
||||
#include "mediapipe/framework/port/opencv_imgproc_inc.h"
|
||||
#endif // !MEDIAPIPE_DISABLE_OPENCV
|
||||
|
||||
#if MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_31
|
||||
#include "tensorflow/lite/delegates/gpu/gl/converters/util.h"
|
||||
#include "tensorflow/lite/delegates/gpu/gl/gl_program.h"
|
||||
@@ -159,9 +162,10 @@ class TensorsToSegmentationCalculator : public CalculatorBase {
|
||||
return options_.gpu_origin() != mediapipe::GpuOrigin_Mode_TOP_LEFT;
|
||||
}
|
||||
|
||||
#if !MEDIAPIPE_DISABLE_OPENCV
|
||||
template <class T>
|
||||
absl::Status ApplyActivation(cv::Mat& tensor_mat, cv::Mat* small_mask_mat);
|
||||
|
||||
#endif // !MEDIAPIPE_DISABLE_OPENCV
|
||||
::mediapipe::TensorsToSegmentationCalculatorOptions options_;
|
||||
|
||||
#if !MEDIAPIPE_DISABLE_GPU
|
||||
@@ -283,7 +287,11 @@ absl::Status TensorsToSegmentationCalculator::Process(CalculatorContext* cc) {
|
||||
RET_CHECK_FAIL() << "GPU processing disabled.";
|
||||
#endif // !MEDIAPIPE_DISABLE_GPU
|
||||
} else {
|
||||
#if !MEDIAPIPE_DISABLE_OPENCV
|
||||
MP_RETURN_IF_ERROR(ProcessCpu(cc));
|
||||
#else
|
||||
RET_CHECK_FAIL() << "OpenCV processing disabled.";
|
||||
#endif // !MEDIAPIPE_DISABLE_OPENCV
|
||||
}
|
||||
|
||||
return absl::OkStatus();
|
||||
@@ -311,6 +319,7 @@ absl::Status TensorsToSegmentationCalculator::Close(CalculatorContext* cc) {
|
||||
|
||||
absl::Status TensorsToSegmentationCalculator::ProcessCpu(
|
||||
CalculatorContext* cc) {
|
||||
#if !MEDIAPIPE_DISABLE_OPENCV
|
||||
// Get input streams, and dimensions.
|
||||
const auto& input_tensors =
|
||||
cc->Inputs().Tag(kTensorsTag).Get<std::vector<Tensor>>();
|
||||
@@ -355,14 +364,17 @@ absl::Status TensorsToSegmentationCalculator::ProcessCpu(
|
||||
std::shared_ptr<ImageFrame> mask_frame = std::make_shared<ImageFrame>(
|
||||
ImageFormat::VEC32F1, output_width, output_height);
|
||||
std::unique_ptr<Image> output_mask = absl::make_unique<Image>(mask_frame);
|
||||
cv::Mat output_mat = formats::MatView(output_mask.get());
|
||||
auto output_mat = formats::MatView(output_mask.get());
|
||||
// Upsample small mask into output.
|
||||
cv::resize(small_mask_mat, output_mat, cv::Size(output_width, output_height));
|
||||
cv::resize(small_mask_mat, *output_mat,
|
||||
cv::Size(output_width, output_height));
|
||||
cc->Outputs().Tag(kMaskTag).Add(output_mask.release(), cc->InputTimestamp());
|
||||
#endif // !MEDIAPIPE_DISABLE_OPENCV
|
||||
|
||||
return absl::OkStatus();
|
||||
}
|
||||
|
||||
#if !MEDIAPIPE_DISABLE_OPENCV
|
||||
template <class T>
|
||||
absl::Status TensorsToSegmentationCalculator::ApplyActivation(
|
||||
cv::Mat& tensor_mat, cv::Mat* small_mask_mat) {
|
||||
@@ -410,6 +422,7 @@ absl::Status TensorsToSegmentationCalculator::ApplyActivation(
|
||||
|
||||
return absl::OkStatus();
|
||||
}
|
||||
#endif // !MEDIAPIPE_DISABLE_OPENCV
|
||||
|
||||
// Steps:
|
||||
// 1. receive tensor
|
||||
|
||||
@@ -334,6 +334,7 @@ cc_library(
|
||||
":image_frame_to_tensor_calculator_cc_proto",
|
||||
"//mediapipe/framework:calculator_framework",
|
||||
"//mediapipe/framework/formats:image_frame",
|
||||
"//mediapipe/framework/port:core_proto",
|
||||
"//mediapipe/framework/port:ret_check",
|
||||
"//mediapipe/framework/port:status",
|
||||
] + select({
|
||||
|
||||
@@ -17,6 +17,7 @@
|
||||
#include "mediapipe/calculators/tensorflow/image_frame_to_tensor_calculator.pb.h"
|
||||
#include "mediapipe/framework/calculator_framework.h"
|
||||
#include "mediapipe/framework/formats/image_frame.h"
|
||||
#include "mediapipe/framework/port/proto_ns.h"
|
||||
#include "mediapipe/framework/port/ret_check.h"
|
||||
#include "mediapipe/framework/port/status.h"
|
||||
#include "mediapipe/framework/port/status_macros.h"
|
||||
@@ -32,7 +33,10 @@ namespace {
|
||||
// Convert the ImageFrame into Tensor with floating point value type.
|
||||
// The value will be normalized based on mean and stddev.
|
||||
std::unique_ptr<tf::Tensor> ImageFrameToNormalizedTensor(
|
||||
const ImageFrame& image_frame, float mean, float stddev) {
|
||||
// const ImageFrame& image_frame, float mean, float stddev) {
|
||||
const ImageFrame& image_frame,
|
||||
const mediapipe::proto_ns::RepeatedField<float>& mean,
|
||||
const mediapipe::proto_ns::RepeatedField<float>& stddev) {
|
||||
const int cols = image_frame.Width();
|
||||
const int rows = image_frame.Height();
|
||||
const int channels = image_frame.NumberOfChannels();
|
||||
@@ -45,7 +49,20 @@ std::unique_ptr<tf::Tensor> ImageFrameToNormalizedTensor(
|
||||
for (int row = 0; row < rows; ++row) {
|
||||
for (int col = 0; col < cols; ++col) {
|
||||
for (int channel = 0; channel < channels; ++channel) {
|
||||
tensor_data(row, col, channel) = (pixel[channel] - mean) / stddev;
|
||||
float mean_value = 0;
|
||||
if (mean.size() > 1) {
|
||||
mean_value = mean[channel];
|
||||
} else if (!mean.empty()) {
|
||||
mean_value = mean[0];
|
||||
}
|
||||
float stddev_value = 1;
|
||||
if (stddev.size() > 1) {
|
||||
stddev_value = stddev[channel];
|
||||
} else if (!stddev.empty()) {
|
||||
stddev_value = stddev[0];
|
||||
}
|
||||
tensor_data(row, col, channel) =
|
||||
(pixel[channel] - mean_value) / stddev_value;
|
||||
}
|
||||
pixel += channels;
|
||||
}
|
||||
@@ -126,7 +143,18 @@ absl::Status ImageFrameToTensorCalculator::Process(CalculatorContext* cc) {
|
||||
const tf::DataType data_type = options_.data_type();
|
||||
RET_CHECK_EQ(data_type, tf::DT_FLOAT)
|
||||
<< "Unsupported data type " << data_type;
|
||||
RET_CHECK_GT(options_.stddev(), 0.0f);
|
||||
RET_CHECK_GT(options_.stddev().size(), 0) << "You must set a stddev.";
|
||||
RET_CHECK_GT(options_.stddev()[0], 0.0f) << "The stddev cannot be zero.";
|
||||
if (options_.stddev().size() > 1) {
|
||||
RET_CHECK_EQ(options_.stddev().size(), video_frame.NumberOfChannels())
|
||||
<< "If specifying multiple stddev normalization values, "
|
||||
<< "the number must match the number of image channels.";
|
||||
}
|
||||
if (options_.mean().size() > 1) {
|
||||
RET_CHECK_EQ(options_.mean().size(), video_frame.NumberOfChannels())
|
||||
<< "If specifying multiple mean normalization values, "
|
||||
<< "the number must match the number of image channels.";
|
||||
}
|
||||
tensor = ImageFrameToNormalizedTensor(video_frame, options_.mean(),
|
||||
options_.stddev());
|
||||
} else {
|
||||
|
||||
@@ -32,6 +32,6 @@ message ImageFrameToTensorCalculatorOptions {
|
||||
// If set, the output tensor T is equal to (F - mean * J) / stddev, where F
|
||||
// and J are the input image frame and the all-ones matrix of the same size,
|
||||
// respectively. Otherwise, T is equal to F.
|
||||
optional float mean = 2;
|
||||
optional float stddev = 3;
|
||||
repeated float mean = 2;
|
||||
repeated float stddev = 3;
|
||||
}
|
||||
|
||||
@@ -454,4 +454,32 @@ TEST_F(ImageFrameToTensorCalculatorTest, FixedRGBFrameWithMeanAndStddev) {
|
||||
EXPECT_EQ(actual[2], 127.0f / 128.0f); // (255 - 128) / 128
|
||||
}
|
||||
|
||||
TEST_F(ImageFrameToTensorCalculatorTest, FixedRGBFrameWithRepeatMeanAndStddev) {
|
||||
runner_ = ::absl::make_unique<CalculatorRunner>(
|
||||
"ImageFrameToTensorCalculator",
|
||||
"[mediapipe.ImageFrameToTensorCalculatorOptions.ext]"
|
||||
"{data_type:DT_FLOAT mean:128.0 mean:128.0 mean:128.0 "
|
||||
" stddev:128.0 stddev:128.0 stddev:128.0}",
|
||||
1, 1, 0);
|
||||
|
||||
// Create a single pixel image of fixed color #0080ff.
|
||||
auto image_frame = ::absl::make_unique<ImageFrame>(ImageFormat::SRGB, 1, 1);
|
||||
const uint8 color[] = {0, 128, 255};
|
||||
SetToColor<uint8>(color, image_frame.get());
|
||||
|
||||
runner_->MutableInputs()->Index(0).packets.push_back(
|
||||
Adopt(image_frame.release()).At(Timestamp(0)));
|
||||
MP_ASSERT_OK(runner_->Run());
|
||||
|
||||
const auto& tensor = runner_->Outputs().Index(0).packets[0].Get<tf::Tensor>();
|
||||
EXPECT_EQ(tensor.dtype(), tf::DT_FLOAT);
|
||||
ASSERT_EQ(tensor.dims(), 3);
|
||||
EXPECT_EQ(tensor.shape().dim_size(0), 1);
|
||||
EXPECT_EQ(tensor.shape().dim_size(1), 1);
|
||||
EXPECT_EQ(tensor.shape().dim_size(2), 3);
|
||||
const float* actual = tensor.flat<float>().data();
|
||||
EXPECT_EQ(actual[0], -1.0f); // ( 0 - 128) / 128
|
||||
EXPECT_EQ(actual[1], 0.0f); // (128 - 128) / 128
|
||||
EXPECT_EQ(actual[2], 127.0f / 128.0f); // (255 - 128) / 128
|
||||
}
|
||||
} // namespace mediapipe
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user