Compare commits

...
19 Commits
Author SHA1 Message Date
MediaPipe Teamandjqtang 3b6d3c4058 Project import generated by Copybara.
GitOrigin-RevId: 4419aaa472eeb91123d1f8576188166ee0e5ea69
2020-03-10 18:14:25 -07:00
MediaPipe Teamandjqtang 252a5713c7 Project import generated by Copybara.
GitOrigin-RevId: 6f964e58d874e47fb6207aa97d060a4cd6428527
2020-03-02 10:35:07 -08:00
MediaPipe TeamandHadon Nash de4fbc10e6 Project import generated by Copybara.
GitOrigin-RevId: 852dfb05d450167899c0dd5ef7c45622a12e865b
2020-02-10 14:13:25 -08:00
MediaPipe Teamandjqtang d144e564d8 Project import generated by Copybara.
GitOrigin-RevId: df2c4ee5ecd342bd88f332389348615b47a0244c
2020-01-17 17:05:43 -08:00
MediaPipe Teamandchris dd02df1dbe Project import generated by Copybara.
GitOrigin-RevId: b695dda274aa3ac3c7d054e150bd9eb5c1285b19
2020-01-17 15:49:22 -08:00
MediaPipe Teamandjqtang 66b377c825 Project import generated by Copybara.
GitOrigin-RevId: 1dd19723270084e701a90f974e35754b3fe20265
2020-01-13 14:42:47 -08:00
MediaPipe Teamandjqtang bf5185f122 Project import generated by Copybara.
GitOrigin-RevId: 72933af9ce469acd89cbf41898dcc06c65df7c8a
2020-01-11 11:16:13 -08:00
MediaPipe Teamandjqtang a2823541e6 Project import generated by Copybara.
GitOrigin-RevId: 1237f560f10007d74f620349f9fe27b492f5faf6
2020-01-10 15:51:18 -08:00
MediaPipe TeamandHadon Nash ae6be10afe Project import generated by Copybara.
GitOrigin-RevId: 0517756260533d374df93679965ca662d0ec6943
2020-01-10 13:13:24 -08:00
MediaPipe Teamandjqtang 38ee2603a7 Project import generated by Copybara.
GitOrigin-RevId: 87e46800807001e01d686fd7bcc2533714556920
2019-12-09 13:11:22 -08:00
MediaPipe Teamandjqtang 86b3283b2f Project import generated by Copybara.
GitOrigin-RevId: 831b7eb6038549a3a5047e7a113d6a11956e2de9
2019-12-06 16:17:14 -08:00
MediaPipe Teamandjqtang 7d470a1335 Project import generated by Copybara.
GitOrigin-RevId: 398d8577074c6e93041c01ed34bd6f27b2773c4f
2019-12-06 16:07:44 -08:00
MediaPipe Teamandmgyong d16cc3be5b Project import generated by Copybara.
GitOrigin-RevId: d91373b4d4d10abef49cab410caa6aadf0875049
2019-12-06 15:57:20 -08:00
MediaPipe Teamandjqtang 137867d088 Project import generated by Copybara.
GitOrigin-RevId: e3566e5029af25b0fc4b1071a49e49ae20aa5df6
2019-12-02 17:54:10 -08:00
MediaPipe Teamandmgyong 446d7cf6b6 Project import generated by Copybara.
GitOrigin-RevId: b02a6442fa6234cd2c15fa19f09accd8767adbee
2019-11-21 14:48:32 -08:00
MediaPipe Teamandmgyong 90f72bd851 Project import generated by Copybara.
GitOrigin-RevId: 5aa039c4a51ab7b4a1c58c17ad13af4c833e25e7
2019-11-21 14:35:46 -08:00
MediaPipe Teamandmgyong 4285aeddfc Project import generated by Copybara.
GitOrigin-RevId: 651ba7a75bb696877570a8a1b4244b34d59088f8
2019-11-21 14:24:17 -08:00
MediaPipe Teamandmgyong 37287925b0 Project import generated by Copybara.
GitOrigin-RevId: ba1d851bc868c2f8037a6fa96ee90e4b8ab9bd40
2019-11-21 14:10:52 -08:00
MediaPipe Teamandmgyong 48bcbb115f Project import generated by Copybara.
GitOrigin-RevId: 50714fe28298d7b707eff7304547d89d6ec34a54
2019-11-21 13:20:47 -08:00
630 changed files with 86388 additions and 2451 deletions
+4 -1
View File
@@ -12,13 +12,16 @@ build --copt='-Wno-comment'
build --copt='-Wno-return-type' build --copt='-Wno-return-type'
build --copt='-Wno-unused-local-typedefs' build --copt='-Wno-unused-local-typedefs'
build --copt='-Wno-ignored-attributes' build --copt='-Wno-ignored-attributes'
# Temporarily set the incompatiblity flag for Bazel 0.27.0 and above # Temporarily set the incompatibility flag for Bazel 0.27.0 and above
build --incompatible_disable_deprecated_attr_params=false build --incompatible_disable_deprecated_attr_params=false
build --incompatible_depset_is_not_iterable=false build --incompatible_depset_is_not_iterable=false
# Sets the default Apple platform to macOS. # Sets the default Apple platform to macOS.
build --apple_platform_type=macos build --apple_platform_type=macos
# Allow debugging with XCODE
build --apple_generate_dsym
# Android configs. # Android configs.
build:android --crosstool_top=//external:android/crosstool build:android --crosstool_top=//external:android/crosstool
build:android --host_crosstool_top=@bazel_tools//tools/cpp:toolchain build:android --host_crosstool_top=@bazel_tools//tools/cpp:toolchain
+5 -1
View File
@@ -30,10 +30,13 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
unzip \ unzip \
python \ python \
python-pip \ python-pip \
python3-pip \
libopencv-core-dev \ libopencv-core-dev \
libopencv-highgui-dev \ libopencv-highgui-dev \
libopencv-imgproc-dev \ libopencv-imgproc-dev \
libopencv-video-dev \ libopencv-video-dev \
libopencv-calib3d-dev \
libopencv-features2d-dev \
software-properties-common && \ software-properties-common && \
add-apt-repository -y ppa:openjdk-r/ppa && \ add-apt-repository -y ppa:openjdk-r/ppa && \
apt-get update && apt-get install -y openjdk-8-jdk && \ apt-get update && apt-get install -y openjdk-8-jdk && \
@@ -42,9 +45,10 @@ RUN apt-get update && apt-get install -y --no-install-recommends \
RUN pip install --upgrade setuptools RUN pip install --upgrade setuptools
RUN pip install future RUN pip install future
RUN pip3 install six
# Install bazel # Install bazel
ARG BAZEL_VERSION=0.26.1 ARG BAZEL_VERSION=1.1.0
RUN mkdir /bazel && \ RUN mkdir /bazel && \
wget --no-check-certificate -O /bazel/installer.sh "https://github.com/bazelbuild/bazel/releases/download/${BAZEL_VERSION}/b\ 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" && \ azel-${BAZEL_VERSION}-installer-linux-x86_64.sh" && \
+25 -10
View File
@@ -1,7 +1,7 @@
![MediaPipe](mediapipe/docs/images/mediapipe_small.png?raw=true "MediaPipe logo") ![MediaPipe](mediapipe/docs/images/mediapipe_small.png?raw=true "MediaPipe logo")
======================================================================= =======================================================================
[MediaPipe](http://mediapipe.dev) is a framework for building multimodal (eg. video, audio, any time series data) applied ML pipelines. With MediaPipe, a perception pipeline can be built as a graph of modular components, including, for instance, inference models (e.g., TensorFlow, TFLite) and media processing functions. [MediaPipe](http://mediapipe.dev) is a framework for building multimodal (eg. video, audio, any time series data), cross platform (i.e Android, iOS, web, edge devices) applied ML pipelines. With MediaPipe, a perception pipeline can be built as a graph of modular components, including, for instance, inference models (e.g., TensorFlow, TFLite) and media processing functions.
![Real-time Face Detection](mediapipe/docs/images/realtime_face_detection.gif) ![Real-time Face Detection](mediapipe/docs/images/realtime_face_detection.gif)
@@ -9,21 +9,28 @@
## ML Solutions in MediaPipe ## ML Solutions in MediaPipe
* [Hand Tracking](mediapipe/docs/hand_tracking_mobile_gpu.md) * [Face Detection](mediapipe/docs/face_detection_mobile_gpu.md) [[Web Demo]](https://viz.mediapipe.dev/runner/demos/face_detection/face_detection.html)
* [Face Detection](mediapipe/docs/face_detection_mobile_gpu.md) * [Multi-hand Tracking](mediapipe/docs/multi_hand_tracking_mobile_gpu.md)
* [Hair Segmentation](mediapipe/docs/hair_segmentation_mobile_gpu.md) * [Hand Tracking](mediapipe/docs/hand_tracking_mobile_gpu.md) [[Web Demo]](https://viz.mediapipe.dev/runner/demos/hand_tracking/hand_tracking.html)
* [Hair Segmentation](mediapipe/docs/hair_segmentation_mobile_gpu.md) [[Web Demo]](https://viz.mediapipe.dev/runner/demos/hair_segmentation/hair_segmentation.html)
* [Object Detection](mediapipe/docs/object_detection_mobile_gpu.md) * [Object Detection](mediapipe/docs/object_detection_mobile_gpu.md)
* [Object Detection and Tracking](mediapipe/docs/object_tracking_mobile_gpu.md)
* [Objectron: 3D Object Detection and Tracking](mediapipe/docs/objectron_mobile_gpu.md)
* [AutoFlip](mediapipe/docs/autoflip.md)
![hand_tracking](mediapipe/docs/images/mobile/hand_tracking_3d_android_gpu_small.gif)
![face_detection](mediapipe/docs/images/mobile/face_detection_android_gpu_small.gif) ![face_detection](mediapipe/docs/images/mobile/face_detection_android_gpu_small.gif)
![multi-hand_tracking](mediapipe/docs/images/mobile/multi_hand_tracking_android_gpu_small.gif)
![hand_tracking](mediapipe/docs/images/mobile/hand_tracking_3d_android_gpu_small.gif)
![hair_segmentation](mediapipe/docs/images/mobile/hair_segmentation_android_gpu_small.gif) ![hair_segmentation](mediapipe/docs/images/mobile/hair_segmentation_android_gpu_small.gif)
![object_detection](mediapipe/docs/images/mobile/object_detection_android_gpu_small.gif) ![object_tracking](mediapipe/docs/images/mobile/object_tracking_android_gpu_small.gif)
## Installation ## Installation
Follow these [instructions](mediapipe/docs/install.md). Follow these [instructions](mediapipe/docs/install.md).
## Getting started ## Getting started
See mobile and desktop [examples](mediapipe/docs/examples.md). See mobile, desktop and Google Coral [examples](mediapipe/docs/examples.md).
Check out some web demos [[Edge detection]](https://viz.mediapipe.dev/runner/demos/edge_detection/edge_detection.html) [[Face detection]](https://viz.mediapipe.dev/runner/demos/face_detection/face_detection.html) [[Hand Tracking]](https://viz.mediapipe.dev/runner/demos/hand_tracking/hand_tracking.html)
## Documentation ## Documentation
[MediaPipe Read-the-Docs](https://mediapipe.readthedocs.io/) or [docs.mediapipe.dev](https://docs.mediapipe.dev) [MediaPipe Read-the-Docs](https://mediapipe.readthedocs.io/) or [docs.mediapipe.dev](https://docs.mediapipe.dev)
@@ -33,14 +40,19 @@ Check out the [Examples page](https://mediapipe.readthedocs.io/en/latest/example
## Visualizing MediaPipe graphs ## Visualizing MediaPipe graphs
A web-based visualizer is hosted on [viz.mediapipe.dev](https://viz.mediapipe.dev/). Please also see instructions [here](mediapipe/docs/visualizer.md). A web-based visualizer is hosted on [viz.mediapipe.dev](https://viz.mediapipe.dev/). Please also see instructions [here](mediapipe/docs/visualizer.md).
## Community forum ## Videos
* [Discuss](https://groups.google.com/forum/#!forum/mediapipe) - General community discussion around MediaPipe * [YouTube Channel](https://www.youtube.com/channel/UCObqmpuSMx-usADtL_qdMAw)
## Publications ## Publications
* [MediaPipe Objectron: Real-time 3D Object Detection on Mobile Devices](https://mediapipe.page.link/objectron-aiblog)
* [AutoFlip: An Open Source Framework for Intelligent Video Reframing](https://mediapipe.page.link/autoflip)
* [Google Developer Blog: MediaPipe on the Web](https://mediapipe.page.link/webdevblog)
* [Google Developer Blog: Object Detection and Tracking using MediaPipe](https://mediapipe.page.link/objecttrackingblog)
* [On-Device, Real-Time Hand Tracking with MediaPipe](https://ai.googleblog.com/2019/08/on-device-real-time-hand-tracking-with.html) * [On-Device, Real-Time Hand Tracking with MediaPipe](https://ai.googleblog.com/2019/08/on-device-real-time-hand-tracking-with.html)
* [MediaPipe: A Framework for Building Perception Pipelines](https://arxiv.org/abs/1906.08172) * [MediaPipe: A Framework for Building Perception Pipelines](https://arxiv.org/abs/1906.08172)
## Events ## Events
* [AI Nextcon 2020, 12-16 Feb 2020, Seattle](http://aisea20.xnextcon.com/)
* [MediaPipe Madrid Meetup, 16 Dec 2019](https://www.meetup.com/Madrid-AI-Developers-Group/events/266329088/) * [MediaPipe Madrid Meetup, 16 Dec 2019](https://www.meetup.com/Madrid-AI-Developers-Group/events/266329088/)
* [MediaPipe London Meetup, Google 123 Building, 12 Dec 2019](https://www.meetup.com/London-AI-Tech-Talk/events/266329038) * [MediaPipe London Meetup, Google 123 Building, 12 Dec 2019](https://www.meetup.com/London-AI-Tech-Talk/events/266329038)
* [ML Conference, Berlin, 11 Dec 2019](https://mlconference.ai/machine-learning-advanced-development/mediapipe-building-real-time-cross-platform-mobile-web-edge-desktop-video-audio-ml-pipelines/) * [ML Conference, Berlin, 11 Dec 2019](https://mlconference.ai/machine-learning-advanced-development/mediapipe-building-real-time-cross-platform-mobile-web-edge-desktop-video-audio-ml-pipelines/)
@@ -50,8 +62,11 @@ A web-based visualizer is hosted on [viz.mediapipe.dev](https://viz.mediapipe.de
* [Google Industry Workshop at ICIP 2019](http://2019.ieeeicip.org/?action=page4&id=14#Google) [Presentation](https://docs.google.com/presentation/d/e/2PACX-1vRIBBbO_LO9v2YmvbHHEt1cwyqH6EjDxiILjuT0foXy1E7g6uyh4CesB2DkkEwlRDO9_lWfuKMZx98T/pub?start=false&loop=false&delayms=3000&slide=id.g556cc1a659_0_5) on Sept 24 in Taipei, Taiwan * [Google Industry Workshop at ICIP 2019](http://2019.ieeeicip.org/?action=page4&id=14#Google) [Presentation](https://docs.google.com/presentation/d/e/2PACX-1vRIBBbO_LO9v2YmvbHHEt1cwyqH6EjDxiILjuT0foXy1E7g6uyh4CesB2DkkEwlRDO9_lWfuKMZx98T/pub?start=false&loop=false&delayms=3000&slide=id.g556cc1a659_0_5) on Sept 24 in Taipei, Taiwan
* [Open sourced at CVPR 2019](https://sites.google.com/corp/view/perception-cv4arvr/mediapipe) on June 17~20 in Long Beach, CA * [Open sourced at CVPR 2019](https://sites.google.com/corp/view/perception-cv4arvr/mediapipe) on June 17~20 in Long Beach, CA
## Community forum
* [Discuss](https://groups.google.com/forum/#!forum/mediapipe) - General community discussion around MediaPipe
## Alpha Disclaimer ## Alpha Disclaimer
MediaPipe is currently in alpha for v0.6. We are still making breaking API changes and expect to get to stable API by v1.0. MediaPipe is currently in alpha for v0.7. We are still making breaking API changes and expect to get to stable API by v1.0.
## Contributing ## Contributing
We welcome contributions. Please follow these [guidelines](./CONTRIBUTING.md). We welcome contributions. Please follow these [guidelines](./CONTRIBUTING.md).
+59 -18
View File
@@ -10,19 +10,25 @@ http_archive(
sha256 = "2ef429f5d7ce7111263289644d233707dba35e39696377ebab8b0bc701f7818e", sha256 = "2ef429f5d7ce7111263289644d233707dba35e39696377ebab8b0bc701f7818e",
) )
load("@bazel_skylib//lib:versions.bzl", "versions") load("@bazel_skylib//lib:versions.bzl", "versions")
versions.check(minimum_bazel_version = "0.24.1") versions.check(minimum_bazel_version = "1.0.0",
maximum_bazel_version = "1.2.1")
# ABSL cpp library.
# ABSL cpp library lts_2020_02_25
http_archive( http_archive(
name = "com_google_absl", name = "com_google_absl",
# Head commit on 2019-04-12.
# TODO: Switch to the latest absl version when the problem gets
# fixed.
urls = [ urls = [
"https://github.com/abseil/abseil-cpp/archive/a02f62f456f2c4a7ecf2be3104fe0c6e16fbad9a.tar.gz", "https://github.com/abseil/abseil-cpp/archive/20200225.tar.gz",
], ],
sha256 = "d437920d1434c766d22e85773b899c77c672b8b4865d5dc2cd61a29fdff3cf03", # Remove after https://github.com/abseil/abseil-cpp/issues/326 is solved.
strip_prefix = "abseil-cpp-a02f62f456f2c4a7ecf2be3104fe0c6e16fbad9a", patches = [
"@//third_party:com_google_absl_f863b622fe13612433fdf43f76547d5edda0c93001.diff"
],
patch_args = [
"-p1",
],
strip_prefix = "abseil-cpp-20200225",
sha256 = "728a813291bdec2aa46eab8356ace9f75ac2ed9dfe2df5ab603c4e6c09f1c353"
) )
http_archive( http_archive(
@@ -72,6 +78,14 @@ http_archive(
], ],
) )
# easyexif
http_archive(
name = "easyexif",
url = "https://github.com/mayanklahiri/easyexif/archive/master.zip",
strip_prefix = "easyexif-master",
build_file = "@//third_party:easyexif.BUILD",
)
# libyuv # libyuv
http_archive( http_archive(
name = "libyuv", name = "libyuv",
@@ -103,15 +117,23 @@ http_archive(
], ],
) )
# 2019-11-12 # 2020-02-12
_TENSORFLOW_GIT_COMMIT = "a5f9bcd64453ff3d1f64cb4da4786db3d2da7f82" # The last commit before TensorFlow switched to Bazel 2.0
_TENSORFLOW_SHA256= "f2b6f2ab2ffe63e86eccd3ce4bea6b7197383d726638dfeeebcdc1e7de73f075" _TENSORFLOW_GIT_COMMIT = "77e9ffb9b2bfb1a4f7056e62d84039626923e328"
_TENSORFLOW_SHA256= "176ccd82f7dd17c5e117b50d353603b129c7a6ccbfebd522ca47cc2a40f33f13"
http_archive( http_archive(
name = "org_tensorflow", name = "org_tensorflow",
urls = [ urls = [
"https://mirror.bazel.build/github.com/tensorflow/tensorflow/archive/%s.tar.gz" % _TENSORFLOW_GIT_COMMIT, "https://mirror.bazel.build/github.com/tensorflow/tensorflow/archive/%s.tar.gz" % _TENSORFLOW_GIT_COMMIT,
"https://github.com/tensorflow/tensorflow/archive/%s.tar.gz" % _TENSORFLOW_GIT_COMMIT, "https://github.com/tensorflow/tensorflow/archive/%s.tar.gz" % _TENSORFLOW_GIT_COMMIT,
], ],
# A compatibility patch
patches = [
"@//third_party:org_tensorflow_528e22eae8bf3206189a066032c66e9e5c9b4a61.diff"
],
patch_args = [
"-p1",
],
strip_prefix = "tensorflow-%s" % _TENSORFLOW_GIT_COMMIT, strip_prefix = "tensorflow-%s" % _TENSORFLOW_GIT_COMMIT,
sha256 = _TENSORFLOW_SHA256, sha256 = _TENSORFLOW_SHA256,
) )
@@ -119,8 +141,22 @@ http_archive(
load("@org_tensorflow//tensorflow:workspace.bzl", "tf_workspace") load("@org_tensorflow//tensorflow:workspace.bzl", "tf_workspace")
tf_workspace(tf_repo_name = "org_tensorflow") tf_workspace(tf_repo_name = "org_tensorflow")
http_archive(
name = "ceres_solver",
url = "https://github.com/ceres-solver/ceres-solver/archive/1.14.0.zip",
patches = [
"@//third_party:ceres_solver_9bf9588988236279e1262f75d7f4d85711dfa172.diff"
],
patch_args = [
"-p1",
],
strip_prefix = "ceres-solver-1.14.0",
sha256 = "5ba6d0db4e784621fda44a50c58bb23b0892684692f0c623e2063f9c19f192f1"
)
# Please run # Please run
# $ sudo apt-get install libopencv-core-dev libopencv-highgui-dev \ # $ sudo apt-get install libopencv-core-dev libopencv-highgui-dev \
# libopencv-calib3d-dev libopencv-features2d-dev \
# libopencv-imgproc-dev libopencv-video-dev # libopencv-imgproc-dev libopencv-video-dev
new_local_repository( new_local_repository(
name = "linux_opencv", name = "linux_opencv",
@@ -149,11 +185,10 @@ new_local_repository(
http_archive( http_archive(
name = "android_opencv", name = "android_opencv",
sha256 = "056b849842e4fa8751d09edbb64530cfa7a63c84ccd232d0ace330e27ba55d0b",
build_file = "@//third_party:opencv_android.BUILD", build_file = "@//third_party:opencv_android.BUILD",
strip_prefix = "OpenCV-android-sdk", strip_prefix = "OpenCV-android-sdk",
type = "zip", type = "zip",
url = "https://github.com/opencv/opencv/releases/download/4.1.0/opencv-4.1.0-android-sdk.zip", url = "https://github.com/opencv/opencv/releases/download/3.4.3/opencv-3.4.3-android-sdk.zip",
) )
# After OpenCV 3.2.0, the pre-compiled opencv2.framework has google protobuf symbols, which will # After OpenCV 3.2.0, the pre-compiled opencv2.framework has google protobuf symbols, which will
@@ -184,13 +219,18 @@ maven_install(
artifacts = [ artifacts = [
"androidx.annotation:annotation:aar:1.1.0", "androidx.annotation:annotation:aar:1.1.0",
"androidx.appcompat:appcompat:aar:1.1.0-rc01", "androidx.appcompat:appcompat:aar:1.1.0-rc01",
"androidx.camera:camera-core:aar:1.0.0-alpha06",
"androidx.camera:camera-camera2:aar:1.0.0-alpha06",
"androidx.constraintlayout:constraintlayout:aar:1.1.3", "androidx.constraintlayout:constraintlayout:aar:1.1.3",
"androidx.core:core:aar:1.1.0-rc03", "androidx.core:core:aar:1.1.0-rc03",
"androidx.legacy:legacy-support-v4:aar:1.0.0", "androidx.legacy:legacy-support-v4:aar:1.0.0",
"androidx.recyclerview:recyclerview:aar:1.1.0-beta02", "androidx.recyclerview:recyclerview:aar:1.1.0-beta02",
"com.google.android.material:material:aar:1.0.0-rc01", "com.google.android.material:material:aar:1.0.0-rc01",
], ],
repositories = ["https://dl.google.com/dl/android/maven2"], repositories = [
"https://dl.google.com/dl/android/maven2",
"https://repo1.maven.org/maven2",
],
) )
maven_server( maven_server(
@@ -206,10 +246,10 @@ maven_jar(
) )
maven_jar( maven_jar(
name = "androidx_concurrent_futures", name = "androidx_concurrent_futures",
artifact = "androidx.concurrent:concurrent-futures:1.0.0-alpha03", artifact = "androidx.concurrent:concurrent-futures:1.0.0-alpha03",
sha1 = "b528df95c7e2fefa2210c0c742bf3e491c1818ae", sha1 = "b528df95c7e2fefa2210c0c742bf3e491c1818ae",
server = "google_server", server = "google_server",
) )
maven_jar( maven_jar(
@@ -284,3 +324,4 @@ http_archive(
strip_prefix = "google-toolbox-for-mac-2.2.1", strip_prefix = "google-toolbox-for-mac-2.2.1",
build_file = "@//third_party:google_toolbox_for_mac.BUILD", build_file = "@//third_party:google_toolbox_for_mac.BUILD",
) )
@@ -9,6 +9,7 @@
"mediapipe/examples/ios/facedetectiongpu/BUILD", "mediapipe/examples/ios/facedetectiongpu/BUILD",
"mediapipe/examples/ios/handdetectiongpu/BUILD", "mediapipe/examples/ios/handdetectiongpu/BUILD",
"mediapipe/examples/ios/handtrackinggpu/BUILD", "mediapipe/examples/ios/handtrackinggpu/BUILD",
"mediapipe/examples/ios/multihandtrackinggpu/BUILD",
"mediapipe/examples/ios/objectdetectioncpu/BUILD", "mediapipe/examples/ios/objectdetectioncpu/BUILD",
"mediapipe/examples/ios/objectdetectiongpu/BUILD" "mediapipe/examples/ios/objectdetectiongpu/BUILD"
], ],
@@ -18,6 +19,7 @@
"//mediapipe/examples/ios/facedetectiongpu:FaceDetectionGpuApp", "//mediapipe/examples/ios/facedetectiongpu:FaceDetectionGpuApp",
"//mediapipe/examples/ios/handdetectiongpu:HandDetectionGpuApp", "//mediapipe/examples/ios/handdetectiongpu:HandDetectionGpuApp",
"//mediapipe/examples/ios/handtrackinggpu:HandTrackingGpuApp", "//mediapipe/examples/ios/handtrackinggpu:HandTrackingGpuApp",
"//mediapipe/examples/ios/multihandtrackinggpu:MultiHandTrackingGpuApp",
"//mediapipe/examples/ios/objectdetectioncpu:ObjectDetectionCpuApp", "//mediapipe/examples/ios/objectdetectioncpu:ObjectDetectionCpuApp",
"//mediapipe/examples/ios/objectdetectiongpu:ObjectDetectionGpuApp", "//mediapipe/examples/ios/objectdetectiongpu:ObjectDetectionGpuApp",
"//mediapipe/objc:mediapipe_framework_ios" "//mediapipe/objc:mediapipe_framework_ios"
@@ -84,6 +86,8 @@
"mediapipe/examples/ios/handdetectiongpu/Base.lproj", "mediapipe/examples/ios/handdetectiongpu/Base.lproj",
"mediapipe/examples/ios/handtrackinggpu", "mediapipe/examples/ios/handtrackinggpu",
"mediapipe/examples/ios/handtrackinggpu/Base.lproj", "mediapipe/examples/ios/handtrackinggpu/Base.lproj",
"mediapipe/examples/ios/multihandtrackinggpu",
"mediapipe/examples/ios/multihandtrackinggpu/Base.lproj",
"mediapipe/examples/ios/objectdetectioncpu", "mediapipe/examples/ios/objectdetectioncpu",
"mediapipe/examples/ios/objectdetectioncpu/Base.lproj", "mediapipe/examples/ios/objectdetectioncpu/Base.lproj",
"mediapipe/examples/ios/objectdetectiongpu", "mediapipe/examples/ios/objectdetectiongpu",
@@ -16,6 +16,7 @@
"mediapipe/examples/ios/facedetectiongpu", "mediapipe/examples/ios/facedetectiongpu",
"mediapipe/examples/ios/handdetectiongpu", "mediapipe/examples/ios/handdetectiongpu",
"mediapipe/examples/ios/handtrackinggpu", "mediapipe/examples/ios/handtrackinggpu",
"mediapipe/examples/ios/multihandtrackinggpu",
"mediapipe/examples/ios/objectdetectioncpu", "mediapipe/examples/ios/objectdetectioncpu",
"mediapipe/examples/ios/objectdetectiongpu" "mediapipe/examples/ios/objectdetectiongpu"
], ],
+119 -5
View File
@@ -47,6 +47,13 @@ proto_library(
deps = ["//mediapipe/framework:calculator_proto"], deps = ["//mediapipe/framework:calculator_proto"],
) )
proto_library(
name = "packet_thinner_calculator_proto",
srcs = ["packet_thinner_calculator.proto"],
visibility = ["//visibility:public"],
deps = ["//mediapipe/framework:calculator_proto"],
)
proto_library( proto_library(
name = "split_vector_calculator_proto", name = "split_vector_calculator_proto",
srcs = ["split_vector_calculator.proto"], srcs = ["split_vector_calculator.proto"],
@@ -79,6 +86,15 @@ proto_library(
], ],
) )
proto_library(
name = "constant_side_packet_calculator_proto",
srcs = ["constant_side_packet_calculator.proto"],
visibility = ["//visibility:public"],
deps = [
"//mediapipe/framework:calculator_proto",
],
)
proto_library( proto_library(
name = "clip_vector_size_calculator_proto", name = "clip_vector_size_calculator_proto",
srcs = ["clip_vector_size_calculator.proto"], srcs = ["clip_vector_size_calculator.proto"],
@@ -102,6 +118,14 @@ mediapipe_cc_proto_library(
deps = [":packet_resampler_calculator_proto"], deps = [":packet_resampler_calculator_proto"],
) )
mediapipe_cc_proto_library(
name = "packet_thinner_calculator_cc_proto",
srcs = ["packet_thinner_calculator.proto"],
cc_deps = ["//mediapipe/framework:calculator_cc_proto"],
visibility = ["//visibility:public"],
deps = [":packet_thinner_calculator_proto"],
)
mediapipe_cc_proto_library( mediapipe_cc_proto_library(
name = "split_vector_calculator_cc_proto", name = "split_vector_calculator_cc_proto",
srcs = ["split_vector_calculator.proto"], srcs = ["split_vector_calculator.proto"],
@@ -158,6 +182,14 @@ mediapipe_cc_proto_library(
deps = [":gate_calculator_proto"], deps = [":gate_calculator_proto"],
) )
mediapipe_cc_proto_library(
name = "constant_side_packet_calculator_cc_proto",
srcs = ["constant_side_packet_calculator.proto"],
cc_deps = ["//mediapipe/framework:calculator_cc_proto"],
visibility = ["//visibility:public"],
deps = [":constant_side_packet_calculator_proto"],
)
cc_library( cc_library(
name = "add_header_calculator", name = "add_header_calculator",
srcs = ["add_header_calculator.cc"], srcs = ["add_header_calculator.cc"],
@@ -165,6 +197,7 @@ cc_library(
deps = [ deps = [
"//mediapipe/framework:calculator_framework", "//mediapipe/framework:calculator_framework",
"//mediapipe/framework/port:logging", "//mediapipe/framework/port:logging",
"//mediapipe/framework/port:status",
], ],
alwayslink = 1, alwayslink = 1,
) )
@@ -283,7 +316,6 @@ cc_test(
srcs = ["concatenate_vector_calculator_test.cc"], srcs = ["concatenate_vector_calculator_test.cc"],
deps = [ deps = [
":concatenate_vector_calculator", ":concatenate_vector_calculator",
"//mediapipe/calculators/core:packet_resampler_calculator_cc_proto",
"//mediapipe/framework:calculator_framework", "//mediapipe/framework:calculator_framework",
"//mediapipe/framework:calculator_runner", "//mediapipe/framework:calculator_runner",
"//mediapipe/framework:timestamp", "//mediapipe/framework:timestamp",
@@ -450,6 +482,37 @@ cc_test(
], ],
) )
cc_library(
name = "packet_thinner_calculator",
srcs = ["packet_thinner_calculator.cc"],
visibility = ["//visibility:public"],
deps = [
"//mediapipe/calculators/core:packet_thinner_calculator_cc_proto",
"//mediapipe/framework:calculator_context",
"//mediapipe/framework:calculator_framework",
"//mediapipe/framework/formats:video_stream_header",
"//mediapipe/framework/port:integral_types",
"//mediapipe/framework/port:logging",
"//mediapipe/framework/port:status",
],
alwayslink = 1,
)
cc_test(
name = "packet_thinner_calculator_test",
srcs = ["packet_thinner_calculator_test.cc"],
deps = [
":packet_thinner_calculator",
"//mediapipe/calculators/core:packet_thinner_calculator_cc_proto",
"//mediapipe/framework:calculator_framework",
"//mediapipe/framework:calculator_runner",
"//mediapipe/framework/formats:video_stream_header",
"//mediapipe/framework/port:gtest_main",
"//mediapipe/framework/port:integral_types",
"@com_google_absl//absl/strings",
],
)
cc_library( cc_library(
name = "pass_through_calculator", name = "pass_through_calculator",
srcs = ["pass_through_calculator.cc"], srcs = ["pass_through_calculator.cc"],
@@ -547,6 +610,22 @@ cc_library(
alwayslink = 1, alwayslink = 1,
) )
cc_test(
name = "side_packet_to_stream_calculator_test",
srcs = ["side_packet_to_stream_calculator_test.cc"],
deps = [
":side_packet_to_stream_calculator",
"//mediapipe/framework:calculator_framework",
"//mediapipe/framework/port:gtest_main",
"//mediapipe/framework/port:integral_types",
"//mediapipe/framework/port:parse_text_proto",
"//mediapipe/framework/port:status",
"//mediapipe/framework/tool:options_util",
"@com_google_absl//absl/memory",
"@com_google_absl//absl/strings",
],
)
cc_test( cc_test(
name = "immediate_mux_calculator_test", name = "immediate_mux_calculator_test",
srcs = ["immediate_mux_calculator_test.cc"], srcs = ["immediate_mux_calculator_test.cc"],
@@ -571,6 +650,7 @@ cc_test(
cc_library( cc_library(
name = "packet_resampler_calculator", name = "packet_resampler_calculator",
srcs = ["packet_resampler_calculator.cc"], srcs = ["packet_resampler_calculator.cc"],
hdrs = ["packet_resampler_calculator.h"],
visibility = [ visibility = [
"//visibility:public", "//visibility:public",
], ],
@@ -594,17 +674,17 @@ cc_library(
cc_test( cc_test(
name = "packet_resampler_calculator_test", name = "packet_resampler_calculator_test",
timeout = "short", timeout = "short",
srcs = ["packet_resampler_calculator_test.cc"], srcs = [
"packet_resampler_calculator_test.cc",
],
deps = [ deps = [
":packet_resampler_calculator", ":packet_resampler_calculator",
"//mediapipe/calculators/core:packet_resampler_calculator_cc_proto", "//mediapipe/calculators/core:packet_resampler_calculator_cc_proto",
"//mediapipe/framework:calculator_framework", "//mediapipe/framework:calculator_framework",
"//mediapipe/framework:calculator_runner", "//mediapipe/framework:calculator_runner",
"//mediapipe/framework:timestamp",
"//mediapipe/framework/formats:video_stream_header", "//mediapipe/framework/formats:video_stream_header",
"//mediapipe/framework/port:gtest_main", "//mediapipe/framework/port:gtest_main",
"//mediapipe/framework/port:parse_text_proto", "//mediapipe/framework/port:parse_text_proto",
"//mediapipe/framework/port:status",
"@com_google_absl//absl/strings", "@com_google_absl//absl/strings",
], ],
) )
@@ -691,12 +771,19 @@ cc_library(
":split_vector_calculator_cc_proto", ":split_vector_calculator_cc_proto",
"//mediapipe/framework:calculator_framework", "//mediapipe/framework:calculator_framework",
"//mediapipe/framework/formats:landmark_cc_proto", "//mediapipe/framework/formats:landmark_cc_proto",
"//mediapipe/framework/formats:rect_cc_proto",
"//mediapipe/framework/port:ret_check", "//mediapipe/framework/port:ret_check",
"//mediapipe/framework/port:status", "//mediapipe/framework/port:status",
"//mediapipe/util:resource_util", "//mediapipe/util:resource_util",
"@org_tensorflow//tensorflow/lite:framework", "@org_tensorflow//tensorflow/lite:framework",
"@org_tensorflow//tensorflow/lite/kernels:builtin_ops", "@org_tensorflow//tensorflow/lite/kernels:builtin_ops",
], ] + select({
"//mediapipe/gpu:disable_gpu": [],
"//mediapipe:ios": [],
"//conditions:default": [
"@org_tensorflow//tensorflow/lite/delegates/gpu/gl:gl_buffer",
],
}),
alwayslink = 1, alwayslink = 1,
) )
@@ -906,3 +993,30 @@ cc_test(
"@com_google_absl//absl/memory", "@com_google_absl//absl/memory",
], ],
) )
cc_library(
name = "constant_side_packet_calculator",
srcs = ["constant_side_packet_calculator.cc"],
visibility = ["//visibility:public"],
deps = [
":constant_side_packet_calculator_cc_proto",
"//mediapipe/framework:calculator_framework",
"//mediapipe/framework:collection_item_id",
"//mediapipe/framework/port:ret_check",
"//mediapipe/framework/port:status",
],
alwayslink = 1,
)
cc_test(
name = "constant_side_packet_calculator_test",
srcs = ["constant_side_packet_calculator_test.cc"],
deps = [
":constant_side_packet_calculator",
"//mediapipe/framework:calculator_framework",
"//mediapipe/framework/port:gtest_main",
"//mediapipe/framework/port:parse_text_proto",
"//mediapipe/framework/port:status",
"@com_google_absl//absl/strings",
],
)
@@ -13,11 +13,12 @@
// limitations under the License. // limitations under the License.
#include "mediapipe/framework/calculator_framework.h" #include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/port/canonical_errors.h"
#include "mediapipe/framework/port/logging.h" #include "mediapipe/framework/port/logging.h"
namespace mediapipe { namespace mediapipe {
// Attach the header from one stream to another stream. // Attach the header from a stream or side input to another stream.
// //
// The header stream (tag HEADER) must not have any packets in it. // The header stream (tag HEADER) must not have any packets in it.
// //
@@ -25,17 +26,53 @@ namespace mediapipe {
// calculator to not need a header or to accept a separate stream with // calculator to not need a header or to accept a separate stream with
// a header, that would be more future proof. // a header, that would be more future proof.
// //
// Example usage 1:
// node {
// calculator: "AddHeaderCalculator"
// input_stream: "DATA:audio"
// input_stream: "HEADER:audio_header"
// output_stream: "audio_with_header"
// }
//
// Example usage 2:
// node {
// calculator: "AddHeaderCalculator"
// input_stream: "DATA:audio"
// input_side_packet: "HEADER:audio_header"
// output_stream: "audio_with_header"
// }
//
class AddHeaderCalculator : public CalculatorBase { class AddHeaderCalculator : public CalculatorBase {
public: public:
static ::mediapipe::Status GetContract(CalculatorContract* cc) { static ::mediapipe::Status GetContract(CalculatorContract* cc) {
cc->Inputs().Tag("HEADER").SetNone(); bool has_side_input = false;
bool has_header_stream = false;
if (cc->InputSidePackets().HasTag("HEADER")) {
cc->InputSidePackets().Tag("HEADER").SetAny();
has_side_input = true;
}
if (cc->Inputs().HasTag("HEADER")) {
cc->Inputs().Tag("HEADER").SetNone();
has_header_stream = true;
}
if (has_side_input == has_header_stream) {
return mediapipe::InvalidArgumentError(
"Header must be provided via exactly one of side input and input "
"stream");
}
cc->Inputs().Tag("DATA").SetAny(); cc->Inputs().Tag("DATA").SetAny();
cc->Outputs().Index(0).SetSameAs(&cc->Inputs().Tag("DATA")); cc->Outputs().Index(0).SetSameAs(&cc->Inputs().Tag("DATA"));
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
::mediapipe::Status Open(CalculatorContext* cc) override { ::mediapipe::Status Open(CalculatorContext* cc) override {
const Packet& header = cc->Inputs().Tag("HEADER").Header(); Packet header;
if (cc->InputSidePackets().HasTag("HEADER")) {
header = cc->InputSidePackets().Tag("HEADER");
}
if (cc->Inputs().HasTag("HEADER")) {
header = cc->Inputs().Tag("HEADER").Header();
}
if (!header.IsEmpty()) { if (!header.IsEmpty()) {
cc->Outputs().Index(0).SetHeader(header); cc->Outputs().Index(0).SetHeader(header);
} }
@@ -14,8 +14,10 @@
#include "mediapipe/framework/calculator_framework.h" #include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/calculator_runner.h" #include "mediapipe/framework/calculator_runner.h"
#include "mediapipe/framework/port/canonical_errors.h"
#include "mediapipe/framework/port/gmock.h" #include "mediapipe/framework/port/gmock.h"
#include "mediapipe/framework/port/gtest.h" #include "mediapipe/framework/port/gtest.h"
#include "mediapipe/framework/port/status.h"
#include "mediapipe/framework/port/status_matchers.h" #include "mediapipe/framework/port/status_matchers.h"
#include "mediapipe/framework/timestamp.h" #include "mediapipe/framework/timestamp.h"
#include "mediapipe/framework/tool/validate_type.h" #include "mediapipe/framework/tool/validate_type.h"
@@ -24,7 +26,7 @@ namespace mediapipe {
class AddHeaderCalculatorTest : public ::testing::Test {}; class AddHeaderCalculatorTest : public ::testing::Test {};
TEST_F(AddHeaderCalculatorTest, Works) { TEST_F(AddHeaderCalculatorTest, HeaderStream) {
CalculatorGraphConfig::Node node; CalculatorGraphConfig::Node node;
node.set_calculator("AddHeaderCalculator"); node.set_calculator("AddHeaderCalculator");
node.add_input_stream("HEADER:header_stream"); node.add_input_stream("HEADER:header_stream");
@@ -96,4 +98,62 @@ TEST_F(AddHeaderCalculatorTest, NoPacketsOnHeaderStream) {
ASSERT_FALSE(runner.Run().ok()); ASSERT_FALSE(runner.Run().ok());
} }
TEST_F(AddHeaderCalculatorTest, InputSidePacket) {
CalculatorGraphConfig::Node node;
node.set_calculator("AddHeaderCalculator");
node.add_input_stream("DATA:data_stream");
node.add_output_stream("merged_stream");
node.add_input_side_packet("HEADER:header");
CalculatorRunner runner(node);
// Set header and add 5 packets.
runner.MutableSidePackets()->Tag("HEADER") =
Adopt(new std::string("my_header"));
for (int i = 0; i < 5; ++i) {
Packet packet = Adopt(new int(i)).At(Timestamp(i * 1000));
runner.MutableInputs()->Tag("DATA").packets.push_back(packet);
}
// Run calculator.
MP_ASSERT_OK(runner.Run());
ASSERT_EQ(1, runner.Outputs().NumEntries());
// Test output.
EXPECT_EQ(std::string("my_header"),
runner.Outputs().Index(0).header.Get<std::string>());
const std::vector<Packet>& output_packets = runner.Outputs().Index(0).packets;
ASSERT_EQ(5, output_packets.size());
for (int i = 0; i < 5; ++i) {
const int val = output_packets[i].Get<int>();
EXPECT_EQ(i, val);
EXPECT_EQ(Timestamp(i * 1000), output_packets[i].Timestamp());
}
}
TEST_F(AddHeaderCalculatorTest, UsingBothSideInputAndStream) {
CalculatorGraphConfig::Node node;
node.set_calculator("AddHeaderCalculator");
node.add_input_stream("HEADER:header_stream");
node.add_input_stream("DATA:data_stream");
node.add_output_stream("merged_stream");
node.add_input_side_packet("HEADER:header");
CalculatorRunner runner(node);
// Set both headers and add 5 packets.
runner.MutableSidePackets()->Tag("HEADER") =
Adopt(new std::string("my_header"));
runner.MutableSidePackets()->Tag("HEADER") =
Adopt(new std::string("my_header"));
for (int i = 0; i < 5; ++i) {
Packet packet = Adopt(new int(i)).At(Timestamp(i * 1000));
runner.MutableInputs()->Tag("DATA").packets.push_back(packet);
}
// Run should fail because header can only be provided one way.
EXPECT_EQ(runner.Run().code(), ::mediapipe::InvalidArgumentError("").code());
}
} // namespace mediapipe } // namespace mediapipe
@@ -21,16 +21,10 @@
namespace mediapipe { namespace mediapipe {
// A calculator to process std::vector<NormalizedLandmark>. // A calculator to process std::vector<NormalizedLandmarkList>.
typedef BeginLoopCalculator<std::vector<::mediapipe::NormalizedLandmark>> typedef BeginLoopCalculator<std::vector<::mediapipe::NormalizedLandmarkList>>
BeginLoopNormalizedLandmarkCalculator; BeginLoopNormalizedLandmarkListVectorCalculator;
REGISTER_CALCULATOR(BeginLoopNormalizedLandmarkCalculator); REGISTER_CALCULATOR(BeginLoopNormalizedLandmarkListVectorCalculator);
// A calculator to process std::vector<std::vector<NormalizedLandmark>>.
typedef BeginLoopCalculator<
std::vector<std::vector<::mediapipe::NormalizedLandmark>>>
BeginLoopNormalizedLandmarksVectorCalculator;
REGISTER_CALCULATOR(BeginLoopNormalizedLandmarksVectorCalculator);
// A calculator to process std::vector<NormalizedRect>. // A calculator to process std::vector<NormalizedRect>.
typedef BeginLoopCalculator<std::vector<::mediapipe::NormalizedRect>> typedef BeginLoopCalculator<std::vector<::mediapipe::NormalizedRect>>
@@ -38,6 +38,8 @@ namespace mediapipe {
// } // }
// } // }
// } // }
// Optionally, you can pass in a side packet that will override `max_vec_size`
// that is specified in the options.
template <typename T> template <typename T>
class ClipVectorSizeCalculator : public CalculatorBase { class ClipVectorSizeCalculator : public CalculatorBase {
public: public:
@@ -53,6 +55,10 @@ class ClipVectorSizeCalculator : public CalculatorBase {
cc->Inputs().Index(0).Set<std::vector<T>>(); cc->Inputs().Index(0).Set<std::vector<T>>();
cc->Outputs().Index(0).Set<std::vector<T>>(); cc->Outputs().Index(0).Set<std::vector<T>>();
// Optional input side packet that determines `max_vec_size`.
if (cc->InputSidePackets().NumEntries() > 0) {
cc->InputSidePackets().Index(0).Set<int>();
}
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -61,6 +67,11 @@ class ClipVectorSizeCalculator : public CalculatorBase {
cc->SetOffset(TimestampDiff(0)); cc->SetOffset(TimestampDiff(0));
max_vec_size_ = cc->Options<::mediapipe::ClipVectorSizeCalculatorOptions>() max_vec_size_ = cc->Options<::mediapipe::ClipVectorSizeCalculatorOptions>()
.max_vec_size(); .max_vec_size();
// Override `max_vec_size` if passed as side packet.
if (cc->InputSidePackets().NumEntries() > 0 &&
!cc->InputSidePackets().Index(0).IsEmpty()) {
max_vec_size_ = cc->InputSidePackets().Index(0).Get<int>();
}
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -176,4 +176,31 @@ TEST(TestClipUniqueIntPtrVectorSizeCalculatorTest, ConsumeOneTimestamp) {
} }
} }
TEST(TestClipIntVectorSizeCalculatorTest, SidePacket) {
CalculatorGraphConfig::Node node_config =
ParseTextProtoOrDie<CalculatorGraphConfig::Node>(R"(
calculator: "TestClipIntVectorSizeCalculator"
input_stream: "input_vector"
input_side_packet: "max_vec_size"
output_stream: "output_vector"
options {
[mediapipe.ClipVectorSizeCalculatorOptions.ext] { max_vec_size: 1 }
}
)");
CalculatorRunner runner(node_config);
// This should override the default of 1 set in the options.
runner.MutableSidePackets()->Index(0) = Adopt(new int(2));
std::vector<int> input = {0, 1, 2, 3};
AddInputVector(input, /*timestamp=*/1, &runner);
MP_ASSERT_OK(runner.Run());
const std::vector<Packet>& outputs = runner.Outputs().Index(0).packets;
EXPECT_EQ(1, outputs.size());
EXPECT_EQ(Timestamp(1), outputs[0].Timestamp());
const std::vector<int>& output = outputs[0].Get<std::vector<int>>();
EXPECT_EQ(2, output.size());
std::vector<int> expected_vector = {0, 1};
EXPECT_EQ(expected_vector, output);
}
} // namespace mediapipe } // namespace mediapipe
@@ -19,7 +19,7 @@
#include "mediapipe/framework/formats/landmark.pb.h" #include "mediapipe/framework/formats/landmark.pb.h"
#include "tensorflow/lite/interpreter.h" #include "tensorflow/lite/interpreter.h"
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__APPLE__) #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
#include "tensorflow/lite/delegates/gpu/gl/gl_buffer.h" #include "tensorflow/lite/delegates/gpu/gl/gl_buffer.h"
#endif // !MEDIAPIPE_DISABLE_GPU #endif // !MEDIAPIPE_DISABLE_GPU
@@ -35,6 +35,16 @@ namespace mediapipe {
typedef ConcatenateVectorCalculator<float> ConcatenateFloatVectorCalculator; typedef ConcatenateVectorCalculator<float> ConcatenateFloatVectorCalculator;
REGISTER_CALCULATOR(ConcatenateFloatVectorCalculator); REGISTER_CALCULATOR(ConcatenateFloatVectorCalculator);
// Example config:
// node {
// calculator: "ConcatenateInt32VectorCalculator"
// input_stream: "int32_vector_1"
// input_stream: "int32_vector_2"
// output_stream: "concatenated_int32_vector"
// }
typedef ConcatenateVectorCalculator<int32> ConcatenateInt32VectorCalculator;
REGISTER_CALCULATOR(ConcatenateInt32VectorCalculator);
// Example config: // Example config:
// node { // node {
// calculator: "ConcatenateTfLiteTensorVectorCalculator" // calculator: "ConcatenateTfLiteTensorVectorCalculator"
@@ -50,7 +60,7 @@ typedef ConcatenateVectorCalculator<::mediapipe::NormalizedLandmark>
ConcatenateLandmarkVectorCalculator; ConcatenateLandmarkVectorCalculator;
REGISTER_CALCULATOR(ConcatenateLandmarkVectorCalculator); REGISTER_CALCULATOR(ConcatenateLandmarkVectorCalculator);
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__APPLE__) #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
typedef ConcatenateVectorCalculator<::tflite::gpu::gl::GlBuffer> typedef ConcatenateVectorCalculator<::tflite::gpu::gl::GlBuffer>
ConcatenateGlBufferVectorCalculator; ConcatenateGlBufferVectorCalculator;
REGISTER_CALCULATOR(ConcatenateGlBufferVectorCalculator); REGISTER_CALCULATOR(ConcatenateGlBufferVectorCalculator);
@@ -0,0 +1,116 @@
// Copyright 2020 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 <string>
#include "mediapipe/calculators/core/constant_side_packet_calculator.pb.h"
#include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/collection_item_id.h"
#include "mediapipe/framework/port/canonical_errors.h"
#include "mediapipe/framework/port/ret_check.h"
#include "mediapipe/framework/port/status.h"
namespace mediapipe {
// Generates an output side packet or multiple output side packets according to
// the specified options.
//
// Example configs:
// node {
// calculator: "ConstantSidePacketCalculator"
// output_side_packet: "PACKET:packet"
// options: {
// [mediapipe.ConstantSidePacketCalculatorOptions.ext]: {
// packet { int_value: 2 }
// }
// }
// }
//
// node {
// calculator: "ConstantSidePacketCalculator"
// output_side_packet: "PACKET:0:int_packet"
// output_side_packet: "PACKET:1:bool_packet"
// options: {
// [mediapipe.ConstantSidePacketCalculatorOptions.ext]: {
// packet { int_value: 2 }
// packet { bool_value: true }
// }
// }
// }
class ConstantSidePacketCalculator : public CalculatorBase {
public:
static ::mediapipe::Status GetContract(CalculatorContract* cc) {
const auto& options = cc->Options().GetExtension(
::mediapipe::ConstantSidePacketCalculatorOptions::ext);
RET_CHECK_EQ(cc->OutputSidePackets().NumEntries(kPacketTag),
options.packet_size())
<< "Number of output side packets has to be same as number of packets "
"configured in options.";
int index = 0;
for (CollectionItemId id = cc->OutputSidePackets().BeginId(kPacketTag);
id != cc->OutputSidePackets().EndId(kPacketTag); ++id, ++index) {
const auto& packet_options = options.packet(index);
auto& packet = cc->OutputSidePackets().Get(id);
if (packet_options.has_int_value()) {
packet.Set<int>();
} else if (packet_options.has_float_value()) {
packet.Set<float>();
} else if (packet_options.has_bool_value()) {
packet.Set<bool>();
} else if (packet_options.has_string_value()) {
packet.Set<std::string>();
} else {
return ::mediapipe::InvalidArgumentError(
"None of supported values were specified in options.");
}
}
return ::mediapipe::OkStatus();
}
::mediapipe::Status Open(CalculatorContext* cc) override {
const auto& options = cc->Options().GetExtension(
::mediapipe::ConstantSidePacketCalculatorOptions::ext);
int index = 0;
for (CollectionItemId id = cc->OutputSidePackets().BeginId(kPacketTag);
id != cc->OutputSidePackets().EndId(kPacketTag); ++id, ++index) {
auto& packet = cc->OutputSidePackets().Get(id);
const auto& packet_options = options.packet(index);
if (packet_options.has_int_value()) {
packet.Set(MakePacket<int>(packet_options.int_value()));
} else if (packet_options.has_float_value()) {
packet.Set(MakePacket<float>(packet_options.float_value()));
} else if (packet_options.has_bool_value()) {
packet.Set(MakePacket<bool>(packet_options.bool_value()));
} else if (packet_options.has_string_value()) {
packet.Set(MakePacket<std::string>(packet_options.string_value()));
} else {
return ::mediapipe::InvalidArgumentError(
"None of supported values were specified in options.");
}
}
return ::mediapipe::OkStatus();
}
::mediapipe::Status Process(CalculatorContext* cc) override {
return ::mediapipe::OkStatus();
}
private:
static constexpr const char* kPacketTag = "PACKET";
};
REGISTER_CALCULATOR(ConstantSidePacketCalculator);
} // namespace mediapipe
@@ -0,0 +1,36 @@
// Copyright 2020 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 ConstantSidePacketCalculatorOptions {
extend CalculatorOptions {
optional ConstantSidePacketCalculatorOptions ext = 291214597;
}
message ConstantSidePacket {
oneof value {
int32 int_value = 1;
float float_value = 2;
bool bool_value = 3;
string string_value = 4;
}
}
repeated ConstantSidePacket packet = 1;
}
@@ -0,0 +1,196 @@
// Copyright 2020 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 <string>
#include "absl/strings/string_view.h"
#include "absl/strings/substitute.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"
namespace mediapipe {
template <typename T>
void DoTestSingleSidePacket(absl::string_view packet_spec,
const T& expected_value) {
static constexpr absl::string_view graph_config_template = R"(
node {
calculator: "ConstantSidePacketCalculator"
output_side_packet: "PACKET:packet"
options: {
[mediapipe.ConstantSidePacketCalculatorOptions.ext]: {
packet $0
}
}
}
)";
CalculatorGraphConfig graph_config =
::mediapipe::ParseTextProtoOrDie<CalculatorGraphConfig>(
absl::Substitute(graph_config_template, packet_spec));
CalculatorGraph graph;
MP_ASSERT_OK(graph.Initialize(graph_config));
MP_ASSERT_OK(graph.StartRun({}));
MP_ASSERT_OK(graph.WaitUntilIdle());
MP_ASSERT_OK(graph.GetOutputSidePacket("packet"));
auto actual_value =
graph.GetOutputSidePacket("packet").ValueOrDie().template Get<T>();
EXPECT_EQ(actual_value, expected_value);
}
TEST(ConstantSidePacketCalculatorTest, EveryPossibleType) {
DoTestSingleSidePacket("{ int_value: 2 }", 2);
DoTestSingleSidePacket("{ float_value: 6.5f }", 6.5f);
DoTestSingleSidePacket("{ bool_value: true }", true);
DoTestSingleSidePacket<std::string>(R"({ string_value: "str" })", "str");
}
TEST(ConstantSidePacketCalculatorTest, MultiplePackets) {
CalculatorGraphConfig graph_config =
::mediapipe::ParseTextProtoOrDie<CalculatorGraphConfig>(R"(
node {
calculator: "ConstantSidePacketCalculator"
output_side_packet: "PACKET:0:int_packet"
output_side_packet: "PACKET:1:float_packet"
output_side_packet: "PACKET:2:bool_packet"
output_side_packet: "PACKET:3:string_packet"
output_side_packet: "PACKET:4:another_string_packet"
output_side_packet: "PACKET:5:another_int_packet"
options: {
[mediapipe.ConstantSidePacketCalculatorOptions.ext]: {
packet { int_value: 256 }
packet { float_value: 0.5f }
packet { bool_value: false }
packet { string_value: "string" }
packet { string_value: "another string" }
packet { int_value: 128 }
}
}
}
)");
CalculatorGraph graph;
MP_ASSERT_OK(graph.Initialize(graph_config));
MP_ASSERT_OK(graph.StartRun({}));
MP_ASSERT_OK(graph.WaitUntilIdle());
MP_ASSERT_OK(graph.GetOutputSidePacket("int_packet"));
EXPECT_EQ(graph.GetOutputSidePacket("int_packet").ValueOrDie().Get<int>(),
256);
MP_ASSERT_OK(graph.GetOutputSidePacket("float_packet"));
EXPECT_EQ(graph.GetOutputSidePacket("float_packet").ValueOrDie().Get<float>(),
0.5f);
MP_ASSERT_OK(graph.GetOutputSidePacket("bool_packet"));
EXPECT_FALSE(
graph.GetOutputSidePacket("bool_packet").ValueOrDie().Get<bool>());
MP_ASSERT_OK(graph.GetOutputSidePacket("string_packet"));
EXPECT_EQ(graph.GetOutputSidePacket("string_packet")
.ValueOrDie()
.Get<std::string>(),
"string");
MP_ASSERT_OK(graph.GetOutputSidePacket("another_string_packet"));
EXPECT_EQ(graph.GetOutputSidePacket("another_string_packet")
.ValueOrDie()
.Get<std::string>(),
"another string");
MP_ASSERT_OK(graph.GetOutputSidePacket("another_int_packet"));
EXPECT_EQ(
graph.GetOutputSidePacket("another_int_packet").ValueOrDie().Get<int>(),
128);
}
TEST(ConstantSidePacketCalculatorTest, ProcessingPacketsWithCorrectTagOnly) {
CalculatorGraphConfig graph_config =
::mediapipe::ParseTextProtoOrDie<CalculatorGraphConfig>(R"(
node {
calculator: "ConstantSidePacketCalculator"
output_side_packet: "PACKET:0:int_packet"
output_side_packet: "no_tag0"
output_side_packet: "PACKET:1:float_packet"
output_side_packet: "INCORRECT_TAG:0:name1"
output_side_packet: "PACKET:2:bool_packet"
output_side_packet: "PACKET:3:string_packet"
output_side_packet: "no_tag2"
output_side_packet: "INCORRECT_TAG:1:name2"
options: {
[mediapipe.ConstantSidePacketCalculatorOptions.ext]: {
packet { int_value: 256 }
packet { float_value: 0.5f }
packet { bool_value: false }
packet { string_value: "string" }
}
}
}
)");
CalculatorGraph graph;
MP_ASSERT_OK(graph.Initialize(graph_config));
MP_ASSERT_OK(graph.StartRun({}));
MP_ASSERT_OK(graph.WaitUntilIdle());
MP_ASSERT_OK(graph.GetOutputSidePacket("int_packet"));
EXPECT_EQ(graph.GetOutputSidePacket("int_packet").ValueOrDie().Get<int>(),
256);
MP_ASSERT_OK(graph.GetOutputSidePacket("float_packet"));
EXPECT_EQ(graph.GetOutputSidePacket("float_packet").ValueOrDie().Get<float>(),
0.5f);
MP_ASSERT_OK(graph.GetOutputSidePacket("bool_packet"));
EXPECT_FALSE(
graph.GetOutputSidePacket("bool_packet").ValueOrDie().Get<bool>());
MP_ASSERT_OK(graph.GetOutputSidePacket("string_packet"));
EXPECT_EQ(graph.GetOutputSidePacket("string_packet")
.ValueOrDie()
.Get<std::string>(),
"string");
}
TEST(ConstantSidePacketCalculatorTest, IncorrectConfig_MoreOptionsThanPackets) {
CalculatorGraphConfig graph_config =
::mediapipe::ParseTextProtoOrDie<CalculatorGraphConfig>(R"(
node {
calculator: "ConstantSidePacketCalculator"
output_side_packet: "PACKET:int_packet"
options: {
[mediapipe.ConstantSidePacketCalculatorOptions.ext]: {
packet { int_value: 256 }
packet { float_value: 0.5f }
}
}
}
)");
CalculatorGraph graph;
EXPECT_FALSE(graph.Initialize(graph_config).ok());
}
TEST(ConstantSidePacketCalculatorTest, IncorrectConfig_MorePacketsThanOptions) {
CalculatorGraphConfig graph_config =
::mediapipe::ParseTextProtoOrDie<CalculatorGraphConfig>(R"(
node {
calculator: "ConstantSidePacketCalculator"
output_side_packet: "PACKET:0:int_packet"
output_side_packet: "PACKET:1:float_packet"
options: {
[mediapipe.ConstantSidePacketCalculatorOptions.ext]: {
packet { int_value: 256 }
}
}
}
)");
CalculatorGraph graph;
EXPECT_FALSE(graph.Initialize(graph_config).ok());
}
} // namespace mediapipe
@@ -26,14 +26,9 @@ typedef EndLoopCalculator<std::vector<::mediapipe::NormalizedRect>>
EndLoopNormalizedRectCalculator; EndLoopNormalizedRectCalculator;
REGISTER_CALCULATOR(EndLoopNormalizedRectCalculator); REGISTER_CALCULATOR(EndLoopNormalizedRectCalculator);
typedef EndLoopCalculator<std::vector<::mediapipe::NormalizedLandmark>> typedef EndLoopCalculator<std::vector<::mediapipe::NormalizedLandmarkList>>
EndLoopNormalizedLandmarkCalculator; EndLoopNormalizedLandmarkListVectorCalculator;
REGISTER_CALCULATOR(EndLoopNormalizedLandmarkCalculator); REGISTER_CALCULATOR(EndLoopNormalizedLandmarkListVectorCalculator);
typedef EndLoopCalculator<
std::vector<std::vector<::mediapipe::NormalizedLandmark>>>
EndLoopNormalizedLandmarksVectorCalculator;
REGISTER_CALCULATOR(EndLoopNormalizedLandmarksVectorCalculator);
typedef EndLoopCalculator<std::vector<bool>> EndLoopBooleanCalculator; typedef EndLoopCalculator<std::vector<bool>> EndLoopBooleanCalculator;
REGISTER_CALCULATOR(EndLoopBooleanCalculator); REGISTER_CALCULATOR(EndLoopBooleanCalculator);
@@ -12,25 +12,17 @@
// See the License for the specific language governing permissions and // See the License for the specific language governing permissions and
// limitations under the License. // limitations under the License.
#include <cstdlib> #include "mediapipe/calculators/core/packet_resampler_calculator.h"
#include <memory>
#include <string>
#include "absl/strings/str_cat.h" #include <memory>
#include "mediapipe/calculators/core/packet_resampler_calculator.pb.h"
#include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/collection_item_id.h"
#include "mediapipe/framework/deps/mathutil.h"
#include "mediapipe/framework/deps/random_base.h"
#include "mediapipe/framework/formats/video_stream_header.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/status.h"
#include "mediapipe/framework/port/status_macros.h"
#include "mediapipe/framework/tool/options_util.h"
namespace { namespace {
// Reflect an integer against the lower and upper bound of an interval.
int64 ReflectBetween(int64 ts, int64 ts_min, int64 ts_max) {
if (ts < ts_min) return 2 * ts_min - ts - 1;
if (ts >= ts_max) return 2 * ts_max - ts - 1;
return ts;
}
// Creates a secure random number generator for use in ProcessWithJitter. // Creates a secure random number generator for use in ProcessWithJitter.
// If no secure random number generator can be constructed, the jitter // If no secure random number generator can be constructed, the jitter
@@ -45,120 +37,7 @@ std::unique_ptr<RandomBase> CreateSecureRandom(const std::string& seed) {
namespace mediapipe { namespace mediapipe {
// This calculator is used to normalize the frequency of the packets
// out of a stream. Given a desired frame rate, packets are going to be
// removed or added to achieve it.
//
// The jitter feature is disabled by default. To enable it, you need to
// implement CreateSecureRandom(const std::string&).
//
// The data stream may be either specified as the only stream (by index)
// or as the stream with tag "DATA".
//
// The input and output streams may be accompanied by a VIDEO_HEADER
// stream. This stream includes a VideoHeader at Timestamp::PreStream().
// The input VideoHeader on the VIDEO_HEADER stream will always be updated
// with the resampler frame rate no matter what the options value for
// output_header is before being output on the output VIDEO_HEADER stream.
// If the input VideoHeader is not available, then only the frame rate
// value will be set in the output.
//
// Related:
// packet_downsampler_calculator.cc: skips packets regardless of timestamps.
class PacketResamplerCalculator : public CalculatorBase {
public:
static ::mediapipe::Status GetContract(CalculatorContract* cc);
::mediapipe::Status Open(CalculatorContext* cc) override;
::mediapipe::Status Close(CalculatorContext* cc) override;
::mediapipe::Status Process(CalculatorContext* cc) override;
private:
// Calculates the first sampled timestamp that incorporates a jittering
// offset.
void InitializeNextOutputTimestampWithJitter();
// Calculates the next sampled timestamp that incorporates a jittering offset.
void UpdateNextOutputTimestampWithJitter();
// Logic for Process() when jitter_ != 0.0.
::mediapipe::Status ProcessWithJitter(CalculatorContext* cc);
// Logic for Process() when jitter_ == 0.0.
::mediapipe::Status ProcessWithoutJitter(CalculatorContext* cc);
// Given the current count of periods that have passed, this returns
// the next valid timestamp of the middle point of the next period:
// if count is 0, it returns the first_timestamp_.
// if count is 1, it returns the first_timestamp_ + period (corresponding
// to the first tick using exact fps)
// e.g. for frame_rate=30 and first_timestamp_=0:
// 0: 0
// 1: 33333
// 2: 66667
// 3: 100000
//
// Can only be used if jitter_ equals zero.
Timestamp PeriodIndexToTimestamp(int64 index) const;
// Given a Timestamp, finds the closest sync Timestamp based on
// first_timestamp_ and the desired fps.
//
// Can only be used if jitter_ equals zero.
int64 TimestampToPeriodIndex(Timestamp timestamp) const;
// Outputs a packet if it is in range (start_time_, end_time_).
void OutputWithinLimits(CalculatorContext* cc, const Packet& packet) const;
// The timestamp of the first packet received.
Timestamp first_timestamp_;
// Number of frames per second (desired output frequency).
double frame_rate_;
// Inverse of frame_rate_.
int64 frame_time_usec_;
// Number of periods that have passed (= #packets sent to the output).
//
// Can only be used if jitter_ equals zero.
int64 period_count_;
// The last packet that was received.
Packet last_packet_;
VideoHeader video_header_;
// The "DATA" input stream.
CollectionItemId input_data_id_;
// The "DATA" output stream.
CollectionItemId output_data_id_;
// Indicator whether to flush last packet even if its timestamp is greater
// than the final stream timestamp. Set to false when jitter_ is non-zero.
bool flush_last_packet_;
// Jitter-related variables.
std::unique_ptr<RandomBase> random_;
double jitter_ = 0.0;
Timestamp next_output_timestamp_;
// If specified, output timestamps are aligned with base_timestamp.
// Otherwise, they are aligned with the first input timestamp.
Timestamp base_timestamp_;
// If specified, only outputs at/after start_time are included.
Timestamp start_time_;
// If specified, only outputs before end_time are included.
Timestamp end_time_;
// If set, the output timestamps nearest to start_time and end_time
// are included in the output, even if the nearest timestamp is not
// between start_time and end_time.
bool round_limits_;
};
REGISTER_CALCULATOR(PacketResamplerCalculator); REGISTER_CALCULATOR(PacketResamplerCalculator);
namespace { namespace {
// Returns a TimestampDiff (assuming microseconds) corresponding to the // Returns a TimestampDiff (assuming microseconds) corresponding to the
// given time in seconds. // given time in seconds.
@@ -209,6 +88,7 @@ TimestampDiff TimestampDiffFromSeconds(double seconds) {
flush_last_packet_ = resampler_options.flush_last_packet(); flush_last_packet_ = resampler_options.flush_last_packet();
jitter_ = resampler_options.jitter(); jitter_ = resampler_options.jitter();
jitter_with_reflection_ = resampler_options.jitter_with_reflection();
input_data_id_ = cc->Inputs().GetId("DATA", 0); input_data_id_ = cc->Inputs().GetId("DATA", 0);
if (!input_data_id_.IsValid()) { if (!input_data_id_.IsValid()) {
@@ -239,6 +119,8 @@ TimestampDiff TimestampDiffFromSeconds(double seconds) {
<< Timestamp::kTimestampUnitsPerSecond; << Timestamp::kTimestampUnitsPerSecond;
frame_time_usec_ = static_cast<int64>(1000000.0 / frame_rate_); frame_time_usec_ = static_cast<int64>(1000000.0 / frame_rate_);
jitter_usec_ = static_cast<int64>(1000000.0 * jitter_ / frame_rate_);
RET_CHECK_LE(jitter_usec_, frame_time_usec_);
video_header_.frame_rate = frame_rate_; video_header_.frame_rate = frame_rate_;
@@ -279,7 +161,10 @@ TimestampDiff TimestampDiffFromSeconds(double seconds) {
"SecureRandom is not available. With \"jitter\" specified, " "SecureRandom is not available. With \"jitter\" specified, "
"PacketResamplerCalculator processing cannot proceed."); "PacketResamplerCalculator processing cannot proceed.");
} }
packet_reservoir_random_ = CreateSecureRandom(seed);
} }
packet_reservoir_ =
std::make_unique<PacketReservoir>(packet_reservoir_random_.get());
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -294,6 +179,14 @@ TimestampDiff TimestampDiffFromSeconds(double seconds) {
} }
} }
if (jitter_ != 0.0 && random_ != nullptr) { if (jitter_ != 0.0 && random_ != nullptr) {
// Packet reservior is used to make sure there's an output for every period,
// e.g. partial period at the end of the stream.
if (packet_reservoir_->IsEnabled() &&
(first_timestamp_ == Timestamp::Unset() ||
(cc->InputTimestamp() - next_output_timestamp_min_).Value() >= 0)) {
auto curr_packet = cc->Inputs().Get(input_data_id_).Value();
packet_reservoir_->AddSample(curr_packet);
}
MP_RETURN_IF_ERROR(ProcessWithJitter(cc)); MP_RETURN_IF_ERROR(ProcessWithJitter(cc));
} else { } else {
MP_RETURN_IF_ERROR(ProcessWithoutJitter(cc)); MP_RETURN_IF_ERROR(ProcessWithoutJitter(cc));
@@ -303,11 +196,34 @@ TimestampDiff TimestampDiffFromSeconds(double seconds) {
} }
void PacketResamplerCalculator::InitializeNextOutputTimestampWithJitter() { void PacketResamplerCalculator::InitializeNextOutputTimestampWithJitter() {
next_output_timestamp_min_ = first_timestamp_;
if (jitter_with_reflection_) {
next_output_timestamp_ =
first_timestamp_ + random_->UnbiasedUniform64(frame_time_usec_);
return;
}
next_output_timestamp_ = next_output_timestamp_ =
first_timestamp_ + frame_time_usec_ * random_->RandFloat(); first_timestamp_ + frame_time_usec_ * random_->RandFloat();
} }
void PacketResamplerCalculator::UpdateNextOutputTimestampWithJitter() { void PacketResamplerCalculator::UpdateNextOutputTimestampWithJitter() {
packet_reservoir_->Clear();
if (jitter_with_reflection_) {
next_output_timestamp_min_ += frame_time_usec_;
Timestamp next_output_timestamp_max_ =
next_output_timestamp_min_ + frame_time_usec_;
next_output_timestamp_ += frame_time_usec_ +
random_->UnbiasedUniform64(2 * jitter_usec_ + 1) -
jitter_usec_;
next_output_timestamp_ = Timestamp(ReflectBetween(
next_output_timestamp_.Value(), next_output_timestamp_min_.Value(),
next_output_timestamp_max_.Value()));
CHECK_GE(next_output_timestamp_, next_output_timestamp_min_);
CHECK_LT(next_output_timestamp_, next_output_timestamp_max_);
return;
}
packet_reservoir_->Disable();
next_output_timestamp_ += next_output_timestamp_ +=
frame_time_usec_ * frame_time_usec_ *
((1.0 - jitter_) + 2.0 * jitter_ * random_->RandFloat()); ((1.0 - jitter_) + 2.0 * jitter_ * random_->RandFloat());
@@ -330,22 +246,27 @@ void PacketResamplerCalculator::UpdateNextOutputTimestampWithJitter() {
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
LOG_IF(WARNING, frame_time_usec_ < if (frame_time_usec_ <
(cc->InputTimestamp() - last_packet_.Timestamp()).Value()) (cc->InputTimestamp() - last_packet_.Timestamp()).Value()) {
<< "Adding jitter is meaningless when upsampling."; LOG_FIRST_N(WARNING, 2)
<< "Adding jitter is not very useful when upsampling.";
}
const int64 curr_diff = while (true) {
(next_output_timestamp_ - cc->InputTimestamp()).Value(); const int64 last_diff =
const int64 last_diff = (next_output_timestamp_ - last_packet_.Timestamp()).Value();
(next_output_timestamp_ - last_packet_.Timestamp()).Value(); RET_CHECK_GT(last_diff, 0);
if (curr_diff * last_diff > 0) { const int64 curr_diff =
return ::mediapipe::OkStatus(); (next_output_timestamp_ - cc->InputTimestamp()).Value();
if (curr_diff > 0) {
break;
}
OutputWithinLimits(cc, (std::abs(curr_diff) > last_diff
? last_packet_
: cc->Inputs().Get(input_data_id_).Value())
.At(next_output_timestamp_));
UpdateNextOutputTimestampWithJitter();
} }
OutputWithinLimits(cc, (std::abs(curr_diff) > std::abs(last_diff)
? last_packet_
: cc->Inputs().Get(input_data_id_).Value())
.At(next_output_timestamp_));
UpdateNextOutputTimestampWithJitter();
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -426,6 +347,9 @@ void PacketResamplerCalculator::UpdateNextOutputTimestampWithJitter() {
OutputWithinLimits(cc, OutputWithinLimits(cc,
last_packet_.At(PeriodIndexToTimestamp(period_count_))); last_packet_.At(PeriodIndexToTimestamp(period_count_)));
} }
if (!packet_reservoir_->IsEmpty()) {
OutputWithinLimits(cc, packet_reservoir_->GetSample());
}
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -0,0 +1,205 @@
#ifndef MEDIAPIPE_CALCULATORS_CORE_PACKET_RESAMPLER_CALCULATOR_H_
#define MEDIAPIPE_CALCULATORS_CORE_PACKET_RESAMPLER_CALCULATOR_H_
#include <cstdlib>
#include <memory>
#include <string>
#include "absl/strings/str_cat.h"
#include "mediapipe/calculators/core/packet_resampler_calculator.pb.h"
#include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/collection_item_id.h"
#include "mediapipe/framework/deps/mathutil.h"
#include "mediapipe/framework/deps/random_base.h"
#include "mediapipe/framework/formats/video_stream_header.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/status.h"
#include "mediapipe/framework/port/status_macros.h"
#include "mediapipe/framework/tool/options_util.h"
namespace mediapipe {
class PacketReservoir {
public:
PacketReservoir(RandomBase* rng) : rng_(rng) {}
// Replace candidate with current packet with 1/count_ probability.
void AddSample(Packet sample) {
if (rng_->UnbiasedUniform(++count_) == 0) {
reservoir_ = sample;
}
}
bool IsEnabled() { return rng_ && enabled_; }
void Disable() {
if (enabled_) enabled_ = false;
}
void Clear() { count_ = 0; }
bool IsEmpty() { return count_ == 0; }
Packet GetSample() { return reservoir_; }
private:
RandomBase* rng_;
bool enabled_ = true;
int32 count_ = 0;
Packet reservoir_;
};
// This calculator is used to normalize the frequency of the packets
// out of a stream. Given a desired frame rate, packets are going to be
// removed or added to achieve it.
//
// If jitter_ is specified:
// - The first packet is chosen randomly (uniform distribution) among frames
// that correspond to timestamps [0, 1/frame_rate). Let the chosen packet
// correspond to timestamp t.
// - The next packet is chosen randomly (uniform distribution) among frames
// that correspond to [t+(1-jitter)/frame_rate, t+(1+jitter)/frame_rate].
// - if jitter_with_reflection_ is true, the timestamp will be reflected
// against the boundaries of [t_0 + (k-1)/frame_rate, t_0 + k/frame_rate)
// so that its marginal distribution is uniform within this interval.
// In the formula, t_0 is the timestamp of the first sampled
// packet, and the k is the packet index.
// See paper (https://arxiv.org/abs/2002.01147) for details.
// - t is updated and the process is repeated.
// - Note that seed is specified as input side packet for reproducibility of
// the resampling. For Cloud ML Video Intelligence API, the hash of the
// input video should serve this purpose. For YouTube, either video ID or
// content hex ID of the input video should do.
//
// If jitter_ is not specified:
// - The first packet defines the first_timestamp of the output stream,
// so it is always emitted.
// - If more packets are emitted, they will have timestamp equal to
// round(first_timestamp + k * period) , where k is a positive
// integer and the period is defined by the frame rate.
// Example: first_timestamp=0, fps=30, then the output stream
// will have timestamps: 0, 33333, 66667, 100000, etc...
// - The packets selected for the output stream are the ones closer
// to the exact middle point (33333.33, 66666.67 in our previous
// example). In case of ties, later packets are chosen.
// - 'Empty' periods happen when there are no packets for a long time
// (greater than a period). In this case, we send a copy of the last
// packet received before the empty period.
// The jitter feature is disabled by default. To enable it, you need to
// implement CreateSecureRandom(const std::string&).
//
// The data stream may be either specified as the only stream (by index)
// or as the stream with tag "DATA".
//
// The input and output streams may be accompanied by a VIDEO_HEADER
// stream. This stream includes a VideoHeader at Timestamp::PreStream().
// The input VideoHeader on the VIDEO_HEADER stream will always be updated
// with the resampler frame rate no matter what the options value for
// output_header is before being output on the output VIDEO_HEADER stream.
// If the input VideoHeader is not available, then only the frame rate
// value will be set in the output.
//
// Related:
// packet_downsampler_calculator.cc: skips packets regardless of timestamps.
class PacketResamplerCalculator : public CalculatorBase {
public:
static ::mediapipe::Status GetContract(CalculatorContract* cc);
::mediapipe::Status Open(CalculatorContext* cc) override;
::mediapipe::Status Close(CalculatorContext* cc) override;
::mediapipe::Status Process(CalculatorContext* cc) override;
private:
// Calculates the first sampled timestamp that incorporates a jittering
// offset.
void InitializeNextOutputTimestampWithJitter();
// Calculates the next sampled timestamp that incorporates a jittering offset.
void UpdateNextOutputTimestampWithJitter();
// Logic for Process() when jitter_ != 0.0.
::mediapipe::Status ProcessWithJitter(CalculatorContext* cc);
// Logic for Process() when jitter_ == 0.0.
::mediapipe::Status ProcessWithoutJitter(CalculatorContext* cc);
// Given the current count of periods that have passed, this returns
// the next valid timestamp of the middle point of the next period:
// if count is 0, it returns the first_timestamp_.
// if count is 1, it returns the first_timestamp_ + period (corresponding
// to the first tick using exact fps)
// e.g. for frame_rate=30 and first_timestamp_=0:
// 0: 0
// 1: 33333
// 2: 66667
// 3: 100000
//
// Can only be used if jitter_ equals zero.
Timestamp PeriodIndexToTimestamp(int64 index) const;
// Given a Timestamp, finds the closest sync Timestamp based on
// first_timestamp_ and the desired fps.
//
// Can only be used if jitter_ equals zero.
int64 TimestampToPeriodIndex(Timestamp timestamp) const;
// Outputs a packet if it is in range (start_time_, end_time_).
void OutputWithinLimits(CalculatorContext* cc, const Packet& packet) const;
// The timestamp of the first packet received.
Timestamp first_timestamp_;
// Number of frames per second (desired output frequency).
double frame_rate_;
// Inverse of frame_rate_.
int64 frame_time_usec_;
// Number of periods that have passed (= #packets sent to the output).
//
// Can only be used if jitter_ equals zero.
int64 period_count_;
// The last packet that was received.
Packet last_packet_;
VideoHeader video_header_;
// The "DATA" input stream.
CollectionItemId input_data_id_;
// The "DATA" output stream.
CollectionItemId output_data_id_;
// Indicator whether to flush last packet even if its timestamp is greater
// than the final stream timestamp. Set to false when jitter_ is non-zero.
bool flush_last_packet_;
// Jitter-related variables.
std::unique_ptr<RandomBase> random_;
double jitter_ = 0.0;
bool jitter_with_reflection_;
int64 jitter_usec_;
Timestamp next_output_timestamp_;
// If jittering_with_reflection_ is true, next_output_timestamp_ will be
// kept within the interval
// [next_output_timestamp_min_, next_output_timestamp_min_ + frame_time_usec_)
Timestamp next_output_timestamp_min_;
// If specified, output timestamps are aligned with base_timestamp.
// Otherwise, they are aligned with the first input timestamp.
Timestamp base_timestamp_;
// If specified, only outputs at/after start_time are included.
Timestamp start_time_;
// If specified, only outputs before end_time are included.
Timestamp end_time_;
// If set, the output timestamps nearest to start_time and end_time
// are included in the output, even if the nearest timestamp is not
// between start_time and end_time.W
bool round_limits_;
// packet reservior used for sampling random packet out of partial
// period when jitter is enabled
std::unique_ptr<PacketReservoir> packet_reservoir_;
// random number generator used in packet_reservior_.
std::unique_ptr<RandomBase> packet_reservoir_random_;
};
} // namespace mediapipe
#endif // MEDIAPIPE_CALCULATORS_CORE_PACKET_RESAMPLER_CALCULATOR_H_
@@ -66,6 +66,7 @@ message PacketResamplerCalculatorOptions {
// pseudo-random number generator does its job and the number of frames is // pseudo-random number generator does its job and the number of frames is
// sufficiently large, the average frame rate will be close to this value. // sufficiently large, the average frame rate will be close to this value.
optional double jitter = 4; optional double jitter = 4;
optional bool jitter_with_reflection = 9 [default = false];
// If specified, output timestamps are aligned with base_timestamp. // If specified, output timestamps are aligned with base_timestamp.
// Otherwise, they are aligned with the first input timestamp. // Otherwise, they are aligned with the first input timestamp.
@@ -12,6 +12,8 @@
// See the License for the specific language governing permissions and // See the License for the specific language governing permissions and
// limitations under the License. // limitations under the License.
#include "mediapipe/calculators/core/packet_resampler_calculator.h"
#include <memory> #include <memory>
#include <string> #include <string>
#include <vector> #include <vector>
@@ -29,7 +31,6 @@
namespace mediapipe { namespace mediapipe {
namespace { namespace {
// A simple version of CalculatorRunner with built-in convenience // A simple version of CalculatorRunner with built-in convenience
// methods for setting inputs from a vector and checking outputs // methods for setting inputs from a vector and checking outputs
// against expected outputs (both timestamps and contents). // against expected outputs (both timestamps and contents).
@@ -0,0 +1,304 @@
// 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.
//
// Declaration of PacketThinnerCalculator.
#include <cmath> // for ceil
#include <memory>
#include "mediapipe/calculators/core/packet_thinner_calculator.pb.h"
#include "mediapipe/framework/calculator_context.h"
#include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/formats/video_stream_header.h"
#include "mediapipe/framework/port/integral_types.h"
#include "mediapipe/framework/port/logging.h"
#include "mediapipe/framework/port/status.h"
namespace mediapipe {
namespace {
const double kTimebaseUs = 1000000; // Microseconds.
const char* const kPeriodTag = "PERIOD";
} // namespace
// This calculator is used to thin an input stream of Packets.
// An example application would be to sample decoded frames of video
// at a coarser temporal resolution. Unless otherwise stated, all
// timestamps are in units of microseconds.
//
// Thinning can be accomplished in one of two ways:
// 1) asynchronous thinning (known below as async):
// Algorithm does not rely on a master clock and is parameterized only
// by a single option -- the period. Once a packet is emitted, the
// thinner will discard subsequent packets for the duration of the period
// [Analogous to a refractory period during which packet emission is
// suppressed.]
// Packets arriving before start_time are discarded, as are packets
// arriving at or after end_time.
// 2) synchronous thinning (known below as sync):
// There are two variants of this algorithm, both parameterized by a
// start_time and a period. As in (1), packets arriving before start_time
// or at/after end_time are discarded. Otherwise, at most one packet is
// emitted during a period, centered at timestamps generated by the
// expression:
// start_time + i * period [where i is a non-negative integer]
// During each period, the packet closest to the generated timestamp is
// emitted (latest in the case of ties). In the first variant
// (sync_output_timestamps = true), the emitted packet is output at the
// generated timestamp. In the second variant, the packet is output at
// its original timestamp. Both variants emit exactly the same packets,
// but at different timestamps.
//
// Thinning period can be provided in the calculator options or via a
// side packet with the tag "PERIOD".
//
// Example config:
// node {
// calculator: "PacketThinnerCalculator"
// input_stream: "signal"
// output_stream: "output"
// options {
// [mediapipe.PacketThinnerCalculatorOptions.ext] {
// thinner_type: SYNC
// period: 10
// sync_output_timestamps: true
// update_frame_rate: false
// }
// }
// }
class PacketThinnerCalculator : public CalculatorBase {
public:
PacketThinnerCalculator() {}
~PacketThinnerCalculator() override {}
static ::mediapipe::Status GetContract(CalculatorContract* cc) {
cc->Inputs().Index(0).SetAny();
cc->Outputs().Index(0).SetSameAs(&cc->Inputs().Index(0));
if (cc->InputSidePackets().HasTag(kPeriodTag)) {
cc->InputSidePackets().Tag(kPeriodTag).Set<int64>();
}
return ::mediapipe::OkStatus();
}
::mediapipe::Status Open(CalculatorContext* cc) override;
::mediapipe::Status Close(CalculatorContext* cc) override;
::mediapipe::Status Process(CalculatorContext* cc) override {
if (cc->InputTimestamp() < start_time_) {
return ::mediapipe::OkStatus(); // Drop packets before start_time_.
} else if (cc->InputTimestamp() >= end_time_) {
if (!cc->Outputs().Index(0).IsClosed()) {
cc->Outputs()
.Index(0)
.Close(); // No more Packets will be output after end_time_.
}
return ::mediapipe::OkStatus();
} else {
return thinner_type_ == PacketThinnerCalculatorOptions::ASYNC
? AsyncThinnerProcess(cc)
: SyncThinnerProcess(cc);
}
}
private:
// Implementation of ASYNC and SYNC versions of thinner algorithm.
::mediapipe::Status AsyncThinnerProcess(CalculatorContext* cc);
::mediapipe::Status SyncThinnerProcess(CalculatorContext* cc);
// Cached option.
PacketThinnerCalculatorOptions::ThinnerType thinner_type_;
// Given a Timestamp, finds the closest sync Timestamp
// based on start_time_ and period_. This can be earlier or
// later than given Timestamp, but is guaranteed to be within
// half a period_.
Timestamp NearestSyncTimestamp(Timestamp now) const;
// Cached option used by both async and sync thinners.
TimestampDiff period_; // Interval during which only one packet is emitted.
Timestamp start_time_; // Cached option - default Timestamp::Min()
Timestamp end_time_; // Cached option - default Timestamp::Max()
// Only used by async thinner:
Timestamp next_valid_timestamp_; // Suppress packets until this timestamp.
// Only used by sync thinner:
Packet saved_packet_; // Best packet not yet emitted.
bool sync_output_timestamps_; // Cached option.
};
REGISTER_CALCULATOR(PacketThinnerCalculator);
namespace {
TimestampDiff abs(TimestampDiff t) { return t < 0 ? -t : t; }
} // namespace
::mediapipe::Status PacketThinnerCalculator::Open(CalculatorContext* cc) {
auto& options = cc->Options<PacketThinnerCalculatorOptions>();
thinner_type_ = options.thinner_type();
// This check enables us to assume only two thinner types exist in Process()
CHECK(thinner_type_ == PacketThinnerCalculatorOptions::ASYNC ||
thinner_type_ == PacketThinnerCalculatorOptions::SYNC)
<< "Unsupported thinner type.";
if (thinner_type_ == PacketThinnerCalculatorOptions::ASYNC) {
// ASYNC thinner outputs packets with the same timestamp as their input so
// its safe to SetOffset(0). SYNC thinner manipulates timestamps of its
// output so we don't do this for that case.
cc->SetOffset(0);
}
if (cc->InputSidePackets().HasTag(kPeriodTag)) {
period_ =
TimestampDiff(cc->InputSidePackets().Tag(kPeriodTag).Get<int64>());
} else {
period_ = TimestampDiff(options.period());
}
CHECK_LT(TimestampDiff(0), period_) << "Specified period must be positive.";
if (options.has_start_time()) {
start_time_ = Timestamp(options.start_time());
} else if (thinner_type_ == PacketThinnerCalculatorOptions::ASYNC) {
start_time_ = Timestamp::Min();
} else {
start_time_ = Timestamp(0);
}
end_time_ =
options.has_end_time() ? Timestamp(options.end_time()) : Timestamp::Max();
CHECK_LT(start_time_, end_time_)
<< "Invalid PacketThinner: start_time must be earlier than end_time";
sync_output_timestamps_ = options.sync_output_timestamps();
next_valid_timestamp_ = start_time_;
// Drop packets until this time.
cc->Outputs().Index(0).SetNextTimestampBound(start_time_);
if (!cc->Inputs().Index(0).Header().IsEmpty()) {
if (options.update_frame_rate()) {
const VideoHeader& video_header =
cc->Inputs().Index(0).Header().Get<VideoHeader>();
double new_frame_rate;
if (thinner_type_ == PacketThinnerCalculatorOptions::ASYNC) {
new_frame_rate =
video_header.frame_rate /
ceil(video_header.frame_rate * options.period() / kTimebaseUs);
} else {
const double sampling_rate = kTimebaseUs / options.period();
new_frame_rate = video_header.frame_rate < sampling_rate
? video_header.frame_rate
: sampling_rate;
}
std::unique_ptr<VideoHeader> header(new VideoHeader);
header->format = video_header.format;
header->width = video_header.width;
header->height = video_header.height;
header->frame_rate = new_frame_rate;
cc->Outputs().Index(0).SetHeader(Adopt(header.release()));
} else {
cc->Outputs().Index(0).SetHeader(cc->Inputs().Index(0).Header());
}
}
return ::mediapipe::OkStatus();
}
::mediapipe::Status PacketThinnerCalculator::Close(CalculatorContext* cc) {
// Emit any saved packets before quitting.
if (!saved_packet_.IsEmpty()) {
// Only sync thinner should have saved packets.
CHECK_EQ(PacketThinnerCalculatorOptions::SYNC, thinner_type_);
if (sync_output_timestamps_) {
cc->Outputs().Index(0).AddPacket(
saved_packet_.At(NearestSyncTimestamp(saved_packet_.Timestamp())));
} else {
cc->Outputs().Index(0).AddPacket(saved_packet_);
}
}
return ::mediapipe::OkStatus();
}
::mediapipe::Status PacketThinnerCalculator::AsyncThinnerProcess(
CalculatorContext* cc) {
if (cc->InputTimestamp() >= next_valid_timestamp_) {
cc->Outputs().Index(0).AddPacket(
cc->Inputs().Index(0).Value()); // Emit current packet.
next_valid_timestamp_ = cc->InputTimestamp() + period_;
// Guaranteed not to emit packets seen during refractory period.
cc->Outputs().Index(0).SetNextTimestampBound(next_valid_timestamp_);
}
return ::mediapipe::OkStatus();
}
::mediapipe::Status PacketThinnerCalculator::SyncThinnerProcess(
CalculatorContext* cc) {
if (saved_packet_.IsEmpty()) {
// If no packet has been saved, store the current packet.
saved_packet_ = cc->Inputs().Index(0).Value();
cc->Outputs().Index(0).SetNextTimestampBound(
sync_output_timestamps_ ? NearestSyncTimestamp(cc->InputTimestamp())
: cc->InputTimestamp());
} else {
// Saved packet exists -- update or emit.
const Timestamp saved = saved_packet_.Timestamp();
const Timestamp saved_sync = NearestSyncTimestamp(saved);
const Timestamp now = cc->InputTimestamp();
const Timestamp now_sync = NearestSyncTimestamp(now);
CHECK_LE(saved_sync, now_sync);
if (saved_sync == now_sync) {
// Saved Packet is in same interval as current packet.
// Replace saved packet with current if it is at least as
// central as the saved packet wrt temporal interval.
// [We break ties in favor of fresher packets]
if (abs(now - now_sync) <= abs(saved - saved_sync)) {
saved_packet_ = cc->Inputs().Index(0).Value();
}
} else {
// Saved packet is the best packet from earlier interval: emit!
if (sync_output_timestamps_) {
cc->Outputs().Index(0).AddPacket(saved_packet_.At(saved_sync));
cc->Outputs().Index(0).SetNextTimestampBound(now_sync);
} else {
cc->Outputs().Index(0).AddPacket(saved_packet_);
cc->Outputs().Index(0).SetNextTimestampBound(now);
}
// Current packet is the first one we've seen from new interval -- save!
saved_packet_ = cc->Inputs().Index(0).Value();
}
}
return ::mediapipe::OkStatus();
}
Timestamp PacketThinnerCalculator::NearestSyncTimestamp(Timestamp now) const {
CHECK_NE(start_time_, Timestamp::Unset())
<< "Method only valid for sync thinner calculator.";
// Computation is done using int64 arithmetic. No easy way to avoid
// since Timestamps don't support div and multiply.
const int64 now64 = now.Value();
const int64 start64 = start_time_.Value();
const int64 period64 = period_.Value();
CHECK_LE(0, period64);
// Round now64 to its closest interval (units of period64).
int64 sync64 =
(now64 - start64 + period64 / 2) / period64 * period64 + start64;
CHECK_LE(abs(now64 - sync64), period64 / 2)
<< "start64: " << start64 << "; now64: " << now64
<< "; sync64: " << sync64;
return Timestamp(sync64);
}
} // namespace mediapipe
@@ -0,0 +1,66 @@
// Copyright 2018 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 PacketThinnerCalculatorOptions {
extend CalculatorOptions {
optional PacketThinnerCalculatorOptions ext = 288533508;
}
enum ThinnerType {
ASYNC = 1; // Asynchronous thinner, described below [default].
SYNC = 2; // Synchronous thinner, also described below.
}
optional ThinnerType thinner_type = 1 [default = ASYNC];
// The period (in microsecond) specifies the temporal interval during which
// only a single packet is emitted in the output stream. Has subtly different
// semantics depending on the thinner type, as follows.
//
// Async thinner: this option is a refractory period -- once a packet is
// emitted, we guarantee that no packets will be emitted for period ticks.
//
// Sync thinner: the period specifies a temporal interval during which
// only one packet is emitted. The emitted packet is guaranteed to be
// the one closest to the center of the temporal interval (no guarantee on
// how ties are broken). More specifically,
// intervals are centered at start_time + i * period
// (for non-negative integers i).
// Thus, each interval extends period/2 ticks before and after its center.
// Additionally, in the sync thinner any packets earlier than start_time
// are discarded and the thinner calls Close() once timestamp equals or
// exceeds end_time.
optional int64 period = 2 [default = 1];
// Packets before start_time and at/after end_time are discarded.
// Additionally, for a sync thinner, start time specifies the center of
// time invervals as described above and therefore should be set explicitly.
optional int64 start_time = 3; // If not specified, set to 0 for SYNC type,
// and set to Timestamp::Min() for ASYNC type.
optional int64 end_time = 4; // Set to Timestamp::Max() if not specified.
// Whether the timestamps of packets emitted by sync thinner should
// correspond to the center of their corresponding temporal interval.
// If false, packets emitted using original timestamp (as in async thinner).
optional bool sync_output_timestamps = 5 [default = true];
// If true, update the frame rate in the header, if it's available, to an
// estimated frame rate due to the sampling.
optional bool update_frame_rate = 6 [default = false];
}
@@ -0,0 +1,357 @@
// 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.
#include <memory>
#include <string>
#include <vector>
#include "absl/strings/str_cat.h"
#include "mediapipe/calculators/core/packet_thinner_calculator.pb.h"
#include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/calculator_runner.h"
#include "mediapipe/framework/formats/video_stream_header.h"
#include "mediapipe/framework/port/gmock.h"
#include "mediapipe/framework/port/gtest.h"
#include "mediapipe/framework/port/integral_types.h"
#include "mediapipe/framework/port/status_matchers.h"
namespace mediapipe {
namespace {
// A simple version of CalculatorRunner with built-in convenience methods for
// setting inputs from a vector and checking outputs against a vector of
// expected outputs.
class SimpleRunner : public CalculatorRunner {
public:
explicit SimpleRunner(const CalculatorOptions& options)
: CalculatorRunner("PacketThinnerCalculator", options) {
SetNumInputs(1);
SetNumOutputs(1);
SetNumInputSidePackets(0);
}
explicit SimpleRunner(const CalculatorGraphConfig::Node& node)
: CalculatorRunner(node) {}
void SetInput(const std::vector<int>& timestamp_list) {
MutableInputs()->Index(0).packets.clear();
for (const int ts : timestamp_list) {
MutableInputs()->Index(0).packets.push_back(
MakePacket<std::string>(absl::StrCat("Frame #", ts))
.At(Timestamp(ts)));
}
}
void SetFrameRate(const double frame_rate) {
auto video_header = absl::make_unique<VideoHeader>();
video_header->frame_rate = frame_rate;
MutableInputs()->Index(0).header = Adopt(video_header.release());
}
std::vector<int64> GetOutputTimestamps() const {
std::vector<int64> timestamps;
for (const Packet& packet : Outputs().Index(0).packets) {
timestamps.emplace_back(packet.Timestamp().Value());
}
return timestamps;
}
double GetFrameRate() const {
CHECK(!Outputs().Index(0).header.IsEmpty());
return Outputs().Index(0).header.Get<VideoHeader>().frame_rate;
}
};
// Check that thinner respects start_time and end_time options.
// We only test with one thinner because the logic for start & end time
// handling is shared across both types of thinner in Process().
TEST(PacketThinnerCalculatorTest, StartAndEndTimeTest) {
CalculatorOptions options;
auto* extension =
options.MutableExtension(PacketThinnerCalculatorOptions::ext);
extension->set_thinner_type(PacketThinnerCalculatorOptions::ASYNC);
extension->set_period(5);
extension->set_start_time(4);
extension->set_end_time(12);
SimpleRunner runner(options);
runner.SetInput({2, 3, 5, 7, 11, 13, 17, 19, 23, 29});
MP_ASSERT_OK(runner.Run());
const std::vector<int64> expected_timestamps = {5, 11};
EXPECT_EQ(expected_timestamps, runner.GetOutputTimestamps());
}
TEST(PacketThinnerCalculatorTest, AsyncUniformStreamThinningTest) {
CalculatorOptions options;
auto* extension =
options.MutableExtension(PacketThinnerCalculatorOptions::ext);
extension->set_thinner_type(PacketThinnerCalculatorOptions::ASYNC);
extension->set_period(5);
SimpleRunner runner(options);
runner.SetInput({2, 4, 6, 8, 10, 12, 14});
MP_ASSERT_OK(runner.Run());
const std::vector<int64> expected_timestamps = {2, 8, 14};
EXPECT_EQ(expected_timestamps, runner.GetOutputTimestamps());
}
TEST(PacketThinnerCalculatorTest, ASyncUniformStreamThinningTestBySidePacket) {
// Note: sync runner but outputting *original* timestamps.
CalculatorGraphConfig::Node node;
node.set_calculator("PacketThinnerCalculator");
node.add_input_side_packet("PERIOD:period");
node.add_input_stream("input_stream");
node.add_output_stream("output_stream");
auto* extension = node.mutable_options()->MutableExtension(
PacketThinnerCalculatorOptions::ext);
extension->set_thinner_type(PacketThinnerCalculatorOptions::ASYNC);
extension->set_start_time(0);
extension->set_sync_output_timestamps(false);
SimpleRunner runner(node);
runner.SetInput({2, 4, 6, 8, 10, 12, 14});
runner.MutableSidePackets()->Tag("PERIOD") = MakePacket<int64>(5);
MP_ASSERT_OK(runner.Run());
const std::vector<int64> expected_timestamps = {2, 8, 14};
EXPECT_EQ(expected_timestamps, runner.GetOutputTimestamps());
}
TEST(PacketThinnerCalculatorTest, SyncUniformStreamThinningTest1) {
// Note: sync runner but outputting *original* timestamps.
CalculatorOptions options;
auto* extension =
options.MutableExtension(PacketThinnerCalculatorOptions::ext);
extension->set_thinner_type(PacketThinnerCalculatorOptions::SYNC);
extension->set_start_time(0);
extension->set_period(5);
extension->set_sync_output_timestamps(false);
SimpleRunner runner(options);
runner.SetInput({2, 4, 6, 8, 10, 12, 14});
MP_ASSERT_OK(runner.Run());
const std::vector<int64> expected_timestamps = {2, 6, 10, 14};
EXPECT_EQ(expected_timestamps, runner.GetOutputTimestamps());
}
TEST(PacketThinnerCalculatorTest, SyncUniformStreamThinningTestBySidePacket1) {
// Note: sync runner but outputting *original* timestamps.
CalculatorGraphConfig::Node node;
node.set_calculator("PacketThinnerCalculator");
node.add_input_side_packet("PERIOD:period");
node.add_input_stream("input_stream");
node.add_output_stream("output_stream");
auto* extension = node.mutable_options()->MutableExtension(
PacketThinnerCalculatorOptions::ext);
extension->set_thinner_type(PacketThinnerCalculatorOptions::SYNC);
extension->set_start_time(0);
extension->set_sync_output_timestamps(false);
SimpleRunner runner(node);
runner.SetInput({2, 4, 6, 8, 10, 12, 14});
runner.MutableSidePackets()->Tag("PERIOD") = MakePacket<int64>(5);
MP_ASSERT_OK(runner.Run());
const std::vector<int64> expected_timestamps = {2, 6, 10, 14};
EXPECT_EQ(expected_timestamps, runner.GetOutputTimestamps());
}
TEST(PacketThinnerCalculatorTest, SyncUniformStreamThinningTest2) {
// Same test but now with synced timestamps.
CalculatorOptions options;
auto* extension =
options.MutableExtension(PacketThinnerCalculatorOptions::ext);
extension->set_thinner_type(PacketThinnerCalculatorOptions::SYNC);
extension->set_start_time(0);
extension->set_period(5);
extension->set_sync_output_timestamps(true);
SimpleRunner runner(options);
runner.SetInput({2, 4, 6, 8, 10, 12, 14});
MP_ASSERT_OK(runner.Run());
const std::vector<int64> expected_timestamps = {0, 5, 10, 15};
EXPECT_EQ(expected_timestamps, runner.GetOutputTimestamps());
}
// Test: Given a stream with timestamps corresponding to first ten prime numbers
// and period of 5, confirm whether timestamps of thinner stream matches
// expectations.
TEST(PacketThinnerCalculatorTest, PrimeStreamThinningTest1) {
// ASYNC thinner.
CalculatorOptions options;
auto* extension =
options.MutableExtension(PacketThinnerCalculatorOptions::ext);
extension->set_thinner_type(PacketThinnerCalculatorOptions::ASYNC);
extension->set_period(5);
SimpleRunner runner(options);
runner.SetInput({2, 3, 5, 7, 11, 13, 17, 19, 23, 29});
MP_ASSERT_OK(runner.Run());
const std::vector<int64> expected_timestamps = {2, 7, 13, 19, 29};
EXPECT_EQ(expected_timestamps, runner.GetOutputTimestamps());
}
TEST(PacketThinnerCalculatorTest, PrimeStreamThinningTest2) {
// SYNC with original timestamps.
CalculatorOptions options;
auto* extension =
options.MutableExtension(PacketThinnerCalculatorOptions::ext);
extension->set_thinner_type(PacketThinnerCalculatorOptions::SYNC);
extension->set_start_time(0);
extension->set_period(5);
extension->set_sync_output_timestamps(false);
SimpleRunner runner(options);
runner.SetInput({2, 3, 5, 7, 11, 13, 17, 19, 23, 29});
MP_ASSERT_OK(runner.Run());
const std::vector<int64> expected_timestamps = {2, 5, 11, 17, 19, 23, 29};
EXPECT_EQ(expected_timestamps, runner.GetOutputTimestamps());
}
// Confirm that Calculator correctly handles boundary cases.
TEST(PacketThinnerCalculatorTest, BoundaryTimestampTest1) {
// Odd period, negative start_time
CalculatorOptions options;
auto* extension =
options.MutableExtension(PacketThinnerCalculatorOptions::ext);
extension->set_thinner_type(PacketThinnerCalculatorOptions::SYNC);
extension->set_start_time(-10);
extension->set_period(5);
extension->set_sync_output_timestamps(true);
SimpleRunner runner(options);
// Two timestamps falling on either side of a period boundary.
runner.SetInput({2, 3});
MP_ASSERT_OK(runner.Run());
const std::vector<int64> expected_timestamps = {0, 5};
EXPECT_EQ(expected_timestamps, runner.GetOutputTimestamps());
}
TEST(PacketThinnerCalculatorTest, BoundaryTimestampTest2) {
// Even period, negative start_time, negative packet timestamps.
CalculatorOptions options;
auto* extension =
options.MutableExtension(PacketThinnerCalculatorOptions::ext);
extension->set_thinner_type(PacketThinnerCalculatorOptions::SYNC);
extension->set_start_time(-144);
extension->set_period(6);
extension->set_sync_output_timestamps(true);
SimpleRunner runner(options);
// Two timestamps falling on either side of a period boundary.
runner.SetInput({-4, -3, 8, 9});
MP_ASSERT_OK(runner.Run());
const std::vector<int64> expected_timestamps = {-6, 0, 6, 12};
EXPECT_EQ(expected_timestamps, runner.GetOutputTimestamps());
}
TEST(PacketThinnerCalculatorTest, FrameRateTest1) {
CalculatorOptions options;
auto* extension =
options.MutableExtension(PacketThinnerCalculatorOptions::ext);
extension->set_thinner_type(PacketThinnerCalculatorOptions::ASYNC);
extension->set_period(5);
extension->set_update_frame_rate(true);
SimpleRunner runner(options);
runner.SetInput({2, 4, 6, 8, 10, 12, 14});
runner.SetFrameRate(1000000.0 / 2);
MP_ASSERT_OK(runner.Run());
const std::vector<int64> expected_timestamps = {2, 8, 14};
EXPECT_EQ(expected_timestamps, runner.GetOutputTimestamps());
// The true sampling period is 6.
EXPECT_DOUBLE_EQ(1000000.0 / 6, runner.GetFrameRate());
}
TEST(PacketThinnerCalculatorTest, FrameRateTest2) {
CalculatorOptions options;
auto* extension =
options.MutableExtension(PacketThinnerCalculatorOptions::ext);
extension->set_thinner_type(PacketThinnerCalculatorOptions::ASYNC);
extension->set_period(5);
extension->set_update_frame_rate(true);
SimpleRunner runner(options);
runner.SetInput({8, 16, 24, 32, 40, 48, 56});
runner.SetFrameRate(1000000.0 / 8);
MP_ASSERT_OK(runner.Run());
const std::vector<int64> expected_timestamps = {8, 16, 24, 32, 40, 48, 56};
EXPECT_EQ(expected_timestamps, runner.GetOutputTimestamps());
// The true sampling period is still 8.
EXPECT_DOUBLE_EQ(1000000.0 / 8, runner.GetFrameRate());
}
TEST(PacketThinnerCalculatorTest, FrameRateTest3) {
// Note: sync runner but outputting *original* timestamps.
CalculatorOptions options;
auto* extension =
options.MutableExtension(PacketThinnerCalculatorOptions::ext);
extension->set_thinner_type(PacketThinnerCalculatorOptions::SYNC);
extension->set_start_time(0);
extension->set_period(5);
extension->set_sync_output_timestamps(false);
extension->set_update_frame_rate(true);
SimpleRunner runner(options);
runner.SetInput({2, 4, 6, 8, 10, 12, 14});
runner.SetFrameRate(1000000.0 / 2);
MP_ASSERT_OK(runner.Run());
const std::vector<int64> expected_timestamps = {2, 6, 10, 14};
EXPECT_EQ(expected_timestamps, runner.GetOutputTimestamps());
// The true (long-run) sampling period is 5.
EXPECT_DOUBLE_EQ(1000000.0 / 5, runner.GetFrameRate());
}
TEST(PacketThinnerCalculatorTest, FrameRateTest4) {
// Same test but now with synced timestamps.
CalculatorOptions options;
auto* extension =
options.MutableExtension(PacketThinnerCalculatorOptions::ext);
extension->set_thinner_type(PacketThinnerCalculatorOptions::SYNC);
extension->set_start_time(0);
extension->set_period(5);
extension->set_sync_output_timestamps(true);
extension->set_update_frame_rate(true);
SimpleRunner runner(options);
runner.SetInput({2, 4, 6, 8, 10, 12, 14});
runner.SetFrameRate(1000000.0 / 2);
MP_ASSERT_OK(runner.Run());
const std::vector<int64> expected_timestamps = {0, 5, 10, 15};
EXPECT_EQ(expected_timestamps, runner.GetOutputTimestamps());
// The true (long-run) sampling period is 5.
EXPECT_DOUBLE_EQ(1000000.0 / 5, runner.GetFrameRate());
}
TEST(PacketThinnerCalculatorTest, FrameRateTest5) {
CalculatorOptions options;
auto* extension =
options.MutableExtension(PacketThinnerCalculatorOptions::ext);
extension->set_thinner_type(PacketThinnerCalculatorOptions::SYNC);
extension->set_start_time(0);
extension->set_period(5);
extension->set_sync_output_timestamps(true);
extension->set_update_frame_rate(true);
SimpleRunner runner(options);
runner.SetInput({8, 16, 24, 32, 40, 48, 56});
runner.SetFrameRate(1000000.0 / 8);
MP_ASSERT_OK(runner.Run());
const std::vector<int64> expected_timestamps = {10, 15, 25, 30, 40, 50, 55};
EXPECT_EQ(expected_timestamps, runner.GetOutputTimestamps());
// The true (long-run) sampling period is 8.
EXPECT_DOUBLE_EQ(1000000.0 / 8, runner.GetFrameRate());
}
} // namespace
} // namespace mediapipe
@@ -17,6 +17,7 @@
#include "mediapipe/framework/calculator_framework.h" #include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/port/ret_check.h" #include "mediapipe/framework/port/ret_check.h"
#include "mediapipe/framework/port/status.h" #include "mediapipe/framework/port/status.h"
#include "mediapipe/framework/timestamp.h"
namespace mediapipe { namespace mediapipe {
@@ -86,6 +87,7 @@ class PreviousLoopbackCalculator : public CalculatorBase {
main_ts_.pop_front(); main_ts_.pop_front();
} }
} }
auto& loop_out = cc->Outputs().Get(loop_out_id_);
while (!main_ts_.empty() && !loopback_packets_.empty()) { while (!main_ts_.empty() && !loopback_packets_.empty()) {
Timestamp main_timestamp = main_ts_.front(); Timestamp main_timestamp = main_ts_.front();
@@ -95,18 +97,31 @@ class PreviousLoopbackCalculator : public CalculatorBase {
if (previous_loopback.IsEmpty()) { if (previous_loopback.IsEmpty()) {
// TODO: SetCompleteTimestampBound would be more useful. // TODO: SetCompleteTimestampBound would be more useful.
cc->Outputs() loop_out.SetNextTimestampBound(main_timestamp + 1);
.Get(loop_out_id_)
.SetNextTimestampBound(main_timestamp + 1);
} else { } else {
cc->Outputs().Get(loop_out_id_).AddPacket(std::move(previous_loopback)); loop_out.AddPacket(std::move(previous_loopback));
}
}
// In case of an empty loopback input, the next timestamp bound for
// loopback input is the loopback timestamp + 1. The next timestamp bound
// for output is set and the main_ts_ vector is truncated accordingly.
if (loopback_packet.IsEmpty() &&
loopback_packet.Timestamp() != Timestamp::Unstarted()) {
Timestamp loopback_bound =
loopback_packet.Timestamp().NextAllowedInStream();
while (!main_ts_.empty() && main_ts_.front() <= loopback_bound) {
main_ts_.pop_front();
}
if (main_ts_.empty()) {
loop_out.SetNextTimestampBound(loopback_bound.NextAllowedInStream());
} }
} }
if (!main_ts_.empty()) { if (!main_ts_.empty()) {
cc->Outputs().Get(loop_out_id_).SetNextTimestampBound(main_ts_.front()); loop_out.SetNextTimestampBound(main_ts_.front());
} }
if (cc->Inputs().Get(main_id_).IsDone() && main_ts_.empty()) { if (cc->Inputs().Get(main_id_).IsDone() && main_ts_.empty()) {
cc->Outputs().Get(loop_out_id_).Close(); loop_out.Close();
} }
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -93,14 +93,19 @@ TEST(PreviousLoopbackCalculator, CorrectTimestamps) {
EXPECT_EQ(TimestampValues(in_prev), (std::vector<int64>{1})); EXPECT_EQ(TimestampValues(in_prev), (std::vector<int64>{1}));
EXPECT_EQ(pair_values(in_prev.back()), std::make_pair(1, -1)); EXPECT_EQ(pair_values(in_prev.back()), std::make_pair(1, -1));
send_packet("in", 2);
MP_EXPECT_OK(graph_.WaitUntilIdle());
EXPECT_EQ(TimestampValues(in_prev), (std::vector<int64>{1, 2}));
EXPECT_EQ(pair_values(in_prev.back()), std::make_pair(2, 1));
send_packet("in", 5); send_packet("in", 5);
MP_EXPECT_OK(graph_.WaitUntilIdle()); MP_EXPECT_OK(graph_.WaitUntilIdle());
EXPECT_EQ(TimestampValues(in_prev), (std::vector<int64>{1, 5})); EXPECT_EQ(TimestampValues(in_prev), (std::vector<int64>{1, 2, 5}));
EXPECT_EQ(pair_values(in_prev.back()), std::make_pair(5, 1)); EXPECT_EQ(pair_values(in_prev.back()), std::make_pair(5, 2));
send_packet("in", 15); send_packet("in", 15);
MP_EXPECT_OK(graph_.WaitUntilIdle()); MP_EXPECT_OK(graph_.WaitUntilIdle());
EXPECT_EQ(TimestampValues(in_prev), (std::vector<int64>{1, 5, 15})); EXPECT_EQ(TimestampValues(in_prev), (std::vector<int64>{1, 2, 5, 15}));
EXPECT_EQ(pair_values(in_prev.back()), std::make_pair(15, 5)); EXPECT_EQ(pair_values(in_prev.back()), std::make_pair(15, 5));
MP_EXPECT_OK(graph_.CloseAllInputStreams()); MP_EXPECT_OK(graph_.CloseAllInputStreams());
@@ -182,21 +187,84 @@ TEST(PreviousLoopbackCalculator, ClosesCorrectly) {
MP_EXPECT_OK(graph_.WaitUntilIdle()); MP_EXPECT_OK(graph_.WaitUntilIdle());
EXPECT_EQ(TimestampValues(outputs), (std::vector<int64>{1})); EXPECT_EQ(TimestampValues(outputs), (std::vector<int64>{1}));
send_packet("in", 2);
MP_EXPECT_OK(graph_.WaitUntilIdle());
EXPECT_EQ(TimestampValues(outputs), (std::vector<int64>{1, 2}));
send_packet("in", 5); send_packet("in", 5);
MP_EXPECT_OK(graph_.WaitUntilIdle()); MP_EXPECT_OK(graph_.WaitUntilIdle());
EXPECT_EQ(TimestampValues(outputs), (std::vector<int64>{1, 5})); EXPECT_EQ(TimestampValues(outputs), (std::vector<int64>{1, 2, 5}));
send_packet("in", 15); send_packet("in", 15);
MP_EXPECT_OK(graph_.WaitUntilIdle()); MP_EXPECT_OK(graph_.WaitUntilIdle());
EXPECT_EQ(TimestampValues(outputs), (std::vector<int64>{1, 5, 15})); EXPECT_EQ(TimestampValues(outputs), (std::vector<int64>{1, 2, 5, 15}));
MP_EXPECT_OK(graph_.CloseAllInputStreams()); MP_EXPECT_OK(graph_.CloseAllInputStreams());
MP_EXPECT_OK(graph_.WaitUntilIdle()); MP_EXPECT_OK(graph_.WaitUntilIdle());
EXPECT_EQ(TimestampValues(outputs), EXPECT_EQ(TimestampValues(outputs),
(std::vector<int64>{1, 5, 15, Timestamp::Max().Value()})); (std::vector<int64>{1, 2, 5, 15, Timestamp::Max().Value()}));
MP_EXPECT_OK(graph_.WaitUntilDone()); MP_EXPECT_OK(graph_.WaitUntilDone());
} }
// Demonstrates that downstream calculators won't be blocked by
// always-empty-LOOP-stream.
TEST(PreviousLoopbackCalculator, EmptyLoopForever) {
std::vector<Packet> outputs;
CalculatorGraphConfig graph_config_ =
ParseTextProtoOrDie<CalculatorGraphConfig>(R"(
input_stream: 'in'
node {
calculator: 'PreviousLoopbackCalculator'
input_stream: 'MAIN:in'
input_stream: 'LOOP:previous'
input_stream_info: { tag_index: 'LOOP' back_edge: true }
output_stream: 'PREV_LOOP:previous'
}
# This calculator synchronizes its inputs as normal, so it is used
# to check that both "in" and "previous" are ready.
node {
calculator: 'PassThroughCalculator'
input_stream: 'in'
input_stream: 'previous'
output_stream: 'out'
output_stream: 'previous2'
}
node {
calculator: 'PacketOnCloseCalculator'
input_stream: 'out'
output_stream: 'close_out'
}
)");
tool::AddVectorSink("close_out", &graph_config_, &outputs);
CalculatorGraph graph_;
MP_ASSERT_OK(graph_.Initialize(graph_config_, {}));
MP_ASSERT_OK(graph_.StartRun({}));
auto send_packet = [&graph_](const std::string& input_name, int n) {
MP_EXPECT_OK(graph_.AddPacketToInputStream(
input_name, MakePacket<int>(n).At(Timestamp(n))));
};
send_packet("in", 0);
MP_EXPECT_OK(graph_.WaitUntilIdle());
EXPECT_EQ(TimestampValues(outputs), (std::vector<int64>{0}));
for (int main_ts = 1; main_ts < 50; ++main_ts) {
send_packet("in", main_ts);
MP_EXPECT_OK(graph_.WaitUntilIdle());
std::vector<int64> ts_values = TimestampValues(outputs);
EXPECT_EQ(ts_values.size(), main_ts + 1);
for (int j = 0; j < main_ts; ++j) {
EXPECT_EQ(ts_values[j], j);
}
}
MP_EXPECT_OK(graph_.CloseAllInputStreams());
MP_EXPECT_OK(graph_.WaitUntilIdle());
MP_EXPECT_OK(graph_.WaitUntilDone());
}
} // anonymous namespace } // anonymous namespace
} // namespace mediapipe } // namespace mediapipe
@@ -28,55 +28,133 @@ using mediapipe::PacketTypeSet;
using mediapipe::Timestamp; using mediapipe::Timestamp;
namespace { namespace {
constexpr char kTagAtPreStream[] = "AT_PRESTREAM";
constexpr char kTagAtPostStream[] = "AT_POSTSTREAM";
constexpr char kTagAtZero[] = "AT_ZERO";
constexpr char kTagAtTick[] = "AT_TICK";
constexpr char kTagTick[] = "TICK";
static std::map<std::string, Timestamp>* kTimestampMap = []() { static std::map<std::string, Timestamp>* kTimestampMap = []() {
auto* res = new std::map<std::string, Timestamp>(); auto* res = new std::map<std::string, Timestamp>();
res->emplace("AT_PRESTREAM", Timestamp::PreStream()); res->emplace(kTagAtPreStream, Timestamp::PreStream());
res->emplace("AT_POSTSTREAM", Timestamp::PostStream()); res->emplace(kTagAtPostStream, Timestamp::PostStream());
res->emplace("AT_ZERO", Timestamp(0)); res->emplace(kTagAtZero, Timestamp(0));
res->emplace(kTagAtTick, Timestamp::Unset());
return res; return res;
}(); }();
template <typename CC>
std::string GetOutputTag(const CC& cc) {
// Single output tag only is required by contract.
return *cc.Outputs().GetTags().begin();
}
} // namespace } // namespace
// Outputs the single input_side_packet at the timestamp specified in the // Outputs side packet(s) in corresponding output stream(s) with a particular
// output_stream tag. Valid tags are AT_PRESTREAM, AT_POSTSTREAM and AT_ZERO. // timestamp, depending on the tag used to define output stream(s). (One tag can
// be used only.)
//
// Valid tags are AT_PRESTREAM, AT_POSTSTREAM, AT_ZERO and AT_TICK and
// corresponding timestamps are Timestamp::PreStream(), Timestamp::PostStream(),
// Timestamp(0) and timestamp of a packet received in TICK input.
//
// Examples:
// node {
// calculator: "SidePacketToStreamCalculator"
// input_side_packet: "side_packet"
// output_stream: "AT_PRESTREAM:packet"
// }
//
// node {
// calculator: "SidePacketToStreamCalculator"
// input_stream: "TICK:tick"
// input_side_packet: "side_packet"
// output_stream: "AT_TICK:packet"
// }
class SidePacketToStreamCalculator : public CalculatorBase { class SidePacketToStreamCalculator : public CalculatorBase {
public: public:
SidePacketToStreamCalculator() = default; SidePacketToStreamCalculator() = default;
~SidePacketToStreamCalculator() override = default; ~SidePacketToStreamCalculator() override = default;
static ::mediapipe::Status GetContract(CalculatorContract* cc); static ::mediapipe::Status GetContract(CalculatorContract* cc);
::mediapipe::Status Open(CalculatorContext* cc) override;
::mediapipe::Status Process(CalculatorContext* cc) override; ::mediapipe::Status Process(CalculatorContext* cc) override;
::mediapipe::Status Close(CalculatorContext* cc) override; ::mediapipe::Status Close(CalculatorContext* cc) override;
private:
bool is_tick_processing_ = false;
std::string output_tag_;
}; };
REGISTER_CALCULATOR(SidePacketToStreamCalculator); REGISTER_CALCULATOR(SidePacketToStreamCalculator);
::mediapipe::Status SidePacketToStreamCalculator::GetContract( ::mediapipe::Status SidePacketToStreamCalculator::GetContract(
CalculatorContract* cc) { CalculatorContract* cc) {
cc->InputSidePackets().Index(0).SetAny(); const auto& tags = cc->Outputs().GetTags();
RET_CHECK(tags.size() == 1 && kTimestampMap->count(*tags.begin()) == 1)
<< "Only one of AT_PRESTREAM, AT_POSTSTREAM, AT_ZERO and AT_TICK tags is "
"allowed and required to specify output stream(s).";
RET_CHECK(
(cc->Outputs().HasTag(kTagAtTick) && cc->Inputs().HasTag(kTagTick)) ||
(!cc->Outputs().HasTag(kTagAtTick) && !cc->Inputs().HasTag(kTagTick)))
<< "Either both of TICK and AT_TICK should be used or none of them.";
const std::string output_tag = GetOutputTag(*cc);
const int num_entries = cc->Outputs().NumEntries(output_tag);
RET_CHECK_EQ(num_entries, cc->InputSidePackets().NumEntries())
<< "Same number of input side packets and output streams is required.";
for (int i = 0; i < num_entries; ++i) {
cc->InputSidePackets().Index(i).SetAny();
cc->Outputs()
.Get(output_tag, i)
.SetSameAs(cc->InputSidePackets().Index(i).GetSameAs());
}
std::set<std::string> tags = cc->Outputs().GetTags(); if (cc->Inputs().HasTag(kTagTick)) {
RET_CHECK_EQ(tags.size(), 1); cc->Inputs().Tag(kTagTick).SetAny();
}
RET_CHECK_EQ(kTimestampMap->count(*tags.begin()), 1); return ::mediapipe::OkStatus();
cc->Outputs().Tag(*tags.begin()).SetAny(); }
::mediapipe::Status SidePacketToStreamCalculator::Open(CalculatorContext* cc) {
output_tag_ = GetOutputTag(*cc);
if (cc->Inputs().HasTag(kTagTick)) {
is_tick_processing_ = true;
// Set offset, so output timestamp bounds are updated in response to TICK
// timestamp bound update.
cc->SetOffset(TimestampDiff(0));
}
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
::mediapipe::Status SidePacketToStreamCalculator::Process( ::mediapipe::Status SidePacketToStreamCalculator::Process(
CalculatorContext* cc) { CalculatorContext* cc) {
return mediapipe::tool::StatusStop(); if (is_tick_processing_) {
// TICK input is guaranteed to be non-empty, as it's the only input stream
// for this calculator.
const auto& timestamp = cc->Inputs().Tag(kTagTick).Value().Timestamp();
for (int i = 0; i < cc->Outputs().NumEntries(output_tag_); ++i) {
cc->Outputs()
.Get(output_tag_, i)
.AddPacket(cc->InputSidePackets().Index(i).At(timestamp));
}
return ::mediapipe::OkStatus();
}
return ::mediapipe::tool::StatusStop();
} }
::mediapipe::Status SidePacketToStreamCalculator::Close(CalculatorContext* cc) { ::mediapipe::Status SidePacketToStreamCalculator::Close(CalculatorContext* cc) {
std::set<std::string> tags = cc->Outputs().GetTags(); if (!cc->Outputs().HasTag(kTagAtTick)) {
RET_CHECK_EQ(tags.size(), 1); const auto& timestamp = kTimestampMap->at(output_tag_);
const std::string& tag = *tags.begin(); for (int i = 0; i < cc->Outputs().NumEntries(output_tag_); ++i) {
RET_CHECK_EQ(kTimestampMap->count(tag), 1); cc->Outputs()
cc->Outputs().Tag(tag).AddPacket( .Get(output_tag_, i)
cc->InputSidePackets().Index(0).At(kTimestampMap->at(tag))); .AddPacket(cc->InputSidePackets().Index(i).At(timestamp));
}
}
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -0,0 +1,275 @@
// Copyright 2020 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 <string>
#include <vector>
#include "absl/memory/memory.h"
#include "absl/strings/match.h"
#include "absl/strings/str_replace.h"
#include "absl/strings/string_view.h"
#include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/port/gtest.h"
#include "mediapipe/framework/port/integral_types.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/tool/options_util.h"
namespace mediapipe {
namespace {
TEST(SidePacketToStreamCalculator, WrongConfig_MissingTick) {
CalculatorGraphConfig graph_config =
ParseTextProtoOrDie<CalculatorGraphConfig>(
R"(
input_stream: "tick"
input_side_packet: "side_packet"
output_stream: "packet"
node {
calculator: "SidePacketToStreamCalculator"
input_side_packet: "side_packet"
output_stream: "AT_TICK:packet"
}
)");
CalculatorGraph graph;
auto status = graph.Initialize(graph_config);
EXPECT_FALSE(status.ok());
EXPECT_PRED2(
absl::StrContains, status.message(),
"Either both of TICK and AT_TICK should be used or none of them.");
}
TEST(SidePacketToStreamCalculator, WrongConfig_NonExistentTag) {
CalculatorGraphConfig graph_config =
ParseTextProtoOrDie<CalculatorGraphConfig>(
R"(
input_stream: "tick"
input_side_packet: "side_packet"
output_stream: "packet"
node {
calculator: "SidePacketToStreamCalculator"
input_side_packet: "side_packet"
output_stream: "DOES_NOT_EXIST:packet"
}
)");
CalculatorGraph graph;
auto status = graph.Initialize(graph_config);
EXPECT_FALSE(status.ok());
EXPECT_PRED2(absl::StrContains, status.message(),
"Only one of AT_PRESTREAM, AT_POSTSTREAM, AT_ZERO and AT_TICK "
"tags is allowed and required to specify output stream(s).");
}
TEST(SidePacketToStreamCalculator, WrongConfig_MixedTags) {
CalculatorGraphConfig graph_config =
ParseTextProtoOrDie<CalculatorGraphConfig>(
R"(
input_stream: "tick"
input_side_packet: "side_packet0"
input_side_packet: "side_packet1"
node {
calculator: "SidePacketToStreamCalculator"
input_side_packet: "side_packet0"
input_side_packet: "side_packet1"
output_stream: "AT_TICK:packet0"
output_stream: "AT_PRE_STREAM:packet1"
}
)");
CalculatorGraph graph;
auto status = graph.Initialize(graph_config);
EXPECT_FALSE(status.ok());
EXPECT_PRED2(absl::StrContains, status.message(),
"Only one of AT_PRESTREAM, AT_POSTSTREAM, AT_ZERO and AT_TICK "
"tags is allowed and required to specify output stream(s).");
}
TEST(SidePacketToStreamCalculator, WrongConfig_NotEnoughSidePackets) {
CalculatorGraphConfig graph_config =
ParseTextProtoOrDie<CalculatorGraphConfig>(
R"(
input_side_packet: "side_packet0"
input_side_packet: "side_packet1"
node {
calculator: "SidePacketToStreamCalculator"
input_side_packet: "side_packet0"
output_stream: "AT_PRESTREAM:0:packet0"
output_stream: "AT_PRESTREAM:1:packet1"
}
)");
CalculatorGraph graph;
auto status = graph.Initialize(graph_config);
EXPECT_FALSE(status.ok());
EXPECT_PRED2(
absl::StrContains, status.message(),
"Same number of input side packets and output streams is required.");
}
TEST(SidePacketToStreamCalculator, WrongConfig_NotEnoughOutputStreams) {
CalculatorGraphConfig graph_config =
ParseTextProtoOrDie<CalculatorGraphConfig>(
R"(
input_side_packet: "side_packet0"
input_side_packet: "side_packet1"
node {
calculator: "SidePacketToStreamCalculator"
input_side_packet: "side_packet0"
input_side_packet: "side_packet1"
output_stream: "AT_PRESTREAM:packet0"
}
)");
CalculatorGraph graph;
auto status = graph.Initialize(graph_config);
EXPECT_FALSE(status.ok());
EXPECT_PRED2(
absl::StrContains, status.message(),
"Same number of input side packets and output streams is required.");
}
void DoTestNonAtTickOutputTag(absl::string_view tag,
Timestamp expected_timestamp) {
CalculatorGraphConfig graph_config =
ParseTextProtoOrDie<CalculatorGraphConfig>(absl::StrReplaceAll(
R"(
input_side_packet: "side_packet"
output_stream: "packet"
node {
calculator: "SidePacketToStreamCalculator"
input_side_packet: "side_packet"
output_stream: "$tag:packet"
}
)",
{{"$tag", tag}}));
CalculatorGraph graph;
MP_ASSERT_OK(graph.Initialize(graph_config));
const int expected_value = 10;
std::vector<Packet> output_packets;
MP_ASSERT_OK(graph.ObserveOutputStream(
"packet", [&output_packets](const Packet& packet) {
output_packets.push_back(packet);
return ::mediapipe::OkStatus();
}));
MP_ASSERT_OK(
graph.StartRun({{"side_packet", MakePacket<int>(expected_value)}}));
MP_ASSERT_OK(graph.WaitForObservedOutput());
ASSERT_FALSE(output_packets.empty());
EXPECT_EQ(expected_timestamp, output_packets.back().Timestamp());
EXPECT_EQ(expected_value, output_packets.back().Get<int>());
}
TEST(SidePacketToStreamCalculator, NoAtTickOutputTags) {
DoTestNonAtTickOutputTag("AT_PRESTREAM", Timestamp::PreStream());
DoTestNonAtTickOutputTag("AT_POSTSTREAM", Timestamp::PostStream());
DoTestNonAtTickOutputTag("AT_ZERO", Timestamp(0));
}
TEST(SidePacketToStreamCalculator, AtTick) {
CalculatorGraphConfig graph_config =
ParseTextProtoOrDie<CalculatorGraphConfig>(
R"(
input_stream: "tick"
input_side_packet: "side_packet"
output_stream: "packet"
node {
calculator: "SidePacketToStreamCalculator"
input_stream: "TICK:tick"
input_side_packet: "side_packet"
output_stream: "AT_TICK:packet"
}
)");
std::vector<Packet> output_packets;
tool::AddVectorSink("packet", &graph_config, &output_packets);
CalculatorGraph graph;
MP_ASSERT_OK(graph.Initialize(graph_config));
const int expected_value = 20;
MP_ASSERT_OK(
graph.StartRun({{"side_packet", MakePacket<int>(expected_value)}}));
auto tick_and_verify = [&graph, &output_packets,
expected_value](int at_timestamp) {
MP_ASSERT_OK(graph.AddPacketToInputStream(
"tick",
MakePacket<int>(/*doesn't matter*/ 1).At(Timestamp(at_timestamp))));
MP_ASSERT_OK(graph.WaitUntilIdle());
ASSERT_FALSE(output_packets.empty());
EXPECT_EQ(Timestamp(at_timestamp), output_packets.back().Timestamp());
EXPECT_EQ(expected_value, output_packets.back().Get<int>());
};
tick_and_verify(/*at_timestamp=*/0);
tick_and_verify(/*at_timestamp=*/1);
tick_and_verify(/*at_timestamp=*/128);
tick_and_verify(/*at_timestamp=*/1024);
tick_and_verify(/*at_timestamp=*/1025);
}
TEST(SidePacketToStreamCalculator, AtTick_MultipleSidePackets) {
CalculatorGraphConfig graph_config =
ParseTextProtoOrDie<CalculatorGraphConfig>(
R"(
input_stream: "tick"
input_side_packet: "side_packet0"
input_side_packet: "side_packet1"
output_stream: "packet0"
output_stream: "packet1"
node {
calculator: "SidePacketToStreamCalculator"
input_stream: "TICK:tick"
input_side_packet: "side_packet0"
input_side_packet: "side_packet1"
output_stream: "AT_TICK:0:packet0"
output_stream: "AT_TICK:1:packet1"
}
)");
std::vector<Packet> output_packets0;
tool::AddVectorSink("packet0", &graph_config, &output_packets0);
std::vector<Packet> output_packets1;
tool::AddVectorSink("packet1", &graph_config, &output_packets1);
CalculatorGraph graph;
MP_ASSERT_OK(graph.Initialize(graph_config));
const int expected_value0 = 20;
const int expected_value1 = 128;
MP_ASSERT_OK(
graph.StartRun({{"side_packet0", MakePacket<int>(expected_value0)},
{"side_packet1", MakePacket<int>(expected_value1)}}));
auto tick_and_verify = [&graph, &output_packets0, &output_packets1,
expected_value0, expected_value1](int at_timestamp) {
MP_ASSERT_OK(graph.AddPacketToInputStream(
"tick",
MakePacket<int>(/*doesn't matter*/ 1).At(Timestamp(at_timestamp))));
MP_ASSERT_OK(graph.WaitUntilIdle());
ASSERT_FALSE(output_packets0.empty());
ASSERT_FALSE(output_packets1.empty());
EXPECT_EQ(Timestamp(at_timestamp), output_packets0.back().Timestamp());
EXPECT_EQ(expected_value0, output_packets0.back().Get<int>());
EXPECT_EQ(Timestamp(at_timestamp), output_packets1.back().Timestamp());
EXPECT_EQ(expected_value1, output_packets1.back().Get<int>());
};
tick_and_verify(/*at_timestamp=*/0);
tick_and_verify(/*at_timestamp=*/1);
tick_and_verify(/*at_timestamp=*/128);
tick_and_verify(/*at_timestamp=*/1024);
tick_and_verify(/*at_timestamp=*/1025);
}
} // namespace
} // namespace mediapipe
@@ -17,8 +17,13 @@
#include <vector> #include <vector>
#include "mediapipe/framework/formats/landmark.pb.h" #include "mediapipe/framework/formats/landmark.pb.h"
#include "mediapipe/framework/formats/rect.pb.h"
#include "tensorflow/lite/interpreter.h" #include "tensorflow/lite/interpreter.h"
#if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
#include "tensorflow/lite/delegates/gpu/gl/gl_buffer.h"
#endif // !MEDIAPIPE_DISABLE_GPU
namespace mediapipe { namespace mediapipe {
// Example config: // Example config:
@@ -35,10 +40,21 @@ namespace mediapipe {
// } // }
// } // }
// } // }
typedef SplitVectorCalculator<TfLiteTensor> SplitTfLiteTensorVectorCalculator; typedef SplitVectorCalculator<TfLiteTensor, false>
SplitTfLiteTensorVectorCalculator;
REGISTER_CALCULATOR(SplitTfLiteTensorVectorCalculator); REGISTER_CALCULATOR(SplitTfLiteTensorVectorCalculator);
typedef SplitVectorCalculator<::mediapipe::NormalizedLandmark> typedef SplitVectorCalculator<::mediapipe::NormalizedLandmark, false>
SplitLandmarkVectorCalculator; SplitLandmarkVectorCalculator;
REGISTER_CALCULATOR(SplitLandmarkVectorCalculator); REGISTER_CALCULATOR(SplitLandmarkVectorCalculator);
typedef SplitVectorCalculator<::mediapipe::NormalizedRect, false>
SplitNormalizedRectVectorCalculator;
REGISTER_CALCULATOR(SplitNormalizedRectVectorCalculator);
#if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
typedef SplitVectorCalculator<::tflite::gpu::gl::GlBuffer, true>
MovableSplitGlBufferVectorCalculator;
REGISTER_CALCULATOR(MovableSplitGlBufferVectorCalculator);
#endif
} // namespace mediapipe } // namespace mediapipe
@@ -15,12 +15,14 @@
#ifndef MEDIAPIPE_CALCULATORS_CORE_SPLIT_VECTOR_CALCULATOR_H_ #ifndef MEDIAPIPE_CALCULATORS_CORE_SPLIT_VECTOR_CALCULATOR_H_
#define MEDIAPIPE_CALCULATORS_CORE_SPLIT_VECTOR_CALCULATOR_H_ #define MEDIAPIPE_CALCULATORS_CORE_SPLIT_VECTOR_CALCULATOR_H_
#include <type_traits>
#include <vector> #include <vector>
#include "mediapipe/calculators/core/split_vector_calculator.pb.h" #include "mediapipe/calculators/core/split_vector_calculator.pb.h"
#include "mediapipe/framework/calculator_framework.h" #include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/port/canonical_errors.h" #include "mediapipe/framework/port/canonical_errors.h"
#include "mediapipe/framework/port/ret_check.h" #include "mediapipe/framework/port/ret_check.h"
#include "mediapipe/framework/port/status.h"
#include "mediapipe/util/resource_util.h" #include "mediapipe/util/resource_util.h"
#include "tensorflow/lite/error_reporter.h" #include "tensorflow/lite/error_reporter.h"
#include "tensorflow/lite/interpreter.h" #include "tensorflow/lite/interpreter.h"
@@ -29,6 +31,20 @@
namespace mediapipe { namespace mediapipe {
template <typename T>
using IsCopyable = std::enable_if_t<std::is_copy_constructible<T>::value, bool>;
template <typename T>
using IsNotCopyable =
std::enable_if_t<!std::is_copy_constructible<T>::value, bool>;
template <typename T>
using IsMovable = std::enable_if_t<std::is_move_constructible<T>::value, bool>;
template <typename T>
using IsNotMovable =
std::enable_if_t<!std::is_move_constructible<T>::value, bool>;
// Splits an input packet with std::vector<T> into multiple std::vector<T> // Splits an input packet with std::vector<T> into multiple std::vector<T>
// output packets using the [begin, end) ranges specified in // output packets using the [begin, end) ranges specified in
// SplitVectorCalculatorOptions. If the option "element_only" is set to true, // SplitVectorCalculatorOptions. If the option "element_only" is set to true,
@@ -39,7 +55,7 @@ namespace mediapipe {
// combined into one vector. // combined into one vector.
// To use this class for a particular type T, register a calculator using // To use this class for a particular type T, register a calculator using
// SplitVectorCalculator<T>. // SplitVectorCalculator<T>.
template <typename T> template <typename T, bool move_elements>
class SplitVectorCalculator : public CalculatorBase { class SplitVectorCalculator : public CalculatorBase {
public: public:
static ::mediapipe::Status GetContract(CalculatorContract* cc) { static ::mediapipe::Status GetContract(CalculatorContract* cc) {
@@ -51,23 +67,16 @@ class SplitVectorCalculator : public CalculatorBase {
const auto& options = const auto& options =
cc->Options<::mediapipe::SplitVectorCalculatorOptions>(); cc->Options<::mediapipe::SplitVectorCalculatorOptions>();
if (!std::is_copy_constructible<T>::value || move_elements) {
// Ranges of elements shouldn't overlap when the vector contains
// non-copyable elements.
RET_CHECK_OK(checkRangesDontOverlap(options));
}
if (options.combine_outputs()) { if (options.combine_outputs()) {
RET_CHECK_EQ(cc->Outputs().NumEntries(), 1); RET_CHECK_EQ(cc->Outputs().NumEntries(), 1);
cc->Outputs().Index(0).Set<std::vector<T>>(); cc->Outputs().Index(0).Set<std::vector<T>>();
for (int i = 0; i < options.ranges_size() - 1; ++i) { RET_CHECK_OK(checkRangesDontOverlap(options));
for (int j = i + 1; j < options.ranges_size(); ++j) {
const auto& range_0 = options.ranges(i);
const auto& range_1 = options.ranges(j);
if ((range_0.begin() >= range_1.begin() &&
range_0.begin() < range_1.end()) ||
(range_1.begin() >= range_0.begin() &&
range_1.begin() < range_0.end())) {
return ::mediapipe::InvalidArgumentError(
"Ranges must be non-overlapping when using combine_outputs "
"option.");
}
}
}
} else { } else {
if (cc->Outputs().NumEntries() != options.ranges_size()) { if (cc->Outputs().NumEntries() != options.ranges_size()) {
return ::mediapipe::InvalidArgumentError( return ::mediapipe::InvalidArgumentError(
@@ -117,14 +126,26 @@ class SplitVectorCalculator : public CalculatorBase {
} }
::mediapipe::Status Process(CalculatorContext* cc) override { ::mediapipe::Status Process(CalculatorContext* cc) override {
const auto& input = cc->Inputs().Index(0).Get<std::vector<T>>(); if (cc->Inputs().Index(0).IsEmpty()) return ::mediapipe::OkStatus();
RET_CHECK_GE(input.size(), max_range_end_);
if (move_elements) {
return ProcessMovableElements<T>(cc);
} else {
return ProcessCopyableElements<T>(cc);
}
}
template <typename U, IsCopyable<U> = true>
::mediapipe::Status ProcessCopyableElements(CalculatorContext* cc) {
// static_assert(std::is_copy_constructible<U>::value,
// "Cannot copy non-copyable elements");
const auto& input = cc->Inputs().Index(0).Get<std::vector<U>>();
RET_CHECK_GE(input.size(), max_range_end_);
if (combine_outputs_) { if (combine_outputs_) {
auto output = absl::make_unique<std::vector<T>>(); auto output = absl::make_unique<std::vector<U>>();
output->reserve(total_elements_); output->reserve(total_elements_);
for (int i = 0; i < ranges_.size(); ++i) { for (int i = 0; i < ranges_.size(); ++i) {
auto elements = absl::make_unique<std::vector<T>>( auto elements = absl::make_unique<std::vector<U>>(
input.begin() + ranges_[i].first, input.begin() + ranges_[i].first,
input.begin() + ranges_[i].second); input.begin() + ranges_[i].second);
output->insert(output->end(), elements->begin(), elements->end()); output->insert(output->end(), elements->begin(), elements->end());
@@ -134,7 +155,7 @@ class SplitVectorCalculator : public CalculatorBase {
if (element_only_) { if (element_only_) {
for (int i = 0; i < ranges_.size(); ++i) { for (int i = 0; i < ranges_.size(); ++i) {
cc->Outputs().Index(i).AddPacket( cc->Outputs().Index(i).AddPacket(
MakePacket<T>(input[ranges_[i].first]).At(cc->InputTimestamp())); MakePacket<U>(input[ranges_[i].first]).At(cc->InputTimestamp()));
} }
} else { } else {
for (int i = 0; i < ranges_.size(); ++i) { for (int i = 0; i < ranges_.size(); ++i) {
@@ -149,7 +170,78 @@ class SplitVectorCalculator : public CalculatorBase {
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
template <typename U, IsNotCopyable<U> = true>
::mediapipe::Status ProcessCopyableElements(CalculatorContext* cc) {
return ::mediapipe::InternalError("Cannot copy non-copyable elements.");
}
template <typename U, IsMovable<U> = true>
::mediapipe::Status ProcessMovableElements(CalculatorContext* cc) {
::mediapipe::StatusOr<std::unique_ptr<std::vector<U>>> input_status =
cc->Inputs().Index(0).Value().Consume<std::vector<U>>();
if (!input_status.ok()) return input_status.status();
std::unique_ptr<std::vector<U>> input_vector =
std::move(input_status).ValueOrDie();
RET_CHECK_GE(input_vector->size(), max_range_end_);
if (combine_outputs_) {
auto output = absl::make_unique<std::vector<U>>();
output->reserve(total_elements_);
for (int i = 0; i < ranges_.size(); ++i) {
output->insert(
output->end(),
std::make_move_iterator(input_vector->begin() + ranges_[i].first),
std::make_move_iterator(input_vector->begin() + ranges_[i].second));
}
cc->Outputs().Index(0).Add(output.release(), cc->InputTimestamp());
} else {
if (element_only_) {
for (int i = 0; i < ranges_.size(); ++i) {
cc->Outputs().Index(i).AddPacket(
MakePacket<U>(std::move(input_vector->at(ranges_[i].first)))
.At(cc->InputTimestamp()));
}
} else {
for (int i = 0; i < ranges_.size(); ++i) {
auto output = absl::make_unique<std::vector<T>>();
output->insert(
output->end(),
std::make_move_iterator(input_vector->begin() + ranges_[i].first),
std::make_move_iterator(input_vector->begin() +
ranges_[i].second));
cc->Outputs().Index(i).Add(output.release(), cc->InputTimestamp());
}
}
}
return ::mediapipe::OkStatus();
}
template <typename U, IsNotMovable<U> = true>
::mediapipe::Status ProcessMovableElements(CalculatorContext* cc) {
return ::mediapipe::InternalError("Cannot move non-movable elements.");
}
private: private:
static ::mediapipe::Status checkRangesDontOverlap(
const ::mediapipe::SplitVectorCalculatorOptions& options) {
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);
const auto& range_1 = options.ranges(j);
if ((range_0.begin() >= range_1.begin() &&
range_0.begin() < range_1.end()) ||
(range_1.begin() >= range_0.begin() &&
range_1.begin() < range_0.end())) {
return ::mediapipe::InvalidArgumentError(
"Ranges must be non-overlapping when using combine_outputs "
"option.");
}
}
}
return ::mediapipe::OkStatus();
}
std::vector<std::pair<int32, int32>> ranges_; std::vector<std::pair<int32, int32>> ranges_;
int32 max_range_end_ = -1; int32 max_range_end_ = -1;
int32 total_elements_ = 0; int32 total_elements_ = 0;
@@ -452,4 +452,243 @@ TEST_F(SplitTfLiteTensorVectorCalculatorTest,
ASSERT_FALSE(graph.Initialize(graph_config).ok()); ASSERT_FALSE(graph.Initialize(graph_config).ok());
} }
typedef SplitVectorCalculator<std::unique_ptr<int>, true>
MovableSplitUniqueIntPtrCalculator;
REGISTER_CALCULATOR(MovableSplitUniqueIntPtrCalculator);
class MovableSplitUniqueIntPtrCalculatorTest : public ::testing::Test {
protected:
void ValidateVectorOutput(std::vector<Packet>& output_packets,
int expected_elements, int input_begin_index) {
ASSERT_EQ(1, output_packets.size());
const std::vector<std::unique_ptr<int>>& output_vec =
output_packets[0].Get<std::vector<std::unique_ptr<int>>>();
ASSERT_EQ(expected_elements, output_vec.size());
for (int i = 0; i < expected_elements; ++i) {
const int expected_value = input_begin_index + i;
const std::unique_ptr<int>& result = output_vec[i];
ASSERT_NE(result, nullptr);
ASSERT_EQ(expected_value, *result);
}
}
void ValidateElementOutput(std::vector<Packet>& output_packets,
int expected_value) {
ASSERT_EQ(1, output_packets.size());
const std::unique_ptr<int>& result =
output_packets[0].Get<std::unique_ptr<int>>();
ASSERT_NE(result, nullptr);
ASSERT_EQ(expected_value, *result);
}
void ValidateCombinedVectorOutput(std::vector<Packet>& output_packets,
int expected_elements,
std::vector<int>& input_begin_indices,
std::vector<int>& input_end_indices) {
ASSERT_EQ(1, output_packets.size());
ASSERT_EQ(input_begin_indices.size(), input_end_indices.size());
const std::vector<std::unique_ptr<int>>& output_vector =
output_packets[0].Get<std::vector<std::unique_ptr<int>>>();
ASSERT_EQ(expected_elements, output_vector.size());
const int num_ranges = input_begin_indices.size();
int element_id = 0;
for (int range_id = 0; range_id < num_ranges; ++range_id) {
for (int i = input_begin_indices[range_id];
i < input_end_indices[range_id]; ++i) {
const int expected_value = i;
const std::unique_ptr<int>& result = output_vector[element_id];
ASSERT_NE(result, nullptr);
ASSERT_EQ(expected_value, *result);
++element_id;
}
}
}
};
TEST_F(MovableSplitUniqueIntPtrCalculatorTest, InvalidOverlappingRangesTest) {
// Prepare a graph to use the TestMovableSplitUniqueIntPtrVectorCalculator.
CalculatorGraphConfig graph_config =
::mediapipe::ParseTextProtoOrDie<CalculatorGraphConfig>(
R"(
input_stream: "input_vector"
node {
calculator: "MovableSplitUniqueIntPtrCalculator"
input_stream: "input_vector"
output_stream: "range_0"
options {
[mediapipe.SplitVectorCalculatorOptions.ext] {
ranges: { begin: 0 end: 3 }
ranges: { begin: 1 end: 4 }
}
}
}
)");
// Run the graph.
CalculatorGraph graph;
// The graph should fail running because there are overlapping ranges.
ASSERT_FALSE(graph.Initialize(graph_config).ok());
}
TEST_F(MovableSplitUniqueIntPtrCalculatorTest, SmokeTest) {
// Prepare a graph to use the TestMovableSplitUniqueIntPtrVectorCalculator.
CalculatorGraphConfig graph_config =
::mediapipe::ParseTextProtoOrDie<CalculatorGraphConfig>(
R"(
input_stream: "input_vector"
node {
calculator: "MovableSplitUniqueIntPtrCalculator"
input_stream: "input_vector"
output_stream: "range_0"
output_stream: "range_1"
output_stream: "range_2"
options {
[mediapipe.SplitVectorCalculatorOptions.ext] {
ranges: { begin: 0 end: 1 }
ranges: { begin: 1 end: 4 }
ranges: { begin: 4 end: 5 }
}
}
}
)");
std::vector<Packet> range_0_packets;
tool::AddVectorSink("range_0", &graph_config, &range_0_packets);
std::vector<Packet> range_1_packets;
tool::AddVectorSink("range_1", &graph_config, &range_1_packets);
std::vector<Packet> range_2_packets;
tool::AddVectorSink("range_2", &graph_config, &range_2_packets);
// Run the graph.
CalculatorGraph graph;
MP_ASSERT_OK(graph.Initialize(graph_config));
MP_ASSERT_OK(graph.StartRun({}));
// input_vector : {0, 1, 2, 3, 4, 5}
std::unique_ptr<std::vector<std::unique_ptr<int>>> input_vector =
absl::make_unique<std::vector<std::unique_ptr<int>>>(6);
for (int i = 0; i < 6; ++i) {
input_vector->at(i) = absl::make_unique<int>(i);
}
MP_ASSERT_OK(graph.AddPacketToInputStream(
"input_vector", Adopt(input_vector.release()).At(Timestamp(1))));
MP_ASSERT_OK(graph.WaitUntilIdle());
MP_ASSERT_OK(graph.CloseAllPacketSources());
MP_ASSERT_OK(graph.WaitUntilDone());
ValidateVectorOutput(range_0_packets, /*expected_elements=*/1,
/*input_begin_index=*/0);
ValidateVectorOutput(range_1_packets, /*expected_elements=*/3,
/*input_begin_index=*/1);
ValidateVectorOutput(range_2_packets, /*expected_elements=*/1,
/*input_begin_index=*/4);
}
TEST_F(MovableSplitUniqueIntPtrCalculatorTest, SmokeTestElementOnly) {
// Prepare a graph to use the TestMovableSplitUniqueIntPtrVectorCalculator.
CalculatorGraphConfig graph_config =
::mediapipe::ParseTextProtoOrDie<CalculatorGraphConfig>(
R"(
input_stream: "input_vector"
node {
calculator: "MovableSplitUniqueIntPtrCalculator"
input_stream: "input_vector"
output_stream: "range_0"
output_stream: "range_1"
output_stream: "range_2"
options {
[mediapipe.SplitVectorCalculatorOptions.ext] {
ranges: { begin: 0 end: 1 }
ranges: { begin: 2 end: 3 }
ranges: { begin: 4 end: 5 }
element_only: true
}
}
}
)");
std::vector<Packet> range_0_packets;
tool::AddVectorSink("range_0", &graph_config, &range_0_packets);
std::vector<Packet> range_1_packets;
tool::AddVectorSink("range_1", &graph_config, &range_1_packets);
std::vector<Packet> range_2_packets;
tool::AddVectorSink("range_2", &graph_config, &range_2_packets);
// Run the graph.
CalculatorGraph graph;
MP_ASSERT_OK(graph.Initialize(graph_config));
MP_ASSERT_OK(graph.StartRun({}));
// input_vector : {0, 1, 2, 3, 4, 5}
std::unique_ptr<std::vector<std::unique_ptr<int>>> input_vector =
absl::make_unique<std::vector<std::unique_ptr<int>>>(6);
for (int i = 0; i < 6; ++i) {
input_vector->at(i) = absl::make_unique<int>(i);
}
MP_ASSERT_OK(graph.AddPacketToInputStream(
"input_vector", Adopt(input_vector.release()).At(Timestamp(1))));
MP_ASSERT_OK(graph.WaitUntilIdle());
MP_ASSERT_OK(graph.CloseAllPacketSources());
MP_ASSERT_OK(graph.WaitUntilDone());
ValidateElementOutput(range_0_packets, /*expected_value=*/0);
ValidateElementOutput(range_1_packets, /*expected_value=*/2);
ValidateElementOutput(range_2_packets, /*expected_value=*/4);
}
TEST_F(MovableSplitUniqueIntPtrCalculatorTest, SmokeTestCombiningOutputs) {
// Prepare a graph to use the TestMovableSplitUniqueIntPtrVectorCalculator.
CalculatorGraphConfig graph_config =
::mediapipe::ParseTextProtoOrDie<CalculatorGraphConfig>(
R"(
input_stream: "input_vector"
node {
calculator: "MovableSplitUniqueIntPtrCalculator"
input_stream: "input_vector"
output_stream: "range_0"
options {
[mediapipe.SplitVectorCalculatorOptions.ext] {
ranges: { begin: 0 end: 1 }
ranges: { begin: 2 end: 3 }
ranges: { begin: 4 end: 5 }
combine_outputs: true
}
}
}
)");
std::vector<Packet> range_0_packets;
tool::AddVectorSink("range_0", &graph_config, &range_0_packets);
// Run the graph.
CalculatorGraph graph;
MP_ASSERT_OK(graph.Initialize(graph_config));
MP_ASSERT_OK(graph.StartRun({}));
// input_vector : {0, 1, 2, 3, 4, 5}
std::unique_ptr<std::vector<std::unique_ptr<int>>> input_vector =
absl::make_unique<std::vector<std::unique_ptr<int>>>(6);
for (int i = 0; i < 6; ++i) {
input_vector->at(i) = absl::make_unique<int>(i);
}
MP_ASSERT_OK(graph.AddPacketToInputStream(
"input_vector", Adopt(input_vector.release()).At(Timestamp(1))));
MP_ASSERT_OK(graph.WaitUntilIdle());
MP_ASSERT_OK(graph.CloseAllPacketSources());
MP_ASSERT_OK(graph.WaitUntilDone());
std::vector<int> input_begin_indices = {0, 2, 4};
std::vector<int> input_end_indices = {1, 3, 5};
ValidateCombinedVectorOutput(range_0_packets, /*expected_elements=*/3,
input_begin_indices, input_end_indices);
}
} // namespace mediapipe } // namespace mediapipe
+22 -5
View File
@@ -80,7 +80,9 @@ mediapipe_cc_proto_library(
name = "opencv_image_encoder_calculator_cc_proto", name = "opencv_image_encoder_calculator_cc_proto",
srcs = ["opencv_image_encoder_calculator.proto"], srcs = ["opencv_image_encoder_calculator.proto"],
cc_deps = ["//mediapipe/framework:calculator_cc_proto"], cc_deps = ["//mediapipe/framework:calculator_cc_proto"],
visibility = ["//visibility:public"], visibility = [
"//visibility:public",
],
deps = [":opencv_image_encoder_calculator_proto"], deps = [":opencv_image_encoder_calculator_proto"],
) )
@@ -330,6 +332,7 @@ cc_library(
cc_library( cc_library(
name = "image_cropping_calculator", name = "image_cropping_calculator",
srcs = ["image_cropping_calculator.cc"], srcs = ["image_cropping_calculator.cc"],
hdrs = ["image_cropping_calculator.h"],
copts = select({ copts = select({
"//mediapipe:apple": [ "//mediapipe:apple": [
"-x objective-c++", "-x objective-c++",
@@ -343,9 +346,7 @@ cc_library(
], ],
"//conditions:default": [], "//conditions:default": [],
}), }),
visibility = [ visibility = ["//visibility:public"],
"//visibility:public",
],
deps = [ deps = [
":image_cropping_calculator_cc_proto", ":image_cropping_calculator_cc_proto",
"//mediapipe/framework:calculator_framework", "//mediapipe/framework:calculator_framework",
@@ -356,19 +357,35 @@ cc_library(
"//mediapipe/framework/port:opencv_imgproc", "//mediapipe/framework/port:opencv_imgproc",
"//mediapipe/framework/port:ret_check", "//mediapipe/framework/port:ret_check",
"//mediapipe/framework/port:status", "//mediapipe/framework/port:status",
"//mediapipe/gpu:gpu_buffer",
] + select({ ] + select({
"//mediapipe/gpu:disable_gpu": [], "//mediapipe/gpu:disable_gpu": [],
"//conditions:default": [ "//conditions:default": [
"//mediapipe/gpu:gl_calculator_helper", "//mediapipe/gpu:gl_calculator_helper",
"//mediapipe/gpu:gl_simple_shaders", "//mediapipe/gpu:gl_simple_shaders",
"//mediapipe/gpu:gl_quad_renderer", "//mediapipe/gpu:gl_quad_renderer",
"//mediapipe/gpu:gpu_buffer",
"//mediapipe/gpu:shader_util", "//mediapipe/gpu:shader_util",
], ],
}), }),
alwayslink = 1, alwayslink = 1,
) )
cc_test(
name = "image_cropping_calculator_test",
srcs = ["image_cropping_calculator_test.cc"],
deps = [
":image_cropping_calculator",
":image_cropping_calculator_cc_proto",
"//mediapipe/framework:calculator_framework",
"//mediapipe/framework/formats:rect_cc_proto",
"//mediapipe/framework/port:gtest_main",
"//mediapipe/framework/port:parse_text_proto",
"//mediapipe/framework/port:status",
"//mediapipe/framework/tool:tag_map",
"//mediapipe/framework/tool:tag_map_helper",
],
)
cc_library( cc_library(
name = "luminance_calculator", name = "luminance_calculator",
srcs = ["luminance_calculator.cc"], srcs = ["luminance_calculator.cc"],
@@ -75,6 +75,11 @@ class ColorConvertCalculator : public CalculatorBase {
static ::mediapipe::Status GetContract(CalculatorContract* cc); static ::mediapipe::Status GetContract(CalculatorContract* cc);
::mediapipe::Status Process(CalculatorContext* cc) override; ::mediapipe::Status Process(CalculatorContext* cc) override;
::mediapipe::Status Open(CalculatorContext* cc) override {
cc->SetOffset(TimestampDiff(0));
return ::mediapipe::OkStatus();
}
private: private:
// Wrangles the appropriate inputs and outputs to perform the color // Wrangles the appropriate inputs and outputs to perform the color
// conversion. The ImageFrame on input_tag is converted using the // conversion. The ImageFrame on input_tag is converted using the
@@ -12,10 +12,10 @@
// See the License for the specific language governing permissions and // See the License for the specific language governing permissions and
// limitations under the License. // limitations under the License.
#include "mediapipe/calculators/image/image_cropping_calculator.h"
#include <cmath> #include <cmath>
#include "mediapipe/calculators/image/image_cropping_calculator.pb.h"
#include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/formats/image_frame.h" #include "mediapipe/framework/formats/image_frame.h"
#include "mediapipe/framework/formats/image_frame_opencv.h" #include "mediapipe/framework/formats/image_frame_opencv.h"
#include "mediapipe/framework/formats/rect.pb.h" #include "mediapipe/framework/formats/rect.pb.h"
@@ -25,7 +25,6 @@
#include "mediapipe/framework/port/status.h" #include "mediapipe/framework/port/status.h"
#if !defined(MEDIAPIPE_DISABLE_GPU) #if !defined(MEDIAPIPE_DISABLE_GPU)
#include "mediapipe/gpu/gl_calculator_helper.h"
#include "mediapipe/gpu/gl_simple_shaders.h" #include "mediapipe/gpu/gl_simple_shaders.h"
#include "mediapipe/gpu/gpu_buffer.h" #include "mediapipe/gpu/gpu_buffer.h"
#include "mediapipe/gpu/shader_util.h" #include "mediapipe/gpu/shader_util.h"
@@ -52,62 +51,6 @@ constexpr char kWidthTag[] = "WIDTH";
} // namespace } // namespace
// Crops the input texture to the given rectangle region. The rectangle can
// be at arbitrary location on the image with rotation. If there's rotation, the
// output texture will have the size of the input rectangle. The rotation should
// be in radian, see rect.proto for detail.
//
// Input:
// One of the following two tags:
// IMAGE - ImageFrame representing the input image.
// IMAGE_GPU - GpuBuffer representing the input image.
// One of the following two tags (optional if WIDTH/HEIGHT is specified):
// RECT - A Rect proto specifying the width/height and location of the
// cropping rectangle.
// NORM_RECT - A NormalizedRect proto specifying the width/height and location
// of the cropping rectangle in normalized coordinates.
// Alternative tags to RECT (optional if RECT/NORM_RECT is specified):
// WIDTH - The desired width of the output cropped image,
// based on image center
// HEIGHT - The desired height of the output cropped image,
// based on image center
//
// Output:
// One of the following two tags:
// IMAGE - Cropped ImageFrame
// IMAGE_GPU - Cropped GpuBuffer.
//
// Note: input_stream values take precedence over options defined in the graph.
//
class ImageCroppingCalculator : public CalculatorBase {
public:
ImageCroppingCalculator() = default;
~ImageCroppingCalculator() override = default;
static ::mediapipe::Status GetContract(CalculatorContract* cc);
::mediapipe::Status Open(CalculatorContext* cc) override;
::mediapipe::Status Process(CalculatorContext* cc) override;
::mediapipe::Status Close(CalculatorContext* cc) override;
private:
::mediapipe::Status RenderCpu(CalculatorContext* cc);
::mediapipe::Status RenderGpu(CalculatorContext* cc);
::mediapipe::Status InitGpu(CalculatorContext* cc);
void GlRender();
void GetOutputDimensions(CalculatorContext* cc, int src_width, int src_height,
int* dst_width, int* dst_height);
mediapipe::ImageCroppingCalculatorOptions options_;
bool use_gpu_ = false;
// Output texture corners (4) after transoformation in normalized coordinates.
float transformed_points_[8];
#if !defined(MEDIAPIPE_DISABLE_GPU)
bool gpu_initialized_ = false;
mediapipe::GlCalculatorHelper gpu_helper_;
GLuint program_ = 0;
#endif // !MEDIAPIPE_DISABLE_GPU
};
REGISTER_CALCULATOR(ImageCroppingCalculator); REGISTER_CALCULATOR(ImageCroppingCalculator);
::mediapipe::Status ImageCroppingCalculator::GetContract( ::mediapipe::Status ImageCroppingCalculator::GetContract(
@@ -132,7 +75,28 @@ REGISTER_CALCULATOR(ImageCroppingCalculator);
} }
#endif // !MEDIAPIPE_DISABLE_GPU #endif // !MEDIAPIPE_DISABLE_GPU
RET_CHECK(cc->Inputs().HasTag(kRectTag) ^ cc->Inputs().HasTag(kNormRectTag)); int flags = 0;
if (cc->Inputs().HasTag(kRectTag)) {
++flags;
}
if (cc->Inputs().HasTag(kWidthTag) && cc->Inputs().HasTag(kHeightTag)) {
++flags;
}
if (cc->Inputs().HasTag(kNormRectTag)) {
++flags;
}
if (cc->Options<mediapipe::ImageCroppingCalculatorOptions>()
.has_norm_width() &&
cc->Options<mediapipe::ImageCroppingCalculatorOptions>()
.has_norm_height()) {
++flags;
}
if (cc->Options<mediapipe::ImageCroppingCalculatorOptions>().has_width() &&
cc->Options<mediapipe::ImageCroppingCalculatorOptions>().has_height()) {
++flags;
}
RET_CHECK(flags == 1) << "Illegal combination of input streams/options.";
if (cc->Inputs().HasTag(kRectTag)) { if (cc->Inputs().HasTag(kRectTag)) {
cc->Inputs().Tag(kRectTag).Set<Rect>(); cc->Inputs().Tag(kRectTag).Set<Rect>();
} }
@@ -172,6 +136,13 @@ REGISTER_CALCULATOR(ImageCroppingCalculator);
#endif // !MEDIAPIPE_DISABLE_GPU #endif // !MEDIAPIPE_DISABLE_GPU
} }
// Validate border mode.
if (use_gpu_) {
MP_RETURN_IF_ERROR(ValidateBorderModeForGPU(cc));
} else {
MP_RETURN_IF_ERROR(ValidateBorderModeForCPU(cc));
}
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -215,6 +186,32 @@ REGISTER_CALCULATOR(ImageCroppingCalculator);
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
::mediapipe::Status ImageCroppingCalculator::ValidateBorderModeForCPU(
CalculatorContext* cc) {
int border_mode;
return GetBorderModeForOpenCV(cc, &border_mode);
}
::mediapipe::Status ImageCroppingCalculator::ValidateBorderModeForGPU(
CalculatorContext* cc) {
mediapipe::ImageCroppingCalculatorOptions options =
cc->Options<mediapipe::ImageCroppingCalculatorOptions>();
switch (options.border_mode()) {
case mediapipe::ImageCroppingCalculatorOptions::BORDER_ZERO:
LOG(WARNING) << "BORDER_ZERO mode is not supported by GPU "
<< "implementation and will fall back into BORDER_REPLICATE";
break;
case mediapipe::ImageCroppingCalculatorOptions::BORDER_REPLICATE:
break;
default:
RET_CHECK_FAIL() << "Unsupported border mode for GPU: "
<< options.border_mode();
}
return ::mediapipe::OkStatus();
}
::mediapipe::Status ImageCroppingCalculator::RenderCpu(CalculatorContext* cc) { ::mediapipe::Status ImageCroppingCalculator::RenderCpu(CalculatorContext* cc) {
if (cc->Inputs().Tag(kImageTag).IsEmpty()) { if (cc->Inputs().Tag(kImageTag).IsEmpty()) {
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
@@ -222,41 +219,12 @@ REGISTER_CALCULATOR(ImageCroppingCalculator);
const auto& input_img = cc->Inputs().Tag(kImageTag).Get<ImageFrame>(); const auto& input_img = cc->Inputs().Tag(kImageTag).Get<ImageFrame>();
cv::Mat input_mat = formats::MatView(&input_img); cv::Mat input_mat = formats::MatView(&input_img);
float rect_center_x = input_img.Width() / 2.0f; auto [target_width, target_height, rect_center_x, rect_center_y, rotation] =
float rect_center_y = input_img.Height() / 2.0f; GetCropSpecs(cc, input_img.Width(), input_img.Height());
float rotation = 0.0f;
int target_width = input_img.Width(); // Get border mode and value for OpenCV.
int target_height = input_img.Height(); int border_mode;
if (cc->Inputs().HasTag(kRectTag)) { MP_RETURN_IF_ERROR(GetBorderModeForOpenCV(cc, &border_mode));
const auto& rect = cc->Inputs().Tag(kRectTag).Get<Rect>();
if (rect.width() > 0 && rect.height() > 0 && rect.x_center() >= 0 &&
rect.y_center() >= 0) {
rect_center_x = rect.x_center();
rect_center_y = rect.y_center();
target_width = rect.width();
target_height = rect.height();
rotation = rect.rotation();
}
} else if (cc->Inputs().HasTag(kNormRectTag)) {
const auto& rect = cc->Inputs().Tag(kNormRectTag).Get<NormalizedRect>();
if (rect.width() > 0.0 && rect.height() > 0.0 && rect.x_center() >= 0.0 &&
rect.y_center() >= 0.0) {
rect_center_x = std::round(rect.x_center() * input_img.Width());
rect_center_y = std::round(rect.y_center() * input_img.Height());
target_width = std::round(rect.width() * input_img.Width());
target_height = std::round(rect.height() * input_img.Height());
rotation = rect.rotation();
}
} else {
if (cc->Inputs().HasTag(kWidthTag) && cc->Inputs().HasTag(kHeightTag)) {
target_width = cc->Inputs().Tag(kWidthTag).Get<int>();
target_height = cc->Inputs().Tag(kHeightTag).Get<int>();
} else if (options_.has_width() && options_.has_height()) {
target_width = options_.width();
target_height = options_.height();
}
rotation = options_.rotation();
}
const cv::RotatedRect min_rect(cv::Point2f(rect_center_x, rect_center_y), const cv::RotatedRect min_rect(cv::Point2f(rect_center_x, rect_center_y),
cv::Size2f(target_width, target_height), cv::Size2f(target_width, target_height),
@@ -277,7 +245,9 @@ REGISTER_CALCULATOR(ImageCroppingCalculator);
cv::getPerspectiveTransform(src_points, dst_points); cv::getPerspectiveTransform(src_points, dst_points);
cv::Mat cropped_image; cv::Mat cropped_image;
cv::warpPerspective(input_mat, cropped_image, projection_matrix, cv::warpPerspective(input_mat, cropped_image, projection_matrix,
cv::Size(min_rect.size.width, min_rect.size.height)); cv::Size(min_rect.size.width, min_rect.size.height),
/* flags = */ 0,
/* borderMode = */ border_mode);
std::unique_ptr<ImageFrame> output_frame(new ImageFrame( std::unique_ptr<ImageFrame> output_frame(new ImageFrame(
input_img.Format(), cropped_image.cols, cropped_image.rows)); input_img.Format(), cropped_image.cols, cropped_image.rows));
@@ -433,46 +403,8 @@ void ImageCroppingCalculator::GetOutputDimensions(CalculatorContext* cc,
int src_width, int src_height, int src_width, int src_height,
int* dst_width, int* dst_width,
int* dst_height) { int* dst_height) {
// Get the size of the cropping box. auto [crop_width, crop_height, x_center, y_center, rotation] =
int crop_width = src_width; GetCropSpecs(cc, src_width, src_height);
int crop_height = src_height;
// Get the center of cropping box. Default is the at the center.
int x_center = src_width / 2;
int y_center = src_height / 2;
// Get the rotation of the cropping box.
float rotation = 0.0f;
if (cc->Inputs().HasTag(kRectTag)) {
const auto& rect = cc->Inputs().Tag(kRectTag).Get<Rect>();
// Only use the rect if it is valid.
if (rect.width() > 0 && rect.height() > 0 && rect.x_center() >= 0 &&
rect.y_center() >= 0) {
x_center = rect.x_center();
y_center = rect.y_center();
crop_width = rect.width();
crop_height = rect.height();
rotation = rect.rotation();
}
} else if (cc->Inputs().HasTag(kNormRectTag)) {
const auto& rect = cc->Inputs().Tag(kNormRectTag).Get<NormalizedRect>();
// Only use the rect if it is valid.
if (rect.width() > 0.0 && rect.height() > 0.0 && rect.x_center() >= 0.0 &&
rect.y_center() >= 0.0) {
x_center = std::round(rect.x_center() * src_width);
y_center = std::round(rect.y_center() * src_height);
crop_width = std::round(rect.width() * src_width);
crop_height = std::round(rect.height() * src_height);
rotation = rect.rotation();
}
} else {
if (cc->Inputs().HasTag(kWidthTag) && cc->Inputs().HasTag(kHeightTag)) {
crop_width = cc->Inputs().Tag(kWidthTag).Get<int>();
crop_height = cc->Inputs().Tag(kHeightTag).Get<int>();
} else if (options_.has_width() && options_.has_height()) {
crop_width = options_.width();
crop_height = options_.height();
}
rotation = options_.rotation();
}
const float half_width = crop_width / 2.0f; const float half_width = crop_width / 2.0f;
const float half_height = crop_height / 2.0f; const float half_height = crop_height / 2.0f;
@@ -501,8 +433,110 @@ void ImageCroppingCalculator::GetOutputDimensions(CalculatorContext* cc,
row_max = std::max(row_max, transformed_points_[i * 2 + 1]); row_max = std::max(row_max, transformed_points_[i * 2 + 1]);
} }
*dst_width = std::round((col_max - col_min) * src_width); int width = static_cast<int>(std::round((col_max - col_min) * src_width));
*dst_height = std::round((row_max - row_min) * src_height); int height = static_cast<int>(std::round((row_max - row_min) * src_height));
// Minimum output dimension 1x1 prevents creation of textures with 0x0.
*dst_width = std::max(1, width);
*dst_height = std::max(1, height);
}
RectSpec ImageCroppingCalculator::GetCropSpecs(const CalculatorContext* cc,
int src_width, int src_height) {
// Get the size of the cropping box.
int crop_width = src_width;
int crop_height = src_height;
// Get the center of cropping box. Default is the at the center.
int x_center = src_width / 2;
int y_center = src_height / 2;
// Get the rotation of the cropping box.
float rotation = 0.0f;
// Get the normalized width and height if specified by the inputs or options.
float normalized_width = 0.0f;
float normalized_height = 0.0f;
mediapipe::ImageCroppingCalculatorOptions options =
cc->Options<mediapipe::ImageCroppingCalculatorOptions>();
// width/height, norm_width/norm_height from input streams take precednece.
if (cc->Inputs().HasTag(kRectTag)) {
const auto& rect = cc->Inputs().Tag(kRectTag).Get<Rect>();
// Only use the rect if it is valid.
if (rect.width() > 0 && rect.height() > 0 && rect.x_center() >= 0 &&
rect.y_center() >= 0) {
x_center = rect.x_center();
y_center = rect.y_center();
crop_width = rect.width();
crop_height = rect.height();
rotation = rect.rotation();
}
} else if (cc->Inputs().HasTag(kNormRectTag)) {
const auto& norm_rect =
cc->Inputs().Tag(kNormRectTag).Get<NormalizedRect>();
if (norm_rect.width() > 0.0 && norm_rect.height() > 0.0) {
normalized_width = norm_rect.width();
normalized_height = norm_rect.height();
x_center = std::round(norm_rect.x_center() * src_width);
y_center = std::round(norm_rect.y_center() * src_height);
rotation = norm_rect.rotation();
}
} else if (cc->Inputs().HasTag(kWidthTag) &&
cc->Inputs().HasTag(kHeightTag)) {
crop_width = cc->Inputs().Tag(kWidthTag).Get<int>();
crop_height = cc->Inputs().Tag(kHeightTag).Get<int>();
} else if (options.has_width() && options.has_height()) {
crop_width = options.width();
crop_height = options.height();
} else if (options.has_norm_width() && options.has_norm_height()) {
normalized_width = options.norm_width();
normalized_height = options.norm_height();
}
// Get the crop width and height from the normalized width and height.
if (normalized_width > 0 && normalized_height > 0) {
crop_width = std::round(normalized_width * src_width);
crop_height = std::round(normalized_height * src_height);
}
// Rotation and center values from input streams take precedence, so only
// look at those values in the options if kRectTag and kNormRectTag are not
// present from the inputs.
if (!cc->Inputs().HasTag(kRectTag) && !cc->Inputs().HasTag(kNormRectTag)) {
if (options.has_norm_center_x() && options.has_norm_center_y()) {
x_center = std::round(options.norm_center_x() * src_width);
y_center = std::round(options.norm_center_y() * src_height);
}
if (options.has_rotation()) {
rotation = options.rotation();
}
}
return {
.width = crop_width,
.height = crop_height,
.center_x = x_center,
.center_y = y_center,
.rotation = rotation,
};
}
::mediapipe::Status ImageCroppingCalculator::GetBorderModeForOpenCV(
CalculatorContext* cc, int* border_mode) {
mediapipe::ImageCroppingCalculatorOptions options =
cc->Options<mediapipe::ImageCroppingCalculatorOptions>();
switch (options.border_mode()) {
case mediapipe::ImageCroppingCalculatorOptions::BORDER_ZERO:
*border_mode = cv::BORDER_CONSTANT;
break;
case mediapipe::ImageCroppingCalculatorOptions::BORDER_REPLICATE:
*border_mode = cv::BORDER_REPLICATE;
break;
default:
RET_CHECK_FAIL() << "Unsupported border mode for CPU: "
<< options.border_mode();
}
return ::mediapipe::OkStatus();
} }
} // namespace mediapipe } // namespace mediapipe
@@ -0,0 +1,91 @@
#ifndef MEDIAPIPE_CALCULATORS_IMAGE_IMAGE_CROPPING_CALCULATOR_H_
#define MEDIAPIPE_CALCULATORS_IMAGE_IMAGE_CROPPING_CALCULATOR_H_
#include "mediapipe/calculators/image/image_cropping_calculator.pb.h"
#include "mediapipe/framework/calculator_framework.h"
#if !defined(MEDIAPIPE_DISABLE_GPU)
#include "mediapipe/gpu/gl_calculator_helper.h"
#endif // !MEDIAPIPE_DISABLE_GPU
// Crops the input texture to the given rectangle region. The rectangle can
// be at arbitrary location on the image with rotation. If there's rotation, the
// output texture will have the size of the input rectangle. The rotation should
// be in radian, see rect.proto for detail.
//
// Input:
// One of the following two tags:
// IMAGE - ImageFrame representing the input image.
// IMAGE_GPU - GpuBuffer representing the input image.
// One of the following two tags (optional if WIDTH/HEIGHT is specified):
// RECT - A Rect proto specifying the width/height and location of the
// cropping rectangle.
// NORM_RECT - A NormalizedRect proto specifying the width/height and location
// of the cropping rectangle in normalized coordinates.
// Alternative tags to RECT (optional if RECT/NORM_RECT is specified):
// WIDTH - The desired width of the output cropped image,
// based on image center
// HEIGHT - The desired height of the output cropped image,
// based on image center
//
// Output:
// One of the following two tags:
// IMAGE - Cropped ImageFrame
// IMAGE_GPU - Cropped GpuBuffer.
//
// Note: input_stream values take precedence over options defined in the graph.
//
namespace mediapipe {
struct RectSpec {
int width;
int height;
int center_x;
int center_y;
float rotation;
bool operator==(const RectSpec& rect) const {
return (width == rect.width && height == rect.height &&
center_x == rect.center_x && center_y == rect.center_y &&
rotation == rect.rotation);
}
};
class ImageCroppingCalculator : public CalculatorBase {
public:
ImageCroppingCalculator() = default;
~ImageCroppingCalculator() override = default;
static ::mediapipe::Status GetContract(CalculatorContract* cc);
::mediapipe::Status Open(CalculatorContext* cc) override;
::mediapipe::Status Process(CalculatorContext* cc) override;
::mediapipe::Status Close(CalculatorContext* cc) override;
static RectSpec GetCropSpecs(const CalculatorContext* cc, int src_width,
int src_height);
private:
::mediapipe::Status ValidateBorderModeForCPU(CalculatorContext* cc);
::mediapipe::Status ValidateBorderModeForGPU(CalculatorContext* cc);
::mediapipe::Status RenderCpu(CalculatorContext* cc);
::mediapipe::Status RenderGpu(CalculatorContext* cc);
::mediapipe::Status InitGpu(CalculatorContext* cc);
void GlRender();
void GetOutputDimensions(CalculatorContext* cc, int src_width, int src_height,
int* dst_width, int* dst_height);
::mediapipe::Status GetBorderModeForOpenCV(CalculatorContext* cc,
int* border_mode);
mediapipe::ImageCroppingCalculatorOptions options_;
bool use_gpu_ = false;
// Output texture corners (4) after transoformation in normalized coordinates.
float transformed_points_[8];
#if !defined(MEDIAPIPE_DISABLE_GPU)
bool gpu_initialized_ = false;
mediapipe::GlCalculatorHelper gpu_helper_;
GLuint program_ = 0;
#endif // !MEDIAPIPE_DISABLE_GPU
};
} // namespace mediapipe
#endif // MEDIAPIPE_CALCULATORS_IMAGE_IMAGE_CROPPING_CALCULATOR_H_
@@ -30,4 +30,25 @@ message ImageCroppingCalculatorOptions {
// Rotation angle is counter-clockwise in radian. // Rotation angle is counter-clockwise in radian.
optional float rotation = 3 [default = 0.0]; optional float rotation = 3 [default = 0.0];
// Normalized width and height of the output rect. Value is within [0, 1].
optional float norm_width = 4;
optional float norm_height = 5;
// Normalized location of the center of the output
// rectangle in image coordinates. Value is within [0, 1].
// The (0, 0) point is at the (top, left) corner.
optional float norm_center_x = 6 [default = 0];
optional float norm_center_y = 7 [default = 0];
enum BorderMode {
// First unspecified value is required by the guideline. See details here:
// https://developers.google.com/protocol-buffers/docs/style#enums
BORDER_UNSPECIFIED = 0;
BORDER_ZERO = 1;
BORDER_REPLICATE = 2;
}
// Specifies behaviour for crops that go beyond image borders.
optional BorderMode border_mode = 8 [default = BORDER_ZERO];
} }
@@ -0,0 +1,216 @@
// Copyright 2020 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/image/image_cropping_calculator.h"
#include <cmath>
#include <memory>
#include "mediapipe/calculators/image/image_cropping_calculator.pb.h"
#include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/formats/rect.pb.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/tool/tag_map.h"
#include "mediapipe/framework/tool/tag_map_helper.h"
namespace mediapipe {
namespace {
constexpr int input_width = 100;
constexpr int input_height = 100;
constexpr char kRectTag[] = "RECT";
constexpr char kHeightTag[] = "HEIGHT";
constexpr char kWidthTag[] = "WIDTH";
// Test normal case, where norm_width and norm_height in options are set.
TEST(ImageCroppingCalculatorTest, GetCroppingDimensionsNormal) {
auto calculator_node =
ParseTextProtoOrDie<mediapipe::CalculatorGraphConfig::Node>(
R"(
calculator: "ImageCroppingCalculator"
input_stream: "IMAGE_GPU:input_frames"
output_stream: "IMAGE_GPU:cropped_output_frames"
options: {
[mediapipe.ImageCroppingCalculatorOptions.ext] {
norm_width: 0.6
norm_height: 0.6
norm_center_x: 0.5
norm_center_y: 0.5
rotation: 0.3
}
}
)");
auto calculator_state = absl::make_unique<CalculatorState>(
"Node", 0, "Calculator", calculator_node, nullptr);
auto cc = absl::make_unique<CalculatorContext>(
calculator_state.get(), tool::CreateTagMap({}).ValueOrDie(),
tool::CreateTagMap({}).ValueOrDie());
RectSpec expectRect = {
.width = 60,
.height = 60,
.center_x = 50,
.center_y = 50,
.rotation = 0.3,
};
EXPECT_EQ(ImageCroppingCalculator::GetCropSpecs(cc.get(), input_width,
input_height),
expectRect);
} // TEST
// Test when (width height) + (norm_width norm_height) are set in options.
// width and height should take precedence.
TEST(ImageCroppingCalculatorTest, RedundantSpecInOptions) {
auto calculator_node =
ParseTextProtoOrDie<mediapipe::CalculatorGraphConfig::Node>(
R"(
calculator: "ImageCroppingCalculator"
input_stream: "IMAGE_GPU:input_frames"
output_stream: "IMAGE_GPU:cropped_output_frames"
options: {
[mediapipe.ImageCroppingCalculatorOptions.ext] {
width: 50
height: 50
norm_width: 0.6
norm_height: 0.6
norm_center_x: 0.5
norm_center_y: 0.5
rotation: 0.3
}
}
)");
auto calculator_state = absl::make_unique<CalculatorState>(
"Node", 0, "Calculator", calculator_node, nullptr);
auto cc = absl::make_unique<CalculatorContext>(
calculator_state.get(), tool::CreateTagMap({}).ValueOrDie(),
tool::CreateTagMap({}).ValueOrDie());
RectSpec expectRect = {
.width = 50,
.height = 50,
.center_x = 50,
.center_y = 50,
.rotation = 0.3,
};
EXPECT_EQ(ImageCroppingCalculator::GetCropSpecs(cc.get(), input_width,
input_height),
expectRect);
} // TEST
// Test when WIDTH HEIGHT are set from input stream,
// and options has norm_width/height set.
// WIDTH HEIGHT from input stream should take precedence.
TEST(ImageCroppingCalculatorTest, RedundantSpectWithInputStream) {
auto calculator_node =
ParseTextProtoOrDie<mediapipe::CalculatorGraphConfig::Node>(
R"(
calculator: "ImageCroppingCalculator"
input_stream: "IMAGE_GPU:input_frames"
input_stream: "WIDTH:crop_width"
input_stream: "HEIGHT:crop_height"
output_stream: "IMAGE_GPU:cropped_output_frames"
options: {
[mediapipe.ImageCroppingCalculatorOptions.ext] {
width: 50
height: 50
norm_width: 0.6
norm_height: 0.6
norm_center_x: 0.5
norm_center_y: 0.5
rotation: 0.3
}
}
)");
auto calculator_state = absl::make_unique<CalculatorState>(
"Node", 0, "Calculator", calculator_node, nullptr);
auto inputTags = tool::CreateTagMap({
"HEIGHT:0:crop_height",
"WIDTH:0:crop_width",
})
.ValueOrDie();
auto cc = absl::make_unique<CalculatorContext>(
calculator_state.get(), inputTags, tool::CreateTagMap({}).ValueOrDie());
auto& inputs = cc->Inputs();
inputs.Tag(kHeightTag).Value() = MakePacket<int>(1);
inputs.Tag(kWidthTag).Value() = MakePacket<int>(1);
RectSpec expectRect = {
.width = 1,
.height = 1,
.center_x = 50,
.center_y = 50,
.rotation = 0.3,
};
EXPECT_EQ(ImageCroppingCalculator::GetCropSpecs(cc.get(), input_width,
input_height),
expectRect);
} // TEST
// Test when RECT is set from input stream,
// and options has norm_width/height set.
// RECT from input stream should take precedence.
TEST(ImageCroppingCalculatorTest, RedundantSpecWithInputStream) {
auto calculator_node =
ParseTextProtoOrDie<mediapipe::CalculatorGraphConfig::Node>(
R"(
calculator: "ImageCroppingCalculator"
input_stream: "IMAGE_GPU:input_frames"
input_stream: "RECT:rect"
output_stream: "IMAGE_GPU:cropped_output_frames"
options: {
[mediapipe.ImageCroppingCalculatorOptions.ext] {
width: 50
height: 50
norm_width: 0.6
norm_height: 0.6
norm_center_x: 0.5
norm_center_y: 0.5
rotation: 0.3
}
}
)");
auto calculator_state = absl::make_unique<CalculatorState>(
"Node", 0, "Calculator", calculator_node, nullptr);
auto inputTags = tool::CreateTagMap({
"RECT:0:rect",
})
.ValueOrDie();
auto cc = absl::make_unique<CalculatorContext>(
calculator_state.get(), inputTags, tool::CreateTagMap({}).ValueOrDie());
auto& inputs = cc->Inputs();
mediapipe::Rect rect = ParseTextProtoOrDie<mediapipe::Rect>(
R"(
width: 1 height: 1 x_center: 40 y_center: 40 rotation: 0.5
)");
inputs.Tag(kRectTag).Value() = MakePacket<mediapipe::Rect>(rect);
RectSpec expectRect = {
.width = 1,
.height = 1,
.center_x = 40,
.center_y = 40,
.rotation = 0.5,
};
EXPECT_EQ(ImageCroppingCalculator::GetCropSpecs(cc.get(), input_width,
input_height),
expectRect);
} // TEST
} // namespace
} // namespace mediapipe
@@ -104,6 +104,14 @@ mediapipe::ScaleMode_Mode ParseScaleMode(
// to be a multiple of 90 degrees. If provided, it overrides the // to be a multiple of 90 degrees. If provided, it overrides the
// ROTATION_DEGREES input side packet. // ROTATION_DEGREES input side packet.
// //
// FLIP_HORIZONTALLY (optional): Whether to flip image horizontally or not. If
// provided, it overrides the FLIP_HORIZONTALLY input side packet and/or
// corresponding field in the calculator options.
//
// FLIP_VERTICALLY (optional): Whether to flip image vertically or not. If
// provided, it overrides the FLIP_VERTICALLY input side packet and/or
// corresponding field in the calculator options.
//
// Output: // Output:
// One of the following two tags: // One of the following two tags:
// IMAGE - ImageFrame representing the output image. // IMAGE - ImageFrame representing the output image.
@@ -129,6 +137,12 @@ mediapipe::ScaleMode_Mode ParseScaleMode(
// degrees. It has to be a multiple of 90 degrees. It overrides the // degrees. It has to be a multiple of 90 degrees. It overrides the
// corresponding field in the calculator options. // corresponding field in the calculator options.
// //
// FLIP_HORIZONTALLY (optional): Whether to flip image horizontally or not.
// It overrides the corresponding field in the calculator options.
//
// FLIP_VERTICALLY (optional): Whether to flip image vertically or not.
// It overrides the corresponding field in the calculator options.
//
// Calculator options (see image_transformation_calculator.proto): // Calculator options (see image_transformation_calculator.proto):
// output_width, output_height - (optional) Desired scaled image size. // output_width, output_height - (optional) Desired scaled image size.
// rotation_mode - (optional) Rotation in multiples of 90 degrees. // rotation_mode - (optional) Rotation in multiples of 90 degrees.
@@ -138,8 +152,7 @@ mediapipe::ScaleMode_Mode ParseScaleMode(
// Note: To enable horizontal or vertical flipping, specify them in the // Note: To enable horizontal or vertical flipping, specify them in the
// calculator options. Flipping is applied after rotation. // calculator options. Flipping is applied after rotation.
// //
// Note: Only scale mode STRETCH is currently supported on CPU, // Note: Only scale mode STRETCH is currently supported on CPU.
// and flipping is not yet supported either.
// //
class ImageTransformationCalculator : public CalculatorBase { class ImageTransformationCalculator : public CalculatorBase {
public: public:
@@ -168,6 +181,8 @@ class ImageTransformationCalculator : public CalculatorBase {
int output_height_ = 0; int output_height_ = 0;
mediapipe::RotationMode_Mode rotation_; mediapipe::RotationMode_Mode rotation_;
mediapipe::ScaleMode_Mode scale_mode_; mediapipe::ScaleMode_Mode scale_mode_;
bool flip_horizontally_ = false;
bool flip_vertically_ = false;
bool use_gpu_ = false; bool use_gpu_ = false;
#if !defined(MEDIAPIPE_DISABLE_GPU) #if !defined(MEDIAPIPE_DISABLE_GPU)
@@ -204,6 +219,12 @@ REGISTER_CALCULATOR(ImageTransformationCalculator);
if (cc->Inputs().HasTag("ROTATION_DEGREES")) { if (cc->Inputs().HasTag("ROTATION_DEGREES")) {
cc->Inputs().Tag("ROTATION_DEGREES").Set<int>(); cc->Inputs().Tag("ROTATION_DEGREES").Set<int>();
} }
if (cc->Inputs().HasTag("FLIP_HORIZONTALLY")) {
cc->Inputs().Tag("FLIP_HORIZONTALLY").Set<bool>();
}
if (cc->Inputs().HasTag("FLIP_VERTICALLY")) {
cc->Inputs().Tag("FLIP_VERTICALLY").Set<bool>();
}
if (cc->InputSidePackets().HasTag("OUTPUT_DIMENSIONS")) { if (cc->InputSidePackets().HasTag("OUTPUT_DIMENSIONS")) {
cc->InputSidePackets().Tag("OUTPUT_DIMENSIONS").Set<DimensionsPacketType>(); cc->InputSidePackets().Tag("OUTPUT_DIMENSIONS").Set<DimensionsPacketType>();
@@ -211,6 +232,12 @@ REGISTER_CALCULATOR(ImageTransformationCalculator);
if (cc->InputSidePackets().HasTag("ROTATION_DEGREES")) { if (cc->InputSidePackets().HasTag("ROTATION_DEGREES")) {
cc->InputSidePackets().Tag("ROTATION_DEGREES").Set<int>(); cc->InputSidePackets().Tag("ROTATION_DEGREES").Set<int>();
} }
if (cc->InputSidePackets().HasTag("FLIP_HORIZONTALLY")) {
cc->InputSidePackets().Tag("FLIP_HORIZONTALLY").Set<bool>();
}
if (cc->InputSidePackets().HasTag("FLIP_VERTICALLY")) {
cc->InputSidePackets().Tag("FLIP_VERTICALLY").Set<bool>();
}
if (cc->Outputs().HasTag("LETTERBOX_PADDING")) { if (cc->Outputs().HasTag("LETTERBOX_PADDING")) {
cc->Outputs().Tag("LETTERBOX_PADDING").Set<std::array<float, 4>>(); cc->Outputs().Tag("LETTERBOX_PADDING").Set<std::array<float, 4>>();
@@ -246,6 +273,7 @@ REGISTER_CALCULATOR(ImageTransformationCalculator);
output_width_ = options_.output_width(); output_width_ = options_.output_width();
output_height_ = options_.output_height(); output_height_ = options_.output_height();
} }
if (cc->InputSidePackets().HasTag("ROTATION_DEGREES")) { if (cc->InputSidePackets().HasTag("ROTATION_DEGREES")) {
rotation_ = DegreesToRotationMode( rotation_ = DegreesToRotationMode(
cc->InputSidePackets().Tag("ROTATION_DEGREES").Get<int>()); cc->InputSidePackets().Tag("ROTATION_DEGREES").Get<int>());
@@ -253,6 +281,20 @@ REGISTER_CALCULATOR(ImageTransformationCalculator);
rotation_ = options_.rotation_mode(); rotation_ = options_.rotation_mode();
} }
if (cc->InputSidePackets().HasTag("FLIP_HORIZONTALLY")) {
flip_horizontally_ =
cc->InputSidePackets().Tag("FLIP_HORIZONTALLY").Get<bool>();
} else {
flip_horizontally_ = options_.flip_horizontally();
}
if (cc->InputSidePackets().HasTag("FLIP_VERTICALLY")) {
flip_vertically_ =
cc->InputSidePackets().Tag("FLIP_VERTICALLY").Get<bool>();
} else {
flip_vertically_ = options_.flip_vertically();
}
scale_mode_ = ParseScaleMode(options_.scale_mode(), DEFAULT_SCALE_MODE); scale_mode_ = ParseScaleMode(options_.scale_mode(), DEFAULT_SCALE_MODE);
if (use_gpu_) { if (use_gpu_) {
@@ -269,12 +311,37 @@ REGISTER_CALCULATOR(ImageTransformationCalculator);
::mediapipe::Status ImageTransformationCalculator::Process( ::mediapipe::Status ImageTransformationCalculator::Process(
CalculatorContext* cc) { CalculatorContext* cc) {
// Override values if specified so.
if (cc->Inputs().HasTag("ROTATION_DEGREES") &&
!cc->Inputs().Tag("ROTATION_DEGREES").IsEmpty()) {
rotation_ =
DegreesToRotationMode(cc->Inputs().Tag("ROTATION_DEGREES").Get<int>());
}
if (cc->Inputs().HasTag("FLIP_HORIZONTALLY") &&
!cc->Inputs().Tag("FLIP_HORIZONTALLY").IsEmpty()) {
flip_horizontally_ = cc->Inputs().Tag("FLIP_HORIZONTALLY").Get<bool>();
}
if (cc->Inputs().HasTag("FLIP_VERTICALLY") &&
!cc->Inputs().Tag("FLIP_VERTICALLY").IsEmpty()) {
flip_vertically_ = cc->Inputs().Tag("FLIP_VERTICALLY").Get<bool>();
}
if (use_gpu_) { if (use_gpu_) {
#if !defined(MEDIAPIPE_DISABLE_GPU) #if !defined(MEDIAPIPE_DISABLE_GPU)
if (cc->Inputs().Tag("IMAGE_GPU").IsEmpty()) {
// Image is missing, hence no way to produce output image. (Timestamp
// bound will be updated automatically.)
return ::mediapipe::OkStatus();
}
return helper_.RunInGlContext( return helper_.RunInGlContext(
[this, cc]() -> ::mediapipe::Status { return RenderGpu(cc); }); [this, cc]() -> ::mediapipe::Status { return RenderGpu(cc); });
#endif // !MEDIAPIPE_DISABLE_GPU #endif // !MEDIAPIPE_DISABLE_GPU
} else { } else {
if (cc->Inputs().Tag("IMAGE").IsEmpty()) {
// Image is missing, hence no way to produce output image. (Timestamp
// bound will be updated automatically.)
return ::mediapipe::OkStatus();
}
return RenderCpu(cc); return RenderCpu(cc);
} }
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
@@ -316,6 +383,11 @@ REGISTER_CALCULATOR(ImageTransformationCalculator);
cv::Mat input_mat = formats::MatView(&input_img); cv::Mat input_mat = formats::MatView(&input_img);
cv::Mat scaled_mat; cv::Mat scaled_mat;
if (!output_height_ || !output_width_) {
output_height_ = input_height;
output_width_ = input_width;
}
if (scale_mode_ == mediapipe::ScaleMode_Mode_STRETCH) { if (scale_mode_ == mediapipe::ScaleMode_Mode_STRETCH) {
cv::resize(input_mat, scaled_mat, cv::Size(output_width_, output_height_)); cv::resize(input_mat, scaled_mat, cv::Size(output_width_, output_height_));
} else { } else {
@@ -356,21 +428,25 @@ REGISTER_CALCULATOR(ImageTransformationCalculator);
.Add(padding.release(), cc->InputTimestamp()); .Add(padding.release(), cc->InputTimestamp());
} }
if (cc->InputSidePackets().HasTag("ROTATION_DEGREES")) {
rotation_ = DegreesToRotationMode(
cc->InputSidePackets().Tag("ROTATION_DEGREES").Get<int>());
}
cv::Mat rotated_mat; cv::Mat rotated_mat;
const int angle = RotationModeToDegrees(rotation_); const int angle = RotationModeToDegrees(rotation_);
cv::Point2f src_center(scaled_mat.cols / 2.0, scaled_mat.rows / 2.0); cv::Point2f src_center(scaled_mat.cols / 2.0, scaled_mat.rows / 2.0);
cv::Mat rotation_mat = cv::getRotationMatrix2D(src_center, angle, 1.0); cv::Mat rotation_mat = cv::getRotationMatrix2D(src_center, angle, 1.0);
cv::warpAffine(scaled_mat, rotated_mat, rotation_mat, scaled_mat.size()); cv::warpAffine(scaled_mat, rotated_mat, rotation_mat, scaled_mat.size());
cv::Mat flipped_mat;
if (flip_horizontally_ || flip_vertically_) {
const int flip_code =
flip_horizontally_ && flip_vertically_ ? -1 : flip_horizontally_;
cv::flip(rotated_mat, flipped_mat, flip_code);
} else {
flipped_mat = rotated_mat;
}
std::unique_ptr<ImageFrame> output_frame( std::unique_ptr<ImageFrame> output_frame(
new ImageFrame(input_img.Format(), output_width, output_height)); new ImageFrame(input_img.Format(), output_width, output_height));
cv::Mat output_mat = formats::MatView(output_frame.get()); cv::Mat output_mat = formats::MatView(output_frame.get());
rotated_mat.copyTo(output_mat); flipped_mat.copyTo(output_mat);
cc->Outputs().Tag("IMAGE").Add(output_frame.release(), cc->InputTimestamp()); cc->Outputs().Tag("IMAGE").Add(output_frame.release(), cc->InputTimestamp());
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
@@ -400,7 +476,7 @@ REGISTER_CALCULATOR(ImageTransformationCalculator);
QuadRenderer* renderer = nullptr; QuadRenderer* renderer = nullptr;
GlTexture src1; GlTexture src1;
#if defined(__APPLE__) && !TARGET_OS_OSX #if defined(MEDIAPIPE_IOS)
if (input.format() == GpuBufferFormat::kBiPlanar420YpCbCr8VideoRange || if (input.format() == GpuBufferFormat::kBiPlanar420YpCbCr8VideoRange ||
input.format() == GpuBufferFormat::kBiPlanar420YpCbCr8FullRange) { input.format() == GpuBufferFormat::kBiPlanar420YpCbCr8FullRange) {
if (!yuv_renderer_) { if (!yuv_renderer_) {
@@ -435,14 +511,8 @@ REGISTER_CALCULATOR(ImageTransformationCalculator);
} }
RET_CHECK(renderer) << "Unsupported input texture type"; RET_CHECK(renderer) << "Unsupported input texture type";
if (cc->InputSidePackets().HasTag("ROTATION_DEGREES")) { mediapipe::FrameScaleMode scale_mode = mediapipe::FrameScaleModeFromProto(
rotation_ = DegreesToRotationMode( scale_mode_, mediapipe::FrameScaleMode::kStretch);
cc->InputSidePackets().Tag("ROTATION_DEGREES").Get<int>());
}
static mediapipe::FrameScaleMode scale_mode =
mediapipe::FrameScaleModeFromProto(scale_mode_,
mediapipe::FrameScaleMode::kStretch);
mediapipe::FrameRotation rotation = mediapipe::FrameRotation rotation =
mediapipe::FrameRotationFromDegrees(RotationModeToDegrees(rotation_)); mediapipe::FrameRotationFromDegrees(RotationModeToDegrees(rotation_));
@@ -455,7 +525,7 @@ REGISTER_CALCULATOR(ImageTransformationCalculator);
MP_RETURN_IF_ERROR(renderer->GlRender( MP_RETURN_IF_ERROR(renderer->GlRender(
src1.width(), src1.height(), dst.width(), dst.height(), scale_mode, src1.width(), src1.height(), dst.width(), dst.height(), scale_mode,
rotation, options_.flip_horizontally(), options_.flip_vertically(), rotation, flip_horizontally_, flip_vertically_,
/*flip_texture=*/false)); /*flip_texture=*/false));
glActiveTexture(GL_TEXTURE1); glActiveTexture(GL_TEXTURE1);
@@ -260,11 +260,11 @@ ScaleImageCalculator::~ScaleImageCalculator() {}
&crop_width_, &crop_height_, // &crop_width_, &crop_height_, //
&col_start_, &row_start_)); &col_start_, &row_start_));
MP_RETURN_IF_ERROR( MP_RETURN_IF_ERROR(
scale_image::FindOutputDimensions(crop_width_, crop_height_, // scale_image::FindOutputDimensions(crop_width_, crop_height_, //
options_.target_width(), // options_.target_width(), //
options_.target_height(), // options_.target_height(), //
options_.preserve_aspect_ratio(), // options_.preserve_aspect_ratio(), //
options_.scale_to_multiple_of_two(), // options_.scale_to_multiple_of(), //
&output_width_, &output_height_)); &output_width_, &output_height_));
MP_RETURN_IF_ERROR(FindInterpolationAlgorithm(options_.algorithm(), MP_RETURN_IF_ERROR(FindInterpolationAlgorithm(options_.algorithm(),
&interpolation_algorithm_)); &interpolation_algorithm_));
@@ -361,17 +361,21 @@ ScaleImageCalculator::~ScaleImageCalculator() {}
output_format_ = input_format_; output_format_ = input_format_;
} }
const bool is_positive_and_even =
(options_.scale_to_multiple_of() >= 1) &&
(options_.scale_to_multiple_of() % 2 == 0);
if (output_format_ == ImageFormat::YCBCR420P) { if (output_format_ == ImageFormat::YCBCR420P) {
RET_CHECK(options_.scale_to_multiple_of_two()) RET_CHECK(is_positive_and_even)
<< "ScaleImageCalculator always outputs width and height that are " << "ScaleImageCalculator always outputs width and height that are "
"divisible by 2 when output format is YCbCr420P. To scale to " "divisible by 2 when output format is YCbCr420P. To scale to "
"width and height of odd numbers, the output format must be SRGB."; "width and height of odd numbers, the output format must be SRGB.";
} else if (options_.preserve_aspect_ratio()) { } else if (options_.preserve_aspect_ratio()) {
RET_CHECK(options_.scale_to_multiple_of_two()) RET_CHECK(options_.scale_to_multiple_of() == 2)
<< "ScaleImageCalculator always outputs width and height that are " << "ScaleImageCalculator always outputs width and height that are "
"divisible by 2 when perserving aspect ratio. To scale to width " "divisible by 2 when preserving aspect ratio. If you'd like to "
"and height of odd numbers, please set " "set scale_to_multiple_of to something other than 2, please "
"preserve_aspect_ratio to false."; "set preserve_aspect_ratio to false.";
} }
if (input_width_ > 0 && input_height_ > 0 && if (input_width_ > 0 && input_height_ > 0 &&
@@ -474,13 +478,20 @@ ScaleImageCalculator::~ScaleImageCalculator() {}
input_width_, "x", input_height_)); input_width_, "x", input_height_));
} }
if (input_format_ != image_frame.Format()) { if (input_format_ != image_frame.Format()) {
std::string image_frame_format_desc, input_format_desc;
#ifdef MEDIAPIPE_MOBILE
image_frame_format_desc = std::to_string(image_frame.Format());
input_format_desc = std::to_string(input_format_);
#else
const proto_ns::EnumDescriptor* desc = ImageFormat::Format_descriptor(); const proto_ns::EnumDescriptor* desc = ImageFormat::Format_descriptor();
image_frame_format_desc =
desc->FindValueByNumber(image_frame.Format())->DebugString();
input_format_desc = desc->FindValueByNumber(input_format_)->DebugString();
#endif // MEDIAPIPE_MOBILE
return tool::StatusFail(absl::StrCat( return tool::StatusFail(absl::StrCat(
"If a header specifies a format, then image frames on " "If a header specifies a format, then image frames on "
"the stream must have that format. Actual format ", "the stream must have that format. Actual format ",
desc->FindValueByNumber(image_frame.Format())->DebugString(), image_frame_format_desc, " but expected ", input_format_desc));
" but expected ",
desc->FindValueByNumber(input_format_)->DebugString()));
} }
} }
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
@@ -11,9 +11,10 @@ import "mediapipe/framework/formats/image_format.proto";
// 2) Scale and convert the image to fit inside target_width x target_height // 2) Scale and convert the image to fit inside target_width x target_height
// using the specified scaling algorithm. (maintaining the aspect // using the specified scaling algorithm. (maintaining the aspect
// ratio if preserve_aspect_ratio is true). // ratio if preserve_aspect_ratio is true).
// The output width and height will be divisible by 2. It is possible to output // The output width and height will be divisible by 2, by default. It is
// width and height that are odd number when the output format is SRGB and not // possible to output width and height that are odd numbers when the output
// perserving the aspect ratio. See scale_to_multiple_of_two option for details. // format is SRGB and the aspect ratio is left unpreserved. See
// scale_to_multiple_of for details.
message ScaleImageCalculatorOptions { message ScaleImageCalculatorOptions {
extend CalculatorOptions { extend CalculatorOptions {
optional ScaleImageCalculatorOptions ext = 66237115; optional ScaleImageCalculatorOptions ext = 66237115;
@@ -23,7 +24,7 @@ message ScaleImageCalculatorOptions {
// depending on the other options below. If unset, use the same width // depending on the other options below. If unset, use the same width
// or height as the input. If only one is set then determine the other // or height as the input. If only one is set then determine the other
// from the aspect ratio (after cropping). The output width and height // from the aspect ratio (after cropping). The output width and height
// will be divisible by 2. // will be divisible by 2, by default.
optional int32 target_width = 1; optional int32 target_width = 1;
optional int32 target_height = 2; optional int32 target_height = 2;
@@ -31,12 +32,14 @@ message ScaleImageCalculatorOptions {
// fits inside the box represented by target_width and target_height. // fits inside the box represented by target_width and target_height.
// Otherwise it is scaled to fit target_width and target_height // Otherwise it is scaled to fit target_width and target_height
// completely. In any case, the aspect ratio that is preserved is // completely. In any case, the aspect ratio that is preserved is
// that after cropping to the minimum/maximum aspect ratio. // that after cropping to the minimum/maximum aspect ratio. Additionally, if
// true, the output width and height will be divisible by 2.
optional bool preserve_aspect_ratio = 3 [default = true]; optional bool preserve_aspect_ratio = 3 [default = true];
// If ratio is positive, crop the image to this minimum and maximum // If ratio is positive, crop the image to this minimum and maximum
// aspect ratio (preserving the center of the frame). This is done // aspect ratio (preserving the center of the frame). This is done
// before scaling. // before scaling. The string must contain "/", so to disable cropping,
// set both to "0/1".
// For example, for a min_aspect_ratio of "9/16" and max of "16/9" the // For example, for a min_aspect_ratio of "9/16" and max of "16/9" the
// following cropping will occur: // following cropping will occur:
// 1920x1080 (which is 16:9) is not cropped // 1920x1080 (which is 16:9) is not cropped
@@ -94,11 +97,13 @@ message ScaleImageCalculatorOptions {
// SRGB or YCBCR420P. // SRGB or YCBCR420P.
optional ImageFormat.Format input_format = 12; optional ImageFormat.Format input_format = 12;
// If true, the output width and height will be divisible by 2. Otherwise it // If set to 2, the target width and height will be rounded-down
// will use the exact specified output width and height, which is only // to the nearest even number. If set to any positive value other than 2,
// supported when the output format is SRGB and preserve_aspect_ratio option // preserve_aspect_ratio must be false and the target width and height will be
// is set to false. // rounded-down to multiples of the given value. If set to any value less than
optional bool scale_to_multiple_of_two = 13 [default = true]; // 1, it will be treated like 1.
// NOTE: If set to an odd number, the output format must be SRGB.
optional int32 scale_to_multiple_of = 13 [default = 2];
// If true, assume the input YUV is BT.709 (this is the HDTV standard, so most // If true, assume the input YUV is BT.709 (this is the HDTV standard, so most
// content is likely using it). If false use the previous assumption of BT.601 // content is likely using it). If false use the previous assumption of BT.601
@@ -88,17 +88,27 @@ double ParseRational(const std::string& rational) {
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
::mediapipe::Status FindOutputDimensions(int input_width, // ::mediapipe::Status FindOutputDimensions(int input_width, //
int input_height, // int input_height, //
int target_width, // int target_width, //
int target_height, // int target_height, //
bool preserve_aspect_ratio, // bool preserve_aspect_ratio, //
bool scale_to_multiple_of_two, // int scale_to_multiple_of, //
int* output_width, int* output_width,
int* output_height) { int* output_height) {
CHECK(output_width); CHECK(output_width);
CHECK(output_height); CHECK(output_height);
if (preserve_aspect_ratio) {
RET_CHECK(scale_to_multiple_of == 2)
<< "FindOutputDimensions always outputs width and height that are "
"divisible by 2 when preserving aspect ratio. If you'd like to "
"set scale_to_multiple_of to something other than 2, please "
"set preserve_aspect_ratio to false.";
}
if (scale_to_multiple_of < 1) scale_to_multiple_of = 1;
if (!preserve_aspect_ratio || (target_width <= 0 && target_height <= 0)) { if (!preserve_aspect_ratio || (target_width <= 0 && target_height <= 0)) {
if (target_width <= 0) { if (target_width <= 0) {
target_width = input_width; target_width = input_width;
@@ -106,13 +116,13 @@ double ParseRational(const std::string& rational) {
if (target_height <= 0) { if (target_height <= 0) {
target_height = input_height; target_height = input_height;
} }
if (scale_to_multiple_of_two) {
*output_width = (target_width / 2) * 2; target_width -= target_width % scale_to_multiple_of;
*output_height = (target_height / 2) * 2; target_height -= target_height % scale_to_multiple_of;
} else {
*output_width = target_width; *output_width = target_width;
*output_height = target_height; *output_height = target_height;
}
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -35,17 +35,19 @@ namespace scale_image {
int* col_start, int* row_start); int* col_start, int* row_start);
// Given an input width and height, a target width and height, whether to // Given an input width and height, a target width and height, whether to
// preserve the aspect ratio, and whether to round down to a multiple of 2, // preserve the aspect ratio, and whether to round-down to the multiple of a
// determine the output width and height. If target_width or target_height is // given number nearest to the targets, determine the output width and height.
// non-positive, then they will be set to the input_width and input_height // If target_width or target_height is non-positive, then they will be set to
// respectively. The output_width and output_height will be reduced as necessary // the input_width and input_height respectively. If scale_to_multiple_of is
// to preserve_aspect_ratio and to scale_to_multipe_of_two if these options are // less than 1, it will be treated like 1. The output_width and
// specified. // output_height will be reduced as necessary to preserve_aspect_ratio if the
// option is specified. If preserving the aspect ratio is desired, you must set
// scale_to_multiple_of to 2.
::mediapipe::Status FindOutputDimensions(int input_width, int input_height, // ::mediapipe::Status FindOutputDimensions(int input_width, int input_height, //
int target_width, int target_width,
int target_height, // int target_height, //
bool preserve_aspect_ratio, // bool preserve_aspect_ratio, //
bool scale_to_multiple_of_two, // int scale_to_multiple_of, //
int* output_width, int* output_height); int* output_width, int* output_height);
} // namespace scale_image } // namespace scale_image
@@ -79,49 +79,49 @@ TEST(ScaleImageUtilsTest, FindOutputDimensionsPreserveRatio) {
int output_width; int output_width;
int output_height; int output_height;
// Not scale. // Not scale.
MP_ASSERT_OK(FindOutputDimensions(200, 100, -1, -1, true, true, &output_width, MP_ASSERT_OK(FindOutputDimensions(200, 100, -1, -1, true, 2, &output_width,
&output_height)); &output_height));
EXPECT_EQ(200, output_width); EXPECT_EQ(200, output_width);
EXPECT_EQ(100, output_height); EXPECT_EQ(100, output_height);
// Not scale with odd input size. // Not scale with odd input size.
MP_ASSERT_OK(FindOutputDimensions(201, 101, -1, -1, false, false, MP_ASSERT_OK(FindOutputDimensions(201, 101, -1, -1, false, 1, &output_width,
&output_width, &output_height)); &output_height));
EXPECT_EQ(201, output_width); EXPECT_EQ(201, output_width);
EXPECT_EQ(101, output_height); EXPECT_EQ(101, output_height);
// Scale down by 1/2. // Scale down by 1/2.
MP_ASSERT_OK(FindOutputDimensions(200, 100, 100, -1, true, true, MP_ASSERT_OK(FindOutputDimensions(200, 100, 100, -1, true, 2, &output_width,
&output_width, &output_height)); &output_height));
EXPECT_EQ(100, output_width); EXPECT_EQ(100, output_width);
EXPECT_EQ(50, output_height); EXPECT_EQ(50, output_height);
// Scale up, doubling dimensions. // Scale up, doubling dimensions.
MP_ASSERT_OK(FindOutputDimensions(200, 100, -1, 200, true, true, MP_ASSERT_OK(FindOutputDimensions(200, 100, -1, 200, true, 2, &output_width,
&output_width, &output_height)); &output_height));
EXPECT_EQ(400, output_width); EXPECT_EQ(400, output_width);
EXPECT_EQ(200, output_height); EXPECT_EQ(200, output_height);
// Fits a 2:1 image into a 150 x 150 box. Output dimensions are always // Fits a 2:1 image into a 150 x 150 box. Output dimensions are always
// visible by 2. // visible by 2.
MP_ASSERT_OK(FindOutputDimensions(200, 100, 150, 150, true, true, MP_ASSERT_OK(FindOutputDimensions(200, 100, 150, 150, true, 2, &output_width,
&output_width, &output_height)); &output_height));
EXPECT_EQ(150, output_width); EXPECT_EQ(150, output_width);
EXPECT_EQ(74, output_height); EXPECT_EQ(74, output_height);
// Fits a 2:1 image into a 400 x 50 box. // Fits a 2:1 image into a 400 x 50 box.
MP_ASSERT_OK(FindOutputDimensions(200, 100, 400, 50, true, true, MP_ASSERT_OK(FindOutputDimensions(200, 100, 400, 50, true, 2, &output_width,
&output_width, &output_height)); &output_height));
EXPECT_EQ(100, output_width); EXPECT_EQ(100, output_width);
EXPECT_EQ(50, output_height); EXPECT_EQ(50, output_height);
// Scale to multiple number with odd targe size. // Scale to multiple number with odd targe size.
MP_ASSERT_OK(FindOutputDimensions(200, 100, 101, -1, true, true, MP_ASSERT_OK(FindOutputDimensions(200, 100, 101, -1, true, 2, &output_width,
&output_width, &output_height)); &output_height));
EXPECT_EQ(100, output_width); EXPECT_EQ(100, output_width);
EXPECT_EQ(50, output_height); EXPECT_EQ(50, output_height);
// Scale to multiple number with odd targe size. // Scale to multiple number with odd targe size.
MP_ASSERT_OK(FindOutputDimensions(200, 100, 101, -1, true, false, MP_ASSERT_OK(FindOutputDimensions(200, 100, 101, -1, true, 2, &output_width,
&output_width, &output_height)); &output_height));
EXPECT_EQ(100, output_width); EXPECT_EQ(100, output_width);
EXPECT_EQ(50, output_height); EXPECT_EQ(50, output_height);
// Scale to odd size. // Scale to odd size.
MP_ASSERT_OK(FindOutputDimensions(200, 100, 151, 101, false, false, MP_ASSERT_OK(FindOutputDimensions(200, 100, 151, 101, false, 1, &output_width,
&output_width, &output_height)); &output_height));
EXPECT_EQ(151, output_width); EXPECT_EQ(151, output_width);
EXPECT_EQ(101, output_height); EXPECT_EQ(101, output_height);
} }
@@ -131,22 +131,62 @@ TEST(ScaleImageUtilsTest, FindOutputDimensionsNoAspectRatio) {
int output_width; int output_width;
int output_height; int output_height;
// Scale width only. // Scale width only.
MP_ASSERT_OK(FindOutputDimensions(200, 100, 100, -1, false, true, MP_ASSERT_OK(FindOutputDimensions(200, 100, 100, -1, false, 2, &output_width,
&output_width, &output_height)); &output_height));
EXPECT_EQ(100, output_width); EXPECT_EQ(100, output_width);
EXPECT_EQ(100, output_height); EXPECT_EQ(100, output_height);
// Scale height only. // Scale height only.
MP_ASSERT_OK(FindOutputDimensions(200, 100, -1, 200, false, true, MP_ASSERT_OK(FindOutputDimensions(200, 100, -1, 200, false, 2, &output_width,
&output_width, &output_height)); &output_height));
EXPECT_EQ(200, output_width); EXPECT_EQ(200, output_width);
EXPECT_EQ(200, output_height); EXPECT_EQ(200, output_height);
// Scale both dimensions. // Scale both dimensions.
MP_ASSERT_OK(FindOutputDimensions(200, 100, 150, 200, false, true, MP_ASSERT_OK(FindOutputDimensions(200, 100, 150, 200, false, 2, &output_width,
&output_width, &output_height)); &output_height));
EXPECT_EQ(150, output_width); EXPECT_EQ(150, output_width);
EXPECT_EQ(200, output_height); EXPECT_EQ(200, output_height);
} }
// Tests scale_to_multiple_of.
TEST(ScaleImageUtilsTest, FindOutputDimensionsDownScaleToMultipleOf) {
int output_width;
int output_height;
// Set no targets, downscale to a multiple of 8.
MP_ASSERT_OK(FindOutputDimensions(100, 100, -1, -1, false, 8, &output_width,
&output_height));
EXPECT_EQ(96, output_width);
EXPECT_EQ(96, output_height);
// Set width target, downscale to a multiple of 8.
MP_ASSERT_OK(FindOutputDimensions(200, 100, 100, -1, false, 8, &output_width,
&output_height));
EXPECT_EQ(96, output_width);
EXPECT_EQ(96, output_height);
// Set height target, downscale to a multiple of 8.
MP_ASSERT_OK(FindOutputDimensions(201, 101, -1, 201, false, 8, &output_width,
&output_height));
EXPECT_EQ(200, output_width);
EXPECT_EQ(200, output_height);
// Set both targets, downscale to a multiple of 8.
MP_ASSERT_OK(FindOutputDimensions(200, 100, 150, 200, false, 8, &output_width,
&output_height));
EXPECT_EQ(144, output_width);
EXPECT_EQ(200, output_height);
// Doesn't throw error if keep aspect is true and downscale multiple is 2.
MP_ASSERT_OK(FindOutputDimensions(200, 100, 400, 200, true, 2, &output_width,
&output_height));
EXPECT_EQ(400, output_width);
EXPECT_EQ(200, output_height);
// Throws error if keep aspect is true, but downscale multiple is not 2.
ASSERT_THAT(FindOutputDimensions(200, 100, 400, 200, true, 4, &output_width,
&output_height),
testing::Not(testing::status::IsOk()));
// Downscaling to multiple ignored if multiple is less than 2.
MP_ASSERT_OK(FindOutputDimensions(200, 100, 401, 201, false, 1, &output_width,
&output_height));
EXPECT_EQ(401, output_width);
EXPECT_EQ(201, output_height);
}
} // namespace } // namespace
} // namespace scale_image } // namespace scale_image
} // namespace mediapipe } // namespace mediapipe
+1
View File
@@ -21,6 +21,7 @@ filegroup(
"dino.jpg", "dino.jpg",
"dino_quality_50.jpg", "dino_quality_50.jpg",
"dino_quality_80.jpg", "dino_quality_80.jpg",
"front_camera_pixel2.jpg",
], ],
visibility = ["//visibility:public"], visibility = ["//visibility:public"],
) )
Binary file not shown.

After

Width:  |  Height:  |  Size: 6.3 MiB

+31 -21
View File
@@ -138,7 +138,7 @@ mediapipe_cc_proto_library(
srcs = ["image_frame_to_tensor_calculator.proto"], srcs = ["image_frame_to_tensor_calculator.proto"],
cc_deps = [ cc_deps = [
"//mediapipe/framework:calculator_cc_proto", "//mediapipe/framework:calculator_cc_proto",
"@org_tensorflow//tensorflow/core:protos_all", "@org_tensorflow//tensorflow/core:protos_all_cc",
], ],
visibility = ["//visibility:public"], visibility = ["//visibility:public"],
deps = [":image_frame_to_tensor_calculator_proto"], deps = [":image_frame_to_tensor_calculator_proto"],
@@ -173,7 +173,7 @@ mediapipe_cc_proto_library(
srcs = ["pack_media_sequence_calculator.proto"], srcs = ["pack_media_sequence_calculator.proto"],
cc_deps = [ cc_deps = [
"//mediapipe/framework:calculator_cc_proto", "//mediapipe/framework:calculator_cc_proto",
"@org_tensorflow//tensorflow/core:protos_all", "@org_tensorflow//tensorflow/core:protos_all_cc",
], ],
visibility = ["//visibility:public"], visibility = ["//visibility:public"],
deps = [":pack_media_sequence_calculator_proto"], deps = [":pack_media_sequence_calculator_proto"],
@@ -192,7 +192,7 @@ mediapipe_cc_proto_library(
srcs = ["tensorflow_session_from_frozen_graph_generator.proto"], srcs = ["tensorflow_session_from_frozen_graph_generator.proto"],
cc_deps = [ cc_deps = [
"//mediapipe/framework:packet_generator_cc_proto", "//mediapipe/framework:packet_generator_cc_proto",
"@org_tensorflow//tensorflow/core:protos_all", "@org_tensorflow//tensorflow/core:protos_all_cc",
], ],
visibility = ["//visibility:public"], visibility = ["//visibility:public"],
deps = [":tensorflow_session_from_frozen_graph_generator_proto"], deps = [":tensorflow_session_from_frozen_graph_generator_proto"],
@@ -203,7 +203,7 @@ mediapipe_cc_proto_library(
srcs = ["tensorflow_session_from_frozen_graph_calculator.proto"], srcs = ["tensorflow_session_from_frozen_graph_calculator.proto"],
cc_deps = [ cc_deps = [
"//mediapipe/framework:calculator_cc_proto", "//mediapipe/framework:calculator_cc_proto",
"@org_tensorflow//tensorflow/core:protos_all", "@org_tensorflow//tensorflow/core:protos_all_cc",
], ],
visibility = ["//visibility:public"], visibility = ["//visibility:public"],
deps = [":tensorflow_session_from_frozen_graph_calculator_proto"], deps = [":tensorflow_session_from_frozen_graph_calculator_proto"],
@@ -277,7 +277,7 @@ mediapipe_cc_proto_library(
srcs = ["vector_int_to_tensor_calculator_options.proto"], srcs = ["vector_int_to_tensor_calculator_options.proto"],
cc_deps = [ cc_deps = [
"//mediapipe/framework:calculator_cc_proto", "//mediapipe/framework:calculator_cc_proto",
"@org_tensorflow//tensorflow/core:protos_all", "@org_tensorflow//tensorflow/core:protos_all_cc",
], ],
visibility = ["//visibility:public"], visibility = ["//visibility:public"],
deps = [":vector_int_to_tensor_calculator_options_proto"], deps = [":vector_int_to_tensor_calculator_options_proto"],
@@ -408,7 +408,7 @@ cc_library(
"//mediapipe/util/sequence:media_sequence", "//mediapipe/util/sequence:media_sequence",
"//mediapipe/util/sequence:media_sequence_util", "//mediapipe/util/sequence:media_sequence_util",
"@com_google_absl//absl/strings", "@com_google_absl//absl/strings",
"@org_tensorflow//tensorflow/core:protos_all", "@org_tensorflow//tensorflow/core:protos_all_cc",
], ],
alwayslink = 1, alwayslink = 1,
) )
@@ -423,7 +423,7 @@ cc_library(
"//mediapipe/framework:calculator_framework", "//mediapipe/framework:calculator_framework",
"//mediapipe/framework/port:ret_check", "//mediapipe/framework/port:ret_check",
"//mediapipe/framework/port:status", "//mediapipe/framework/port:status",
"@org_tensorflow//tensorflow/core:protos_all", "@org_tensorflow//tensorflow/core:protos_all_cc",
], ],
alwayslink = 1, alwayslink = 1,
) )
@@ -654,7 +654,7 @@ cc_library(
"//mediapipe/framework/port:ret_check", "//mediapipe/framework/port:ret_check",
"//mediapipe/framework/port:status", "//mediapipe/framework/port:status",
"@org_tensorflow//tensorflow/core:lib", "@org_tensorflow//tensorflow/core:lib",
"@org_tensorflow//tensorflow/core:protos_all", "@org_tensorflow//tensorflow/core:protos_all_cc",
], ],
alwayslink = 1, alwayslink = 1,
) )
@@ -695,7 +695,7 @@ cc_library(
"//mediapipe/util:audio_decoder_cc_proto", "//mediapipe/util:audio_decoder_cc_proto",
"//mediapipe/util/sequence:media_sequence", "//mediapipe/util/sequence:media_sequence",
"@com_google_absl//absl/strings", "@com_google_absl//absl/strings",
"@org_tensorflow//tensorflow/core:protos_all", "@org_tensorflow//tensorflow/core:protos_all_cc",
], ],
alwayslink = 1, alwayslink = 1,
) )
@@ -737,7 +737,7 @@ cc_library(
"//mediapipe/framework:calculator_framework", "//mediapipe/framework:calculator_framework",
"//mediapipe/framework:packet", "//mediapipe/framework:packet",
"//mediapipe/framework/port:status", "//mediapipe/framework/port:status",
"@org_tensorflow//tensorflow/core:protos_all", "@org_tensorflow//tensorflow/core:protos_all_cc",
], ],
alwayslink = 1, alwayslink = 1,
) )
@@ -745,6 +745,7 @@ cc_library(
cc_test( cc_test(
name = "graph_tensors_packet_generator_test", name = "graph_tensors_packet_generator_test",
srcs = ["graph_tensors_packet_generator_test.cc"], srcs = ["graph_tensors_packet_generator_test.cc"],
linkstatic = 1,
deps = [ deps = [
":graph_tensors_packet_generator", ":graph_tensors_packet_generator",
":graph_tensors_packet_generator_cc_proto", ":graph_tensors_packet_generator_cc_proto",
@@ -761,6 +762,7 @@ cc_test(
name = "image_frame_to_tensor_calculator_test", name = "image_frame_to_tensor_calculator_test",
size = "small", size = "small",
srcs = ["image_frame_to_tensor_calculator_test.cc"], srcs = ["image_frame_to_tensor_calculator_test.cc"],
linkstatic = 1,
deps = [ deps = [
":image_frame_to_tensor_calculator", ":image_frame_to_tensor_calculator",
"//mediapipe/framework:calculator_framework", "//mediapipe/framework:calculator_framework",
@@ -777,6 +779,7 @@ cc_test(
name = "matrix_to_tensor_calculator_test", name = "matrix_to_tensor_calculator_test",
size = "small", size = "small",
srcs = ["matrix_to_tensor_calculator_test.cc"], srcs = ["matrix_to_tensor_calculator_test.cc"],
linkstatic = 1,
deps = [ deps = [
":matrix_to_tensor_calculator", ":matrix_to_tensor_calculator",
":matrix_to_tensor_calculator_options_cc_proto", ":matrix_to_tensor_calculator_options_cc_proto",
@@ -793,6 +796,7 @@ cc_test(
name = "lapped_tensor_buffer_calculator_test", name = "lapped_tensor_buffer_calculator_test",
size = "small", size = "small",
srcs = ["lapped_tensor_buffer_calculator_test.cc"], srcs = ["lapped_tensor_buffer_calculator_test.cc"],
linkstatic = 1,
deps = [ deps = [
":lapped_tensor_buffer_calculator", ":lapped_tensor_buffer_calculator",
":lapped_tensor_buffer_calculator_cc_proto", ":lapped_tensor_buffer_calculator_cc_proto",
@@ -801,7 +805,7 @@ cc_test(
"//mediapipe/framework/port:gtest_main", "//mediapipe/framework/port:gtest_main",
"@com_google_absl//absl/memory", "@com_google_absl//absl/memory",
"@org_tensorflow//tensorflow/core:framework", "@org_tensorflow//tensorflow/core:framework",
"@org_tensorflow//tensorflow/core:protos_all", "@org_tensorflow//tensorflow/core:protos_all_cc",
], ],
) )
@@ -840,7 +844,7 @@ cc_test(
"//mediapipe/util/sequence:media_sequence", "//mediapipe/util/sequence:media_sequence",
"@com_google_absl//absl/memory", "@com_google_absl//absl/memory",
"@com_google_absl//absl/strings", "@com_google_absl//absl/strings",
"@org_tensorflow//tensorflow/core:protos_all", "@org_tensorflow//tensorflow/core:protos_all_cc",
], ],
) )
@@ -867,7 +871,7 @@ cc_test(
"@com_google_absl//absl/strings", "@com_google_absl//absl/strings",
"@org_tensorflow//tensorflow/core:direct_session", "@org_tensorflow//tensorflow/core:direct_session",
"@org_tensorflow//tensorflow/core:framework", "@org_tensorflow//tensorflow/core:framework",
"@org_tensorflow//tensorflow/core:protos_all", "@org_tensorflow//tensorflow/core:protos_all_cc",
"@org_tensorflow//tensorflow/core:testlib", "@org_tensorflow//tensorflow/core:testlib",
"@org_tensorflow//tensorflow/core/kernels:conv_ops", "@org_tensorflow//tensorflow/core/kernels:conv_ops",
"@org_tensorflow//tensorflow/core/kernels:math", "@org_tensorflow//tensorflow/core/kernels:math",
@@ -897,7 +901,7 @@ cc_test(
"@com_google_absl//absl/strings", "@com_google_absl//absl/strings",
"@org_tensorflow//tensorflow/core:direct_session", "@org_tensorflow//tensorflow/core:direct_session",
"@org_tensorflow//tensorflow/core:framework", "@org_tensorflow//tensorflow/core:framework",
"@org_tensorflow//tensorflow/core:protos_all", "@org_tensorflow//tensorflow/core:protos_all_cc",
"@org_tensorflow//tensorflow/core:testlib", "@org_tensorflow//tensorflow/core:testlib",
"@org_tensorflow//tensorflow/core/kernels:conv_ops", "@org_tensorflow//tensorflow/core/kernels:conv_ops",
"@org_tensorflow//tensorflow/core/kernels:math", "@org_tensorflow//tensorflow/core/kernels:math",
@@ -956,6 +960,7 @@ cc_test(
cc_test( cc_test(
name = "tensor_squeeze_dimensions_calculator_test", name = "tensor_squeeze_dimensions_calculator_test",
srcs = ["tensor_squeeze_dimensions_calculator_test.cc"], srcs = ["tensor_squeeze_dimensions_calculator_test.cc"],
linkstatic = 1,
deps = [ deps = [
":tensor_squeeze_dimensions_calculator", ":tensor_squeeze_dimensions_calculator",
":tensor_squeeze_dimensions_calculator_cc_proto", ":tensor_squeeze_dimensions_calculator_cc_proto",
@@ -963,7 +968,7 @@ cc_test(
"//mediapipe/framework:calculator_runner", "//mediapipe/framework:calculator_runner",
"//mediapipe/framework/port:gtest_main", "//mediapipe/framework/port:gtest_main",
"@org_tensorflow//tensorflow/core:framework", "@org_tensorflow//tensorflow/core:framework",
"@org_tensorflow//tensorflow/core:protos_all", "@org_tensorflow//tensorflow/core:protos_all_cc",
], ],
) )
@@ -971,6 +976,7 @@ cc_test(
name = "tensor_to_image_frame_calculator_test", name = "tensor_to_image_frame_calculator_test",
size = "small", size = "small",
srcs = ["tensor_to_image_frame_calculator_test.cc"], srcs = ["tensor_to_image_frame_calculator_test.cc"],
linkstatic = 1,
deps = [ deps = [
":tensor_to_image_frame_calculator", ":tensor_to_image_frame_calculator",
":tensor_to_image_frame_calculator_cc_proto", ":tensor_to_image_frame_calculator_cc_proto",
@@ -979,7 +985,7 @@ cc_test(
"//mediapipe/framework/formats:image_frame", "//mediapipe/framework/formats:image_frame",
"//mediapipe/framework/port:gtest_main", "//mediapipe/framework/port:gtest_main",
"@org_tensorflow//tensorflow/core:framework", "@org_tensorflow//tensorflow/core:framework",
"@org_tensorflow//tensorflow/core:protos_all", "@org_tensorflow//tensorflow/core:protos_all_cc",
], ],
) )
@@ -987,6 +993,7 @@ cc_test(
name = "tensor_to_matrix_calculator_test", name = "tensor_to_matrix_calculator_test",
size = "small", size = "small",
srcs = ["tensor_to_matrix_calculator_test.cc"], srcs = ["tensor_to_matrix_calculator_test.cc"],
linkstatic = 1,
deps = [ deps = [
":tensor_to_matrix_calculator", ":tensor_to_matrix_calculator",
":tensor_to_matrix_calculator_cc_proto", ":tensor_to_matrix_calculator_cc_proto",
@@ -996,13 +1003,14 @@ cc_test(
"//mediapipe/framework/formats:time_series_header_cc_proto", "//mediapipe/framework/formats:time_series_header_cc_proto",
"//mediapipe/framework/port:gtest_main", "//mediapipe/framework/port:gtest_main",
"@org_tensorflow//tensorflow/core:framework", "@org_tensorflow//tensorflow/core:framework",
"@org_tensorflow//tensorflow/core:protos_all", "@org_tensorflow//tensorflow/core:protos_all_cc",
], ],
) )
cc_test( cc_test(
name = "tensor_to_vector_float_calculator_test", name = "tensor_to_vector_float_calculator_test",
srcs = ["tensor_to_vector_float_calculator_test.cc"], srcs = ["tensor_to_vector_float_calculator_test.cc"],
linkstatic = 1,
deps = [ deps = [
":tensor_to_vector_float_calculator", ":tensor_to_vector_float_calculator",
":tensor_to_vector_float_calculator_options_cc_proto", ":tensor_to_vector_float_calculator_options_cc_proto",
@@ -1010,7 +1018,7 @@ cc_test(
"//mediapipe/framework:calculator_runner", "//mediapipe/framework:calculator_runner",
"//mediapipe/framework/port:gtest_main", "//mediapipe/framework/port:gtest_main",
"@org_tensorflow//tensorflow/core:framework", "@org_tensorflow//tensorflow/core:framework",
"@org_tensorflow//tensorflow/core:protos_all", "@org_tensorflow//tensorflow/core:protos_all_cc",
], ],
) )
@@ -1030,13 +1038,14 @@ cc_test(
"//mediapipe/util/sequence:media_sequence", "//mediapipe/util/sequence:media_sequence",
"@com_google_absl//absl/memory", "@com_google_absl//absl/memory",
"@com_google_absl//absl/strings", "@com_google_absl//absl/strings",
"@org_tensorflow//tensorflow/core:protos_all", "@org_tensorflow//tensorflow/core:protos_all_cc",
], ],
) )
cc_test( cc_test(
name = "vector_int_to_tensor_calculator_test", name = "vector_int_to_tensor_calculator_test",
srcs = ["vector_int_to_tensor_calculator_test.cc"], srcs = ["vector_int_to_tensor_calculator_test.cc"],
linkstatic = 1,
deps = [ deps = [
":vector_int_to_tensor_calculator", ":vector_int_to_tensor_calculator",
":vector_int_to_tensor_calculator_options_cc_proto", ":vector_int_to_tensor_calculator_options_cc_proto",
@@ -1044,13 +1053,14 @@ cc_test(
"//mediapipe/framework:calculator_runner", "//mediapipe/framework:calculator_runner",
"//mediapipe/framework/port:gtest_main", "//mediapipe/framework/port:gtest_main",
"@org_tensorflow//tensorflow/core:framework", "@org_tensorflow//tensorflow/core:framework",
"@org_tensorflow//tensorflow/core:protos_all", "@org_tensorflow//tensorflow/core:protos_all_cc",
], ],
) )
cc_test( cc_test(
name = "vector_float_to_tensor_calculator_test", name = "vector_float_to_tensor_calculator_test",
srcs = ["vector_float_to_tensor_calculator_test.cc"], srcs = ["vector_float_to_tensor_calculator_test.cc"],
linkstatic = 1,
deps = [ deps = [
":vector_float_to_tensor_calculator", ":vector_float_to_tensor_calculator",
":vector_float_to_tensor_calculator_options_cc_proto", ":vector_float_to_tensor_calculator_options_cc_proto",
@@ -1058,7 +1068,7 @@ cc_test(
"//mediapipe/framework:calculator_runner", "//mediapipe/framework:calculator_runner",
"//mediapipe/framework/port:gtest_main", "//mediapipe/framework/port:gtest_main",
"@org_tensorflow//tensorflow/core:framework", "@org_tensorflow//tensorflow/core:framework",
"@org_tensorflow//tensorflow/core:protos_all", "@org_tensorflow//tensorflow/core:protos_all_cc",
], ],
) )
@@ -34,6 +34,7 @@ namespace mediapipe {
const char kSequenceExampleTag[] = "SEQUENCE_EXAMPLE"; const char kSequenceExampleTag[] = "SEQUENCE_EXAMPLE";
const char kImageTag[] = "IMAGE"; const char kImageTag[] = "IMAGE";
const char kFloatContextFeaturePrefixTag[] = "FLOAT_CONTEXT_FEATURE_";
const char kFloatFeaturePrefixTag[] = "FLOAT_FEATURE_"; const char kFloatFeaturePrefixTag[] = "FLOAT_FEATURE_";
const char kForwardFlowEncodedTag[] = "FORWARD_FLOW_ENCODED"; const char kForwardFlowEncodedTag[] = "FORWARD_FLOW_ENCODED";
const char kBBoxTag[] = "BBOX"; const char kBBoxTag[] = "BBOX";
@@ -145,6 +146,9 @@ class PackMediaSequenceCalculator : public CalculatorBase {
} }
cc->Inputs().Tag(tag).Set<std::vector<Detection>>(); cc->Inputs().Tag(tag).Set<std::vector<Detection>>();
} }
if (absl::StartsWith(tag, kFloatContextFeaturePrefixTag)) {
cc->Inputs().Tag(tag).Set<std::vector<float>>();
}
if (absl::StartsWith(tag, kFloatFeaturePrefixTag)) { if (absl::StartsWith(tag, kFloatFeaturePrefixTag)) {
cc->Inputs().Tag(tag).Set<std::vector<float>>(); cc->Inputs().Tag(tag).Set<std::vector<float>>();
} }
@@ -264,7 +268,7 @@ class PackMediaSequenceCalculator : public CalculatorBase {
if (options.output_only_if_all_present()) { if (options.output_only_if_all_present()) {
::mediapipe::Status status = VerifySequence(); ::mediapipe::Status status = VerifySequence();
if (!status.ok()) { if (!status.ok()) {
cc->GetCounter(status.error_message())->Increment(); cc->GetCounter(status.ToString())->Increment();
return status; return status;
} }
} }
@@ -344,6 +348,17 @@ class PackMediaSequenceCalculator : public CalculatorBase {
sequence_.get()); sequence_.get());
} }
} }
if (absl::StartsWith(tag, kFloatContextFeaturePrefixTag) &&
!cc->Inputs().Tag(tag).IsEmpty()) {
std::string key =
tag.substr(sizeof(kFloatContextFeaturePrefixTag) /
sizeof(*kFloatContextFeaturePrefixTag) -
1);
RET_CHECK_EQ(cc->InputTimestamp(), Timestamp::PostStream());
mpms::SetContextFeatureFloats(
key, cc->Inputs().Tag(tag).Get<std::vector<float>>(),
sequence_.get());
}
if (absl::StartsWith(tag, kFloatFeaturePrefixTag) && if (absl::StartsWith(tag, kFloatFeaturePrefixTag) &&
!cc->Inputs().Tag(tag).IsEmpty()) { !cc->Inputs().Tag(tag).IsEmpty()) {
std::string key = tag.substr(sizeof(kFloatFeaturePrefixTag) / std::string key = tag.substr(sizeof(kFloatFeaturePrefixTag) /
@@ -194,6 +194,38 @@ TEST_F(PackMediaSequenceCalculatorTest, PacksTwoFloatLists) {
} }
} }
TEST_F(PackMediaSequenceCalculatorTest, PacksTwoContextFloatLists) {
SetUpCalculator(
{"FLOAT_CONTEXT_FEATURE_TEST:test", "FLOAT_CONTEXT_FEATURE_OTHER:test2"},
{}, false, true);
auto input_sequence = absl::make_unique<tf::SequenceExample>();
auto vf_ptr = absl::make_unique<std::vector<float>>(2, 3);
runner_->MutableInputs()
->Tag("FLOAT_CONTEXT_FEATURE_TEST")
.packets.push_back(Adopt(vf_ptr.release()).At(Timestamp::PostStream()));
vf_ptr = absl::make_unique<std::vector<float>>(2, 4);
runner_->MutableInputs()
->Tag("FLOAT_CONTEXT_FEATURE_OTHER")
.packets.push_back(Adopt(vf_ptr.release()).At(Timestamp::PostStream()));
runner_->MutableSidePackets()->Tag("SEQUENCE_EXAMPLE") =
Adopt(input_sequence.release());
MP_ASSERT_OK(runner_->Run());
const std::vector<Packet>& output_packets =
runner_->Outputs().Tag("SEQUENCE_EXAMPLE").packets;
ASSERT_EQ(1, output_packets.size());
const tf::SequenceExample& output_sequence =
output_packets[0].Get<tf::SequenceExample>();
ASSERT_THAT(mpms::GetContextFeatureFloats("TEST", output_sequence),
testing::ElementsAre(3, 3));
ASSERT_THAT(mpms::GetContextFeatureFloats("OTHER", output_sequence),
testing::ElementsAre(4, 4));
}
TEST_F(PackMediaSequenceCalculatorTest, PacksAdditionalContext) { TEST_F(PackMediaSequenceCalculatorTest, PacksAdditionalContext) {
tf::Features context; tf::Features context;
(*context.mutable_feature())["TEST"].mutable_bytes_list()->add_value("YES"); (*context.mutable_feature())["TEST"].mutable_bytes_list()->add_value("YES");
@@ -34,7 +34,7 @@
#include "tensorflow/core/framework/tensor_shape.h" #include "tensorflow/core/framework/tensor_shape.h"
#include "tensorflow/core/framework/tensor_util.h" #include "tensorflow/core/framework/tensor_util.h"
#if !defined(__ANDROID__) && !defined(__APPLE__) #if !defined(MEDIAPIPE_MOBILE) && !defined(__APPLE__)
#include "tensorflow/core/profiler/lib/traceme.h" #include "tensorflow/core/profiler/lib/traceme.h"
#endif #endif
@@ -441,7 +441,7 @@ class TensorFlowInferenceCalculator : public CalculatorBase {
const int64 run_start_time = absl::ToUnixMicros(clock_->TimeNow()); const int64 run_start_time = absl::ToUnixMicros(clock_->TimeNow());
tf::Status tf_status; tf::Status tf_status;
{ {
#if !defined(__ANDROID__) && !defined(__APPLE__) #if !defined(MEDIAPIPE_MOBILE) && !defined(__APPLE__)
tensorflow::profiler::TraceMe trace(absl::string_view(cc->NodeName())); tensorflow::profiler::TraceMe trace(absl::string_view(cc->NodeName()));
#endif #endif
tf_status = session_->Run(input_tensors, output_tensor_names, tf_status = session_->Run(input_tensors, output_tensor_names,
@@ -454,7 +454,7 @@ class TensorFlowInferenceCalculator : public CalculatorBase {
// RET_CHECK on the tf::Status object itself in order to print an // RET_CHECK on the tf::Status object itself in order to print an
// informative error message. // informative error message.
RET_CHECK(tf_status.ok()) << "Run failed: " << tf_status.error_message(); RET_CHECK(tf_status.ok()) << "Run failed: " << tf_status.ToString();
const int64 run_end_time = absl::ToUnixMicros(clock_->TimeNow()); const int64 run_end_time = absl::ToUnixMicros(clock_->TimeNow());
cc->GetCounter(kTotalSessionRunsTimeUsecsCounterSuffix) cc->GetCounter(kTotalSessionRunsTimeUsecsCounterSuffix)
@@ -31,8 +31,7 @@
#include "mediapipe/framework/tool/status_util.h" #include "mediapipe/framework/tool/status_util.h"
#include "tensorflow/core/public/session_options.h" #include "tensorflow/core/public/session_options.h"
#if defined(MEDIAPIPE_LITE) || defined(__ANDROID__) || \ #if defined(MEDIAPIPE_MOBILE)
defined(__APPLE__) && !TARGET_OS_OSX
#include "mediapipe/util/android/file/base/helpers.h" #include "mediapipe/util/android/file/base/helpers.h"
#else #else
#include "mediapipe/framework/port/file_helpers.h" #include "mediapipe/framework/port/file_helpers.h"
@@ -110,7 +109,7 @@ class TensorFlowSessionFromFrozenGraphCalculator : public CalculatorBase {
RET_CHECK(graph_def.ParseFromString(graph_def_serialized)); RET_CHECK(graph_def.ParseFromString(graph_def_serialized));
const tf::Status tf_status = session->session->Create(graph_def); const tf::Status tf_status = session->session->Create(graph_def);
RET_CHECK(tf_status.ok()) << "Create failed: " << tf_status.error_message(); RET_CHECK(tf_status.ok()) << "Create failed: " << tf_status.ToString();
for (const auto& key_value : options.tag_to_tensor_names()) { for (const auto& key_value : options.tag_to_tensor_names()) {
session->tag_to_tensor_map[key_value.first] = key_value.second; session->tag_to_tensor_map[key_value.first] = key_value.second;
@@ -120,7 +119,7 @@ class TensorFlowSessionFromFrozenGraphCalculator : public CalculatorBase {
session->session->Run({}, {}, initialization_op_names, {}); session->session->Run({}, {}, initialization_op_names, {});
// RET_CHECK on the tf::Status object itself in order to print an // RET_CHECK on the tf::Status object itself in order to print an
// informative error message. // informative error message.
RET_CHECK(tf_status.ok()) << "Run failed: " << tf_status.error_message(); RET_CHECK(tf_status.ok()) << "Run failed: " << tf_status.ToString();
} }
cc->OutputSidePackets().Tag("SESSION").Set(Adopt(session.release())); cc->OutputSidePackets().Tag("SESSION").Set(Adopt(session.release()));
@@ -109,7 +109,7 @@ class TensorFlowSessionFromFrozenGraphGenerator : public PacketGenerator {
RET_CHECK(graph_def.ParseFromString(graph_def_serialized)); RET_CHECK(graph_def.ParseFromString(graph_def_serialized));
const tf::Status tf_status = session->session->Create(graph_def); const tf::Status tf_status = session->session->Create(graph_def);
RET_CHECK(tf_status.ok()) << "Create failed: " << tf_status.error_message(); RET_CHECK(tf_status.ok()) << "Create failed: " << tf_status.ToString();
for (const auto& key_value : options.tag_to_tensor_names()) { for (const auto& key_value : options.tag_to_tensor_names()) {
session->tag_to_tensor_map[key_value.first] = key_value.second; session->tag_to_tensor_map[key_value.first] = key_value.second;
@@ -119,7 +119,7 @@ class TensorFlowSessionFromFrozenGraphGenerator : public PacketGenerator {
session->session->Run({}, {}, initialization_op_names, {}); session->session->Run({}, {}, initialization_op_names, {});
// RET_CHECK on the tf::Status object itself in order to print an // RET_CHECK on the tf::Status object itself in order to print an
// informative error message. // informative error message.
RET_CHECK(tf_status.ok()) << "Run failed: " << tf_status.error_message(); RET_CHECK(tf_status.ok()) << "Run failed: " << tf_status.ToString();
} }
output_side_packets->Tag("SESSION") = Adopt(session.release()); output_side_packets->Tag("SESSION") = Adopt(session.release());
@@ -17,7 +17,7 @@
#if !defined(__ANDROID__) #if !defined(__ANDROID__)
#include "mediapipe/framework/port/file_helpers.h" #include "mediapipe/framework/port/file_helpers.h"
#endif #endif
#include "absl/strings/substitute.h" #include "absl/strings/str_replace.h"
#include "mediapipe/calculators/tensorflow/tensorflow_session.h" #include "mediapipe/calculators/tensorflow/tensorflow_session.h"
#include "mediapipe/calculators/tensorflow/tensorflow_session_from_saved_model_calculator.pb.h" #include "mediapipe/calculators/tensorflow/tensorflow_session_from_saved_model_calculator.pb.h"
#include "mediapipe/framework/calculator_framework.h" #include "mediapipe/framework/calculator_framework.h"
@@ -63,7 +63,7 @@ const std::string MaybeConvertSignatureToTag(
output.resize(name.length()); output.resize(name.length());
std::transform(name.begin(), name.end(), output.begin(), std::transform(name.begin(), name.end(), output.begin(),
[](unsigned char c) { return std::toupper(c); }); [](unsigned char c) { return std::toupper(c); });
output = absl::Substitute(output, "/", "_"); output = absl::StrReplaceAll(output, {{"/", "_"}});
return output; return output;
} else { } else {
return name; return name;
@@ -140,7 +140,7 @@ class TensorFlowSessionFromSavedModelCalculator : public CalculatorBase {
if (!status.ok()) { if (!status.ok()) {
return ::mediapipe::Status( return ::mediapipe::Status(
static_cast<::mediapipe::StatusCode>(status.code()), static_cast<::mediapipe::StatusCode>(status.code()),
status.error_message()); status.ToString());
} }
auto session = absl::make_unique<TensorFlowSession>(); auto session = absl::make_unique<TensorFlowSession>();
@@ -12,7 +12,7 @@
// See the License for the specific language governing permissions and // See the License for the specific language governing permissions and
// limitations under the License. // limitations under the License.
#include "absl/strings/substitute.h" #include "absl/strings/str_replace.h"
#include "mediapipe/calculators/tensorflow/tensorflow_session.h" #include "mediapipe/calculators/tensorflow/tensorflow_session.h"
#include "mediapipe/calculators/tensorflow/tensorflow_session_from_saved_model_calculator.pb.h" #include "mediapipe/calculators/tensorflow/tensorflow_session_from_saved_model_calculator.pb.h"
#include "mediapipe/framework/calculator.pb.h" #include "mediapipe/framework/calculator.pb.h"
@@ -17,7 +17,7 @@
#if !defined(__ANDROID__) #if !defined(__ANDROID__)
#include "mediapipe/framework/port/file_helpers.h" #include "mediapipe/framework/port/file_helpers.h"
#endif #endif
#include "absl/strings/substitute.h" #include "absl/strings/str_replace.h"
#include "mediapipe/calculators/tensorflow/tensorflow_session.h" #include "mediapipe/calculators/tensorflow/tensorflow_session.h"
#include "mediapipe/calculators/tensorflow/tensorflow_session_from_saved_model_generator.pb.h" #include "mediapipe/calculators/tensorflow/tensorflow_session_from_saved_model_generator.pb.h"
#include "mediapipe/framework/deps/file_path.h" #include "mediapipe/framework/deps/file_path.h"
@@ -65,7 +65,7 @@ const std::string MaybeConvertSignatureToTag(
output.resize(name.length()); output.resize(name.length());
std::transform(name.begin(), name.end(), output.begin(), std::transform(name.begin(), name.end(), output.begin(),
[](unsigned char c) { return std::toupper(c); }); [](unsigned char c) { return std::toupper(c); });
output = absl::Substitute(output, "/", "_"); output = absl::StrReplaceAll(output, {{"/", "_"}});
return output; return output;
} else { } else {
return name; return name;
@@ -135,7 +135,7 @@ class TensorFlowSessionFromSavedModelGenerator : public PacketGenerator {
if (!status.ok()) { if (!status.ok()) {
return ::mediapipe::Status( return ::mediapipe::Status(
static_cast<::mediapipe::StatusCode>(status.code()), static_cast<::mediapipe::StatusCode>(status.code()),
status.error_message()); status.ToString());
} }
auto session = absl::make_unique<TensorFlowSession>(); auto session = absl::make_unique<TensorFlowSession>();
@@ -12,7 +12,7 @@
// See the License for the specific language governing permissions and // See the License for the specific language governing permissions and
// limitations under the License. // limitations under the License.
#include "absl/strings/substitute.h" #include "absl/strings/str_replace.h"
#include "mediapipe/calculators/tensorflow/tensorflow_session.h" #include "mediapipe/calculators/tensorflow/tensorflow_session.h"
#include "mediapipe/calculators/tensorflow/tensorflow_session_from_saved_model_generator.pb.h" #include "mediapipe/calculators/tensorflow/tensorflow_session_from_saved_model_generator.pb.h"
#include "mediapipe/framework/calculator_framework.h" #include "mediapipe/framework/calculator_framework.h"
@@ -81,11 +81,11 @@ class TFRecordReaderCalculator : public CalculatorBase {
auto tf_status = tensorflow::Env::Default()->NewRandomAccessFile( auto tf_status = tensorflow::Env::Default()->NewRandomAccessFile(
cc->InputSidePackets().Tag(kTFRecordPath).Get<std::string>(), &file); cc->InputSidePackets().Tag(kTFRecordPath).Get<std::string>(), &file);
RET_CHECK(tf_status.ok()) RET_CHECK(tf_status.ok())
<< "Failed to open tfrecord file: " << tf_status.error_message(); << "Failed to open tfrecord file: " << tf_status.ToString();
tensorflow::io::RecordReader reader(file.get(), tensorflow::io::RecordReader reader(file.get(),
tensorflow::io::RecordReaderOptions()); tensorflow::io::RecordReaderOptions());
tensorflow::uint64 offset = 0; tensorflow::uint64 offset = 0;
std::string example_str; tensorflow::tstring example_str;
const int target_idx = const int target_idx =
cc->InputSidePackets().HasTag(kRecordIndex) cc->InputSidePackets().HasTag(kRecordIndex)
? cc->InputSidePackets().Tag(kRecordIndex).Get<int>() ? cc->InputSidePackets().Tag(kRecordIndex).Get<int>()
@@ -94,11 +94,11 @@ class TFRecordReaderCalculator : public CalculatorBase {
while (current_idx <= target_idx) { while (current_idx <= target_idx) {
tf_status = reader.ReadRecord(&offset, &example_str); tf_status = reader.ReadRecord(&offset, &example_str);
RET_CHECK(tf_status.ok()) RET_CHECK(tf_status.ok())
<< "Failed to read tfrecord: " << tf_status.error_message(); << "Failed to read tfrecord: " << tf_status.ToString();
if (current_idx == target_idx) { if (current_idx == target_idx) {
if (cc->OutputSidePackets().HasTag(kExampleTag)) { if (cc->OutputSidePackets().HasTag(kExampleTag)) {
tensorflow::Example tf_example; tensorflow::Example tf_example;
tf_example.ParseFromString(example_str); tf_example.ParseFromArray(example_str.data(), example_str.size());
cc->OutputSidePackets() cc->OutputSidePackets()
.Tag(kExampleTag) .Tag(kExampleTag)
.Set(MakePacket<tensorflow::Example>(std::move(tf_example))); .Set(MakePacket<tensorflow::Example>(std::move(tf_example)));
+8 -2
View File
@@ -13,12 +13,12 @@
# limitations under the License. # limitations under the License.
# #
load("//mediapipe/framework/port:build_config.bzl", "mediapipe_cc_proto_library")
licenses(["notice"]) # Apache 2.0 licenses(["notice"]) # Apache 2.0
package(default_visibility = ["//visibility:private"]) package(default_visibility = ["//visibility:private"])
load("//mediapipe/framework/port:build_config.bzl", "mediapipe_cc_proto_library")
proto_library( proto_library(
name = "ssd_anchors_calculator_proto", name = "ssd_anchors_calculator_proto",
srcs = ["ssd_anchors_calculator.proto"], srcs = ["ssd_anchors_calculator.proto"],
@@ -249,6 +249,11 @@ cc_library(
"@org_tensorflow//tensorflow/lite/delegates/gpu/gl:gl_program", "@org_tensorflow//tensorflow/lite/delegates/gpu/gl:gl_program",
"@org_tensorflow//tensorflow/lite/delegates/gpu/gl:gl_shader", "@org_tensorflow//tensorflow/lite/delegates/gpu/gl:gl_shader",
], ],
}) + select({
"//conditions:default": [],
"//mediapipe:android": [
"@org_tensorflow//tensorflow/lite/delegates/nnapi:nnapi_delegate",
],
}), }),
alwayslink = 1, alwayslink = 1,
) )
@@ -480,6 +485,7 @@ cc_test(
"//mediapipe/framework/port:integral_types", "//mediapipe/framework/port:integral_types",
"//mediapipe/framework/port:parse_text_proto", "//mediapipe/framework/port:parse_text_proto",
"//mediapipe/framework/tool:validate_type", "//mediapipe/framework/tool:validate_type",
"@com_google_absl//absl/strings",
"@org_tensorflow//tensorflow/lite:framework", "@org_tensorflow//tensorflow/lite:framework",
"@org_tensorflow//tensorflow/lite/kernels:builtin_ops", "@org_tensorflow//tensorflow/lite/kernels:builtin_ops",
], ],
@@ -25,8 +25,7 @@
#include "tensorflow/lite/error_reporter.h" #include "tensorflow/lite/error_reporter.h"
#include "tensorflow/lite/interpreter.h" #include "tensorflow/lite/interpreter.h"
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
#include "mediapipe/gpu/gl_calculator_helper.h" #include "mediapipe/gpu/gl_calculator_helper.h"
#include "mediapipe/gpu/gpu_buffer.h" #include "mediapipe/gpu/gpu_buffer.h"
#include "tensorflow/lite/delegates/gpu/gl/gl_buffer.h" #include "tensorflow/lite/delegates/gpu/gl/gl_buffer.h"
@@ -35,7 +34,7 @@
#include "tensorflow/lite/delegates/gpu/gl_delegate.h" #include "tensorflow/lite/delegates/gpu/gl_delegate.h"
#endif // !MEDIAPIPE_DISABLE_GPU #endif // !MEDIAPIPE_DISABLE_GPU
#if defined(__APPLE__) && !TARGET_OS_OSX // iOS #if defined(MEDIAPIPE_IOS)
#import <CoreVideo/CoreVideo.h> #import <CoreVideo/CoreVideo.h>
#import <Metal/Metal.h> #import <Metal/Metal.h>
#import <MetalKit/MetalKit.h> #import <MetalKit/MetalKit.h>
@@ -46,10 +45,9 @@
#include "tensorflow/lite/delegates/gpu/metal_delegate.h" #include "tensorflow/lite/delegates/gpu/metal_delegate.h"
#endif // iOS #endif // iOS
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
typedef ::tflite::gpu::gl::GlBuffer GpuTensor; typedef ::tflite::gpu::gl::GlBuffer GpuTensor;
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS #elif defined(MEDIAPIPE_IOS)
typedef id<MTLBuffer> GpuTensor; typedef id<MTLBuffer> GpuTensor;
#endif #endif
@@ -69,8 +67,7 @@ typedef Eigen::Matrix<float, Eigen::Dynamic, Eigen::Dynamic, Eigen::ColMajor>
namespace mediapipe { namespace mediapipe {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
using ::tflite::gpu::gl::CreateReadWriteShaderStorageBuffer; using ::tflite::gpu::gl::CreateReadWriteShaderStorageBuffer;
using ::tflite::gpu::gl::GlProgram; using ::tflite::gpu::gl::GlProgram;
using ::tflite::gpu::gl::GlShader; using ::tflite::gpu::gl::GlShader;
@@ -80,7 +77,7 @@ struct GPUData {
GlShader shader; GlShader shader;
GlProgram program; GlProgram program;
}; };
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS #elif defined(MEDIAPIPE_IOS)
struct GPUData { struct GPUData {
int elements = 1; int elements = 1;
GpuTensor buffer; GpuTensor buffer;
@@ -149,11 +146,10 @@ class TfLiteConverterCalculator : public CalculatorBase {
std::unique_ptr<tflite::Interpreter> interpreter_ = nullptr; std::unique_ptr<tflite::Interpreter> interpreter_ = nullptr;
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
mediapipe::GlCalculatorHelper gpu_helper_; mediapipe::GlCalculatorHelper gpu_helper_;
std::unique_ptr<GPUData> gpu_data_out_; std::unique_ptr<GPUData> gpu_data_out_;
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS #elif defined(MEDIAPIPE_IOS)
MPPMetalHelper* gpu_helper_ = nullptr; MPPMetalHelper* gpu_helper_ = nullptr;
std::unique_ptr<GPUData> gpu_data_out_; std::unique_ptr<GPUData> gpu_data_out_;
#endif #endif
@@ -202,10 +198,9 @@ REGISTER_CALCULATOR(TfLiteConverterCalculator);
#endif // !MEDIAPIPE_DISABLE_GPU #endif // !MEDIAPIPE_DISABLE_GPU
if (use_gpu) { if (use_gpu) {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
MP_RETURN_IF_ERROR(mediapipe::GlCalculatorHelper::UpdateContract(cc)); MP_RETURN_IF_ERROR(mediapipe::GlCalculatorHelper::UpdateContract(cc));
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS #elif defined(MEDIAPIPE_IOS)
MP_RETURN_IF_ERROR([MPPMetalHelper updateContract:cc]); MP_RETURN_IF_ERROR([MPPMetalHelper updateContract:cc]);
#endif #endif
} }
@@ -236,10 +231,9 @@ REGISTER_CALCULATOR(TfLiteConverterCalculator);
cc->Outputs().HasTag("TENSORS_GPU")); cc->Outputs().HasTag("TENSORS_GPU"));
// Cannot use quantization. // Cannot use quantization.
use_quantized_tensors_ = false; use_quantized_tensors_ = false;
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
MP_RETURN_IF_ERROR(gpu_helper_.Open(cc)); MP_RETURN_IF_ERROR(gpu_helper_.Open(cc));
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS #elif defined(MEDIAPIPE_IOS)
gpu_helper_ = [[MPPMetalHelper alloc] initWithCalculatorContext:cc]; gpu_helper_ = [[MPPMetalHelper alloc] initWithCalculatorContext:cc];
RET_CHECK(gpu_helper_); RET_CHECK(gpu_helper_);
#endif #endif
@@ -270,11 +264,10 @@ REGISTER_CALCULATOR(TfLiteConverterCalculator);
} }
::mediapipe::Status TfLiteConverterCalculator::Close(CalculatorContext* cc) { ::mediapipe::Status TfLiteConverterCalculator::Close(CalculatorContext* cc) {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
gpu_helper_.RunInGlContext([this] { gpu_data_out_.reset(); }); gpu_helper_.RunInGlContext([this] { gpu_data_out_.reset(); });
#endif #endif
#if defined(__APPLE__) && !TARGET_OS_OSX // iOS #if defined(MEDIAPIPE_IOS)
gpu_data_out_.reset(); gpu_data_out_.reset();
#endif #endif
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
@@ -301,11 +294,15 @@ REGISTER_CALCULATOR(TfLiteConverterCalculator);
if (use_quantized_tensors_) { if (use_quantized_tensors_) {
RET_CHECK(image_frame.Format() != mediapipe::ImageFormat::VEC32F1) RET_CHECK(image_frame.Format() != mediapipe::ImageFormat::VEC32F1)
<< "Only 8-bit input images are supported for quantization."; << "Only 8-bit input images are supported for quantization.";
quant.type = kTfLiteAffineQuantization;
quant.params = nullptr;
// Optional: Set 'quant' quantization params here if needed. // Optional: Set 'quant' quantization params here if needed.
interpreter_->SetTensorParametersReadWrite(0, kTfLiteUInt8, "", interpreter_->SetTensorParametersReadWrite(0, kTfLiteUInt8, "",
{channels_preserved}, quant); {channels_preserved}, quant);
} else { } else {
// Default TfLiteQuantization used for no quantization. // Initialize structure for no quantization.
quant.type = kTfLiteNoQuantization;
quant.params = nullptr;
interpreter_->SetTensorParametersReadWrite(0, kTfLiteFloat32, "", interpreter_->SetTensorParametersReadWrite(0, kTfLiteFloat32, "",
{channels_preserved}, quant); {channels_preserved}, quant);
} }
@@ -390,8 +387,7 @@ REGISTER_CALCULATOR(TfLiteConverterCalculator);
::mediapipe::Status TfLiteConverterCalculator::ProcessGPU( ::mediapipe::Status TfLiteConverterCalculator::ProcessGPU(
CalculatorContext* cc) { CalculatorContext* cc) {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
// GpuBuffer to tflite::gpu::GlBuffer conversion. // GpuBuffer to tflite::gpu::GlBuffer conversion.
const auto& input = cc->Inputs().Tag("IMAGE_GPU").Get<mediapipe::GpuBuffer>(); const auto& input = cc->Inputs().Tag("IMAGE_GPU").Get<mediapipe::GpuBuffer>();
MP_RETURN_IF_ERROR( MP_RETURN_IF_ERROR(
@@ -427,43 +423,38 @@ REGISTER_CALCULATOR(TfLiteConverterCalculator);
cc->Outputs() cc->Outputs()
.Tag("TENSORS_GPU") .Tag("TENSORS_GPU")
.Add(output_tensors.release(), cc->InputTimestamp()); .Add(output_tensors.release(), cc->InputTimestamp());
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS #elif defined(MEDIAPIPE_IOS)
// GpuBuffer to id<MTLBuffer> conversion. // GpuBuffer to id<MTLBuffer> conversion.
const auto& input = cc->Inputs().Tag("IMAGE_GPU").Get<mediapipe::GpuBuffer>(); const auto& input = cc->Inputs().Tag("IMAGE_GPU").Get<mediapipe::GpuBuffer>();
{ id<MTLCommandBuffer> command_buffer = [gpu_helper_ commandBuffer];
id<MTLTexture> src_texture = [gpu_helper_ metalTextureWithGpuBuffer:input];
id<MTLCommandBuffer> command_buffer = [gpu_helper_ commandBuffer]; id<MTLTexture> src_texture = [gpu_helper_ metalTextureWithGpuBuffer:input];
command_buffer.label = @"TfLiteConverterCalculatorConvert"; command_buffer.label = @"TfLiteConverterCalculatorConvertAndBlit";
id<MTLComputeCommandEncoder> compute_encoder = id<MTLComputeCommandEncoder> compute_encoder =
[command_buffer computeCommandEncoder]; [command_buffer computeCommandEncoder];
[compute_encoder setComputePipelineState:gpu_data_out_->pipeline_state]; [compute_encoder setComputePipelineState:gpu_data_out_->pipeline_state];
[compute_encoder setTexture:src_texture atIndex:0]; [compute_encoder setTexture:src_texture atIndex:0];
[compute_encoder setBuffer:gpu_data_out_->buffer offset:0 atIndex:1]; [compute_encoder setBuffer:gpu_data_out_->buffer offset:0 atIndex:1];
MTLSize threads_per_group = MTLSizeMake(kWorkgroupSize, kWorkgroupSize, 1); MTLSize threads_per_group = MTLSizeMake(kWorkgroupSize, kWorkgroupSize, 1);
MTLSize threadgroups = MTLSize threadgroups =
MTLSizeMake(NumGroups(input.width(), kWorkgroupSize), MTLSizeMake(NumGroups(input.width(), kWorkgroupSize),
NumGroups(input.height(), kWorkgroupSize), 1); NumGroups(input.height(), kWorkgroupSize), 1);
[compute_encoder dispatchThreadgroups:threadgroups [compute_encoder dispatchThreadgroups:threadgroups
threadsPerThreadgroup:threads_per_group]; threadsPerThreadgroup:threads_per_group];
[compute_encoder endEncoding]; [compute_encoder endEncoding];
[command_buffer commit];
[command_buffer waitUntilCompleted];
}
// Copy into outputs. // Copy into outputs.
// TODO Avoid this copy. // TODO Avoid this copy.
auto output_tensors = absl::make_unique<std::vector<GpuTensor>>(); auto output_tensors = absl::make_unique<std::vector<GpuTensor>>();
output_tensors->resize(1); output_tensors->resize(1);
{ id<MTLDevice> device = gpu_helper_.mtlDevice;
id<MTLDevice> device = gpu_helper_.mtlDevice; output_tensors->at(0) =
output_tensors->at(0) = [device newBufferWithLength:gpu_data_out_->elements * sizeof(float)
[device newBufferWithLength:gpu_data_out_->elements * sizeof(float) options:MTLResourceStorageModeShared];
options:MTLResourceStorageModeShared]; [MPPMetalUtil blitMetalBufferTo:output_tensors->at(0)
[MPPMetalUtil blitMetalBufferTo:output_tensors->at(0) from:gpu_data_out_->buffer
from:gpu_data_out_->buffer blocking:false
blocking:true commandBuffer:command_buffer];
commandBuffer:[gpu_helper_ commandBuffer]];
}
cc->Outputs() cc->Outputs()
.Tag("TENSORS_GPU") .Tag("TENSORS_GPU")
@@ -493,8 +484,7 @@ REGISTER_CALCULATOR(TfLiteConverterCalculator);
RET_CHECK_FAIL() << "Num input channels is less than desired output."; RET_CHECK_FAIL() << "Num input channels is less than desired output.";
#endif // !MEDIAPIPE_DISABLE_GPU #endif // !MEDIAPIPE_DISABLE_GPU
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext( MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext(
[this, &include_alpha, &input, &single_channel]() -> ::mediapipe::Status { [this, &include_alpha, &input, &single_channel]() -> ::mediapipe::Status {
// Device memory. // Device memory.
@@ -538,7 +528,9 @@ REGISTER_CALCULATOR(TfLiteConverterCalculator);
&gpu_data_out_->program)); &gpu_data_out_->program));
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
})); }));
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS
#elif defined(MEDIAPIPE_IOS)
RET_CHECK(include_alpha) RET_CHECK(include_alpha)
<< "iOS GPU inference currently accepts only RGBA input."; << "iOS GPU inference currently accepts only RGBA input.";
@@ -619,7 +611,7 @@ REGISTER_CALCULATOR(TfLiteConverterCalculator);
CHECK_GE(max_num_channels_, 1); CHECK_GE(max_num_channels_, 1);
CHECK_LE(max_num_channels_, 4); CHECK_LE(max_num_channels_, 4);
CHECK_NE(max_num_channels_, 2); CHECK_NE(max_num_channels_, 2);
#if defined(__APPLE__) && !TARGET_OS_OSX // iOS #if defined(MEDIAPIPE_IOS)
if (cc->Inputs().HasTag("IMAGE_GPU")) if (cc->Inputs().HasTag("IMAGE_GPU"))
// Currently on iOS, tflite gpu input tensor must be 4 channels, // Currently on iOS, tflite gpu input tensor must be 4 channels,
// so input image must be 4 channels also (checked in InitGpu). // so input image must be 4 channels also (checked in InitGpu).
@@ -27,7 +27,7 @@
#include "tensorflow/lite/kernels/register.h" #include "tensorflow/lite/kernels/register.h"
#include "tensorflow/lite/model.h" #include "tensorflow/lite/model.h"
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(MEDIAPIPE_DISABLE_GL_COMPUTE) #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
#include "mediapipe/gpu/gl_calculator_helper.h" #include "mediapipe/gpu/gl_calculator_helper.h"
#include "mediapipe/gpu/gpu_buffer.h" #include "mediapipe/gpu/gpu_buffer.h"
#include "tensorflow/lite/delegates/gpu/common/shape.h" #include "tensorflow/lite/delegates/gpu/common/shape.h"
@@ -35,9 +35,9 @@
#include "tensorflow/lite/delegates/gpu/gl/gl_program.h" #include "tensorflow/lite/delegates/gpu/gl/gl_program.h"
#include "tensorflow/lite/delegates/gpu/gl/gl_shader.h" #include "tensorflow/lite/delegates/gpu/gl/gl_shader.h"
#include "tensorflow/lite/delegates/gpu/gl_delegate.h" #include "tensorflow/lite/delegates/gpu/gl_delegate.h"
#endif // !MEDIAPIPE_DISABLE_GPU #endif // !MEDIAPIPE_DISABLE_GL_COMPUTE
#if defined(__APPLE__) && !TARGET_OS_OSX // iOS #if defined(MEDIAPIPE_IOS)
#import <CoreVideo/CoreVideo.h> #import <CoreVideo/CoreVideo.h>
#import <Metal/Metal.h> #import <Metal/Metal.h>
#import <MetalKit/MetalKit.h> #import <MetalKit/MetalKit.h>
@@ -51,12 +51,19 @@
#include "tensorflow/lite/delegates/gpu/metal_delegate_internal.h" #include "tensorflow/lite/delegates/gpu/metal_delegate_internal.h"
#endif // iOS #endif // iOS
namespace { #if defined(MEDIAPIPE_ANDROID)
#include "tensorflow/lite/delegates/nnapi/nnapi_delegate.h"
#endif // ANDROID
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ namespace {
!defined(__APPLE__) // Commonly used to compute the number of blocks to launch in a kernel.
int NumGroups(const int size, const int group_size) { // NOLINT
return (size + group_size - 1) / group_size;
}
#if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
typedef ::tflite::gpu::gl::GlBuffer GpuTensor; typedef ::tflite::gpu::gl::GlBuffer GpuTensor;
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS #elif defined(MEDIAPIPE_IOS)
typedef id<MTLBuffer> GpuTensor; typedef id<MTLBuffer> GpuTensor;
#endif #endif
@@ -64,14 +71,35 @@ typedef id<MTLBuffer> GpuTensor;
size_t RoundUp(size_t n, size_t m) { return ((n + m - 1) / m) * m; } // NOLINT size_t RoundUp(size_t n, size_t m) { return ((n + m - 1) / m) * m; } // NOLINT
} // namespace } // namespace
#if defined(MEDIAPIPE_EDGE_TPU)
#include "edgetpu.h"
// Creates and returns an Edge TPU interpreter to run the given edgetpu model.
std::unique_ptr<tflite::Interpreter> BuildEdgeTpuInterpreter(
const tflite::FlatBufferModel& model,
tflite::ops::builtin::BuiltinOpResolver* resolver,
edgetpu::EdgeTpuContext* edgetpu_context) {
resolver->AddCustom(edgetpu::kCustomOp, edgetpu::RegisterCustomOp());
std::unique_ptr<tflite::Interpreter> interpreter;
if (tflite::InterpreterBuilder(model, *resolver)(&interpreter) != kTfLiteOk) {
std::cerr << "Failed to build edge TPU interpreter." << std::endl;
}
interpreter->SetExternalContext(kTfLiteEdgeTpuContext, edgetpu_context);
interpreter->SetNumThreads(1);
if (interpreter->AllocateTensors() != kTfLiteOk) {
std::cerr << "Failed to allocate edge TPU tensors." << std::endl;
}
return interpreter;
}
#endif // MEDIAPIPE_EDGE_TPU
// TfLiteInferenceCalculator File Layout: // TfLiteInferenceCalculator File Layout:
// * Header // * Header
// * Core // * Core
// * Aux // * Aux
namespace mediapipe { namespace mediapipe {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
using ::tflite::gpu::gl::CopyBuffer; using ::tflite::gpu::gl::CopyBuffer;
using ::tflite::gpu::gl::CreateReadWriteShaderStorageBuffer; using ::tflite::gpu::gl::CreateReadWriteShaderStorageBuffer;
using ::tflite::gpu::gl::GlBuffer; using ::tflite::gpu::gl::GlBuffer;
@@ -120,7 +148,7 @@ struct GPUData {
// options: { // options: {
// [mediapipe.TfLiteInferenceCalculatorOptions.ext] { // [mediapipe.TfLiteInferenceCalculatorOptions.ext] {
// model_path: "modelname.tflite" // model_path: "modelname.tflite"
// use_gpu: true // delegate { gpu {} }
// } // }
// } // }
// } // }
@@ -135,6 +163,9 @@ struct GPUData {
// //
class TfLiteInferenceCalculator : public CalculatorBase { class TfLiteInferenceCalculator : public CalculatorBase {
public: public:
using TfLiteDelegatePtr =
std::unique_ptr<TfLiteDelegate, std::function<void(TfLiteDelegate*)>>;
static ::mediapipe::Status GetContract(CalculatorContract* cc); static ::mediapipe::Status GetContract(CalculatorContract* cc);
::mediapipe::Status Open(CalculatorContext* cc) override; ::mediapipe::Status Open(CalculatorContext* cc) override;
@@ -148,20 +179,25 @@ class TfLiteInferenceCalculator : public CalculatorBase {
std::unique_ptr<tflite::Interpreter> interpreter_; std::unique_ptr<tflite::Interpreter> interpreter_;
std::unique_ptr<tflite::FlatBufferModel> model_; std::unique_ptr<tflite::FlatBufferModel> model_;
TfLiteDelegate* delegate_ = nullptr; TfLiteDelegatePtr delegate_;
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
mediapipe::GlCalculatorHelper gpu_helper_; mediapipe::GlCalculatorHelper gpu_helper_;
std::unique_ptr<GPUData> gpu_data_in_; std::vector<std::unique_ptr<GPUData>> gpu_data_in_;
std::vector<std::unique_ptr<GPUData>> gpu_data_out_; std::vector<std::unique_ptr<GPUData>> gpu_data_out_;
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS #elif defined(MEDIAPIPE_IOS)
MPPMetalHelper* gpu_helper_ = nullptr; MPPMetalHelper* gpu_helper_ = nullptr;
std::unique_ptr<GPUData> gpu_data_in_; std::vector<std::unique_ptr<GPUData>> gpu_data_in_;
std::vector<std::unique_ptr<GPUData>> gpu_data_out_; std::vector<std::unique_ptr<GPUData>> gpu_data_out_;
id<MTLComputePipelineState> fp32_to_fp16_program_;
TFLBufferConvert* converter_from_BPHWC4_ = nil; TFLBufferConvert* converter_from_BPHWC4_ = nil;
#endif #endif
#if defined(MEDIAPIPE_EDGE_TPU)
std::shared_ptr<edgetpu::EdgeTpuContext> edgetpu_context_ =
edgetpu::EdgeTpuManager::GetSingleton()->OpenDevice();
#endif
std::string model_path_ = ""; std::string model_path_ = "";
bool gpu_inference_ = false; bool gpu_inference_ = false;
bool gpu_input_ = false; bool gpu_input_ = false;
@@ -179,12 +215,18 @@ REGISTER_CALCULATOR(TfLiteInferenceCalculator);
RET_CHECK(cc->Outputs().HasTag("TENSORS") ^ RET_CHECK(cc->Outputs().HasTag("TENSORS") ^
cc->Outputs().HasTag("TENSORS_GPU")); cc->Outputs().HasTag("TENSORS_GPU"));
bool use_gpu = false; const auto& options =
cc->Options<::mediapipe::TfLiteInferenceCalculatorOptions>();
bool use_gpu =
options.has_delegate() ? options.delegate().has_gpu() : options.use_gpu();
if (cc->Inputs().HasTag("TENSORS")) if (cc->Inputs().HasTag("TENSORS"))
cc->Inputs().Tag("TENSORS").Set<std::vector<TfLiteTensor>>(); cc->Inputs().Tag("TENSORS").Set<std::vector<TfLiteTensor>>();
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) #if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__)
if (cc->Inputs().HasTag("TENSORS_GPU")) { if (cc->Inputs().HasTag("TENSORS_GPU")) {
RET_CHECK(!options.has_delegate() || options.delegate().has_gpu())
<< "GPU input is compatible with GPU delegate only.";
cc->Inputs().Tag("TENSORS_GPU").Set<std::vector<GpuTensor>>(); cc->Inputs().Tag("TENSORS_GPU").Set<std::vector<GpuTensor>>();
use_gpu |= true; use_gpu |= true;
} }
@@ -194,6 +236,9 @@ REGISTER_CALCULATOR(TfLiteInferenceCalculator);
cc->Outputs().Tag("TENSORS").Set<std::vector<TfLiteTensor>>(); cc->Outputs().Tag("TENSORS").Set<std::vector<TfLiteTensor>>();
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) #if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__)
if (cc->Outputs().HasTag("TENSORS_GPU")) { if (cc->Outputs().HasTag("TENSORS_GPU")) {
RET_CHECK(!options.has_delegate() || options.delegate().has_gpu())
<< "GPU output is compatible with GPU delegate only.";
cc->Outputs().Tag("TENSORS_GPU").Set<std::vector<GpuTensor>>(); cc->Outputs().Tag("TENSORS_GPU").Set<std::vector<GpuTensor>>();
use_gpu |= true; use_gpu |= true;
} }
@@ -205,15 +250,10 @@ REGISTER_CALCULATOR(TfLiteInferenceCalculator);
.Set<tflite::ops::builtin::BuiltinOpResolver>(); .Set<tflite::ops::builtin::BuiltinOpResolver>();
} }
const auto& options =
cc->Options<::mediapipe::TfLiteInferenceCalculatorOptions>();
use_gpu |= options.use_gpu();
if (use_gpu) { if (use_gpu) {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
MP_RETURN_IF_ERROR(mediapipe::GlCalculatorHelper::UpdateContract(cc)); MP_RETURN_IF_ERROR(mediapipe::GlCalculatorHelper::UpdateContract(cc));
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS #elif defined(MEDIAPIPE_IOS)
MP_RETURN_IF_ERROR([MPPMetalHelper updateContract:cc]); MP_RETURN_IF_ERROR([MPPMetalHelper updateContract:cc]);
#endif #endif
} }
@@ -253,26 +293,24 @@ REGISTER_CALCULATOR(TfLiteInferenceCalculator);
MP_RETURN_IF_ERROR(LoadModel(cc)); MP_RETURN_IF_ERROR(LoadModel(cc));
if (gpu_inference_) { if (gpu_inference_) {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
MP_RETURN_IF_ERROR(gpu_helper_.Open(cc)); MP_RETURN_IF_ERROR(gpu_helper_.Open(cc));
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS #elif defined(MEDIAPIPE_IOS)
gpu_helper_ = [[MPPMetalHelper alloc] initWithCalculatorContext:cc]; gpu_helper_ = [[MPPMetalHelper alloc] initWithCalculatorContext:cc];
RET_CHECK(gpu_helper_); RET_CHECK(gpu_helper_);
#endif #endif
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \
!defined(__APPLE__) #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext( MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext(
[this, &cc]() -> ::mediapipe::Status { return LoadDelegate(cc); })); [this, &cc]() -> ::mediapipe::Status { return LoadDelegate(cc); }));
#else #else
MP_RETURN_IF_ERROR(LoadDelegate(cc)); MP_RETURN_IF_ERROR(LoadDelegate(cc));
#endif #endif
} else {
#if defined(__EMSCRIPTEN__) || defined(MEDIAPIPE_ANDROID)
MP_RETURN_IF_ERROR(LoadDelegate(cc));
#endif // __EMSCRIPTEN__ || ANDROID
} }
#if defined(__EMSCRIPTEN__)
MP_RETURN_IF_ERROR(LoadDelegate(cc));
#endif // __EMSCRIPTEN__
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -280,26 +318,44 @@ REGISTER_CALCULATOR(TfLiteInferenceCalculator);
// 1. Receive pre-processed tensor inputs. // 1. Receive pre-processed tensor inputs.
if (gpu_input_) { if (gpu_input_) {
// Read GPU input into SSBO. // Read GPU input into SSBO.
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
const auto& input_tensors = const auto& input_tensors =
cc->Inputs().Tag("TENSORS_GPU").Get<std::vector<GpuTensor>>(); cc->Inputs().Tag("TENSORS_GPU").Get<std::vector<GpuTensor>>();
RET_CHECK_EQ(input_tensors.size(), 1); RET_CHECK_GT(input_tensors.size(), 0);
MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext( MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext(
[this, &input_tensors]() -> ::mediapipe::Status { [this, &input_tensors]() -> ::mediapipe::Status {
// Explicit copy input. // Explicit copy input.
RET_CHECK_CALL(CopyBuffer(input_tensors[0], gpu_data_in_->buffer)); gpu_data_in_.resize(input_tensors.size());
for (int i = 0; i < input_tensors.size(); ++i) {
RET_CHECK_CALL(
CopyBuffer(input_tensors[i], gpu_data_in_[i]->buffer));
}
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
})); }));
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS #elif defined(MEDIAPIPE_IOS)
const auto& input_tensors = const auto& input_tensors =
cc->Inputs().Tag("TENSORS_GPU").Get<std::vector<GpuTensor>>(); cc->Inputs().Tag("TENSORS_GPU").Get<std::vector<GpuTensor>>();
RET_CHECK_EQ(input_tensors.size(), 1); RET_CHECK_GT(input_tensors.size(), 0);
// Explicit copy input. // Explicit copy input with conversion float 32 bits to 16 bits.
[MPPMetalUtil blitMetalBufferTo:gpu_data_in_->buffer gpu_data_in_.resize(input_tensors.size());
from:input_tensors[0] id<MTLCommandBuffer> command_buffer = [gpu_helper_ commandBuffer];
blocking:true command_buffer.label = @"TfLiteInferenceCalculatorConvert";
commandBuffer:[gpu_helper_ commandBuffer]]; id<MTLComputeCommandEncoder> compute_encoder =
[command_buffer computeCommandEncoder];
[compute_encoder setComputePipelineState:fp32_to_fp16_program_];
for (int i = 0; i < input_tensors.size(); ++i) {
[compute_encoder setBuffer:input_tensors[i] offset:0 atIndex:0];
[compute_encoder setBuffer:gpu_data_in_[i]->buffer offset:0 atIndex:1];
constexpr int kWorkgroupSize = 64; // Block size for GPU shader.
MTLSize threads_per_group = MTLSizeMake(kWorkgroupSize, 1, 1);
const int threadgroups =
NumGroups(gpu_data_in_[i]->elements, kWorkgroupSize);
[compute_encoder dispatchThreadgroups:MTLSizeMake(threadgroups, 1, 1)
threadsPerThreadgroup:threads_per_group];
}
[compute_encoder endEncoding];
[command_buffer commit];
#else #else
RET_CHECK_FAIL() << "GPU processing not enabled."; RET_CHECK_FAIL() << "GPU processing not enabled.";
#endif #endif
@@ -327,14 +383,13 @@ REGISTER_CALCULATOR(TfLiteInferenceCalculator);
// 2. Run inference. // 2. Run inference.
if (gpu_inference_) { if (gpu_inference_) {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
MP_RETURN_IF_ERROR( MP_RETURN_IF_ERROR(
gpu_helper_.RunInGlContext([this]() -> ::mediapipe::Status { gpu_helper_.RunInGlContext([this]() -> ::mediapipe::Status {
RET_CHECK_EQ(interpreter_->Invoke(), kTfLiteOk); RET_CHECK_EQ(interpreter_->Invoke(), kTfLiteOk);
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
})); }));
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS #elif defined(MEDIAPIPE_IOS)
RET_CHECK_EQ(interpreter_->Invoke(), kTfLiteOk); RET_CHECK_EQ(interpreter_->Invoke(), kTfLiteOk);
#endif #endif
} else { } else {
@@ -343,8 +398,7 @@ REGISTER_CALCULATOR(TfLiteInferenceCalculator);
// 3. Output processed tensors. // 3. Output processed tensors.
if (gpu_output_) { if (gpu_output_) {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
// Output result tensors (GPU). // Output result tensors (GPU).
auto output_tensors = absl::make_unique<std::vector<GpuTensor>>(); auto output_tensors = absl::make_unique<std::vector<GpuTensor>>();
MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext( MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext(
@@ -361,7 +415,7 @@ REGISTER_CALCULATOR(TfLiteInferenceCalculator);
cc->Outputs() cc->Outputs()
.Tag("TENSORS_GPU") .Tag("TENSORS_GPU")
.Add(output_tensors.release(), cc->InputTimestamp()); .Add(output_tensors.release(), cc->InputTimestamp());
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS #elif defined(MEDIAPIPE_IOS)
// Output result tensors (GPU). // Output result tensors (GPU).
auto output_tensors = absl::make_unique<std::vector<GpuTensor>>(); auto output_tensors = absl::make_unique<std::vector<GpuTensor>>();
output_tensors->resize(gpu_data_out_.size()); output_tensors->resize(gpu_data_out_.size());
@@ -382,7 +436,6 @@ REGISTER_CALCULATOR(TfLiteInferenceCalculator);
} }
[convert_command endEncoding]; [convert_command endEncoding];
[command_buffer commit]; [command_buffer commit];
[command_buffer waitUntilCompleted];
cc->Outputs() cc->Outputs()
.Tag("TENSORS_GPU") .Tag("TENSORS_GPU")
.Add(output_tensors.release(), cc->InputTimestamp()); .Add(output_tensors.release(), cc->InputTimestamp());
@@ -406,25 +459,34 @@ REGISTER_CALCULATOR(TfLiteInferenceCalculator);
::mediapipe::Status TfLiteInferenceCalculator::Close(CalculatorContext* cc) { ::mediapipe::Status TfLiteInferenceCalculator::Close(CalculatorContext* cc) {
if (delegate_) { if (delegate_) {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ if (gpu_inference_) {
!defined(__APPLE__) #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext([this]() -> Status { MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext([this]() -> Status {
TfLiteGpuDelegateDelete(delegate_); delegate_ = nullptr;
gpu_data_in_.reset(); for (int i = 0; i < gpu_data_in_.size(); ++i) {
gpu_data_in_[i].reset();
}
for (int i = 0; i < gpu_data_out_.size(); ++i) {
gpu_data_out_[i].reset();
}
return ::mediapipe::OkStatus();
}));
#elif defined(MEDIAPIPE_IOS)
delegate_ = nullptr;
for (int i = 0; i < gpu_data_in_.size(); ++i) {
gpu_data_in_[i].reset();
}
for (int i = 0; i < gpu_data_out_.size(); ++i) { for (int i = 0; i < gpu_data_out_.size(); ++i) {
gpu_data_out_[i].reset(); gpu_data_out_[i].reset();
} }
return ::mediapipe::OkStatus();
}));
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS
TFLGpuDelegateDelete(delegate_);
gpu_data_in_.reset();
for (int i = 0; i < gpu_data_out_.size(); ++i) {
gpu_data_out_[i].reset();
}
#endif #endif
delegate_ = nullptr; } else {
delegate_ = nullptr;
}
} }
#if defined(MEDIAPIPE_EDGE_TPU)
edgetpu_context_.reset();
#endif
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -438,7 +500,7 @@ REGISTER_CALCULATOR(TfLiteInferenceCalculator);
// Get model name. // Get model name.
if (!options.model_path().empty()) { if (!options.model_path().empty()) {
auto model_path = options.model_path(); std::string model_path = options.model_path();
ASSIGN_OR_RETURN(model_path_, mediapipe::PathToResourceAsFile(model_path)); ASSIGN_OR_RETURN(model_path_, mediapipe::PathToResourceAsFile(model_path));
} else { } else {
@@ -448,7 +510,8 @@ REGISTER_CALCULATOR(TfLiteInferenceCalculator);
} }
// Get execution modes. // Get execution modes.
gpu_inference_ = options.use_gpu(); gpu_inference_ =
options.has_delegate() ? options.delegate().has_gpu() : options.use_gpu();
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -458,21 +521,27 @@ REGISTER_CALCULATOR(TfLiteInferenceCalculator);
model_ = tflite::FlatBufferModel::BuildFromFile(model_path_.c_str()); model_ = tflite::FlatBufferModel::BuildFromFile(model_path_.c_str());
RET_CHECK(model_); RET_CHECK(model_);
tflite::ops::builtin::BuiltinOpResolver op_resolver;
if (cc->InputSidePackets().HasTag("CUSTOM_OP_RESOLVER")) { if (cc->InputSidePackets().HasTag("CUSTOM_OP_RESOLVER")) {
const auto& op_resolver = op_resolver = cc->InputSidePackets()
cc->InputSidePackets() .Tag("CUSTOM_OP_RESOLVER")
.Tag("CUSTOM_OP_RESOLVER") .Get<tflite::ops::builtin::BuiltinOpResolver>();
.Get<tflite::ops::builtin::BuiltinOpResolver>();
tflite::InterpreterBuilder(*model_, op_resolver)(&interpreter_);
} else {
const tflite::ops::builtin::BuiltinOpResolver op_resolver;
tflite::InterpreterBuilder(*model_, op_resolver)(&interpreter_);
} }
#if defined(MEDIAPIPE_EDGE_TPU)
interpreter_ =
BuildEdgeTpuInterpreter(*model_, &op_resolver, edgetpu_context_.get());
#else
tflite::InterpreterBuilder(*model_, op_resolver)(&interpreter_);
#endif // MEDIAPIPE_EDGE_TPU
RET_CHECK(interpreter_); RET_CHECK(interpreter_);
#if defined(__EMSCRIPTEN__) #if defined(__EMSCRIPTEN__) || defined(MEDIAPIPE_EDGE_TPU)
interpreter_->SetNumThreads(1); interpreter_->SetNumThreads(1);
#else
interpreter_->SetNumThreads(
cc->Options<mediapipe::TfLiteInferenceCalculatorOptions>()
.cpu_num_thread());
#endif // __EMSCRIPTEN__ #endif // __EMSCRIPTEN__
if (gpu_output_) { if (gpu_output_) {
@@ -490,8 +559,39 @@ REGISTER_CALCULATOR(TfLiteInferenceCalculator);
::mediapipe::Status TfLiteInferenceCalculator::LoadDelegate( ::mediapipe::Status TfLiteInferenceCalculator::LoadDelegate(
CalculatorContext* cc) { CalculatorContext* cc) {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ const auto& calculator_opts =
!defined(__APPLE__) cc->Options<mediapipe::TfLiteInferenceCalculatorOptions>();
if (calculator_opts.has_delegate() &&
calculator_opts.delegate().has_tflite()) {
// Default tflite inference requeqsted - no need to modify graph.
return ::mediapipe::OkStatus();
}
if (!gpu_inference_) {
#if defined(MEDIAPIPE_ANDROID)
const bool nnapi_requested = calculator_opts.has_delegate()
? calculator_opts.delegate().has_nnapi()
: calculator_opts.use_nnapi();
if (nnapi_requested) {
// Attempt to use NNAPI.
// If not supported, the default CPU delegate will be created and used.
interpreter_->SetAllowFp16PrecisionForFp32(1);
delegate_ =
TfLiteDelegatePtr(tflite::NnApiDelegate(), [](TfLiteDelegate*) {
// No need to free according to tflite::NnApiDelegate()
// documentation.
});
RET_CHECK_EQ(interpreter_->ModifyGraphWithDelegate(delegate_.get()),
kTfLiteOk);
return ::mediapipe::OkStatus();
}
#endif // MEDIAPIPE_ANDROID
// Return, no need for GPU delegate below.
return ::mediapipe::OkStatus();
}
#if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
// Configure and create the delegate. // Configure and create the delegate.
TfLiteGpuDelegateOptions options = TfLiteGpuDelegateOptionsDefault(); TfLiteGpuDelegateOptions options = TfLiteGpuDelegateOptionsDefault();
options.compile_options.precision_loss_allowed = 1; options.compile_options.precision_loss_allowed = 1;
@@ -499,28 +599,30 @@ REGISTER_CALCULATOR(TfLiteInferenceCalculator);
TFLITE_GL_OBJECT_TYPE_FASTEST; TFLITE_GL_OBJECT_TYPE_FASTEST;
options.compile_options.dynamic_batch_enabled = 0; options.compile_options.dynamic_batch_enabled = 0;
options.compile_options.inline_parameters = 1; options.compile_options.inline_parameters = 1;
if (!delegate_) delegate_ = TfLiteGpuDelegateCreate(&options); if (!delegate_)
delegate_ = TfLiteDelegatePtr(TfLiteGpuDelegateCreate(&options),
&TfLiteGpuDelegateDelete);
if (gpu_input_) { if (gpu_input_) {
// Get input image sizes. // Get input image sizes.
gpu_data_in_ = absl::make_unique<GPUData>();
const auto& input_indices = interpreter_->inputs(); const auto& input_indices = interpreter_->inputs();
RET_CHECK_EQ(input_indices.size(), 1); // TODO accept > 1. gpu_data_in_.resize(input_indices.size());
const TfLiteTensor* tensor = interpreter_->tensor(input_indices[0]); for (int i = 0; i < input_indices.size(); ++i) {
gpu_data_in_->elements = 1; const TfLiteTensor* tensor = interpreter_->tensor(input_indices[0]);
for (int d = 0; d < tensor->dims->size; ++d) { gpu_data_in_[i] = absl::make_unique<GPUData>();
gpu_data_in_->elements *= tensor->dims->data[d]; gpu_data_in_[i]->elements = 1;
for (int d = 0; d < tensor->dims->size; ++d) {
gpu_data_in_[i]->elements *= tensor->dims->data[d];
}
// Create and bind input buffer.
RET_CHECK_CALL(
::tflite::gpu::gl::CreateReadWriteShaderStorageBuffer<float>(
gpu_data_in_[i]->elements, &gpu_data_in_[i]->buffer));
RET_CHECK_EQ(TfLiteGpuDelegateBindBufferToTensor(
delegate_.get(), gpu_data_in_[i]->buffer.id(),
interpreter_->inputs()[i]),
kTfLiteOk);
} }
CHECK_GE(tensor->dims->data[3], 1);
CHECK_LE(tensor->dims->data[3], 4);
CHECK_NE(tensor->dims->data[3], 2);
// Create and bind input buffer.
RET_CHECK_CALL(::tflite::gpu::gl::CreateReadWriteShaderStorageBuffer<float>(
gpu_data_in_->elements, &gpu_data_in_->buffer));
RET_CHECK_EQ(TfLiteGpuDelegateBindBufferToTensor(
delegate_, gpu_data_in_->buffer.id(),
interpreter_->inputs()[0]), // First tensor only
kTfLiteOk);
} }
if (gpu_output_) { if (gpu_output_) {
// Get output image sizes. // Get output image sizes.
@@ -540,53 +642,85 @@ REGISTER_CALCULATOR(TfLiteInferenceCalculator);
for (int i = 0; i < gpu_data_out_.size(); ++i) { for (int i = 0; i < gpu_data_out_.size(); ++i) {
RET_CHECK_CALL(CreateReadWriteShaderStorageBuffer<float>( RET_CHECK_CALL(CreateReadWriteShaderStorageBuffer<float>(
gpu_data_out_[i]->elements, &gpu_data_out_[i]->buffer)); gpu_data_out_[i]->elements, &gpu_data_out_[i]->buffer));
RET_CHECK_EQ( RET_CHECK_EQ(TfLiteGpuDelegateBindBufferToTensor(
TfLiteGpuDelegateBindBufferToTensor( delegate_.get(), gpu_data_out_[i]->buffer.id(),
delegate_, gpu_data_out_[i]->buffer.id(), output_indices[i]), output_indices[i]),
kTfLiteOk); kTfLiteOk);
} }
} }
// Must call this last. // Must call this last.
RET_CHECK_EQ(interpreter_->ModifyGraphWithDelegate(delegate_), kTfLiteOk); RET_CHECK_EQ(interpreter_->ModifyGraphWithDelegate(delegate_.get()),
kTfLiteOk);
#endif // OpenGL #endif // OpenGL
#if defined(__APPLE__) && !TARGET_OS_OSX // iOS #if defined(MEDIAPIPE_IOS)
const int kHalfSize = 2; // sizeof(half)
// Configure and create the delegate. // Configure and create the delegate.
TFLGpuDelegateOptions options; TFLGpuDelegateOptions options;
options.allow_precision_loss = false; // Must match converter, F=float/T=half options.allow_precision_loss = true;
options.wait_type = TFLGpuDelegateWaitType::TFLGpuDelegateWaitTypePassive; options.wait_type = TFLGpuDelegateWaitType::TFLGpuDelegateWaitTypePassive;
if (!delegate_) delegate_ = TFLGpuDelegateCreate(&options); if (!delegate_)
delegate_ = TfLiteDelegatePtr(TFLGpuDelegateCreate(&options),
&TFLGpuDelegateDelete);
id<MTLDevice> device = gpu_helper_.mtlDevice; id<MTLDevice> device = gpu_helper_.mtlDevice;
if (gpu_input_) { if (gpu_input_) {
// Get input image sizes. // Get input image sizes.
gpu_data_in_ = absl::make_unique<GPUData>();
const auto& input_indices = interpreter_->inputs(); const auto& input_indices = interpreter_->inputs();
RET_CHECK_EQ(input_indices.size(), 1); gpu_data_in_.resize(input_indices.size());
const TfLiteTensor* tensor = interpreter_->tensor(input_indices[0]); for (int i = 0; i < input_indices.size(); ++i) {
gpu_data_in_->elements = 1; const TfLiteTensor* tensor = interpreter_->tensor(input_indices[i]);
// On iOS GPU, input must be 4 channels, regardless of what model expects. gpu_data_in_[i] = absl::make_unique<GPUData>();
{ gpu_data_in_[i]->shape.b = tensor->dims->data[0];
gpu_data_in_->elements *= tensor->dims->data[0]; // batch gpu_data_in_[i]->shape.h = tensor->dims->data[1];
gpu_data_in_->elements *= tensor->dims->data[1]; // height gpu_data_in_[i]->shape.w = tensor->dims->data[2];
gpu_data_in_->elements *= tensor->dims->data[2]; // width // On iOS GPU, input must be 4 channels, regardless of what model expects.
gpu_data_in_->elements *= 4; // channels gpu_data_in_[i]->shape.c = 4;
gpu_data_in_[i]->elements =
gpu_data_in_[i]->shape.b * gpu_data_in_[i]->shape.h *
gpu_data_in_[i]->shape.w * gpu_data_in_[i]->shape.c;
// Input to model can be RGBA only.
if (tensor->dims->data[3] != 4) {
LOG(WARNING) << "Please ensure input GPU tensor is 4 channels.";
}
const std::string shader_source =
absl::Substitute(R"(#include <metal_stdlib>
using namespace metal;
kernel void convertKernel(device float4* const input_buffer [[buffer(0)]],
device half4* output_buffer [[buffer(1)]],
uint gid [[thread_position_in_grid]]) {
if (gid >= $0) return;
output_buffer[gid] = half4(input_buffer[gid]);
})",
gpu_data_in_[i]->elements / 4);
NSString* library_source =
[NSString stringWithUTF8String:shader_source.c_str()];
NSError* error = nil;
id<MTLLibrary> library =
[device newLibraryWithSource:library_source options:nil error:&error];
RET_CHECK(library != nil) << "Couldn't create shader library "
<< [[error localizedDescription] UTF8String];
id<MTLFunction> kernel_func = nil;
kernel_func = [library newFunctionWithName:@"convertKernel"];
RET_CHECK(kernel_func != nil) << "Couldn't create kernel function.";
fp32_to_fp16_program_ =
[device newComputePipelineStateWithFunction:kernel_func error:&error];
RET_CHECK(fp32_to_fp16_program_ != nil)
<< "Couldn't create pipeline state "
<< [[error localizedDescription] UTF8String];
// Create and bind input buffer.
gpu_data_in_[i]->buffer =
[device newBufferWithLength:gpu_data_in_[i]->elements * kHalfSize
options:MTLResourceStorageModeShared];
RET_CHECK_EQ(interpreter_->ModifyGraphWithDelegate(delegate_.get()),
kTfLiteOk);
RET_CHECK_EQ(
TFLGpuDelegateBindMetalBufferToTensor(
delegate_.get(), input_indices[i], gpu_data_in_[i]->buffer),
true);
} }
// Input to model can be RGBA only.
if (tensor->dims->data[3] != 4) {
LOG(WARNING) << "Please ensure input GPU tensor is 4 channels.";
}
// Create and bind input buffer.
gpu_data_in_->buffer =
[device newBufferWithLength:gpu_data_in_->elements * sizeof(float)
options:MTLResourceStorageModeShared];
RET_CHECK_EQ(interpreter_->ModifyGraphWithDelegate(delegate_), kTfLiteOk);
RET_CHECK_EQ(TFLGpuDelegateBindMetalBufferToTensor(
delegate_,
input_indices[0], // First tensor only
gpu_data_in_->buffer),
true);
} }
if (gpu_output_) { if (gpu_output_) {
// Get output image sizes. // Get output image sizes.
@@ -627,15 +761,17 @@ REGISTER_CALCULATOR(TfLiteInferenceCalculator);
interpreter_->SetAllowBufferHandleOutput(true); interpreter_->SetAllowBufferHandleOutput(true);
for (int i = 0; i < gpu_data_out_.size(); ++i) { for (int i = 0; i < gpu_data_out_.size(); ++i) {
gpu_data_out_[i]->buffer = gpu_data_out_[i]->buffer =
[device newBufferWithLength:gpu_data_out_[i]->elements * sizeof(float) [device newBufferWithLength:gpu_data_out_[i]->elements * kHalfSize
options:MTLResourceStorageModeShared]; options:MTLResourceStorageModeShared];
RET_CHECK_EQ(TFLGpuDelegateBindMetalBufferToTensor( RET_CHECK_EQ(
delegate_, output_indices[i], gpu_data_out_[i]->buffer), TFLGpuDelegateBindMetalBufferToTensor(
true); delegate_.get(), output_indices[i], gpu_data_out_[i]->buffer),
true);
} }
// Create converter for GPU output. // Create converter for GPU output.
converter_from_BPHWC4_ = [[TFLBufferConvert alloc] initWithDevice:device converter_from_BPHWC4_ = [[TFLBufferConvert alloc] initWithDevice:device
isFloat16:false isFloat16:true
convertToPBHWC4:false]; convertToPBHWC4:false];
if (converter_from_BPHWC4_ == nil) { if (converter_from_BPHWC4_ == nil) {
return mediapipe::InternalError( return mediapipe::InternalError(
@@ -27,7 +27,7 @@ import "mediapipe/framework/calculator.proto";
// options { // options {
// [mediapipe.TfLiteInferenceCalculatorOptions.ext] { // [mediapipe.TfLiteInferenceCalculatorOptions.ext] {
// model_path: "model.tflite" // model_path: "model.tflite"
// use_gpu: true // delegate { gpu {} }
// } // }
// } // }
// } // }
@@ -37,6 +37,22 @@ message TfLiteInferenceCalculatorOptions {
optional TfLiteInferenceCalculatorOptions ext = 233867213; optional TfLiteInferenceCalculatorOptions ext = 233867213;
} }
message Delegate {
// Default inference provided by tflite.
message TfLite {}
// Delegate to run GPU inference depending on the device.
// (Can use OpenGl, OpenCl, Metal depending on the device.)
message Gpu {}
// Android only.
message Nnapi {}
oneof delegate {
TfLite tflite = 1;
Gpu gpu = 2;
Nnapi nnapi = 3;
}
}
// Path to the TF Lite model (ex: /path/to/modelname.tflite). // Path to the TF Lite model (ex: /path/to/modelname.tflite).
// On mobile, this is generally just modelname.tflite. // On mobile, this is generally just modelname.tflite.
optional string model_path = 1; optional string model_path = 1;
@@ -44,5 +60,22 @@ message TfLiteInferenceCalculatorOptions {
// Whether the TF Lite GPU or CPU backend should be used. Effective only when // Whether the TF Lite GPU or CPU backend should be used. Effective only when
// input tensors are on CPU. For input tensors on GPU, GPU backend is always // input tensors are on CPU. For input tensors on GPU, GPU backend is always
// used. // used.
optional bool use_gpu = 2 [default = false]; // DEPRECATED: configure "delegate" instead.
optional bool use_gpu = 2 [deprecated = true, default = false];
// Android only. When true, an NNAPI delegate will be used for inference.
// If NNAPI is not available, then the default CPU delegate will be used
// automatically.
// DEPRECATED: configure "delegate" instead.
optional bool use_nnapi = 3 [deprecated = true, default = false];
// The number of threads available to the interpreter. Effective only when
// input tensors are on CPU and 'use_gpu' is false.
optional int32 cpu_num_thread = 4 [default = -1];
// TfLite delegate to run inference.
// NOTE: calculator is free to choose delegate if not specified explicitly.
// NOTE: use_gpu/use_nnapi are ignored if specified. (Delegate takes
// precedence over use_* deprecated options.)
optional Delegate delegate = 5;
} }
@@ -16,6 +16,8 @@
#include <string> #include <string>
#include <vector> #include <vector>
#include "absl/strings/str_replace.h"
#include "absl/strings/string_view.h"
#include "mediapipe/calculators/tflite/tflite_inference_calculator.pb.h" #include "mediapipe/calculators/tflite/tflite_inference_calculator.pb.h"
#include "mediapipe/framework/calculator_framework.h" #include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/calculator_runner.h" #include "mediapipe/framework/calculator_runner.h"
@@ -39,13 +41,7 @@ namespace mediapipe {
using ::tflite::Interpreter; using ::tflite::Interpreter;
class TfLiteInferenceCalculatorTest : public ::testing::Test { void DoSmokeTest(absl::string_view delegate) {
protected:
std::unique_ptr<CalculatorRunner> runner_ = nullptr;
};
// Tests a simple add model that adds an input tensor to itself.
TEST_F(TfLiteInferenceCalculatorTest, SmokeTest) {
const int width = 8; const int width = 8;
const int height = 8; const int height = 8;
const int channels = 3; const int channels = 3;
@@ -73,23 +69,24 @@ TEST_F(TfLiteInferenceCalculatorTest, SmokeTest) {
auto input_vec = absl::make_unique<std::vector<TfLiteTensor>>(); auto input_vec = absl::make_unique<std::vector<TfLiteTensor>>();
input_vec->emplace_back(*tensor); input_vec->emplace_back(*tensor);
std::string graph_proto = R"(
input_stream: "tensor_in"
node {
calculator: "TfLiteInferenceCalculator"
input_stream: "TENSORS:tensor_in"
output_stream: "TENSORS:tensor_out"
options {
[mediapipe.TfLiteInferenceCalculatorOptions.ext] {
model_path: "mediapipe/calculators/tflite/testdata/add.bin"
$delegate
}
}
}
)";
ASSERT_EQ(absl::StrReplaceAll({{"$delegate", delegate}}, &graph_proto), 1);
// Prepare single calculator graph to and wait for packets. // Prepare single calculator graph to and wait for packets.
CalculatorGraphConfig graph_config = CalculatorGraphConfig graph_config =
::mediapipe::ParseTextProtoOrDie<CalculatorGraphConfig>( ::mediapipe::ParseTextProtoOrDie<CalculatorGraphConfig>(graph_proto);
R"(
input_stream: "tensor_in"
node {
calculator: "TfLiteInferenceCalculator"
input_stream: "TENSORS:tensor_in"
output_stream: "TENSORS:tensor_out"
options {
[mediapipe.TfLiteInferenceCalculatorOptions.ext] {
use_gpu: false
model_path: "mediapipe/calculators/tflite/testdata/add.bin"
}
}
}
)");
std::vector<Packet> output_packets; std::vector<Packet> output_packets;
tool::AddVectorSink("tensor_out", &graph_config, &output_packets); tool::AddVectorSink("tensor_out", &graph_config, &output_packets);
CalculatorGraph graph(graph_config); CalculatorGraph graph(graph_config);
@@ -120,4 +117,10 @@ TEST_F(TfLiteInferenceCalculatorTest, SmokeTest) {
MP_ASSERT_OK(graph.WaitUntilDone()); MP_ASSERT_OK(graph.WaitUntilDone());
} }
// Tests a simple add model that adds an input tensor to itself.
TEST(TfLiteInferenceCalculatorTest, SmokeTest) {
DoSmokeTest(/*delegate=*/"");
DoSmokeTest(/*delegate=*/"delegate { tflite {} }");
}
} // namespace mediapipe } // namespace mediapipe
@@ -24,8 +24,7 @@
#include "mediapipe/framework/port/ret_check.h" #include "mediapipe/framework/port/ret_check.h"
#include "mediapipe/util/resource_util.h" #include "mediapipe/util/resource_util.h"
#include "tensorflow/lite/interpreter.h" #include "tensorflow/lite/interpreter.h"
#if defined(__EMSCRIPTEN__) || defined(__ANDROID__) || \ #if defined(MEDIAPIPE_MOBILE)
(defined(__APPLE__) && !TARGET_OS_OSX)
#include "mediapipe/util/android/file/base/file.h" #include "mediapipe/util/android/file/base/file.h"
#include "mediapipe/util/android/file/base/helpers.h" #include "mediapipe/util/android/file/base/helpers.h"
#else #else
@@ -27,8 +27,7 @@
#include "mediapipe/framework/port/ret_check.h" #include "mediapipe/framework/port/ret_check.h"
#include "tensorflow/lite/interpreter.h" #include "tensorflow/lite/interpreter.h"
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
#include "mediapipe/gpu/gl_calculator_helper.h" #include "mediapipe/gpu/gl_calculator_helper.h"
#include "tensorflow/lite/delegates/gpu/gl/gl_buffer.h" #include "tensorflow/lite/delegates/gpu/gl/gl_buffer.h"
#include "tensorflow/lite/delegates/gpu/gl/gl_program.h" #include "tensorflow/lite/delegates/gpu/gl/gl_program.h"
@@ -36,7 +35,7 @@
#include "tensorflow/lite/delegates/gpu/gl_delegate.h" #include "tensorflow/lite/delegates/gpu/gl_delegate.h"
#endif // !MEDIAPIPE_DISABLE_GPU #endif // !MEDIAPIPE_DISABLE_GPU
#if defined(__APPLE__) && !TARGET_OS_OSX // iOS #if defined(MEDIAPIPE_IOS)
#import <CoreVideo/CoreVideo.h> #import <CoreVideo/CoreVideo.h>
#import <Metal/Metal.h> #import <Metal/Metal.h>
#import <MetalKit/MetalKit.h> #import <MetalKit/MetalKit.h>
@@ -56,17 +55,15 @@ constexpr int kNumCoordsPerBox = 4;
namespace mediapipe { namespace mediapipe {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
using ::tflite::gpu::gl::CreateReadWriteShaderStorageBuffer; using ::tflite::gpu::gl::CreateReadWriteShaderStorageBuffer;
using ::tflite::gpu::gl::GlShader; using ::tflite::gpu::gl::GlShader;
#endif #endif
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
typedef ::tflite::gpu::gl::GlBuffer GpuTensor; typedef ::tflite::gpu::gl::GlBuffer GpuTensor;
typedef ::tflite::gpu::gl::GlProgram GpuProgram; typedef ::tflite::gpu::gl::GlProgram GpuProgram;
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS #elif defined(MEDIAPIPE_IOS)
typedef id<MTLBuffer> GpuTensor; typedef id<MTLBuffer> GpuTensor;
typedef id<MTLComputePipelineState> GpuProgram; typedef id<MTLComputePipelineState> GpuProgram;
#endif #endif
@@ -183,11 +180,10 @@ class TfLiteTensorsToDetectionsCalculator : public CalculatorBase {
std::vector<Anchor> anchors_; std::vector<Anchor> anchors_;
bool side_packet_anchors_{}; bool side_packet_anchors_{};
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
mediapipe::GlCalculatorHelper gpu_helper_; mediapipe::GlCalculatorHelper gpu_helper_;
std::unique_ptr<GPUData> gpu_data_; std::unique_ptr<GPUData> gpu_data_;
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS #elif defined(MEDIAPIPE_IOS)
MPPMetalHelper* gpu_helper_ = nullptr; MPPMetalHelper* gpu_helper_ = nullptr;
std::unique_ptr<GPUData> gpu_data_; std::unique_ptr<GPUData> gpu_data_;
#endif #endif
@@ -226,10 +222,9 @@ REGISTER_CALCULATOR(TfLiteTensorsToDetectionsCalculator);
} }
if (use_gpu) { if (use_gpu) {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
MP_RETURN_IF_ERROR(mediapipe::GlCalculatorHelper::UpdateContract(cc)); MP_RETURN_IF_ERROR(mediapipe::GlCalculatorHelper::UpdateContract(cc));
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS #elif defined(MEDIAPIPE_IOS)
MP_RETURN_IF_ERROR([MPPMetalHelper updateContract:cc]); MP_RETURN_IF_ERROR([MPPMetalHelper updateContract:cc]);
#endif #endif
} }
@@ -243,10 +238,9 @@ REGISTER_CALCULATOR(TfLiteTensorsToDetectionsCalculator);
if (cc->Inputs().HasTag("TENSORS_GPU")) { if (cc->Inputs().HasTag("TENSORS_GPU")) {
gpu_input_ = true; gpu_input_ = true;
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
MP_RETURN_IF_ERROR(gpu_helper_.Open(cc)); MP_RETURN_IF_ERROR(gpu_helper_.Open(cc));
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS #elif defined(MEDIAPIPE_IOS)
gpu_helper_ = [[MPPMetalHelper alloc] initWithCalculatorContext:cc]; gpu_helper_ = [[MPPMetalHelper alloc] initWithCalculatorContext:cc];
RET_CHECK(gpu_helper_); RET_CHECK(gpu_helper_);
#endif #endif
@@ -406,8 +400,7 @@ REGISTER_CALCULATOR(TfLiteTensorsToDetectionsCalculator);
} }
::mediapipe::Status TfLiteTensorsToDetectionsCalculator::ProcessGPU( ::mediapipe::Status TfLiteTensorsToDetectionsCalculator::ProcessGPU(
CalculatorContext* cc, std::vector<Detection>* output_detections) { CalculatorContext* cc, std::vector<Detection>* output_detections) {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
const auto& input_tensors = const auto& input_tensors =
cc->Inputs().Tag("TENSORS_GPU").Get<std::vector<GpuTensor>>(); cc->Inputs().Tag("TENSORS_GPU").Get<std::vector<GpuTensor>>();
RET_CHECK_GE(input_tensors.size(), 2); RET_CHECK_GE(input_tensors.size(), 2);
@@ -470,7 +463,7 @@ REGISTER_CALCULATOR(TfLiteTensorsToDetectionsCalculator);
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
})); }));
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS #elif defined(MEDIAPIPE_IOS)
const auto& input_tensors = const auto& input_tensors =
cc->Inputs().Tag("TENSORS_GPU").Get<std::vector<GpuTensor>>(); cc->Inputs().Tag("TENSORS_GPU").Get<std::vector<GpuTensor>>();
@@ -479,11 +472,11 @@ REGISTER_CALCULATOR(TfLiteTensorsToDetectionsCalculator);
// Copy inputs. // Copy inputs.
[MPPMetalUtil blitMetalBufferTo:gpu_data_->raw_boxes_buffer [MPPMetalUtil blitMetalBufferTo:gpu_data_->raw_boxes_buffer
from:input_tensors[0] from:input_tensors[0]
blocking:true blocking:false
commandBuffer:[gpu_helper_ commandBuffer]]; commandBuffer:[gpu_helper_ commandBuffer]];
[MPPMetalUtil blitMetalBufferTo:gpu_data_->raw_scores_buffer [MPPMetalUtil blitMetalBufferTo:gpu_data_->raw_scores_buffer
from:input_tensors[1] from:input_tensors[1]
blocking:true blocking:false
commandBuffer:[gpu_helper_ commandBuffer]]; commandBuffer:[gpu_helper_ commandBuffer]];
if (!anchors_init_) { if (!anchors_init_) {
if (side_packet_anchors_) { if (side_packet_anchors_) {
@@ -498,48 +491,37 @@ REGISTER_CALCULATOR(TfLiteTensorsToDetectionsCalculator);
RET_CHECK_EQ(input_tensors.size(), kNumInputTensorsWithAnchors); RET_CHECK_EQ(input_tensors.size(), kNumInputTensorsWithAnchors);
[MPPMetalUtil blitMetalBufferTo:gpu_data_->raw_anchors_buffer [MPPMetalUtil blitMetalBufferTo:gpu_data_->raw_anchors_buffer
from:input_tensors[2] from:input_tensors[2]
blocking:true blocking:false
commandBuffer:[gpu_helper_ commandBuffer]]; commandBuffer:[gpu_helper_ commandBuffer]];
} }
anchors_init_ = true; anchors_init_ = true;
} }
// Run shaders. // Run shaders.
{ id<MTLCommandBuffer> command_buffer = [gpu_helper_ commandBuffer];
id<MTLCommandBuffer> command_buffer = [gpu_helper_ commandBuffer]; command_buffer.label = @"TfLiteDecodeAndScoreBoxes";
command_buffer.label = @"TfLiteDecodeBoxes"; id<MTLComputeCommandEncoder> command_encoder =
id<MTLComputeCommandEncoder> decode_command = [command_buffer computeCommandEncoder];
[command_buffer computeCommandEncoder]; [command_encoder setComputePipelineState:gpu_data_->decode_program];
[decode_command setComputePipelineState:gpu_data_->decode_program]; [command_encoder setBuffer:gpu_data_->decoded_boxes_buffer
[decode_command setBuffer:gpu_data_->decoded_boxes_buffer offset:0
offset:0 atIndex:0];
atIndex:0]; [command_encoder setBuffer:gpu_data_->raw_boxes_buffer offset:0 atIndex:1];
[decode_command setBuffer:gpu_data_->raw_boxes_buffer offset:0 atIndex:1]; [command_encoder setBuffer:gpu_data_->raw_anchors_buffer offset:0 atIndex:2];
[decode_command setBuffer:gpu_data_->raw_anchors_buffer offset:0 atIndex:2]; MTLSize decode_threads_per_group = MTLSizeMake(1, 1, 1);
MTLSize decode_threads_per_group = MTLSizeMake(1, 1, 1); MTLSize decode_threadgroups = MTLSizeMake(num_boxes_, 1, 1);
MTLSize decode_threadgroups = MTLSizeMake(num_boxes_, 1, 1); [command_encoder dispatchThreadgroups:decode_threadgroups
[decode_command dispatchThreadgroups:decode_threadgroups threadsPerThreadgroup:decode_threads_per_group];
threadsPerThreadgroup:decode_threads_per_group];
[decode_command endEncoding]; [command_encoder setComputePipelineState:gpu_data_->score_program];
[command_buffer commit]; [command_encoder setBuffer:gpu_data_->scored_boxes_buffer offset:0 atIndex:0];
[command_buffer waitUntilCompleted]; [command_encoder setBuffer:gpu_data_->raw_scores_buffer offset:0 atIndex:1];
} MTLSize score_threads_per_group = MTLSizeMake(1, num_classes_, 1);
{ MTLSize score_threadgroups = MTLSizeMake(num_boxes_, 1, 1);
id<MTLCommandBuffer> command_buffer = [gpu_helper_ commandBuffer]; [command_encoder dispatchThreadgroups:score_threadgroups
command_buffer.label = @"TfLiteScoreBoxes";
id<MTLComputeCommandEncoder> score_command =
[command_buffer computeCommandEncoder];
[score_command setComputePipelineState:gpu_data_->score_program];
[score_command setBuffer:gpu_data_->scored_boxes_buffer offset:0 atIndex:0];
[score_command setBuffer:gpu_data_->raw_scores_buffer offset:0 atIndex:1];
MTLSize score_threads_per_group = MTLSizeMake(1, num_classes_, 1);
MTLSize score_threadgroups = MTLSizeMake(num_boxes_, 1, 1);
[score_command dispatchThreadgroups:score_threadgroups
threadsPerThreadgroup:score_threads_per_group]; threadsPerThreadgroup:score_threads_per_group];
[score_command endEncoding]; [command_encoder endEncoding];
[command_buffer commit]; [MPPMetalUtil commitCommandBufferAndWait:command_buffer];
[command_buffer waitUntilCompleted];
}
// Copy decoded boxes from GPU to CPU. // Copy decoded boxes from GPU to CPU.
std::vector<float> boxes(num_boxes_ * num_coords_); std::vector<float> boxes(num_boxes_ * num_coords_);
@@ -569,12 +551,11 @@ REGISTER_CALCULATOR(TfLiteTensorsToDetectionsCalculator);
::mediapipe::Status TfLiteTensorsToDetectionsCalculator::Close( ::mediapipe::Status TfLiteTensorsToDetectionsCalculator::Close(
CalculatorContext* cc) { CalculatorContext* cc) {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
gpu_helper_.RunInGlContext([this] { gpu_data_.reset(); }); gpu_helper_.RunInGlContext([this] { gpu_data_.reset(); });
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS #elif defined(MEDIAPIPE_IOS)
gpu_data_.reset(); gpu_data_.reset();
#endif // !MEDIAPIPE_DISABLE_GPU #endif
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -723,8 +704,7 @@ Detection TfLiteTensorsToDetectionsCalculator::ConvertToDetection(
::mediapipe::Status TfLiteTensorsToDetectionsCalculator::GpuInit( ::mediapipe::Status TfLiteTensorsToDetectionsCalculator::GpuInit(
CalculatorContext* cc) { CalculatorContext* cc) {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext([this]() MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext([this]()
-> ::mediapipe::Status { -> ::mediapipe::Status {
gpu_data_ = absl::make_unique<GPUData>(); gpu_data_ = absl::make_unique<GPUData>();
@@ -937,8 +917,7 @@ void main() {
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
})); }));
#elif defined(__APPLE__) && !TARGET_OS_OSX // iOS #elif defined(MEDIAPIPE_IOS)
// TODO consolidate Metal and OpenGL shaders via vulkan.
gpu_data_ = absl::make_unique<GPUData>(); gpu_data_ = absl::make_unique<GPUData>();
id<MTLDevice> device = gpu_helper_.mtlDevice; id<MTLDevice> device = gpu_helper_.mtlDevice;
@@ -1168,7 +1147,7 @@ kernel void scoreKernel(
CHECK_LT(num_classes_, max_wg_size) << "# classes must be <" << max_wg_size; CHECK_LT(num_classes_, max_wg_size) << "# classes must be <" << max_wg_size;
} }
#endif // __ANDROID__ or iOS #endif // !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -28,6 +28,21 @@ namespace mediapipe {
// TENSORS - Vector of TfLiteTensor of type kTfLiteFloat32. Only the first // TENSORS - Vector of TfLiteTensor of type kTfLiteFloat32. Only the first
// tensor will be used. The size of the values must be // tensor will be used. The size of the values must be
// (num_dimension x num_landmarks). // (num_dimension x num_landmarks).
//
// FLIP_HORIZONTALLY (optional): Whether to flip landmarks horizontally or
// not. Overrides corresponding side packet and/or field in the calculator
// options.
//
// FLIP_VERTICALLY (optional): Whether to flip landmarks vertically or not.
// Overrides corresponding side packet and/or field in the calculator options.
//
// Input side packet:
// FLIP_HORIZONTALLY (optional): Whether to flip landmarks horizontally or
// not. Overrides the corresponding field in the calculator options.
//
// FLIP_VERTICALLY (optional): Whether to flip landmarks vertically or not.
// Overrides the corresponding field in the calculator options.
//
// Output: // Output:
// LANDMARKS(optional) - Result MediaPipe landmarks. // LANDMARKS(optional) - Result MediaPipe landmarks.
// NORM_LANDMARKS(optional) - Result MediaPipe normalized landmarks. // NORM_LANDMARKS(optional) - Result MediaPipe normalized landmarks.
@@ -61,6 +76,8 @@ class TfLiteTensorsToLandmarksCalculator : public CalculatorBase {
private: private:
::mediapipe::Status LoadOptions(CalculatorContext* cc); ::mediapipe::Status LoadOptions(CalculatorContext* cc);
int num_landmarks_ = 0; int num_landmarks_ = 0;
bool flip_vertically_ = false;
bool flip_horizontally_ = false;
::mediapipe::TfLiteTensorsToLandmarksCalculatorOptions options_; ::mediapipe::TfLiteTensorsToLandmarksCalculatorOptions options_;
}; };
@@ -75,12 +92,28 @@ REGISTER_CALCULATOR(TfLiteTensorsToLandmarksCalculator);
cc->Inputs().Tag("TENSORS").Set<std::vector<TfLiteTensor>>(); cc->Inputs().Tag("TENSORS").Set<std::vector<TfLiteTensor>>();
} }
if (cc->Inputs().HasTag("FLIP_HORIZONTALLY")) {
cc->Inputs().Tag("FLIP_HORIZONTALLY").Set<bool>();
}
if (cc->Inputs().HasTag("FLIP_VERTICALLY")) {
cc->Inputs().Tag("FLIP_VERTICALLY").Set<bool>();
}
if (cc->InputSidePackets().HasTag("FLIP_HORIZONTALLY")) {
cc->InputSidePackets().Tag("FLIP_HORIZONTALLY").Set<bool>();
}
if (cc->InputSidePackets().HasTag("FLIP_VERTICALLY")) {
cc->InputSidePackets().Tag("FLIP_VERTICALLY").Set<bool>();
}
if (cc->Outputs().HasTag("LANDMARKS")) { if (cc->Outputs().HasTag("LANDMARKS")) {
cc->Outputs().Tag("LANDMARKS").Set<std::vector<Landmark>>(); cc->Outputs().Tag("LANDMARKS").Set<LandmarkList>();
} }
if (cc->Outputs().HasTag("NORM_LANDMARKS")) { if (cc->Outputs().HasTag("NORM_LANDMARKS")) {
cc->Outputs().Tag("NORM_LANDMARKS").Set<std::vector<NormalizedLandmark>>(); cc->Outputs().Tag("NORM_LANDMARKS").Set<NormalizedLandmarkList>();
} }
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
@@ -98,17 +131,40 @@ REGISTER_CALCULATOR(TfLiteTensorsToLandmarksCalculator);
<< "Must provide input with/height for getting normalized landmarks."; << "Must provide input with/height for getting normalized landmarks.";
} }
if (cc->Outputs().HasTag("LANDMARKS") && if (cc->Outputs().HasTag("LANDMARKS") &&
(options_.flip_vertically() || options_.flip_horizontally())) { (options_.flip_vertically() || options_.flip_horizontally() ||
cc->InputSidePackets().HasTag("FLIP_HORIZONTALLY") ||
cc->InputSidePackets().HasTag("FLIP_VERTICALLY"))) {
RET_CHECK(options_.has_input_image_height() && RET_CHECK(options_.has_input_image_height() &&
options_.has_input_image_width()) options_.has_input_image_width())
<< "Must provide input with/height for using flip_vertically option " << "Must provide input with/height for using flip_vertically option "
"when outputing landmarks in absolute coordinates."; "when outputing landmarks in absolute coordinates.";
} }
flip_horizontally_ =
cc->InputSidePackets().HasTag("FLIP_HORIZONTALLY")
? cc->InputSidePackets().Tag("FLIP_HORIZONTALLY").Get<bool>()
: options_.flip_horizontally();
flip_horizontally_ =
cc->InputSidePackets().HasTag("FLIP_VERTICALLY")
? cc->InputSidePackets().Tag("FLIP_VERTICALLY").Get<bool>()
: options_.flip_vertically();
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
::mediapipe::Status TfLiteTensorsToLandmarksCalculator::Process( ::mediapipe::Status TfLiteTensorsToLandmarksCalculator::Process(
CalculatorContext* cc) { CalculatorContext* cc) {
// Override values if specified so.
if (cc->Inputs().HasTag("FLIP_HORIZONTALLY") &&
!cc->Inputs().Tag("FLIP_HORIZONTALLY").IsEmpty()) {
flip_horizontally_ = cc->Inputs().Tag("FLIP_HORIZONTALLY").Get<bool>();
}
if (cc->Inputs().HasTag("FLIP_VERTICALLY") &&
!cc->Inputs().Tag("FLIP_VERTICALLY").IsEmpty()) {
flip_vertically_ = cc->Inputs().Tag("FLIP_VERTICALLY").Get<bool>();
}
if (cc->Inputs().Tag("TENSORS").IsEmpty()) { if (cc->Inputs().Tag("TENSORS").IsEmpty()) {
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -127,54 +183,55 @@ REGISTER_CALCULATOR(TfLiteTensorsToLandmarksCalculator);
const float* raw_landmarks = raw_tensor->data.f; const float* raw_landmarks = raw_tensor->data.f;
auto output_landmarks = absl::make_unique<std::vector<Landmark>>(); LandmarkList output_landmarks;
for (int ld = 0; ld < num_landmarks_; ++ld) { for (int ld = 0; ld < num_landmarks_; ++ld) {
const int offset = ld * num_dimensions; const int offset = ld * num_dimensions;
Landmark landmark; Landmark* landmark = output_landmarks.add_landmark();
if (options_.flip_horizontally()) { if (flip_horizontally_) {
landmark.set_x(options_.input_image_width() - raw_landmarks[offset]); landmark->set_x(options_.input_image_width() - raw_landmarks[offset]);
} else { } else {
landmark.set_x(raw_landmarks[offset]); landmark->set_x(raw_landmarks[offset]);
} }
if (num_dimensions > 1) { if (num_dimensions > 1) {
if (options_.flip_vertically()) { if (flip_vertically_) {
landmark.set_y(options_.input_image_height() - landmark->set_y(options_.input_image_height() -
raw_landmarks[offset + 1]); raw_landmarks[offset + 1]);
} else { } else {
landmark.set_y(raw_landmarks[offset + 1]); landmark->set_y(raw_landmarks[offset + 1]);
} }
} }
if (num_dimensions > 2) { if (num_dimensions > 2) {
landmark.set_z(raw_landmarks[offset + 2]); landmark->set_z(raw_landmarks[offset + 2]);
} }
output_landmarks->push_back(landmark);
} }
// Output normalized landmarks if required. // Output normalized landmarks if required.
if (cc->Outputs().HasTag("NORM_LANDMARKS")) { if (cc->Outputs().HasTag("NORM_LANDMARKS")) {
auto output_norm_landmarks = NormalizedLandmarkList output_norm_landmarks;
absl::make_unique<std::vector<NormalizedLandmark>>(); // for (const auto& landmark : output_landmarks) {
for (const auto& landmark : *output_landmarks) { for (int i = 0; i < output_landmarks.landmark_size(); ++i) {
NormalizedLandmark norm_landmark; const Landmark& landmark = output_landmarks.landmark(i);
norm_landmark.set_x(static_cast<float>(landmark.x()) / NormalizedLandmark* norm_landmark = output_norm_landmarks.add_landmark();
options_.input_image_width()); norm_landmark->set_x(static_cast<float>(landmark.x()) /
norm_landmark.set_y(static_cast<float>(landmark.y()) / options_.input_image_width());
options_.input_image_height()); norm_landmark->set_y(static_cast<float>(landmark.y()) /
norm_landmark.set_z(landmark.z() / options_.normalize_z()); options_.input_image_height());
norm_landmark->set_z(landmark.z() / options_.normalize_z());
output_norm_landmarks->push_back(norm_landmark);
} }
cc->Outputs() cc->Outputs()
.Tag("NORM_LANDMARKS") .Tag("NORM_LANDMARKS")
.Add(output_norm_landmarks.release(), cc->InputTimestamp()); .AddPacket(MakePacket<NormalizedLandmarkList>(output_norm_landmarks)
.At(cc->InputTimestamp()));
} }
// Output absolute landmarks. // Output absolute landmarks.
if (cc->Outputs().HasTag("LANDMARKS")) { if (cc->Outputs().HasTag("LANDMARKS")) {
cc->Outputs() cc->Outputs()
.Tag("LANDMARKS") .Tag("LANDMARKS")
.Add(output_landmarks.release(), cc->InputTimestamp()); .AddPacket(MakePacket<LandmarkList>(output_landmarks)
.At(cc->InputTimestamp()));
} }
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
@@ -28,8 +28,7 @@
#include "mediapipe/util/resource_util.h" #include "mediapipe/util/resource_util.h"
#include "tensorflow/lite/interpreter.h" #include "tensorflow/lite/interpreter.h"
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
#include "mediapipe/gpu/gl_calculator_helper.h" #include "mediapipe/gpu/gl_calculator_helper.h"
#include "mediapipe/gpu/gl_simple_shaders.h" #include "mediapipe/gpu/gl_simple_shaders.h"
#include "mediapipe/gpu/shader_util.h" #include "mediapipe/gpu/shader_util.h"
@@ -54,8 +53,7 @@ float Clamp(float val, float min, float max) {
namespace mediapipe { namespace mediapipe {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
using ::tflite::gpu::gl::CopyBuffer; using ::tflite::gpu::gl::CopyBuffer;
using ::tflite::gpu::gl::CreateReadWriteRgbaImageTexture; using ::tflite::gpu::gl::CreateReadWriteRgbaImageTexture;
using ::tflite::gpu::gl::CreateReadWriteShaderStorageBuffer; using ::tflite::gpu::gl::CreateReadWriteShaderStorageBuffer;
@@ -131,8 +129,7 @@ class TfLiteTensorsToSegmentationCalculator : public CalculatorBase {
int tensor_channels_ = 0; int tensor_channels_ = 0;
bool use_gpu_ = false; bool use_gpu_ = false;
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
mediapipe::GlCalculatorHelper gpu_helper_; mediapipe::GlCalculatorHelper gpu_helper_;
std::unique_ptr<GlProgram> mask_program_with_prev_; std::unique_ptr<GlProgram> mask_program_with_prev_;
std::unique_ptr<GlProgram> mask_program_no_prev_; std::unique_ptr<GlProgram> mask_program_no_prev_;
@@ -162,8 +159,7 @@ REGISTER_CALCULATOR(TfLiteTensorsToSegmentationCalculator);
} }
// Inputs GPU. // Inputs GPU.
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
if (cc->Inputs().HasTag("TENSORS_GPU")) { if (cc->Inputs().HasTag("TENSORS_GPU")) {
cc->Inputs().Tag("TENSORS_GPU").Set<std::vector<GlBuffer>>(); cc->Inputs().Tag("TENSORS_GPU").Set<std::vector<GlBuffer>>();
use_gpu |= true; use_gpu |= true;
@@ -182,8 +178,7 @@ REGISTER_CALCULATOR(TfLiteTensorsToSegmentationCalculator);
if (cc->Outputs().HasTag("MASK")) { if (cc->Outputs().HasTag("MASK")) {
cc->Outputs().Tag("MASK").Set<ImageFrame>(); cc->Outputs().Tag("MASK").Set<ImageFrame>();
} }
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
if (cc->Outputs().HasTag("MASK_GPU")) { if (cc->Outputs().HasTag("MASK_GPU")) {
cc->Outputs().Tag("MASK_GPU").Set<mediapipe::GpuBuffer>(); cc->Outputs().Tag("MASK_GPU").Set<mediapipe::GpuBuffer>();
use_gpu |= true; use_gpu |= true;
@@ -191,8 +186,7 @@ REGISTER_CALCULATOR(TfLiteTensorsToSegmentationCalculator);
#endif // !MEDIAPIPE_DISABLE_GPU #endif // !MEDIAPIPE_DISABLE_GPU
if (use_gpu) { if (use_gpu) {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
MP_RETURN_IF_ERROR(mediapipe::GlCalculatorHelper::UpdateContract(cc)); MP_RETURN_IF_ERROR(mediapipe::GlCalculatorHelper::UpdateContract(cc));
#endif // !MEDIAPIPE_DISABLE_GPU #endif // !MEDIAPIPE_DISABLE_GPU
} }
@@ -205,8 +199,7 @@ REGISTER_CALCULATOR(TfLiteTensorsToSegmentationCalculator);
if (cc->Inputs().HasTag("TENSORS_GPU")) { if (cc->Inputs().HasTag("TENSORS_GPU")) {
use_gpu_ = true; use_gpu_ = true;
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
MP_RETURN_IF_ERROR(gpu_helper_.Open(cc)); MP_RETURN_IF_ERROR(gpu_helper_.Open(cc));
#endif // !MEDIAPIPE_DISABLE_GPU #endif // !MEDIAPIPE_DISABLE_GPU
} }
@@ -214,8 +207,7 @@ REGISTER_CALCULATOR(TfLiteTensorsToSegmentationCalculator);
MP_RETURN_IF_ERROR(LoadOptions(cc)); MP_RETURN_IF_ERROR(LoadOptions(cc));
if (use_gpu_) { if (use_gpu_) {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
MP_RETURN_IF_ERROR( MP_RETURN_IF_ERROR(
gpu_helper_.RunInGlContext([this, cc]() -> ::mediapipe::Status { gpu_helper_.RunInGlContext([this, cc]() -> ::mediapipe::Status {
MP_RETURN_IF_ERROR(InitGpu(cc)); MP_RETURN_IF_ERROR(InitGpu(cc));
@@ -232,8 +224,7 @@ REGISTER_CALCULATOR(TfLiteTensorsToSegmentationCalculator);
::mediapipe::Status TfLiteTensorsToSegmentationCalculator::Process( ::mediapipe::Status TfLiteTensorsToSegmentationCalculator::Process(
CalculatorContext* cc) { CalculatorContext* cc) {
if (use_gpu_) { if (use_gpu_) {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
MP_RETURN_IF_ERROR( MP_RETURN_IF_ERROR(
gpu_helper_.RunInGlContext([this, cc]() -> ::mediapipe::Status { gpu_helper_.RunInGlContext([this, cc]() -> ::mediapipe::Status {
MP_RETURN_IF_ERROR(ProcessGpu(cc)); MP_RETURN_IF_ERROR(ProcessGpu(cc));
@@ -249,8 +240,7 @@ REGISTER_CALCULATOR(TfLiteTensorsToSegmentationCalculator);
::mediapipe::Status TfLiteTensorsToSegmentationCalculator::Close( ::mediapipe::Status TfLiteTensorsToSegmentationCalculator::Close(
CalculatorContext* cc) { CalculatorContext* cc) {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
gpu_helper_.RunInGlContext([this] { gpu_helper_.RunInGlContext([this] {
if (upsample_program_) glDeleteProgram(upsample_program_); if (upsample_program_) glDeleteProgram(upsample_program_);
upsample_program_ = 0; upsample_program_ = 0;
@@ -377,8 +367,7 @@ REGISTER_CALCULATOR(TfLiteTensorsToSegmentationCalculator);
if (cc->Inputs().Tag("TENSORS_GPU").IsEmpty()) { if (cc->Inputs().Tag("TENSORS_GPU").IsEmpty()) {
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
// Get input streams. // Get input streams.
const auto& input_tensors = const auto& input_tensors =
cc->Inputs().Tag("TENSORS_GPU").Get<std::vector<GlBuffer>>(); cc->Inputs().Tag("TENSORS_GPU").Get<std::vector<GlBuffer>>();
@@ -464,8 +453,7 @@ REGISTER_CALCULATOR(TfLiteTensorsToSegmentationCalculator);
} }
void TfLiteTensorsToSegmentationCalculator::GlRender() { void TfLiteTensorsToSegmentationCalculator::GlRender() {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
static const GLfloat square_vertices[] = { static const GLfloat square_vertices[] = {
-1.0f, -1.0f, // bottom left -1.0f, -1.0f, // bottom left
1.0f, -1.0f, // bottom right 1.0f, -1.0f, // bottom right
@@ -537,8 +525,7 @@ void TfLiteTensorsToSegmentationCalculator::GlRender() {
::mediapipe::Status TfLiteTensorsToSegmentationCalculator::InitGpu( ::mediapipe::Status TfLiteTensorsToSegmentationCalculator::InitGpu(
CalculatorContext* cc) { CalculatorContext* cc) {
#if !defined(MEDIAPIPE_DISABLE_GPU) && !defined(__EMSCRIPTEN__) && \ #if !defined(MEDIAPIPE_DISABLE_GL_COMPUTE)
!defined(__APPLE__)
MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext([this]() MP_RETURN_IF_ERROR(gpu_helper_.RunInGlContext([this]()
-> ::mediapipe::Status { -> ::mediapipe::Status {
// A shader to process a segmentation tensor into an output mask, // A shader to process a segmentation tensor into an output mask,
+133
View File
@@ -39,6 +39,15 @@ proto_library(
], ],
) )
proto_library(
name = "timed_box_list_id_to_label_calculator_proto",
srcs = ["timed_box_list_id_to_label_calculator.proto"],
visibility = ["//visibility:public"],
deps = [
"//mediapipe/framework:calculator_proto",
],
)
proto_library( proto_library(
name = "latency_proto", name = "latency_proto",
srcs = ["latency.proto"], srcs = ["latency.proto"],
@@ -113,6 +122,18 @@ mediapipe_cc_proto_library(
], ],
) )
mediapipe_cc_proto_library(
name = "timed_box_list_id_to_label_calculator_cc_proto",
srcs = ["timed_box_list_id_to_label_calculator.proto"],
cc_deps = [
"//mediapipe/framework:calculator_cc_proto",
],
visibility = ["//visibility:public"],
deps = [
":timed_box_list_id_to_label_calculator_proto",
],
)
mediapipe_cc_proto_library( mediapipe_cc_proto_library(
name = "latency_cc_proto", name = "latency_cc_proto",
srcs = ["latency.proto"], srcs = ["latency.proto"],
@@ -313,6 +334,34 @@ cc_library(
alwayslink = 1, alwayslink = 1,
) )
cc_library(
name = "timed_box_list_id_to_label_calculator",
srcs = ["timed_box_list_id_to_label_calculator.cc"],
visibility = ["//visibility:public"],
deps = [
":timed_box_list_id_to_label_calculator_cc_proto",
"//mediapipe/framework/port:status",
"//mediapipe/framework:calculator_framework",
"//mediapipe/framework:packet",
"//mediapipe/util/tracking:box_tracker_cc_proto",
"//mediapipe/util:resource_util",
] + select({
"//mediapipe:android": [
"//mediapipe/util/android/file/base",
],
"//mediapipe:apple": [
"//mediapipe/util/android/file/base",
],
"//mediapipe:macos": [
"//mediapipe/framework/port:file_helpers",
],
"//conditions:default": [
"//mediapipe/framework/port:file_helpers",
],
}),
alwayslink = 1,
)
cc_library( cc_library(
name = "non_max_suppression_calculator", name = "non_max_suppression_calculator",
srcs = ["non_max_suppression_calculator.cc"], srcs = ["non_max_suppression_calculator.cc"],
@@ -437,6 +486,7 @@ cc_library(
"//mediapipe/framework/formats:rect_cc_proto", "//mediapipe/framework/formats:rect_cc_proto",
"//mediapipe/framework/port:ret_check", "//mediapipe/framework/port:ret_check",
"//mediapipe/framework/port:status", "//mediapipe/framework/port:status",
"@com_google_absl//absl/types:optional",
], ],
alwayslink = 1, alwayslink = 1,
) )
@@ -508,6 +558,17 @@ proto_library(
], ],
) )
proto_library(
name = "timed_box_list_to_render_data_calculator_proto",
srcs = ["timed_box_list_to_render_data_calculator.proto"],
visibility = ["//visibility:public"],
deps = [
"//mediapipe/framework:calculator_proto",
"//mediapipe/util:color_proto",
"//mediapipe/util:render_data_proto",
],
)
proto_library( proto_library(
name = "labels_to_render_data_calculator_proto", name = "labels_to_render_data_calculator_proto",
srcs = ["labels_to_render_data_calculator.proto"], srcs = ["labels_to_render_data_calculator.proto"],
@@ -651,6 +712,37 @@ cc_library(
alwayslink = 1, alwayslink = 1,
) )
mediapipe_cc_proto_library(
name = "timed_box_list_to_render_data_calculator_cc_proto",
srcs = ["timed_box_list_to_render_data_calculator.proto"],
cc_deps = [
"//mediapipe/framework:calculator_cc_proto",
"//mediapipe/util:color_cc_proto",
"//mediapipe/util:render_data_cc_proto",
],
visibility = ["//visibility:public"],
deps = [":timed_box_list_to_render_data_calculator_proto"],
)
cc_library(
name = "timed_box_list_to_render_data_calculator",
srcs = ["timed_box_list_to_render_data_calculator.cc"],
visibility = ["//visibility:public"],
deps = [
":timed_box_list_to_render_data_calculator_cc_proto",
"//mediapipe/framework:calculator_framework",
"//mediapipe/framework:calculator_options_cc_proto",
"//mediapipe/framework/port:ret_check",
"//mediapipe/util:color_cc_proto",
"//mediapipe/util:render_data_cc_proto",
"//mediapipe/util/tracking:box_tracker_cc_proto",
"//mediapipe/util/tracking:tracking_cc_proto",
"@com_google_absl//absl/memory",
"@com_google_absl//absl/strings",
],
alwayslink = 1,
)
cc_library( cc_library(
name = "labels_to_render_data_calculator", name = "labels_to_render_data_calculator",
srcs = ["labels_to_render_data_calculator.cc"], srcs = ["labels_to_render_data_calculator.cc"],
@@ -886,6 +978,19 @@ cc_library(
alwayslink = 1, alwayslink = 1,
) )
cc_library(
name = "local_file_pattern_contents_calculator",
srcs = ["local_file_pattern_contents_calculator.cc"],
visibility = ["//visibility:public"],
deps = [
"//mediapipe/framework:calculator_framework",
"//mediapipe/framework/port:file_helpers",
"//mediapipe/framework/port:ret_check",
"//mediapipe/framework/port:status",
],
alwayslink = 1,
)
cc_library( cc_library(
name = "filter_collection_calculator", name = "filter_collection_calculator",
srcs = ["filter_collection_calculator.cc"], srcs = ["filter_collection_calculator.cc"],
@@ -983,3 +1088,31 @@ cc_test(
"//mediapipe/framework/port:parse_text_proto", "//mediapipe/framework/port:parse_text_proto",
], ],
) )
cc_library(
name = "detections_to_timed_box_list_calculator",
srcs = ["detections_to_timed_box_list_calculator.cc"],
visibility = ["//visibility:public"],
deps = [
"//mediapipe/framework:calculator_framework",
"//mediapipe/framework/formats:detection_cc_proto",
"//mediapipe/framework/formats:location_data_cc_proto",
"//mediapipe/framework/port:ret_check",
"//mediapipe/framework/port:status",
"//mediapipe/util/tracking:box_tracker",
],
alwayslink = 1,
)
cc_library(
name = "detection_unique_id_calculator",
srcs = ["detection_unique_id_calculator.cc"],
visibility = ["//visibility:public"],
deps = [
"//mediapipe/framework:calculator_framework",
"//mediapipe/framework/formats:detection_cc_proto",
"//mediapipe/framework/port:ret_check",
"//mediapipe/framework/port:status",
],
alwayslink = 1,
)
@@ -39,13 +39,13 @@ namespace mediapipe {
namespace { namespace {
constexpr char kInputFrameTag[] = "INPUT_FRAME"; constexpr char kInputFrameTag[] = "IMAGE";
constexpr char kOutputFrameTag[] = "OUTPUT_FRAME"; constexpr char kOutputFrameTag[] = "IMAGE";
constexpr char kInputVectorTag[] = "VECTOR"; constexpr char kInputVectorTag[] = "VECTOR";
constexpr char kInputFrameTagGpu[] = "INPUT_FRAME_GPU"; constexpr char kInputFrameTagGpu[] = "IMAGE_GPU";
constexpr char kOutputFrameTagGpu[] = "OUTPUT_FRAME_GPU"; constexpr char kOutputFrameTagGpu[] = "IMAGE_GPU";
enum { ATTRIB_VERTEX, ATTRIB_TEXTURE_POSITION, NUM_ATTRIBUTES }; enum { ATTRIB_VERTEX, ATTRIB_TEXTURE_POSITION, NUM_ATTRIBUTES };
@@ -61,7 +61,7 @@ constexpr int kAnnotationBackgroundColor[] = {100, 101, 102};
// A calculator for rendering data on images. // A calculator for rendering data on images.
// //
// Inputs: // Inputs:
// 1. INPUT_FRAME or INPUT_FRAME_GPU (optional): An ImageFrame (or GpuBuffer) // 1. IMAGE or IMAGE_GPU (optional): An ImageFrame (or GpuBuffer)
// containing the input image. // containing the input image.
// If output is CPU, and input isn't provided, the renderer creates a // If output is CPU, and input isn't provided, the renderer creates a
// blank canvas with the width, height and color provided in the options. // blank canvas with the width, height and color provided in the options.
@@ -73,7 +73,7 @@ constexpr int kAnnotationBackgroundColor[] = {100, 101, 102};
// input vector items. These input streams are tagged with "VECTOR". // input vector items. These input streams are tagged with "VECTOR".
// //
// Output: // Output:
// 1. OUTPUT_FRAME or OUTPUT_FRAME_GPU: A rendered ImageFrame (or GpuBuffer). // 1. IMAGE or IMAGE_GPU: A rendered ImageFrame (or GpuBuffer).
// //
// For CPU input frames, only SRGBA, SRGB and GRAY8 format are supported. The // For CPU input frames, only SRGBA, SRGB and GRAY8 format are supported. The
// output format is the same as input except for GRAY8 where the output is in // output format is the same as input except for GRAY8 where the output is in
@@ -87,13 +87,13 @@ constexpr int kAnnotationBackgroundColor[] = {100, 101, 102};
// Example config (CPU): // Example config (CPU):
// node { // node {
// calculator: "AnnotationOverlayCalculator" // calculator: "AnnotationOverlayCalculator"
// input_stream: "INPUT_FRAME:image_frames" // input_stream: "IMAGE:image_frames"
// input_stream: "render_data_1" // input_stream: "render_data_1"
// input_stream: "render_data_2" // input_stream: "render_data_2"
// input_stream: "render_data_3" // input_stream: "render_data_3"
// input_stream: "VECTOR:0:render_data_vec_0" // input_stream: "VECTOR:0:render_data_vec_0"
// input_stream: "VECTOR:1:render_data_vec_1" // input_stream: "VECTOR:1:render_data_vec_1"
// output_stream: "OUTPUT_FRAME:decorated_frames" // output_stream: "IMAGE:decorated_frames"
// options { // options {
// [mediapipe.AnnotationOverlayCalculatorOptions.ext] { // [mediapipe.AnnotationOverlayCalculatorOptions.ext] {
// } // }
@@ -103,13 +103,13 @@ constexpr int kAnnotationBackgroundColor[] = {100, 101, 102};
// Example config (GPU): // Example config (GPU):
// node { // node {
// calculator: "AnnotationOverlayCalculator" // calculator: "AnnotationOverlayCalculator"
// input_stream: "INPUT_FRAME_GPU:image_frames" // input_stream: "IMAGE_GPU:image_frames"
// input_stream: "render_data_1" // input_stream: "render_data_1"
// input_stream: "render_data_2" // input_stream: "render_data_2"
// input_stream: "render_data_3" // input_stream: "render_data_3"
// input_stream: "VECTOR:0:render_data_vec_0" // input_stream: "VECTOR:0:render_data_vec_0"
// input_stream: "VECTOR:1:render_data_vec_1" // input_stream: "VECTOR:1:render_data_vec_1"
// output_stream: "OUTPUT_FRAME_GPU:decorated_frames" // output_stream: "IMAGE_GPU:decorated_frames"
// options { // options {
// [mediapipe.AnnotationOverlayCalculatorOptions.ext] { // [mediapipe.AnnotationOverlayCalculatorOptions.ext] {
// } // }
@@ -37,6 +37,8 @@ namespace mediapipe {
// } // }
// } // }
// } // }
// Optionally, uses a side packet to override `min_size` specified in the
// calculator options.
template <typename IterableT> template <typename IterableT>
class CollectionHasMinSizeCalculator : public CalculatorBase { class CollectionHasMinSizeCalculator : public CalculatorBase {
public: public:
@@ -54,6 +56,10 @@ class CollectionHasMinSizeCalculator : public CalculatorBase {
cc->Inputs().Tag("ITERABLE").Set<IterableT>(); cc->Inputs().Tag("ITERABLE").Set<IterableT>();
cc->Outputs().Index(0).Set<bool>(); cc->Outputs().Index(0).Set<bool>();
// Optional input side packet that determines `min_size_`.
if (cc->InputSidePackets().NumEntries() > 0) {
cc->InputSidePackets().Index(0).Set<int>();
}
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -62,6 +68,11 @@ class CollectionHasMinSizeCalculator : public CalculatorBase {
min_size_ = min_size_ =
cc->Options<::mediapipe::CollectionHasMinSizeCalculatorOptions>() cc->Options<::mediapipe::CollectionHasMinSizeCalculatorOptions>()
.min_size(); .min_size();
// Override `min_size` if passed as side packet.
if (cc->InputSidePackets().NumEntries() > 0 &&
!cc->InputSidePackets().Index(0).IsEmpty()) {
min_size_ = cc->InputSidePackets().Index(0).Get<int>();
}
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -12,15 +12,14 @@
// See the License for the specific language governing permissions and // See the License for the specific language governing permissions and
// limitations under the License. // limitations under the License.
#include "mediapipe//framework/packet.h"
#include "mediapipe/calculators/util/detection_label_id_to_text_calculator.pb.h" #include "mediapipe/calculators/util/detection_label_id_to_text_calculator.pb.h"
#include "mediapipe/framework/calculator_framework.h" #include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/formats/detection.pb.h" #include "mediapipe/framework/formats/detection.pb.h"
#include "mediapipe/framework/packet.h"
#include "mediapipe/framework/port/status.h" #include "mediapipe/framework/port/status.h"
#include "mediapipe/util/resource_util.h" #include "mediapipe/util/resource_util.h"
#if defined(MEDIAPIPE_LITE) || defined(__EMSCRIPTEN__) || \ #if defined(MEDIAPIPE_MOBILE)
defined(__ANDROID__) || (defined(__APPLE__) && !TARGET_OS_OSX)
#include "mediapipe/util/android/file/base/file.h" #include "mediapipe/util/android/file/base/file.h"
#include "mediapipe/util/android/file/base/helpers.h" #include "mediapipe/util/android/file/base/helpers.h"
#else #else
@@ -0,0 +1,110 @@
// 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.
#include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/formats/detection.pb.h"
#include "mediapipe/framework/port/status.h"
namespace mediapipe {
namespace {
constexpr char kDetectionsTag[] = "DETECTIONS";
constexpr char kDetectionListTag[] = "DETECTION_LIST";
// Each detection processed by DetectionUniqueIDCalculator will be assigned an
// unique id that starts from 1. If a detection already has an ID other than 0,
// the ID will be overwritten.
static int64 detection_id = 0;
inline int GetNextDetectionId() { return ++detection_id; }
} // namespace
// Assign a unique id to detections.
// Note that the calculator will consume the input vector of Detection or
// DetectionList. So the input stream can not be connected to other calculators.
//
// Example config:
// node {
// calculator: "DetectionUniqueIdCalculator"
// input_stream: "DETECTIONS:detections"
// output_stream: "DETECTIONS:output_detections"
// }
class DetectionUniqueIdCalculator : public CalculatorBase {
public:
static ::mediapipe::Status GetContract(CalculatorContract* cc) {
RET_CHECK(cc->Inputs().HasTag(kDetectionListTag) ||
cc->Inputs().HasTag(kDetectionsTag))
<< "None of the input streams are provided.";
if (cc->Inputs().HasTag(kDetectionListTag)) {
RET_CHECK(cc->Outputs().HasTag(kDetectionListTag));
cc->Inputs().Tag(kDetectionListTag).Set<DetectionList>();
cc->Outputs().Tag(kDetectionListTag).Set<DetectionList>();
}
if (cc->Inputs().HasTag(kDetectionsTag)) {
RET_CHECK(cc->Outputs().HasTag(kDetectionsTag));
cc->Inputs().Tag(kDetectionsTag).Set<std::vector<Detection>>();
cc->Outputs().Tag(kDetectionsTag).Set<std::vector<Detection>>();
}
return ::mediapipe::OkStatus();
}
::mediapipe::Status Open(CalculatorContext* cc) override {
cc->SetOffset(::mediapipe::TimestampDiff(0));
return ::mediapipe::OkStatus();
}
::mediapipe::Status Process(CalculatorContext* cc) override;
};
REGISTER_CALCULATOR(DetectionUniqueIdCalculator);
::mediapipe::Status DetectionUniqueIdCalculator::Process(
CalculatorContext* cc) {
if (cc->Inputs().HasTag(kDetectionListTag) &&
!cc->Inputs().Tag(kDetectionListTag).IsEmpty()) {
auto result =
cc->Inputs().Tag(kDetectionListTag).Value().Consume<DetectionList>();
if (result.ok()) {
auto detection_list = std::move(result).ValueOrDie();
for (Detection& detection : *detection_list->mutable_detection()) {
detection.set_detection_id(GetNextDetectionId());
}
cc->Outputs()
.Tag(kDetectionListTag)
.Add(detection_list.release(), cc->InputTimestamp());
}
}
if (cc->Inputs().HasTag(kDetectionsTag) &&
!cc->Inputs().Tag(kDetectionsTag).IsEmpty()) {
auto result = cc->Inputs()
.Tag(kDetectionsTag)
.Value()
.Consume<std::vector<Detection>>();
if (result.ok()) {
auto detections = std::move(result).ValueOrDie();
for (Detection& detection : *detections) {
detection.set_detection_id(GetNextDetectionId());
}
cc->Outputs()
.Tag(kDetectionsTag)
.Add(detections.release(), cc->InputTimestamp());
}
}
return ::mediapipe::OkStatus();
}
} // namespace mediapipe
@@ -39,7 +39,8 @@ constexpr char kNormRectsTag[] = "NORM_RECTS";
} // namespace } // namespace
::mediapipe::Status DetectionsToRectsCalculator::DetectionToRect( ::mediapipe::Status DetectionsToRectsCalculator::DetectionToRect(
const Detection& detection, Rect* rect) { const Detection& detection, const DetectionSpec& detection_spec,
Rect* rect) {
const LocationData location_data = detection.location_data(); const LocationData location_data = detection.location_data();
RET_CHECK(location_data.format() == LocationData::BOUNDING_BOX) RET_CHECK(location_data.format() == LocationData::BOUNDING_BOX)
<< "Only Detection with formats of BOUNDING_BOX can be converted to Rect"; << "Only Detection with formats of BOUNDING_BOX can be converted to Rect";
@@ -52,7 +53,8 @@ constexpr char kNormRectsTag[] = "NORM_RECTS";
} }
::mediapipe::Status DetectionsToRectsCalculator::DetectionToNormalizedRect( ::mediapipe::Status DetectionsToRectsCalculator::DetectionToNormalizedRect(
const Detection& detection, NormalizedRect* rect) { const Detection& detection, const DetectionSpec& detection_spec,
NormalizedRect* rect) {
const LocationData location_data = detection.location_data(); const LocationData location_data = detection.location_data();
RET_CHECK(location_data.format() == LocationData::RELATIVE_BOUNDING_BOX) RET_CHECK(location_data.format() == LocationData::RELATIVE_BOUNDING_BOX)
<< "Only Detection with formats of RELATIVE_BOUNDING_BOX can be " << "Only Detection with formats of RELATIVE_BOUNDING_BOX can be "
@@ -174,27 +176,31 @@ constexpr char kNormRectsTag[] = "NORM_RECTS";
} }
} }
std::pair<int, int> image_size; // Get dynamic calculator options (e.g. `image_size`).
if (rotate_) { const DetectionSpec detection_spec = GetDetectionSpec(cc);
RET_CHECK(!cc->Inputs().Tag(kImageSizeTag).IsEmpty());
image_size = cc->Inputs().Tag(kImageSizeTag).Get<std::pair<int, int>>();
}
if (cc->Outputs().HasTag(kRectTag)) { if (cc->Outputs().HasTag(kRectTag)) {
auto output_rect = absl::make_unique<Rect>(); auto output_rect = absl::make_unique<Rect>();
MP_RETURN_IF_ERROR(DetectionToRect(detections[0], output_rect.get())); MP_RETURN_IF_ERROR(
DetectionToRect(detections[0], detection_spec, output_rect.get()));
if (rotate_) { if (rotate_) {
output_rect->set_rotation(ComputeRotation(detections[0], image_size)); float rotation;
MP_RETURN_IF_ERROR(
ComputeRotation(detections[0], detection_spec, &rotation));
output_rect->set_rotation(rotation);
} }
cc->Outputs().Tag(kRectTag).Add(output_rect.release(), cc->Outputs().Tag(kRectTag).Add(output_rect.release(),
cc->InputTimestamp()); cc->InputTimestamp());
} }
if (cc->Outputs().HasTag(kNormRectTag)) { if (cc->Outputs().HasTag(kNormRectTag)) {
auto output_rect = absl::make_unique<NormalizedRect>(); auto output_rect = absl::make_unique<NormalizedRect>();
MP_RETURN_IF_ERROR( MP_RETURN_IF_ERROR(DetectionToNormalizedRect(detections[0], detection_spec,
DetectionToNormalizedRect(detections[0], output_rect.get())); output_rect.get()));
if (rotate_) { if (rotate_) {
output_rect->set_rotation(ComputeRotation(detections[0], image_size)); float rotation;
MP_RETURN_IF_ERROR(
ComputeRotation(detections[0], detection_spec, &rotation));
output_rect->set_rotation(rotation);
} }
cc->Outputs() cc->Outputs()
.Tag(kNormRectTag) .Tag(kNormRectTag)
@@ -203,11 +209,13 @@ constexpr char kNormRectsTag[] = "NORM_RECTS";
if (cc->Outputs().HasTag(kRectsTag)) { if (cc->Outputs().HasTag(kRectsTag)) {
auto output_rects = absl::make_unique<std::vector<Rect>>(detections.size()); auto output_rects = absl::make_unique<std::vector<Rect>>(detections.size());
for (int i = 0; i < detections.size(); ++i) { for (int i = 0; i < detections.size(); ++i) {
MP_RETURN_IF_ERROR( MP_RETURN_IF_ERROR(DetectionToRect(detections[i], detection_spec,
DetectionToRect(detections[i], &(output_rects->at(i)))); &(output_rects->at(i))));
if (rotate_) { if (rotate_) {
output_rects->at(i).set_rotation( float rotation;
ComputeRotation(detections[i], image_size)); MP_RETURN_IF_ERROR(
ComputeRotation(detections[i], detection_spec, &rotation));
output_rects->at(i).set_rotation(rotation);
} }
} }
cc->Outputs().Tag(kRectsTag).Add(output_rects.release(), cc->Outputs().Tag(kRectsTag).Add(output_rects.release(),
@@ -217,11 +225,13 @@ constexpr char kNormRectsTag[] = "NORM_RECTS";
auto output_rects = auto output_rects =
absl::make_unique<std::vector<NormalizedRect>>(detections.size()); absl::make_unique<std::vector<NormalizedRect>>(detections.size());
for (int i = 0; i < detections.size(); ++i) { for (int i = 0; i < detections.size(); ++i) {
MP_RETURN_IF_ERROR( MP_RETURN_IF_ERROR(DetectionToNormalizedRect(
DetectionToNormalizedRect(detections[i], &(output_rects->at(i)))); detections[i], detection_spec, &(output_rects->at(i))));
if (rotate_) { if (rotate_) {
output_rects->at(i).set_rotation( float rotation;
ComputeRotation(detections[i], image_size)); MP_RETURN_IF_ERROR(
ComputeRotation(detections[i], detection_spec, &rotation));
output_rects->at(i).set_rotation(rotation);
} }
} }
cc->Outputs() cc->Outputs()
@@ -232,21 +242,35 @@ constexpr char kNormRectsTag[] = "NORM_RECTS";
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
float DetectionsToRectsCalculator::ComputeRotation( ::mediapipe::Status DetectionsToRectsCalculator::ComputeRotation(
const Detection& detection, const std::pair<int, int> image_size) { const Detection& detection, const DetectionSpec& detection_spec,
float* rotation) {
const auto& location_data = detection.location_data(); const auto& location_data = detection.location_data();
const auto& image_size = detection_spec.image_size;
RET_CHECK(image_size) << "Image size is required to calculate rotation";
const float x0 = location_data.relative_keypoints(start_keypoint_index_).x() * const float x0 = location_data.relative_keypoints(start_keypoint_index_).x() *
image_size.first; image_size->first;
const float y0 = location_data.relative_keypoints(start_keypoint_index_).y() * const float y0 = location_data.relative_keypoints(start_keypoint_index_).y() *
image_size.second; image_size->second;
const float x1 = location_data.relative_keypoints(end_keypoint_index_).x() * const float x1 = location_data.relative_keypoints(end_keypoint_index_).x() *
image_size.first; image_size->first;
const float y1 = location_data.relative_keypoints(end_keypoint_index_).y() * const float y1 = location_data.relative_keypoints(end_keypoint_index_).y() *
image_size.second; image_size->second;
float rotation = target_angle_ - std::atan2(-(y1 - y0), x1 - x0); *rotation = NormalizeRadians(target_angle_ - std::atan2(-(y1 - y0), x1 - x0));
return NormalizeRadians(rotation); return ::mediapipe::OkStatus();
}
DetectionSpec DetectionsToRectsCalculator::GetDetectionSpec(
const CalculatorContext* cc) {
absl::optional<std::pair<int, int>> image_size;
if (cc->Inputs().HasTag(kImageSizeTag)) {
image_size = cc->Inputs().Tag(kImageSizeTag).Get<std::pair<int, int>>();
}
return {image_size};
} }
REGISTER_CALCULATOR(DetectionsToRectsCalculator); REGISTER_CALCULATOR(DetectionsToRectsCalculator);
@@ -16,6 +16,7 @@
#include <cmath> #include <cmath>
#include "absl/types/optional.h"
#include "mediapipe/calculators/util/detections_to_rects_calculator.pb.h" #include "mediapipe/calculators/util/detections_to_rects_calculator.pb.h"
#include "mediapipe/framework/calculator_framework.h" #include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/calculator_options.pb.h" #include "mediapipe/framework/calculator_options.pb.h"
@@ -27,6 +28,13 @@
namespace mediapipe { namespace mediapipe {
// Dynamic options passed as calculator `input_stream` that can be used for
// calculation of rectangle or rotation for given detection. Does not include
// static calculator options which are available via private fields.
struct DetectionSpec {
absl::optional<std::pair<int, int>> image_size;
};
// A calculator that converts Detection proto to Rect proto. // A calculator that converts Detection proto to Rect proto.
// //
// Detection is the format for encoding one or more detections in an image. // Detection is the format for encoding one or more detections in an image.
@@ -81,13 +89,16 @@ class DetectionsToRectsCalculator : public CalculatorBase {
::mediapipe::Status Process(CalculatorContext* cc) override; ::mediapipe::Status Process(CalculatorContext* cc) override;
protected: protected:
virtual float ComputeRotation(const ::mediapipe::Detection& detection,
const std::pair<int, int> image_size);
virtual ::mediapipe::Status DetectionToRect( virtual ::mediapipe::Status DetectionToRect(
const ::mediapipe::Detection& detection, ::mediapipe::Rect* rect); const ::mediapipe::Detection& detection,
const DetectionSpec& detection_spec, ::mediapipe::Rect* rect);
virtual ::mediapipe::Status DetectionToNormalizedRect( virtual ::mediapipe::Status DetectionToNormalizedRect(
const ::mediapipe::Detection& detection, const ::mediapipe::Detection& detection,
::mediapipe::NormalizedRect* rect); const DetectionSpec& detection_spec, ::mediapipe::NormalizedRect* rect);
virtual ::mediapipe::Status ComputeRotation(
const ::mediapipe::Detection& detection,
const DetectionSpec& detection_spec, float* rotation);
virtual DetectionSpec GetDetectionSpec(const CalculatorContext* cc);
static inline float NormalizeRadians(float angle) { static inline float NormalizeRadians(float angle) {
return angle - 2 * M_PI * std::floor((angle - (-M_PI)) / (2 * M_PI)); return angle - 2 * M_PI * std::floor((angle - (-M_PI)) / (2 * M_PI));
@@ -39,6 +39,8 @@ constexpr char kKeypointLabel[] = "KEYPOINT";
// The ratio of detection label font height to the height of detection bounding // The ratio of detection label font height to the height of detection bounding
// box. // box.
constexpr double kLabelToBoundingBoxRatio = 0.1; constexpr double kLabelToBoundingBoxRatio = 0.1;
// Perserve 2 decimal digits.
constexpr float kNumScoreDecimalDigitsMultipler = 100;
} // namespace } // namespace
@@ -235,18 +237,26 @@ void DetectionsToRenderDataCalculator::AddLabels(
std::string label_str = detection.label().empty() std::string label_str = detection.label().empty()
? absl::StrCat(detection.label_id(i)) ? absl::StrCat(detection.label_id(i))
: detection.label(i); : detection.label(i);
const float rounded_score =
std::round(detection.score(i) * kNumScoreDecimalDigitsMultipler) /
kNumScoreDecimalDigitsMultipler;
std::string label_and_score = std::string label_and_score =
absl::StrCat(label_str, options.text_delimiter(), detection.score(i), absl::StrCat(label_str, options.text_delimiter(), rounded_score,
options.text_delimiter()); options.text_delimiter());
label_and_scores.push_back(label_and_score); label_and_scores.push_back(label_and_score);
} }
std::vector<std::string> labels; std::vector<std::string> labels;
if (options.render_detection_id()) {
const std::string detection_id_str =
absl::StrCat("Id: ", detection.detection_id());
labels.push_back(detection_id_str);
}
if (options.one_label_per_line()) { if (options.one_label_per_line()) {
labels.swap(label_and_scores); labels.insert(labels.end(), label_and_scores.begin(),
label_and_scores.end());
} else { } else {
labels.push_back(absl::StrJoin(label_and_scores, "")); labels.push_back(absl::StrJoin(label_and_scores, ""));
} }
// Add the render annotations for "label(_id),score". // Add the render annotations for "label(_id),score".
for (int i = 0; i < labels.size(); ++i) { for (int i = 0; i < labels.size(); ++i) {
auto label = labels.at(i); auto label = labels.at(i);
@@ -53,4 +53,7 @@ message DetectionsToRenderDataCalculatorOptions {
// instances of this calculator are present in the graph, this value // instances of this calculator are present in the graph, this value
// should be unique among them. // should be unique among them.
optional string scene_class = 7 [default = "DETECTION"]; optional string scene_class = 7 [default = "DETECTION"];
// If true, renders the detection id in the first line before the labels.
optional bool render_detection_id = 8 [default = false];
} }
@@ -0,0 +1,110 @@
// 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.
#include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/formats/detection.pb.h"
#include "mediapipe/framework/formats/location_data.pb.h"
#include "mediapipe/framework/port/status.h"
#include "mediapipe/util/tracking/box_tracker.h"
namespace mediapipe {
namespace {
constexpr char kDetectionsTag[] = "DETECTIONS";
constexpr char kDetectionListTag[] = "DETECTION_LIST";
constexpr char kBoxesTag[] = "BOXES";
} // namespace
// A calculator that converts Detection proto to TimedBoxList proto for
// tracking.
//
// Please note that only Location Data formats of RELATIVE_BOUNDING_BOX are
// supported.
//
// Example config:
// node {
// calculator: "DetectionsToTimedBoxListCalculator"
// input_stream: "DETECTIONS:detections"
// output_stream: "BOXES:boxes"
// }
class DetectionsToTimedBoxListCalculator : public CalculatorBase {
public:
static ::mediapipe::Status GetContract(CalculatorContract* cc) {
RET_CHECK(cc->Inputs().HasTag(kDetectionListTag) ||
cc->Inputs().HasTag(kDetectionsTag))
<< "None of the input streams are provided.";
if (cc->Inputs().HasTag(kDetectionListTag)) {
cc->Inputs().Tag(kDetectionListTag).Set<DetectionList>();
}
if (cc->Inputs().HasTag(kDetectionsTag)) {
cc->Inputs().Tag(kDetectionsTag).Set<std::vector<Detection>>();
}
cc->Outputs().Tag(kBoxesTag).Set<TimedBoxProtoList>();
return ::mediapipe::OkStatus();
}
::mediapipe::Status Open(CalculatorContext* cc) override {
cc->SetOffset(TimestampDiff(0));
return ::mediapipe::OkStatus();
}
::mediapipe::Status Process(CalculatorContext* cc) override;
private:
void ConvertDetectionToTimedBox(const Detection& detection,
TimedBoxProto* box, CalculatorContext* cc);
};
REGISTER_CALCULATOR(DetectionsToTimedBoxListCalculator);
::mediapipe::Status DetectionsToTimedBoxListCalculator::Process(
CalculatorContext* cc) {
auto output_timed_box_list = absl::make_unique<TimedBoxProtoList>();
if (cc->Inputs().HasTag(kDetectionListTag)) {
const auto& detection_list =
cc->Inputs().Tag(kDetectionListTag).Get<DetectionList>();
for (const auto& detection : detection_list.detection()) {
TimedBoxProto* box = output_timed_box_list->add_box();
ConvertDetectionToTimedBox(detection, box, cc);
}
}
if (cc->Inputs().HasTag(kDetectionsTag)) {
const auto& detections =
cc->Inputs().Tag(kDetectionsTag).Get<std::vector<Detection>>();
for (const auto& detection : detections) {
TimedBoxProto* box = output_timed_box_list->add_box();
ConvertDetectionToTimedBox(detection, box, cc);
}
}
cc->Outputs().Tag(kBoxesTag).Add(output_timed_box_list.release(),
cc->InputTimestamp());
return ::mediapipe::OkStatus();
}
void DetectionsToTimedBoxListCalculator::ConvertDetectionToTimedBox(
const Detection& detection, TimedBoxProto* box, CalculatorContext* cc) {
const auto& relative_bounding_box =
detection.location_data().relative_bounding_box();
box->set_left(relative_bounding_box.xmin());
box->set_right(relative_bounding_box.xmin() + relative_bounding_box.width());
box->set_top(relative_bounding_box.ymin());
box->set_bottom(relative_bounding_box.ymin() +
relative_bounding_box.height());
box->set_id(detection.detection_id());
box->set_time_msec(cc->InputTimestamp().Microseconds() / 1000);
}
} // namespace mediapipe
@@ -27,8 +27,8 @@ typedef FilterCollectionCalculator<std::vector<::mediapipe::NormalizedRect>>
REGISTER_CALCULATOR(FilterNormalizedRectCollectionCalculator); REGISTER_CALCULATOR(FilterNormalizedRectCollectionCalculator);
typedef FilterCollectionCalculator< typedef FilterCollectionCalculator<
std::vector<std::vector<::mediapipe::NormalizedLandmark>>> std::vector<::mediapipe::NormalizedLandmarkList>>
FilterLandmarksCollectionCalculator; FilterLandmarkListCollectionCalculator;
REGISTER_CALCULATOR(FilterLandmarksCollectionCalculator); REGISTER_CALCULATOR(FilterLandmarkListCollectionCalculator);
} // namespace mediapipe } // namespace mediapipe
@@ -93,6 +93,7 @@ REGISTER_CALCULATOR(LabelsToRenderDataCalculator);
} }
::mediapipe::Status LabelsToRenderDataCalculator::Open(CalculatorContext* cc) { ::mediapipe::Status LabelsToRenderDataCalculator::Open(CalculatorContext* cc) {
cc->SetOffset(TimestampDiff(0));
options_ = cc->Options<LabelsToRenderDataCalculatorOptions>(); options_ = cc->Options<LabelsToRenderDataCalculatorOptions>();
num_colors_ = options_.color_size(); num_colors_ = options_.color_size();
label_height_px_ = std::ceil(options_.font_height_px() * kFontHeightScale); label_height_px_ = std::ceil(options_.font_height_px() * kFontHeightScale);
@@ -12,20 +12,6 @@
// See the License for the specific language governing permissions and // See the License for the specific language governing permissions and
// limitations under the License. // limitations under the License.
// 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.
#include <cmath> #include <cmath>
#include <vector> #include <vector>
@@ -49,7 +35,7 @@ constexpr char kLetterboxPaddingTag[] = "LETTERBOX_PADDING";
// corresponding input image before letterboxing. // corresponding input image before letterboxing.
// //
// Input: // Input:
// LANDMARKS: An std::vector<NormalizedLandmark> representing landmarks on an // LANDMARKS: A NormalizedLandmarkList representing landmarks on an
// letterboxed image. // letterboxed image.
// //
// LETTERBOX_PADDING: An std::array<float, 4> representing the letterbox // LETTERBOX_PADDING: An std::array<float, 4> representing the letterbox
@@ -57,7 +43,7 @@ constexpr char kLetterboxPaddingTag[] = "LETTERBOX_PADDING";
// image, normalized to [0.f, 1.f] by the letterboxed image dimensions. // image, normalized to [0.f, 1.f] by the letterboxed image dimensions.
// //
// Output: // Output:
// LANDMARKS: An std::vector<NormalizedLandmark> representing landmarks with // LANDMARKS: An NormalizedLandmarkList proto representing landmarks with
// their locations adjusted to the letterbox-removed (non-padded) image. // their locations adjusted to the letterbox-removed (non-padded) image.
// //
// Usage example: // Usage example:
@@ -67,6 +53,15 @@ constexpr char kLetterboxPaddingTag[] = "LETTERBOX_PADDING";
// input_stream: "LETTERBOX_PADDING:letterbox_padding" // input_stream: "LETTERBOX_PADDING:letterbox_padding"
// output_stream: "LANDMARKS:adjusted_landmarks" // output_stream: "LANDMARKS:adjusted_landmarks"
// } // }
//
// node {
// calculator: "LandmarkLetterboxRemovalCalculator"
// input_stream: "LANDMARKS:0:landmarks_0"
// input_stream: "LANDMARKS:1:landmarks_1"
// input_stream: "LETTERBOX_PADDING:letterbox_padding"
// output_stream: "LANDMARKS:0:adjusted_landmarks_0"
// output_stream: "LANDMARKS:1:adjusted_landmarks_1"
// }
class LandmarkLetterboxRemovalCalculator : public CalculatorBase { class LandmarkLetterboxRemovalCalculator : public CalculatorBase {
public: public:
static ::mediapipe::Status GetContract(CalculatorContract* cc) { static ::mediapipe::Status GetContract(CalculatorContract* cc) {
@@ -74,10 +69,20 @@ class LandmarkLetterboxRemovalCalculator : public CalculatorBase {
cc->Inputs().HasTag(kLetterboxPaddingTag)) cc->Inputs().HasTag(kLetterboxPaddingTag))
<< "Missing one or more input streams."; << "Missing one or more input streams.";
cc->Inputs().Tag(kLandmarksTag).Set<std::vector<NormalizedLandmark>>(); RET_CHECK_EQ(cc->Inputs().NumEntries(kLandmarksTag),
cc->Outputs().NumEntries(kLandmarksTag))
<< "Same number of input and output landmarks is required.";
for (CollectionItemId id = cc->Inputs().BeginId(kLandmarksTag);
id != cc->Inputs().EndId(kLandmarksTag); ++id) {
cc->Inputs().Get(id).Set<NormalizedLandmarkList>();
}
cc->Inputs().Tag(kLetterboxPaddingTag).Set<std::array<float, 4>>(); cc->Inputs().Tag(kLetterboxPaddingTag).Set<std::array<float, 4>>();
cc->Outputs().Tag(kLandmarksTag).Set<std::vector<NormalizedLandmark>>(); for (CollectionItemId id = cc->Outputs().BeginId(kLandmarksTag);
id != cc->Outputs().EndId(kLandmarksTag); ++id) {
cc->Outputs().Get(id).Set<NormalizedLandmarkList>();
}
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -89,39 +94,45 @@ class LandmarkLetterboxRemovalCalculator : public CalculatorBase {
} }
::mediapipe::Status Process(CalculatorContext* cc) override { ::mediapipe::Status Process(CalculatorContext* cc) override {
// Only process if there's input landmarks. if (cc->Inputs().Tag(kLetterboxPaddingTag).IsEmpty()) {
if (cc->Inputs().Tag(kLandmarksTag).IsEmpty()) {
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
const auto& input_landmarks =
cc->Inputs().Tag(kLandmarksTag).Get<std::vector<NormalizedLandmark>>();
const auto& letterbox_padding = const auto& letterbox_padding =
cc->Inputs().Tag(kLetterboxPaddingTag).Get<std::array<float, 4>>(); cc->Inputs().Tag(kLetterboxPaddingTag).Get<std::array<float, 4>>();
const float left = letterbox_padding[0]; const float left = letterbox_padding[0];
const float top = letterbox_padding[1]; const float top = letterbox_padding[1];
const float left_and_right = letterbox_padding[0] + letterbox_padding[2]; const float left_and_right = letterbox_padding[0] + letterbox_padding[2];
const float top_and_bottom = letterbox_padding[1] + letterbox_padding[3]; const float top_and_bottom = letterbox_padding[1] + letterbox_padding[3];
auto output_landmarks = CollectionItemId input_id = cc->Inputs().BeginId(kLandmarksTag);
absl::make_unique<std::vector<NormalizedLandmark>>(); CollectionItemId output_id = cc->Outputs().BeginId(kLandmarksTag);
for (const auto& landmark : input_landmarks) { // Number of inputs and outpus is the same according to the contract.
NormalizedLandmark new_landmark; for (; input_id != cc->Inputs().EndId(kLandmarksTag);
const float new_x = (landmark.x() - left) / (1.0f - left_and_right); ++input_id, ++output_id) {
const float new_y = (landmark.y() - top) / (1.0f - top_and_bottom); const auto& input_packet = cc->Inputs().Get(input_id);
if (input_packet.IsEmpty()) {
continue;
}
new_landmark.set_x(new_x); const NormalizedLandmarkList& input_landmarks =
new_landmark.set_y(new_y); input_packet.Get<NormalizedLandmarkList>();
// Keep z-coord as is. NormalizedLandmarkList output_landmarks;
new_landmark.set_z(landmark.z()); for (int i = 0; i < input_landmarks.landmark_size(); ++i) {
const NormalizedLandmark& landmark = input_landmarks.landmark(i);
NormalizedLandmark* new_landmark = output_landmarks.add_landmark();
const float new_x = (landmark.x() - left) / (1.0f - left_and_right);
const float new_y = (landmark.y() - top) / (1.0f - top_and_bottom);
output_landmarks->emplace_back(new_landmark); new_landmark->set_x(new_x);
new_landmark->set_y(new_y);
// Keep z-coord as is.
new_landmark->set_z(landmark.z());
}
cc->Outputs().Get(output_id).AddPacket(
MakePacket<NormalizedLandmarkList>(output_landmarks)
.At(cc->InputTimestamp()));
} }
cc->Outputs()
.Tag(kLandmarksTag)
.Add(output_landmarks.release(), cc->InputTimestamp());
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
}; };
@@ -43,10 +43,10 @@ CalculatorGraphConfig::Node GetDefaultNode() {
TEST(LandmarkLetterboxRemovalCalculatorTest, PaddingLeftRight) { TEST(LandmarkLetterboxRemovalCalculatorTest, PaddingLeftRight) {
CalculatorRunner runner(GetDefaultNode()); CalculatorRunner runner(GetDefaultNode());
auto landmarks = absl::make_unique<std::vector<NormalizedLandmark>>(); auto landmarks = absl::make_unique<NormalizedLandmarkList>();
landmarks->push_back(CreateLandmark(0.5f, 0.5f)); *landmarks->add_landmark() = CreateLandmark(0.5f, 0.5f);
landmarks->push_back(CreateLandmark(0.2f, 0.2f)); *landmarks->add_landmark() = CreateLandmark(0.2f, 0.2f);
landmarks->push_back(CreateLandmark(0.7f, 0.7f)); *landmarks->add_landmark() = CreateLandmark(0.7f, 0.7f);
runner.MutableInputs() runner.MutableInputs()
->Tag("LANDMARKS") ->Tag("LANDMARKS")
.packets.push_back( .packets.push_back(
@@ -61,26 +61,28 @@ TEST(LandmarkLetterboxRemovalCalculatorTest, PaddingLeftRight) {
MP_ASSERT_OK(runner.Run()) << "Calculator execution failed."; MP_ASSERT_OK(runner.Run()) << "Calculator execution failed.";
const std::vector<Packet>& output = runner.Outputs().Tag("LANDMARKS").packets; const std::vector<Packet>& output = runner.Outputs().Tag("LANDMARKS").packets;
ASSERT_EQ(1, output.size()); ASSERT_EQ(1, output.size());
const auto& output_landmarks = const auto& output_landmarks = output[0].Get<NormalizedLandmarkList>();
output[0].Get<std::vector<NormalizedLandmark>>();
EXPECT_EQ(output_landmarks.size(), 3); EXPECT_EQ(output_landmarks.landmark_size(), 3);
EXPECT_THAT(output_landmarks[0].x(), testing::FloatNear(0.6f, 1e-5)); EXPECT_THAT(output_landmarks.landmark(0).x(), testing::FloatNear(0.6f, 1e-5));
EXPECT_THAT(output_landmarks[0].y(), testing::FloatNear(0.5f, 1e-5)); EXPECT_THAT(output_landmarks.landmark(0).y(), testing::FloatNear(0.5f, 1e-5));
EXPECT_THAT(output_landmarks[1].x(), testing::FloatNear(0.0f, 1e-5)); EXPECT_THAT(output_landmarks.landmark(1).x(), testing::FloatNear(0.0f, 1e-5));
EXPECT_THAT(output_landmarks[1].y(), testing::FloatNear(0.2f, 1e-5)); EXPECT_THAT(output_landmarks.landmark(1).y(), testing::FloatNear(0.2f, 1e-5));
EXPECT_THAT(output_landmarks[2].x(), testing::FloatNear(1.0f, 1e-5)); EXPECT_THAT(output_landmarks.landmark(2).x(), testing::FloatNear(1.0f, 1e-5));
EXPECT_THAT(output_landmarks[2].y(), testing::FloatNear(0.7f, 1e-5)); EXPECT_THAT(output_landmarks.landmark(2).y(), testing::FloatNear(0.7f, 1e-5));
} }
TEST(LandmarkLetterboxRemovalCalculatorTest, PaddingTopBottom) { TEST(LandmarkLetterboxRemovalCalculatorTest, PaddingTopBottom) {
CalculatorRunner runner(GetDefaultNode()); CalculatorRunner runner(GetDefaultNode());
auto landmarks = absl::make_unique<std::vector<NormalizedLandmark>>(); auto landmarks = absl::make_unique<NormalizedLandmarkList>();
landmarks->push_back(CreateLandmark(0.5f, 0.5f)); NormalizedLandmark* landmark = landmarks->add_landmark();
landmarks->push_back(CreateLandmark(0.2f, 0.2f)); *landmark = CreateLandmark(0.5f, 0.5f);
landmarks->push_back(CreateLandmark(0.7f, 0.7f)); landmark = landmarks->add_landmark();
*landmark = CreateLandmark(0.2f, 0.2f);
landmark = landmarks->add_landmark();
*landmark = CreateLandmark(0.7f, 0.7f);
runner.MutableInputs() runner.MutableInputs()
->Tag("LANDMARKS") ->Tag("LANDMARKS")
.packets.push_back( .packets.push_back(
@@ -95,17 +97,16 @@ TEST(LandmarkLetterboxRemovalCalculatorTest, PaddingTopBottom) {
MP_ASSERT_OK(runner.Run()) << "Calculator execution failed."; MP_ASSERT_OK(runner.Run()) << "Calculator execution failed.";
const std::vector<Packet>& output = runner.Outputs().Tag("LANDMARKS").packets; const std::vector<Packet>& output = runner.Outputs().Tag("LANDMARKS").packets;
ASSERT_EQ(1, output.size()); ASSERT_EQ(1, output.size());
const auto& output_landmarks = const auto& output_landmarks = output[0].Get<NormalizedLandmarkList>();
output[0].Get<std::vector<NormalizedLandmark>>();
EXPECT_EQ(output_landmarks.size(), 3); EXPECT_EQ(output_landmarks.landmark_size(), 3);
EXPECT_THAT(output_landmarks[0].x(), testing::FloatNear(0.5f, 1e-5)); EXPECT_THAT(output_landmarks.landmark(0).x(), testing::FloatNear(0.5f, 1e-5));
EXPECT_THAT(output_landmarks[0].y(), testing::FloatNear(0.6f, 1e-5)); EXPECT_THAT(output_landmarks.landmark(0).y(), testing::FloatNear(0.6f, 1e-5));
EXPECT_THAT(output_landmarks[1].x(), testing::FloatNear(0.2f, 1e-5)); EXPECT_THAT(output_landmarks.landmark(1).x(), testing::FloatNear(0.2f, 1e-5));
EXPECT_THAT(output_landmarks[1].y(), testing::FloatNear(0.0f, 1e-5)); EXPECT_THAT(output_landmarks.landmark(1).y(), testing::FloatNear(0.0f, 1e-5));
EXPECT_THAT(output_landmarks[2].x(), testing::FloatNear(0.7f, 1e-5)); EXPECT_THAT(output_landmarks.landmark(2).x(), testing::FloatNear(0.7f, 1e-5));
EXPECT_THAT(output_landmarks[2].y(), testing::FloatNear(1.0f, 1e-5)); EXPECT_THAT(output_landmarks.landmark(2).y(), testing::FloatNear(1.0f, 1e-5));
} }
} // namespace mediapipe } // namespace mediapipe
@@ -12,20 +12,6 @@
// See the License for the specific language governing permissions and // See the License for the specific language governing permissions and
// limitations under the License. // limitations under the License.
// 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.
#include <cmath> #include <cmath>
#include <vector> #include <vector>
@@ -47,13 +33,13 @@ constexpr char kRectTag[] = "NORM_RECT";
// Projects normalized landmarks in a rectangle to its original coordinates. The // Projects normalized landmarks in a rectangle to its original coordinates. The
// rectangle must also be in normalized coordinates. // rectangle must also be in normalized coordinates.
// Input: // Input:
// NORM_LANDMARKS: An std::vector<NormalizedLandmark> representing landmarks // NORM_LANDMARKS: A NormalizedLandmarkList representing landmarks
// in a normalized rectangle. // in a normalized rectangle.
// NORM_RECT: An NormalizedRect representing a normalized rectangle in image // NORM_RECT: An NormalizedRect representing a normalized rectangle in image
// coordinates. // coordinates.
// //
// Output: // Output:
// NORM_LANDMARKS: An std::vector<NormalizedLandmark> representing landmarks // NORM_LANDMARKS: A NormalizedLandmarkList representing landmarks
// with their locations adjusted to the image. // with their locations adjusted to the image.
// //
// Usage example: // Usage example:
@@ -63,6 +49,15 @@ constexpr char kRectTag[] = "NORM_RECT";
// input_stream: "NORM_RECT:rect" // input_stream: "NORM_RECT:rect"
// output_stream: "NORM_LANDMARKS:projected_landmarks" // output_stream: "NORM_LANDMARKS:projected_landmarks"
// } // }
//
// node {
// calculator: "LandmarkProjectionCalculator"
// input_stream: "NORM_LANDMARKS:0:landmarks_0"
// input_stream: "NORM_LANDMARKS:1:landmarks_1"
// input_stream: "NORM_RECT:rect"
// output_stream: "NORM_LANDMARKS:0:projected_landmarks_0"
// output_stream: "NORM_LANDMARKS:1:projected_landmarks_1"
// }
class LandmarkProjectionCalculator : public CalculatorBase { class LandmarkProjectionCalculator : public CalculatorBase {
public: public:
static ::mediapipe::Status GetContract(CalculatorContract* cc) { static ::mediapipe::Status GetContract(CalculatorContract* cc) {
@@ -70,10 +65,20 @@ class LandmarkProjectionCalculator : public CalculatorBase {
cc->Inputs().HasTag(kRectTag)) cc->Inputs().HasTag(kRectTag))
<< "Missing one or more input streams."; << "Missing one or more input streams.";
cc->Inputs().Tag(kLandmarksTag).Set<std::vector<NormalizedLandmark>>(); RET_CHECK_EQ(cc->Inputs().NumEntries(kLandmarksTag),
cc->Outputs().NumEntries(kLandmarksTag))
<< "Same number of input and output landmarks is required.";
for (CollectionItemId id = cc->Inputs().BeginId(kLandmarksTag);
id != cc->Inputs().EndId(kLandmarksTag); ++id) {
cc->Inputs().Get(id).Set<NormalizedLandmarkList>();
}
cc->Inputs().Tag(kRectTag).Set<NormalizedRect>(); cc->Inputs().Tag(kRectTag).Set<NormalizedRect>();
cc->Outputs().Tag(kLandmarksTag).Set<std::vector<NormalizedLandmark>>(); for (CollectionItemId id = cc->Outputs().BeginId(kLandmarksTag);
id != cc->Outputs().EndId(kLandmarksTag); ++id) {
cc->Outputs().Get(id).Set<NormalizedLandmarkList>();
}
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -85,42 +90,50 @@ class LandmarkProjectionCalculator : public CalculatorBase {
} }
::mediapipe::Status Process(CalculatorContext* cc) override { ::mediapipe::Status Process(CalculatorContext* cc) override {
const auto& options = if (cc->Inputs().Tag(kRectTag).IsEmpty()) {
cc->Options<::mediapipe::LandmarkProjectionCalculatorOptions>();
// Only process if there's input landmarks.
if (cc->Inputs().Tag(kLandmarksTag).IsEmpty()) {
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
const auto& input_landmarks =
cc->Inputs().Tag(kLandmarksTag).Get<std::vector<NormalizedLandmark>>();
const auto& input_rect = cc->Inputs().Tag(kRectTag).Get<NormalizedRect>(); const auto& input_rect = cc->Inputs().Tag(kRectTag).Get<NormalizedRect>();
auto output_landmarks = const auto& options =
absl::make_unique<std::vector<NormalizedLandmark>>(); cc->Options<::mediapipe::LandmarkProjectionCalculatorOptions>();
for (const auto& landmark : input_landmarks) {
NormalizedLandmark new_landmark;
const float x = landmark.x() - 0.5f; CollectionItemId input_id = cc->Inputs().BeginId(kLandmarksTag);
const float y = landmark.y() - 0.5f; CollectionItemId output_id = cc->Outputs().BeginId(kLandmarksTag);
const float angle = options.ignore_rotation() ? 0 : input_rect.rotation(); // Number of inputs and outpus is the same according to the contract.
float new_x = std::cos(angle) * x - std::sin(angle) * y; for (; input_id != cc->Inputs().EndId(kLandmarksTag);
float new_y = std::sin(angle) * x + std::cos(angle) * y; ++input_id, ++output_id) {
const auto& input_packet = cc->Inputs().Get(input_id);
if (input_packet.IsEmpty()) {
continue;
}
new_x = new_x * input_rect.width() + input_rect.x_center(); const auto& input_landmarks = input_packet.Get<NormalizedLandmarkList>();
new_y = new_y * input_rect.height() + input_rect.y_center(); NormalizedLandmarkList output_landmarks;
for (int i = 0; i < input_landmarks.landmark_size(); ++i) {
const NormalizedLandmark& landmark = input_landmarks.landmark(i);
NormalizedLandmark* new_landmark = output_landmarks.add_landmark();
new_landmark.set_x(new_x); const float x = landmark.x() - 0.5f;
new_landmark.set_y(new_y); const float y = landmark.y() - 0.5f;
// Keep z-coord as is. const float angle =
new_landmark.set_z(landmark.z()); options.ignore_rotation() ? 0 : input_rect.rotation();
float new_x = std::cos(angle) * x - std::sin(angle) * y;
float new_y = std::sin(angle) * x + std::cos(angle) * y;
output_landmarks->emplace_back(new_landmark); new_x = new_x * input_rect.width() + input_rect.x_center();
new_y = new_y * input_rect.height() + input_rect.y_center();
new_landmark->set_x(new_x);
new_landmark->set_y(new_y);
// Keep z-coord as is.
new_landmark->set_z(landmark.z());
}
cc->Outputs().Get(output_id).AddPacket(
MakePacket<NormalizedLandmarkList>(output_landmarks)
.At(cc->InputTimestamp()));
} }
cc->Outputs()
.Tag(kLandmarksTag)
.Add(output_landmarks.release(), cc->InputTimestamp());
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
}; };
@@ -28,8 +28,7 @@ namespace {
constexpr char kDetectionTag[] = "DETECTION"; constexpr char kDetectionTag[] = "DETECTION";
constexpr char kNormalizedLandmarksTag[] = "NORM_LANDMARKS"; constexpr char kNormalizedLandmarksTag[] = "NORM_LANDMARKS";
Detection ConvertLandmarksToDetection( Detection ConvertLandmarksToDetection(const NormalizedLandmarkList& landmarks) {
const std::vector<NormalizedLandmark>& landmarks) {
Detection detection; Detection detection;
LocationData* location_data = detection.mutable_location_data(); LocationData* location_data = detection.mutable_location_data();
@@ -37,7 +36,8 @@ Detection ConvertLandmarksToDetection(
float x_max = std::numeric_limits<float>::min(); float x_max = std::numeric_limits<float>::min();
float y_min = std::numeric_limits<float>::max(); float y_min = std::numeric_limits<float>::max();
float y_max = std::numeric_limits<float>::min(); float y_max = std::numeric_limits<float>::min();
for (const auto& landmark : landmarks) { for (int i = 0; i < landmarks.landmark_size(); ++i) {
const NormalizedLandmark& landmark = landmarks.landmark(i);
x_min = std::min(x_min, landmark.x()); x_min = std::min(x_min, landmark.x());
x_max = std::max(x_max, landmark.x()); x_max = std::max(x_max, landmark.x());
y_min = std::min(y_min, landmark.y()); y_min = std::min(y_min, landmark.y());
@@ -67,7 +67,7 @@ Detection ConvertLandmarksToDetection(
// to specify a subset of landmarks for creating the detection. // to specify a subset of landmarks for creating the detection.
// //
// Input: // Input:
// NOMR_LANDMARKS: A vector of NormalizedLandmark. // NOMR_LANDMARKS: A NormalizedLandmarkList proto.
// //
// Output: // Output:
// DETECTION: A Detection proto. // DETECTION: A Detection proto.
@@ -95,9 +95,7 @@ REGISTER_CALCULATOR(LandmarksToDetectionCalculator);
RET_CHECK(cc->Inputs().HasTag(kNormalizedLandmarksTag)); RET_CHECK(cc->Inputs().HasTag(kNormalizedLandmarksTag));
RET_CHECK(cc->Outputs().HasTag(kDetectionTag)); RET_CHECK(cc->Outputs().HasTag(kDetectionTag));
// TODO: Also support converting Landmark to Detection. // TODO: Also support converting Landmark to Detection.
cc->Inputs() cc->Inputs().Tag(kNormalizedLandmarksTag).Set<NormalizedLandmarkList>();
.Tag(kNormalizedLandmarksTag)
.Set<std::vector<NormalizedLandmark>>();
cc->Outputs().Tag(kDetectionTag).Set<Detection>(); cc->Outputs().Tag(kDetectionTag).Set<Detection>();
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
@@ -113,19 +111,20 @@ REGISTER_CALCULATOR(LandmarksToDetectionCalculator);
::mediapipe::Status LandmarksToDetectionCalculator::Process( ::mediapipe::Status LandmarksToDetectionCalculator::Process(
CalculatorContext* cc) { CalculatorContext* cc) {
const auto& landmarks = cc->Inputs() const auto& landmarks =
.Tag(kNormalizedLandmarksTag) cc->Inputs().Tag(kNormalizedLandmarksTag).Get<NormalizedLandmarkList>();
.Get<std::vector<NormalizedLandmark>>(); RET_CHECK_GT(landmarks.landmark_size(), 0)
RET_CHECK_GT(landmarks.size(), 0) << "Input landmark vector is empty."; << "Input landmark vector is empty.";
auto detection = absl::make_unique<Detection>(); auto detection = absl::make_unique<Detection>();
if (options_.selected_landmark_indices_size()) { if (options_.selected_landmark_indices_size()) {
std::vector<NormalizedLandmark> subset_landmarks( NormalizedLandmarkList subset_landmarks;
options_.selected_landmark_indices_size()); for (int i = 0; i < options_.selected_landmark_indices_size(); ++i) {
for (int i = 0; i < subset_landmarks.size(); ++i) { RET_CHECK_LT(options_.selected_landmark_indices(i),
RET_CHECK_LT(options_.selected_landmark_indices(i), landmarks.size()) landmarks.landmark_size())
<< "Index of landmark subset is out of range."; << "Index of landmark subset is out of range.";
subset_landmarks[i] = landmarks[options_.selected_landmark_indices(i)]; *subset_landmarks.add_landmark() =
landmarks.landmark(options_.selected_landmark_indices(i));
} }
*detection = ConvertLandmarksToDetection(subset_landmarks); *detection = ConvertLandmarksToDetection(subset_landmarks);
} else { } else {
@@ -48,7 +48,7 @@ constexpr char kMatrixTag[] = "MATRIX";
// Converts a vector of landmarks to a vector of floats or a matrix. // Converts a vector of landmarks to a vector of floats or a matrix.
// Input: // Input:
// NORM_LANDMARKS: An std::vector<NormalizedLandmark>. // NORM_LANDMARKS: A NormalizedLandmarkList proto.
// //
// Output: // Output:
// FLOATS(optional): A vector of floats from flattened landmarks. // FLOATS(optional): A vector of floats from flattened landmarks.
@@ -63,7 +63,7 @@ constexpr char kMatrixTag[] = "MATRIX";
class LandmarksToFloatsCalculator : public CalculatorBase { class LandmarksToFloatsCalculator : public CalculatorBase {
public: public:
static ::mediapipe::Status GetContract(CalculatorContract* cc) { static ::mediapipe::Status GetContract(CalculatorContract* cc) {
cc->Inputs().Tag(kLandmarksTag).Set<std::vector<NormalizedLandmark>>(); cc->Inputs().Tag(kLandmarksTag).Set<NormalizedLandmarkList>();
RET_CHECK(cc->Outputs().HasTag(kFloatsTag) || RET_CHECK(cc->Outputs().HasTag(kFloatsTag) ||
cc->Outputs().HasTag(kMatrixTag)); cc->Outputs().HasTag(kMatrixTag));
if (cc->Outputs().HasTag(kFloatsTag)) { if (cc->Outputs().HasTag(kFloatsTag)) {
@@ -94,11 +94,12 @@ class LandmarksToFloatsCalculator : public CalculatorBase {
} }
const auto& input_landmarks = const auto& input_landmarks =
cc->Inputs().Tag(kLandmarksTag).Get<std::vector<NormalizedLandmark>>(); cc->Inputs().Tag(kLandmarksTag).Get<NormalizedLandmarkList>();
if (cc->Outputs().HasTag(kFloatsTag)) { if (cc->Outputs().HasTag(kFloatsTag)) {
auto output_floats = absl::make_unique<std::vector<float>>(); auto output_floats = absl::make_unique<std::vector<float>>();
for (const auto& landmark : input_landmarks) { for (int i = 0; i < input_landmarks.landmark_size(); ++i) {
const NormalizedLandmark& landmark = input_landmarks.landmark(i);
output_floats->emplace_back(landmark.x()); output_floats->emplace_back(landmark.x());
if (num_dimensions_ > 1) { if (num_dimensions_ > 1) {
output_floats->emplace_back(landmark.y()); output_floats->emplace_back(landmark.y());
@@ -113,14 +114,14 @@ class LandmarksToFloatsCalculator : public CalculatorBase {
.Add(output_floats.release(), cc->InputTimestamp()); .Add(output_floats.release(), cc->InputTimestamp());
} else { } else {
auto output_matrix = absl::make_unique<Matrix>(); auto output_matrix = absl::make_unique<Matrix>();
output_matrix->setZero(num_dimensions_, input_landmarks.size()); output_matrix->setZero(num_dimensions_, input_landmarks.landmark_size());
for (int i = 0; i < input_landmarks.size(); ++i) { for (int i = 0; i < input_landmarks.landmark_size(); ++i) {
(*output_matrix)(0, i) = input_landmarks[i].x(); (*output_matrix)(0, i) = input_landmarks.landmark(i).x();
if (num_dimensions_ > 1) { if (num_dimensions_ > 1) {
(*output_matrix)(1, i) = input_landmarks[i].y(); (*output_matrix)(1, i) = input_landmarks.landmark(i).y();
} }
if (num_dimensions_ > 2) { if (num_dimensions_ > 2) {
(*output_matrix)(2, i) = input_landmarks[i].z(); (*output_matrix)(2, i) = input_landmarks.landmark(i).z();
} }
} }
cc->Outputs() cc->Outputs()
@@ -46,12 +46,13 @@ inline float Remap(float x, float lo, float hi, float scale) {
return (x - lo) / (hi - lo + 1e-6) * scale; return (x - lo) / (hi - lo + 1e-6) * scale;
} }
template <class LandmarkType> template <class LandmarkListType, class LandmarkType>
inline void GetMinMaxZ(const std::vector<LandmarkType>& landmarks, float* z_min, inline void GetMinMaxZ(const LandmarkListType& landmarks, float* z_min,
float* z_max) { float* z_max) {
*z_min = std::numeric_limits<float>::max(); *z_min = std::numeric_limits<float>::max();
*z_max = std::numeric_limits<float>::min(); *z_max = std::numeric_limits<float>::min();
for (const auto& landmark : landmarks) { for (int i = 0; i < landmarks.landmark_size(); ++i) {
const LandmarkType& landmark = landmarks.landmark(i);
*z_min = std::min(landmark.z(), *z_min); *z_min = std::min(landmark.z(), *z_min);
*z_max = std::max(landmark.z(), *z_max); *z_max = std::max(landmark.z(), *z_max);
} }
@@ -73,7 +74,7 @@ void SetColorSizeValueFromZ(float z, float z_min, float z_max,
} // namespace } // namespace
// A calculator that converts Landmark proto to RenderData proto for // A calculator that converts Landmark proto to RenderData proto for
// visualization. The input should be std::vector<Landmark>. It is also possible // visualization. The input should be LandmarkList proto. It is also possible
// to specify the connections between landmarks. // to specify the connections between landmarks.
// //
// Example config: // Example config:
@@ -121,11 +122,11 @@ class LandmarksToRenderDataCalculator : public CalculatorBase {
const LandmarksToRenderDataCalculatorOptions& options, bool normalized, const LandmarksToRenderDataCalculatorOptions& options, bool normalized,
int gray_val1, int gray_val2, RenderData* render_data); int gray_val1, int gray_val2, RenderData* render_data);
template <class LandmarkType> template <class LandmarkListType>
void AddConnections(const std::vector<LandmarkType>& landmarks, void AddConnections(const LandmarkListType& landmarks, bool normalized,
bool normalized, RenderData* render_data); RenderData* render_data);
template <class LandmarkType> template <class LandmarkListType>
void AddConnectionsWithDepth(const std::vector<LandmarkType>& landmarks, void AddConnectionsWithDepth(const LandmarkListType& landmarks,
bool normalized, float min_z, float max_z, bool normalized, float min_z, float max_z,
RenderData* render_data); RenderData* render_data);
@@ -144,10 +145,10 @@ REGISTER_CALCULATOR(LandmarksToRenderDataCalculator);
"normalized landmarks."; "normalized landmarks.";
if (cc->Inputs().HasTag(kLandmarksTag)) { if (cc->Inputs().HasTag(kLandmarksTag)) {
cc->Inputs().Tag(kLandmarksTag).Set<std::vector<Landmark>>(); cc->Inputs().Tag(kLandmarksTag).Set<LandmarkList>();
} }
if (cc->Inputs().HasTag(kNormLandmarksTag)) { if (cc->Inputs().HasTag(kNormLandmarksTag)) {
cc->Inputs().Tag(kNormLandmarksTag).Set<std::vector<NormalizedLandmark>>(); cc->Inputs().Tag(kNormLandmarksTag).Set<NormalizedLandmarkList>();
} }
cc->Outputs().Tag(kRenderDataTag).Set<RenderData>(); cc->Outputs().Tag(kRenderDataTag).Set<RenderData>();
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
@@ -169,16 +170,17 @@ REGISTER_CALCULATOR(LandmarksToRenderDataCalculator);
float z_max = 0.f; float z_max = 0.f;
if (cc->Inputs().HasTag(kLandmarksTag)) { if (cc->Inputs().HasTag(kLandmarksTag)) {
const auto& landmarks = const LandmarkList& landmarks =
cc->Inputs().Tag(kLandmarksTag).Get<std::vector<Landmark>>(); cc->Inputs().Tag(kLandmarksTag).Get<LandmarkList>();
RET_CHECK_EQ(options_.landmark_connections_size() % 2, 0) RET_CHECK_EQ(options_.landmark_connections_size() % 2, 0)
<< "Number of entries in landmark connections must be a multiple of 2"; << "Number of entries in landmark connections must be a multiple of 2";
if (visualize_depth) { if (visualize_depth) {
GetMinMaxZ(landmarks, &z_min, &z_max); GetMinMaxZ<LandmarkList, Landmark>(landmarks, &z_min, &z_max);
} }
// Only change rendering if there are actually z values other than 0. // Only change rendering if there are actually z values other than 0.
visualize_depth &= ((z_max - z_min) > 1e-3); visualize_depth &= ((z_max - z_min) > 1e-3);
for (const auto& landmark : landmarks) { for (int i = 0; i < landmarks.landmark_size(); ++i) {
const Landmark& landmark = landmarks.landmark(i);
auto* landmark_data_render = auto* landmark_data_render =
AddPointRenderData(options_, render_data.get()); AddPointRenderData(options_, render_data.get());
if (visualize_depth) { if (visualize_depth) {
@@ -191,25 +193,27 @@ REGISTER_CALCULATOR(LandmarksToRenderDataCalculator);
landmark_data->set_y(landmark.y()); landmark_data->set_y(landmark.y());
} }
if (visualize_depth) { if (visualize_depth) {
AddConnectionsWithDepth(landmarks, /*normalized=*/false, z_min, z_max, AddConnectionsWithDepth<LandmarkList>(landmarks, /*normalized=*/false,
render_data.get()); z_min, z_max, render_data.get());
} else { } else {
AddConnections(landmarks, /*normalized=*/false, render_data.get()); AddConnections<LandmarkList>(landmarks, /*normalized=*/false,
render_data.get());
} }
} }
if (cc->Inputs().HasTag(kNormLandmarksTag)) { if (cc->Inputs().HasTag(kNormLandmarksTag)) {
const auto& landmarks = cc->Inputs() const NormalizedLandmarkList& landmarks =
.Tag(kNormLandmarksTag) cc->Inputs().Tag(kNormLandmarksTag).Get<NormalizedLandmarkList>();
.Get<std::vector<NormalizedLandmark>>();
RET_CHECK_EQ(options_.landmark_connections_size() % 2, 0) RET_CHECK_EQ(options_.landmark_connections_size() % 2, 0)
<< "Number of entries in landmark connections must be a multiple of 2"; << "Number of entries in landmark connections must be a multiple of 2";
if (visualize_depth) { if (visualize_depth) {
GetMinMaxZ(landmarks, &z_min, &z_max); GetMinMaxZ<NormalizedLandmarkList, NormalizedLandmark>(landmarks, &z_min,
&z_max);
} }
// Only change rendering if there are actually z values other than 0. // Only change rendering if there are actually z values other than 0.
visualize_depth &= ((z_max - z_min) > 1e-3); visualize_depth &= ((z_max - z_min) > 1e-3);
for (const auto& landmark : landmarks) { for (int i = 0; i < landmarks.landmark_size(); ++i) {
const NormalizedLandmark& landmark = landmarks.landmark(i);
auto* landmark_data_render = auto* landmark_data_render =
AddPointRenderData(options_, render_data.get()); AddPointRenderData(options_, render_data.get());
if (visualize_depth) { if (visualize_depth) {
@@ -222,10 +226,11 @@ REGISTER_CALCULATOR(LandmarksToRenderDataCalculator);
landmark_data->set_y(landmark.y()); landmark_data->set_y(landmark.y());
} }
if (visualize_depth) { if (visualize_depth) {
AddConnectionsWithDepth(landmarks, /*normalized=*/true, z_min, z_max, AddConnectionsWithDepth<NormalizedLandmarkList>(
render_data.get()); landmarks, /*normalized=*/true, z_min, z_max, render_data.get());
} else { } else {
AddConnections(landmarks, /*normalized=*/true, render_data.get()); AddConnections<NormalizedLandmarkList>(landmarks, /*normalized=*/true,
render_data.get());
} }
} }
@@ -235,13 +240,13 @@ REGISTER_CALCULATOR(LandmarksToRenderDataCalculator);
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
template <class LandmarkType> template <class LandmarkListType>
void LandmarksToRenderDataCalculator::AddConnectionsWithDepth( void LandmarksToRenderDataCalculator::AddConnectionsWithDepth(
const std::vector<LandmarkType>& landmarks, bool normalized, float min_z, const LandmarkListType& landmarks, bool normalized, float min_z,
float max_z, RenderData* render_data) { float max_z, RenderData* render_data) {
for (int i = 0; i < options_.landmark_connections_size(); i += 2) { for (int i = 0; i < options_.landmark_connections_size(); i += 2) {
const auto& ld0 = landmarks[options_.landmark_connections(i)]; const auto& ld0 = landmarks.landmark(options_.landmark_connections(i));
const auto& ld1 = landmarks[options_.landmark_connections(i + 1)]; const auto& ld1 = landmarks.landmark(options_.landmark_connections(i + 1));
const int gray_val1 = const int gray_val1 =
255 - static_cast<int>(Remap(ld0.z(), min_z, max_z, 255)); 255 - static_cast<int>(Remap(ld0.z(), min_z, max_z, 255));
const int gray_val2 = const int gray_val2 =
@@ -272,13 +277,13 @@ void LandmarksToRenderDataCalculator::AddConnectionToRenderData(
connection_annotation->set_thickness(options.thickness()); connection_annotation->set_thickness(options.thickness());
} }
template <class LandmarkType> template <class LandmarkListType>
void LandmarksToRenderDataCalculator::AddConnections( void LandmarksToRenderDataCalculator::AddConnections(
const std::vector<LandmarkType>& landmarks, bool normalized, const LandmarkListType& landmarks, bool normalized,
RenderData* render_data) { RenderData* render_data) {
for (int i = 0; i < options_.landmark_connections_size(); i += 2) { for (int i = 0; i < options_.landmark_connections_size(); i += 2) {
const auto& ld0 = landmarks[options_.landmark_connections(i)]; const auto& ld0 = landmarks.landmark(options_.landmark_connections(i));
const auto& ld1 = landmarks[options_.landmark_connections(i + 1)]; const auto& ld1 = landmarks.landmark(options_.landmark_connections(i + 1));
AddConnectionToRenderData(ld0.x(), ld0.y(), ld1.x(), ld1.y(), options_, AddConnectionToRenderData(ld0.x(), ld0.y(), ld1.x(), ld1.y(), options_,
normalized, render_data); normalized, render_data);
} }
@@ -0,0 +1,75 @@
// 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.
#include <memory>
#include <string>
#include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/port/file_helpers.h"
#include "mediapipe/framework/port/status.h"
namespace mediapipe {
// The calculator takes the path to local directory and desired file suffix to
// mach as input side packets, and outputs the contents of those files that
// match the pattern. Those matched files will be sent sequentially through the
// output stream with incremental timestamp difference by 1.
//
// Example config:
// node {
// calculator: "LocalFilePatternContentsCalculator"
// input_side_packet: "FILE_DIRECTORY:file_directory"
// input_side_packet: "FILE_SUFFIX:file_suffix"
// output_stream: "CONTENTS:contents"
// }
class LocalFilePatternContentsCalculator : public CalculatorBase {
public:
static ::mediapipe::Status GetContract(CalculatorContract* cc) {
cc->InputSidePackets().Tag("FILE_DIRECTORY").Set<std::string>();
cc->InputSidePackets().Tag("FILE_SUFFIX").Set<std::string>();
cc->Outputs().Tag("CONTENTS").Set<std::string>();
return ::mediapipe::OkStatus();
}
::mediapipe::Status Open(CalculatorContext* cc) override {
MP_RETURN_IF_ERROR(::mediapipe::file::MatchFileTypeInDirectory(
cc->InputSidePackets().Tag("FILE_DIRECTORY").Get<std::string>(),
cc->InputSidePackets().Tag("FILE_SUFFIX").Get<std::string>(),
&filenames_));
return ::mediapipe::OkStatus();
}
::mediapipe::Status Process(CalculatorContext* cc) override {
if (current_output_ < filenames_.size()) {
auto contents = absl::make_unique<std::string>();
LOG(INFO) << filenames_[current_output_];
MP_RETURN_IF_ERROR(mediapipe::file::GetContents(
filenames_[current_output_], contents.get()));
++current_output_;
cc->Outputs()
.Tag("CONTENTS")
.Add(contents.release(), Timestamp(current_output_));
} else {
return tool::StatusStop();
}
return ::mediapipe::OkStatus();
}
private:
std::vector<std::string> filenames_;
int current_output_ = 0;
};
REGISTER_CALCULATOR(LocalFilePatternContentsCalculator);
} // namespace mediapipe
@@ -45,15 +45,17 @@ RenderAnnotation::Rectangle* NewRect(
void SetRect(bool normalized, double xmin, double ymin, double width, void SetRect(bool normalized, double xmin, double ymin, double width,
double height, double rotation, double height, double rotation,
RenderAnnotation::Rectangle* rect) { RenderAnnotation::Rectangle* rect) {
if (xmin + width < 0.0 || ymin + height < 0.0) return; if (rotation == 0.0) {
if (normalized) { if (xmin + width < 0.0 || ymin + height < 0.0) return;
if (xmin > 1.0 || ymin > 1.0) return; if (normalized) {
if (xmin > 1.0 || ymin > 1.0) return;
}
} }
rect->set_normalized(normalized); rect->set_normalized(normalized);
rect->set_left(normalized ? std::max(xmin, 0.0) : xmin); rect->set_left(xmin);
rect->set_top(normalized ? std::max(ymin, 0.0) : ymin); rect->set_top(ymin);
rect->set_right(normalized ? std::min(xmin + width, 1.0) : xmin + width); rect->set_right(xmin + width);
rect->set_bottom(normalized ? std::min(ymin + height, 1.0) : ymin + height); rect->set_bottom(ymin + height);
rect->set_rotation(rotation); rect->set_rotation(rotation);
} }
@@ -0,0 +1,105 @@
// 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.
#include "mediapipe/calculators/util/timed_box_list_id_to_label_calculator.pb.h"
#include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/packet.h"
#include "mediapipe/framework/port/status.h"
#include "mediapipe/util/resource_util.h"
#include "mediapipe/util/tracking/box_tracker.pb.h"
#if defined(MEDIAPIPE_MOBILE)
#include "mediapipe/util/android/file/base/file.h"
#include "mediapipe/util/android/file/base/helpers.h"
#else
#include "mediapipe/framework/port/file_helpers.h"
#endif
namespace mediapipe {
using mediapipe::TimedBoxProto;
using mediapipe::TimedBoxProtoList;
// Takes a label map (from label IDs to names), and populate the label field in
// TimedBoxProto according to it's ID.
//
// Example usage:
// node {
// calculator: "TimedBoxListIdToLabelCalculator"
// input_stream: "input_timed_box_list"
// output_stream: "output_timed_box_list"
// node_options: {
// [mediapipe.TimedBoxListIdToLabelCalculatorOptions] {
// label_map_path: "labelmap.txt"
// }
// }
// }
class TimedBoxListIdToLabelCalculator : public CalculatorBase {
public:
static ::mediapipe::Status GetContract(CalculatorContract* cc);
::mediapipe::Status Open(CalculatorContext* cc) override;
::mediapipe::Status Process(CalculatorContext* cc) override;
private:
std::unordered_map<int, std::string> label_map_;
};
REGISTER_CALCULATOR(TimedBoxListIdToLabelCalculator);
::mediapipe::Status TimedBoxListIdToLabelCalculator::GetContract(
CalculatorContract* cc) {
cc->Inputs().Index(0).Set<TimedBoxProtoList>();
cc->Outputs().Index(0).Set<TimedBoxProtoList>();
return ::mediapipe::OkStatus();
}
::mediapipe::Status TimedBoxListIdToLabelCalculator::Open(
CalculatorContext* cc) {
cc->SetOffset(TimestampDiff(0));
const auto& options =
cc->Options<::mediapipe::TimedBoxListIdToLabelCalculatorOptions>();
std::string string_path;
ASSIGN_OR_RETURN(string_path, PathToResourceAsFile(options.label_map_path()));
std::string label_map_string;
MP_RETURN_IF_ERROR(file::GetContents(string_path, &label_map_string));
std::istringstream stream(label_map_string);
std::string line;
int i = 0;
while (std::getline(stream, line)) {
label_map_[i++] = line;
}
return ::mediapipe::OkStatus();
}
::mediapipe::Status TimedBoxListIdToLabelCalculator::Process(
CalculatorContext* cc) {
const auto& input_list = cc->Inputs().Index(0).Get<TimedBoxProtoList>();
auto output_list = absl::make_unique<TimedBoxProtoList>();
for (const auto& input_box : input_list.box()) {
TimedBoxProto* box_ptr = output_list->add_box();
*box_ptr = input_box;
if (label_map_.find(input_box.id()) != label_map_.end()) {
box_ptr->set_label(label_map_[input_box.id()]);
}
}
cc->Outputs().Index(0).Add(output_list.release(), cc->InputTimestamp());
return ::mediapipe::OkStatus();
}
} // namespace mediapipe
@@ -0,0 +1,28 @@
// 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.
syntax = "proto2";
package mediapipe;
import "mediapipe/framework/calculator.proto";
message TimedBoxListIdToLabelCalculatorOptions {
extend mediapipe.CalculatorOptions {
optional TimedBoxListIdToLabelCalculatorOptions ext = 297701606;
}
// Path to a label map file for getting the actual name of detected classes.
optional string label_map_path = 1;
}
@@ -0,0 +1,166 @@
// 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.
#include "absl/memory/memory.h"
#include "absl/strings/str_cat.h"
#include "absl/strings/str_join.h"
#include "mediapipe/calculators/util/timed_box_list_to_render_data_calculator.pb.h"
#include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/calculator_options.pb.h"
#include "mediapipe/framework/port/ret_check.h"
#include "mediapipe/util/color.pb.h"
#include "mediapipe/util/render_data.pb.h"
#include "mediapipe/util/tracking/box_tracker.pb.h"
#include "mediapipe/util/tracking/tracking.pb.h"
namespace mediapipe {
namespace {
constexpr char kTimedBoxListTag[] = "BOX_LIST";
constexpr char kRenderDataTag[] = "RENDER_DATA";
void AddTimedBoxProtoToRenderData(
const TimedBoxProto& box_proto,
const TimedBoxListToRenderDataCalculatorOptions& options,
RenderData* render_data) {
if (box_proto.has_quad() && box_proto.quad().vertices_size() > 0 &&
box_proto.quad().vertices_size() % 2 == 0) {
const int num_corners = box_proto.quad().vertices_size() / 2;
for (int i = 0; i < num_corners; ++i) {
const int next_corner = (i + 1) % num_corners;
auto* line_annotation = render_data->add_render_annotations();
line_annotation->mutable_color()->set_r(options.box_color().r());
line_annotation->mutable_color()->set_g(options.box_color().g());
line_annotation->mutable_color()->set_b(options.box_color().b());
line_annotation->set_thickness(options.thickness());
RenderAnnotation::Line* line = line_annotation->mutable_line();
line->set_normalized(true);
line->set_x_start(box_proto.quad().vertices(i * 2));
line->set_y_start(box_proto.quad().vertices(i * 2 + 1));
line->set_x_end(box_proto.quad().vertices(next_corner * 2));
line->set_y_end(box_proto.quad().vertices(next_corner * 2 + 1));
}
} else {
auto* rect_annotation = render_data->add_render_annotations();
rect_annotation->mutable_color()->set_r(options.box_color().r());
rect_annotation->mutable_color()->set_g(options.box_color().g());
rect_annotation->mutable_color()->set_b(options.box_color().b());
rect_annotation->set_thickness(options.thickness());
RenderAnnotation::Rectangle* rect = rect_annotation->mutable_rectangle();
rect->set_normalized(true);
rect->set_left(box_proto.left());
rect->set_right(box_proto.right());
rect->set_top(box_proto.top());
rect->set_bottom(box_proto.bottom());
rect->set_rotation(box_proto.rotation());
}
if (box_proto.has_label()) {
auto* label_annotation = render_data->add_render_annotations();
label_annotation->mutable_color()->set_r(options.box_color().r());
label_annotation->mutable_color()->set_g(options.box_color().g());
label_annotation->mutable_color()->set_b(options.box_color().b());
label_annotation->set_thickness(options.thickness());
RenderAnnotation::Text* text = label_annotation->mutable_text();
text->set_display_text(box_proto.label());
text->set_normalized(true);
constexpr float text_left_start = 0.3f;
text->set_left((1.0f - text_left_start) * box_proto.left() +
text_left_start * box_proto.right());
constexpr float text_baseline = 0.6f;
text->set_baseline(text_baseline * box_proto.bottom() +
(1.0f - text_baseline) * box_proto.top());
constexpr float text_height = 0.2f;
text->set_font_height((box_proto.bottom() - box_proto.top()) * text_height);
}
}
} // namespace
// A calculator that converts TimedBoxProtoList proto to RenderData proto for
// visualization. If the input TimedBoxProto contains `quad` field, this
// calculator will draw a quadrilateral based on it. Otherwise this calculator
// will draw a rotated rectangle based on `top`, `bottom`, `left`, `right` and
// `rotation` fields
//
// Example config:
// node {
// calculator: "TimedBoxListToRenderDataCalculator"
// input_stream: "BOX_LIST:landmarks"
// output_stream: "RENDER_DATA:render_data"
// options {
// [TimedBoxListToRenderDataCalculatorOptions.ext] {
// box_color { r: 0 g: 255 b: 0 }
// thickness: 4.0
// }
// }
// }
class TimedBoxListToRenderDataCalculator : public CalculatorBase {
public:
TimedBoxListToRenderDataCalculator() {}
~TimedBoxListToRenderDataCalculator() override {}
TimedBoxListToRenderDataCalculator(
const TimedBoxListToRenderDataCalculator&) = delete;
TimedBoxListToRenderDataCalculator& operator=(
const TimedBoxListToRenderDataCalculator&) = delete;
static ::mediapipe::Status GetContract(CalculatorContract* cc);
::mediapipe::Status Open(CalculatorContext* cc) override;
::mediapipe::Status Process(CalculatorContext* cc) override;
private:
TimedBoxListToRenderDataCalculatorOptions options_;
};
REGISTER_CALCULATOR(TimedBoxListToRenderDataCalculator);
::mediapipe::Status TimedBoxListToRenderDataCalculator::GetContract(
CalculatorContract* cc) {
if (cc->Inputs().HasTag(kTimedBoxListTag)) {
cc->Inputs().Tag(kTimedBoxListTag).Set<TimedBoxProtoList>();
}
cc->Outputs().Tag(kRenderDataTag).Set<RenderData>();
return ::mediapipe::OkStatus();
}
::mediapipe::Status TimedBoxListToRenderDataCalculator::Open(
CalculatorContext* cc) {
cc->SetOffset(TimestampDiff(0));
options_ = cc->Options<TimedBoxListToRenderDataCalculatorOptions>();
return ::mediapipe::OkStatus();
}
::mediapipe::Status TimedBoxListToRenderDataCalculator::Process(
CalculatorContext* cc) {
auto render_data = absl::make_unique<RenderData>();
if (cc->Inputs().HasTag(kTimedBoxListTag)) {
const auto& box_list =
cc->Inputs().Tag(kTimedBoxListTag).Get<TimedBoxProtoList>();
for (const auto& box : box_list.box()) {
AddTimedBoxProtoToRenderData(box, options_, render_data.get());
}
}
cc->Outputs()
.Tag(kRenderDataTag)
.Add(render_data.release(), cc->InputTimestamp());
return ::mediapipe::OkStatus();
}
} // namespace mediapipe
@@ -0,0 +1,32 @@
// 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.
syntax = "proto2";
package mediapipe;
import "mediapipe/framework/calculator.proto";
import "mediapipe/util/color.proto";
message TimedBoxListToRenderDataCalculatorOptions {
extend CalculatorOptions {
optional TimedBoxListToRenderDataCalculatorOptions ext = 289899854;
}
// Color of boxes.
optional Color box_color = 1;
// Thickness of the drawing of boxes.
optional double thickness = 2 [default = 1.0];
}
@@ -29,8 +29,7 @@
#include "mediapipe/framework/port/statusor.h" #include "mediapipe/framework/port/statusor.h"
#include "mediapipe/util/resource_util.h" #include "mediapipe/util/resource_util.h"
#if defined(MEDIAPIPE_LITE) || defined(__EMSCRIPTEN__) || \ #if defined(MEDIAPIPE_MOBILE)
defined(__ANDROID__) || (defined(__APPLE__) && !TARGET_OS_OSX)
#include "mediapipe/util/android/file/base/file.h" #include "mediapipe/util/android/file/base/file.h"
#include "mediapipe/util/android/file/base/helpers.h" #include "mediapipe/util/android/file/base/helpers.h"
#else #else
+257
View File
@@ -37,6 +37,86 @@ proto_library(
deps = ["//mediapipe/framework:calculator_proto"], deps = ["//mediapipe/framework:calculator_proto"],
) )
proto_library(
name = "motion_analysis_calculator_proto",
srcs = ["motion_analysis_calculator.proto"],
deps = [
"//mediapipe/framework:calculator_proto",
"//mediapipe/util/tracking:motion_analysis_proto",
],
)
proto_library(
name = "flow_packager_calculator_proto",
srcs = ["flow_packager_calculator.proto"],
deps = [
"//mediapipe/framework:calculator_proto",
"//mediapipe/util/tracking:flow_packager_proto",
],
)
proto_library(
name = "box_tracker_calculator_proto",
srcs = ["box_tracker_calculator.proto"],
visibility = ["//visibility:public"],
deps = [
"//mediapipe/framework:calculator_proto",
"//mediapipe/util/tracking:box_tracker_proto",
],
)
proto_library(
name = "video_pre_stream_calculator_proto",
srcs = ["video_pre_stream_calculator.proto"],
visibility = ["//visibility:public"],
deps = [
"//mediapipe/framework:calculator_proto",
],
)
mediapipe_cc_proto_library(
name = "motion_analysis_calculator_cc_proto",
srcs = ["motion_analysis_calculator.proto"],
cc_deps = [
"//mediapipe/framework:calculator_cc_proto",
"//mediapipe/util/tracking:motion_analysis_cc_proto",
],
visibility = ["//visibility:public"],
deps = [":motion_analysis_calculator_proto"],
)
mediapipe_cc_proto_library(
name = "flow_packager_calculator_cc_proto",
srcs = ["flow_packager_calculator.proto"],
cc_deps = [
"//mediapipe/framework:calculator_cc_proto",
"//mediapipe/util/tracking:flow_packager_cc_proto",
],
visibility = ["//visibility:public"],
deps = [":flow_packager_calculator_proto"],
)
mediapipe_cc_proto_library(
name = "box_tracker_calculator_cc_proto",
srcs = ["box_tracker_calculator.proto"],
cc_deps = [
"//mediapipe/framework:calculator_cc_proto",
"//mediapipe/util/tracking:box_tracker_cc_proto",
],
visibility = ["//visibility:public"],
deps = [":box_tracker_calculator_proto"],
)
mediapipe_cc_proto_library(
name = "video_pre_stream_calculator_cc_proto",
srcs = ["video_pre_stream_calculator.proto"],
cc_deps = [
"//mediapipe/framework:calculator_cc_proto",
],
visibility = ["//visibility:public"],
deps = [":video_pre_stream_calculator_proto"],
)
mediapipe_cc_proto_library( mediapipe_cc_proto_library(
name = "flow_to_image_calculator_cc_proto", name = "flow_to_image_calculator_cc_proto",
srcs = ["flow_to_image_calculator.proto"], srcs = ["flow_to_image_calculator.proto"],
@@ -131,6 +211,107 @@ cc_library(
alwayslink = 1, alwayslink = 1,
) )
cc_library(
name = "motion_analysis_calculator",
srcs = ["motion_analysis_calculator.cc"],
visibility = ["//visibility:public"],
deps = [
":motion_analysis_calculator_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:integral_types",
"//mediapipe/framework/port:logging",
"//mediapipe/framework/port:ret_check",
"//mediapipe/framework/port:status",
"//mediapipe/util/tracking:camera_motion",
"//mediapipe/util/tracking:camera_motion_cc_proto",
"//mediapipe/util/tracking:frame_selection_cc_proto",
"//mediapipe/util/tracking:motion_analysis",
"//mediapipe/util/tracking:motion_estimation",
"//mediapipe/util/tracking:motion_models",
"//mediapipe/util/tracking:region_flow_cc_proto",
"@com_google_absl//absl/strings",
],
alwayslink = 1,
)
cc_library(
name = "flow_packager_calculator",
srcs = ["flow_packager_calculator.cc"],
visibility = ["//visibility:public"],
deps = [
":flow_packager_calculator_cc_proto",
"//mediapipe/framework:calculator_framework",
"//mediapipe/framework/port:integral_types",
"//mediapipe/framework/port:logging",
"//mediapipe/util/tracking:camera_motion_cc_proto",
"//mediapipe/util/tracking:flow_packager",
"//mediapipe/util/tracking:region_flow_cc_proto",
"@com_google_absl//absl/strings",
"@com_google_absl//absl/strings:str_format",
],
alwayslink = 1,
)
cc_library(
name = "box_tracker_calculator",
srcs = ["box_tracker_calculator.cc"],
visibility = ["//visibility:public"],
deps = [
":box_tracker_calculator_cc_proto",
"//mediapipe/framework:calculator_framework",
"//mediapipe/framework/formats:image_frame",
"//mediapipe/framework/formats:image_frame_opencv",
"//mediapipe/framework/formats:video_stream_header", # fixdeps: keep -- required for exobazel build.
"//mediapipe/framework/port:integral_types",
"//mediapipe/framework/port:logging",
"//mediapipe/framework/port:parse_text_proto",
"//mediapipe/framework/port:ret_check",
"//mediapipe/framework/port:status",
"//mediapipe/framework/tool:options_util",
"//mediapipe/util/tracking",
"//mediapipe/util/tracking:box_tracker",
"//mediapipe/util/tracking:tracking_visualization_utilities",
"@com_google_absl//absl/strings",
],
alwayslink = 1,
)
cc_library(
name = "tracked_detection_manager_calculator",
srcs = ["tracked_detection_manager_calculator.cc"],
visibility = ["//visibility:public"],
deps = [
"//mediapipe/framework:calculator_framework",
"//mediapipe/framework/formats:detection_cc_proto",
"//mediapipe/framework/formats:location_data_cc_proto",
"//mediapipe/framework/formats:rect_cc_proto",
"//mediapipe/framework/port:status",
"//mediapipe/util/tracking",
"//mediapipe/util/tracking:box_tracker",
"//mediapipe/util/tracking:tracked_detection",
"//mediapipe/util/tracking:tracked_detection_manager",
"//mediapipe/util/tracking:tracking_visualization_utilities",
"@com_google_absl//absl/container:node_hash_map",
],
alwayslink = 1,
)
cc_library(
name = "video_pre_stream_calculator",
srcs = ["video_pre_stream_calculator.cc"],
visibility = ["//visibility:public"],
deps = [
":video_pre_stream_calculator_cc_proto",
"//mediapipe/framework:calculator_framework",
"//mediapipe/framework/formats:image_frame",
"//mediapipe/framework/formats:video_stream_header",
],
alwayslink = 1,
)
filegroup( filegroup(
name = "test_videos", name = "test_videos",
srcs = [ srcs = [
@@ -187,6 +368,7 @@ cc_test(
cc_test( cc_test(
name = "tvl1_optical_flow_calculator_test", name = "tvl1_optical_flow_calculator_test",
srcs = ["tvl1_optical_flow_calculator_test.cc"], srcs = ["tvl1_optical_flow_calculator_test.cc"],
linkstatic = 1,
deps = [ deps = [
":tvl1_optical_flow_calculator", ":tvl1_optical_flow_calculator",
"//mediapipe/framework:calculator_framework", "//mediapipe/framework:calculator_framework",
@@ -201,3 +383,78 @@ cc_test(
"//mediapipe/framework/port:parse_text_proto", "//mediapipe/framework/port:parse_text_proto",
], ],
) )
MEDIAPIPE_DEPS = [
"//mediapipe/calculators/video:box_tracker_calculator",
"//mediapipe/calculators/video:flow_packager_calculator",
"//mediapipe/calculators/video:motion_analysis_calculator",
"//mediapipe/framework/stream_handler:fixed_size_input_stream_handler",
"//mediapipe/framework/stream_handler:sync_set_input_stream_handler",
]
mediapipe_binary_graph(
name = "parallel_tracker_binarypb",
graph = "testdata/parallel_tracker_graph.pbtxt",
output_name = "testdata/parallel_tracker.binarypb",
visibility = ["//visibility:public"],
deps = MEDIAPIPE_DEPS,
)
mediapipe_binary_graph(
name = "tracker_binarypb",
graph = "testdata/tracker_graph.pbtxt",
output_name = "testdata/tracker.binarypb",
visibility = ["//visibility:public"],
deps = MEDIAPIPE_DEPS,
)
cc_test(
name = "tracking_graph_test",
size = "small",
srcs = ["tracking_graph_test.cc"],
copts = ["-DPARALLEL_INVOKER_ACTIVE"] + select({
"//mediapipe:apple": [],
"//mediapipe:android": [],
"//conditions:default": [],
}),
data = [
":testdata/lenna.png",
":testdata/parallel_tracker.binarypb",
":testdata/tracker.binarypb",
],
deps = [
":box_tracker_calculator",
":box_tracker_calculator_cc_proto",
":flow_packager_calculator",
":motion_analysis_calculator",
"//mediapipe/framework:calculator_cc_proto",
"//mediapipe/framework:calculator_framework",
"//mediapipe/framework:packet",
"//mediapipe/framework/deps:file_path",
"//mediapipe/framework/formats:image_frame",
"//mediapipe/framework/port:advanced_proto",
"//mediapipe/framework/port:core_proto",
"//mediapipe/framework/port:file_helpers",
"//mediapipe/framework/port:gtest_main",
"//mediapipe/framework/port:opencv_highgui",
"//mediapipe/framework/port:status",
"//mediapipe/framework/stream_handler:fixed_size_input_stream_handler",
"//mediapipe/framework/stream_handler:sync_set_input_stream_handler",
"//mediapipe/util/tracking:box_tracker_cc_proto",
"//mediapipe/util/tracking:tracking_cc_proto",
],
)
cc_test(
name = "video_pre_stream_calculator_test",
srcs = ["video_pre_stream_calculator_test.cc"],
deps = [
":video_pre_stream_calculator",
"//mediapipe/framework:calculator_framework",
"//mediapipe/framework/formats:image_frame",
"//mediapipe/framework/formats:video_stream_header",
"//mediapipe/framework/port:gtest_main",
"//mediapipe/framework/port:parse_text_proto",
"//mediapipe/framework/port:status",
],
)
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,55 @@
// 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.
syntax = "proto2";
package mediapipe;
import "mediapipe/framework/calculator.proto";
import "mediapipe/util/tracking/box_tracker.proto";
message BoxTrackerCalculatorOptions {
extend CalculatorOptions {
optional BoxTrackerCalculatorOptions ext = 268767860;
}
optional BoxTrackerOptions tracker_options = 1;
// Initial position to be tracked. Can also be supplied as side packet or
// as input stream.
optional TimedBoxProtoList initial_position = 2;
// If set and VIZ stream is present, renders tracking data into the
// visualization.
optional bool visualize_tracking_data = 3 [default = false];
// If set and VIZ stream is present, renders the box state
// into the visualization.
optional bool visualize_state = 4 [default = false];
// If set and VIZ stream is present, renders the internal box state
// into the visualization.
optional bool visualize_internal_state = 5 [default = false];
// Size of the track data cache during streaming mode. This allows to buffer
// track_data's for fast forward tracking, i.e. any TimedBox received
// via input stream START_POS can be tracked towards the current track head
// (i.e. last received TrackingData). Measured in number of frames.
optional int32 streaming_track_data_cache_size = 6 [default = 0];
// Add a transition period of N frames to smooth the jump from original
// tracking to reset start pos with motion compensation. The transition will
// be a linear decay of original tracking result. 0 means no transition.
optional int32 start_pos_transition_frames = 7 [default = 0];
}
@@ -0,0 +1,281 @@
// 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.
#include <stdio.h>
#include <fstream>
#include <memory>
#include "absl/strings/str_format.h"
#include "absl/strings/string_view.h"
#include "mediapipe/calculators/video/flow_packager_calculator.pb.h"
#include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/port/integral_types.h"
#include "mediapipe/framework/port/logging.h"
#include "mediapipe/util/tracking/camera_motion.pb.h"
#include "mediapipe/util/tracking/flow_packager.h"
#include "mediapipe/util/tracking/region_flow.pb.h"
namespace mediapipe {
using mediapipe::CameraMotion;
using mediapipe::FlowPackager;
using mediapipe::RegionFlowFeatureList;
using mediapipe::TrackingData;
using mediapipe::TrackingDataChunk;
// A calculator that packages input CameraMotion and RegionFlowFeatureList
// into a TrackingData and optionally writes TrackingDataChunks to file.
//
// Input stream:
// FLOW: Input region flow (proto RegionFlowFeatureList).
// CAMERA: Input camera stream (proto CameraMotion, optional).
//
// Input side packets:
// CACHE_DIR: Optional caching directory tracking files are written to.
//
// Output streams.
// TRACKING: Output tracking data (proto TrackingData, per frame
// optional).
// TRACKING_CHUNK: Output tracking chunks (proto TrackingDataChunk,
// per chunk, optional), output at the first timestamp
// of each chunk.
// COMPLETE: Optional output packet sent on PreStream to
// to signal downstream calculators that all data has been
// processed and calculator is closed. Can be used to indicate
// that all data as been written to CACHE_DIR.
class FlowPackagerCalculator : public CalculatorBase {
public:
~FlowPackagerCalculator() override = default;
static ::mediapipe::Status GetContract(CalculatorContract* cc);
::mediapipe::Status Open(CalculatorContext* cc) override;
::mediapipe::Status Process(CalculatorContext* cc) override;
::mediapipe::Status Close(CalculatorContext* cc) override;
// Writes passed chunk to disk.
void WriteChunk(const TrackingDataChunk& chunk) const;
// Initializes next chunk for tracking beginning from last frame of
// current chunk (Chunking is design with one frame overlap).
void PrepareCurrentForNextChunk(TrackingDataChunk* chunk);
private:
FlowPackagerCalculatorOptions options_;
// Caching options.
bool use_caching_ = false;
bool build_chunk_ = false;
std::string cache_dir_;
int chunk_idx_ = -1;
TrackingDataChunk tracking_chunk_;
int frame_idx_ = 0;
Timestamp prev_timestamp_;
std::unique_ptr<FlowPackager> flow_packager_;
};
REGISTER_CALCULATOR(FlowPackagerCalculator);
::mediapipe::Status FlowPackagerCalculator::GetContract(
CalculatorContract* cc) {
if (!cc->Inputs().HasTag("FLOW")) {
return tool::StatusFail("No input flow was specified.");
}
cc->Inputs().Tag("FLOW").Set<RegionFlowFeatureList>();
if (cc->Inputs().HasTag("CAMERA")) {
cc->Inputs().Tag("CAMERA").Set<CameraMotion>();
}
if (cc->Outputs().HasTag("TRACKING")) {
cc->Outputs().Tag("TRACKING").Set<TrackingData>();
}
if (cc->Outputs().HasTag("TRACKING_CHUNK")) {
cc->Outputs().Tag("TRACKING_CHUNK").Set<TrackingDataChunk>();
}
if (cc->Outputs().HasTag("COMPLETE")) {
cc->Outputs().Tag("COMPLETE").Set<bool>();
}
if (cc->InputSidePackets().HasTag("CACHE_DIR")) {
cc->InputSidePackets().Tag("CACHE_DIR").Set<std::string>();
}
return ::mediapipe::OkStatus();
}
::mediapipe::Status FlowPackagerCalculator::Open(CalculatorContext* cc) {
options_ = cc->Options<FlowPackagerCalculatorOptions>();
flow_packager_.reset(new FlowPackager(options_.flow_packager_options()));
use_caching_ = cc->InputSidePackets().HasTag("CACHE_DIR");
build_chunk_ = use_caching_ || cc->Outputs().HasTag("TRACKING_CHUNK");
if (use_caching_) {
cache_dir_ = cc->InputSidePackets().Tag("CACHE_DIR").Get<std::string>();
}
return ::mediapipe::OkStatus();
}
::mediapipe::Status FlowPackagerCalculator::Process(CalculatorContext* cc) {
InputStream* flow_stream = &(cc->Inputs().Tag("FLOW"));
const RegionFlowFeatureList& flow = flow_stream->Get<RegionFlowFeatureList>();
const Timestamp timestamp = flow_stream->Value().Timestamp();
const CameraMotion* camera_motion = nullptr;
if (cc->Inputs().HasTag("CAMERA")) {
InputStream* camera_stream = &(cc->Inputs().Tag("CAMERA"));
camera_motion = &camera_stream->Get<CameraMotion>();
}
std::unique_ptr<TrackingData> tracking_data(new TrackingData());
flow_packager_->PackFlow(flow, camera_motion, tracking_data.get());
if (build_chunk_) {
if (chunk_idx_ < 0) { // Lazy init, determine first start.
chunk_idx_ =
timestamp.Value() / 1000 / options_.caching_chunk_size_msec();
tracking_chunk_.set_first_chunk(true);
}
CHECK_GE(chunk_idx_, 0);
TrackingDataChunk::Item* item = tracking_chunk_.add_item();
item->set_frame_idx(frame_idx_);
item->set_timestamp_usec(timestamp.Value());
if (frame_idx_ > 0) {
item->set_prev_timestamp_usec(prev_timestamp_.Value());
}
if (cc->Outputs().HasTag("TRACKING")) {
// Need to copy as output is requested.
*item->mutable_tracking_data() = *tracking_data;
} else {
item->mutable_tracking_data()->Swap(tracking_data.get());
}
const int next_chunk_msec =
options_.caching_chunk_size_msec() * (chunk_idx_ + 1);
if (timestamp.Value() / 1000 >= next_chunk_msec) {
if (cc->Outputs().HasTag("TRACKING_CHUNK")) {
cc->Outputs()
.Tag("TRACKING_CHUNK")
.Add(new TrackingDataChunk(tracking_chunk_),
Timestamp(tracking_chunk_.item(0).timestamp_usec()));
}
if (use_caching_) {
WriteChunk(tracking_chunk_);
}
PrepareCurrentForNextChunk(&tracking_chunk_);
}
}
if (cc->Outputs().HasTag("TRACKING")) {
cc->Outputs()
.Tag("TRACKING")
.Add(tracking_data.release(), flow_stream->Value().Timestamp());
}
prev_timestamp_ = timestamp;
++frame_idx_;
return ::mediapipe::OkStatus();
}
::mediapipe::Status FlowPackagerCalculator::Close(CalculatorContext* cc) {
if (frame_idx_ > 0) {
tracking_chunk_.set_last_chunk(true);
if (cc->Outputs().HasTag("TRACKING_CHUNK")) {
cc->Outputs()
.Tag("TRACKING_CHUNK")
.Add(new TrackingDataChunk(tracking_chunk_),
Timestamp(tracking_chunk_.item(0).timestamp_usec()));
}
if (use_caching_) {
WriteChunk(tracking_chunk_);
}
}
if (cc->Outputs().HasTag("COMPLETE")) {
cc->Outputs().Tag("COMPLETE").Add(new bool(true), Timestamp::PreStream());
}
return ::mediapipe::OkStatus();
}
void FlowPackagerCalculator::WriteChunk(const TrackingDataChunk& chunk) const {
if (chunk.item_size() == 0) {
LOG(ERROR) << "Write chunk called with empty tracking data."
<< "This can only occur if the spacing between frames "
<< "is larger than the requested chunk size. Try increasing "
<< "the chunk size";
return;
}
auto format_runtime =
absl::ParsedFormat<'d'>::New(options_.cache_file_format());
std::string chunk_file;
if (format_runtime) {
chunk_file =
cache_dir_ + "/" + absl::StrFormat(*format_runtime, chunk_idx_);
} else {
LOG(ERROR) << "chache_file_format wrong. fall back to chunk_%04d.";
chunk_file = cache_dir_ + "/" + absl::StrFormat("chunk_%04d", chunk_idx_);
}
std::string data;
chunk.SerializeToString(&data);
const char* temp_filename = tempnam(cache_dir_.c_str(), nullptr);
std::ofstream out_file(temp_filename);
if (!out_file) {
LOG(ERROR) << "Could not open " << temp_filename;
} else {
out_file.write(data.data(), data.size());
}
if (rename(temp_filename, chunk_file.c_str()) != 0) {
LOG(ERROR) << "Failed to rename to " << chunk_file;
}
LOG(INFO) << "Wrote chunk : " << chunk_file;
}
void FlowPackagerCalculator::PrepareCurrentForNextChunk(
TrackingDataChunk* chunk) {
CHECK(chunk);
if (chunk->item_size() == 0) {
LOG(ERROR) << "Called with empty chunk. Unexpected.";
return;
}
chunk->set_first_chunk(false);
// Buffer last item for next chunk.
TrackingDataChunk::Item last_item;
last_item.Swap(chunk->mutable_item(chunk->item_size() - 1));
chunk->Clear();
chunk->add_item()->Swap(&last_item);
++chunk_idx_;
}
} // namespace mediapipe
@@ -0,0 +1,36 @@
// 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.
syntax = "proto2";
package mediapipe;
import "mediapipe/framework/calculator.proto";
import "mediapipe/util/tracking/flow_packager.proto";
message FlowPackagerCalculatorOptions {
extend CalculatorOptions {
optional FlowPackagerCalculatorOptions ext = 271236147;
}
optional mediapipe.FlowPackagerOptions flow_packager_options = 1;
// Chunk size for caching files that are written to the externally specified
// caching directory. Specified in msec.
// Note that each chunk always contains at its end the first frame of the
// next chunk (to enable forward tracking across chunk boundaries).
optional int32 caching_chunk_size_msec = 2 [default = 2500];
optional string cache_file_format = 3 [default = "chunk_%04d"];
}
@@ -0,0 +1,988 @@
// 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.
#include <cmath>
#include <fstream>
#include <memory>
#include "absl/strings/numbers.h"
#include "absl/strings/str_split.h"
#include "absl/strings/string_view.h"
#include "mediapipe/calculators/video/motion_analysis_calculator.pb.h"
#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/port/integral_types.h"
#include "mediapipe/framework/port/logging.h"
#include "mediapipe/framework/port/ret_check.h"
#include "mediapipe/framework/port/status.h"
#include "mediapipe/util/tracking/camera_motion.h"
#include "mediapipe/util/tracking/camera_motion.pb.h"
#include "mediapipe/util/tracking/frame_selection.pb.h"
#include "mediapipe/util/tracking/motion_analysis.h"
#include "mediapipe/util/tracking/motion_estimation.h"
#include "mediapipe/util/tracking/motion_models.h"
#include "mediapipe/util/tracking/region_flow.pb.h"
namespace mediapipe {
using mediapipe::AffineAdapter;
using mediapipe::CameraMotion;
using mediapipe::FrameSelectionResult;
using mediapipe::Homography;
using mediapipe::HomographyAdapter;
using mediapipe::LinearSimilarityModel;
using mediapipe::MixtureHomography;
using mediapipe::MixtureRowWeights;
using mediapipe::MotionAnalysis;
using mediapipe::ProjectViaFit;
using mediapipe::RegionFlowComputationOptions;
using mediapipe::RegionFlowFeatureList;
using mediapipe::SalientPointFrame;
using mediapipe::TranslationModel;
const char kOptionsTag[] = "OPTIONS";
// A calculator that performs motion analysis on an incoming video stream.
//
// Input streams: (at least one of them is required).
// VIDEO: The input video stream (ImageFrame, sRGB, sRGBA or GRAY8).
// SELECTION: Optional input stream to perform analysis only on selected
// frames. If present needs to contain camera motion
// and features.
//
// Input side packets:
// CSV_FILE: Read motion models as homographies from CSV file. Expected
// to be defined in the frame domain (un-normalized).
// Should store 9 floats per row.
// Specify number of homographies per frames via option
// meta_models_per_frame. For values > 1, MixtureHomographies
// are created, for value == 1, a single Homography is used.
// DOWNSAMPLE: Optionally specify downsampling factor via input side packet
// overriding value in the graph settings.
// Output streams (all are optional).
// FLOW: Sparse feature tracks in form of proto RegionFlowFeatureList.
// CAMERA: Camera motion as proto CameraMotion describing the per frame-
// pair motion. Has VideoHeader from input video.
// SALIENCY: Foreground saliency (objects moving different from the
// background) as proto SalientPointFrame.
// VIZ: Visualization stream as ImageFrame, sRGB, visualizing
// features and saliency (set via
// analysis_options().visualization_options())
// DENSE_FG: Dense foreground stream, describing per-pixel foreground-
// ness as confidence between 0 (background) and 255
// (foreground). Output is ImageFrame (GRAY8).
// VIDEO_OUT: Optional output stream when SELECTION is used. Output is input
// VIDEO at the selected frames. Required VIDEO to be present.
// GRAY_VIDEO_OUT: Optional output stream for downsampled, grayscale video.
// Requires VIDEO to be present and SELECTION to not be used.
class MotionAnalysisCalculator : public CalculatorBase {
// TODO: Activate once leakr approval is ready.
// typedef com::google::android::libraries::micro::proto::Data HomographyData;
public:
~MotionAnalysisCalculator() override = default;
static ::mediapipe::Status GetContract(CalculatorContract* cc);
::mediapipe::Status Open(CalculatorContext* cc) override;
::mediapipe::Status Process(CalculatorContext* cc) override;
::mediapipe::Status Close(CalculatorContext* cc) override;
private:
// Outputs results to Outputs() if MotionAnalysis buffered sufficient results.
// Otherwise no-op. Set flush to true to force output of all buffered data.
void OutputMotionAnalyzedFrames(bool flush, CalculatorContext* cc);
// Lazy init function to be called on Process.
::mediapipe::Status InitOnProcess(InputStream* video_stream,
InputStream* selection_stream);
// Parses CSV file contents to homographies.
bool ParseModelCSV(const std::string& contents,
std::deque<Homography>* homographies);
// Turns list of 9-tuple floating values into set of homographies.
bool HomographiesFromValues(const std::vector<float>& homog_values,
std::deque<Homography>* homographies);
// Appends CameraMotions and features from homographies.
// Set append_identity to true to add an identity transform to the beginning
// of the each list *in addition* to the motions derived from homographies.
void AppendCameraMotionsFromHomographies(
const std::deque<Homography>& homographies, bool append_identity,
std::deque<CameraMotion>* camera_motions,
std::deque<RegionFlowFeatureList>* features);
// Helper function to subtract current metadata motion from features. Used
// for hybrid estimation case.
void SubtractMetaMotion(const CameraMotion& meta_motion,
RegionFlowFeatureList* features);
// Inverse of above function to add back meta motion and replace
// feature location with originals after estimation.
void AddMetaMotion(const CameraMotion& meta_motion,
const RegionFlowFeatureList& meta_features,
RegionFlowFeatureList* features, CameraMotion* motion);
MotionAnalysisCalculatorOptions options_;
int frame_width_ = -1;
int frame_height_ = -1;
int frame_idx_ = 0;
// Buffers incoming video frame packets (if visualization output is requested)
std::vector<Packet> packet_buffer_;
// Buffers incoming timestamps until MotionAnalysis is ready to output via
// above OutputMotionAnalyzedFrames.
std::vector<Timestamp> timestamp_buffer_;
// Input indicators for each stream.
bool selection_input_ = false;
bool video_input_ = false;
// Output indicators for each stream.
bool region_flow_feature_output_ = false;
bool camera_motion_output_ = false;
bool saliency_output_ = false;
bool visualize_output_ = false;
bool dense_foreground_output_ = false;
bool video_output_ = false;
bool grayscale_output_ = false;
bool csv_file_input_ = false;
// Inidicates if saliency should be computed.
bool with_saliency_ = false;
// Set if hybrid meta analysis - see proto for details.
bool hybrid_meta_analysis_ = false;
// Concatenated motions for each selected frame. Used in case
// hybrid estimation is requested to fallback to valid models.
std::deque<CameraMotion> selected_motions_;
// Normalized homographies from CSV file or metadata.
std::deque<Homography> meta_homographies_;
std::deque<CameraMotion> meta_motions_;
std::deque<RegionFlowFeatureList> meta_features_;
// Offset into above meta_motions_ and features_ when using
// hybrid meta analysis.
int hybrid_meta_offset_ = 0;
std::unique_ptr<MotionAnalysis> motion_analysis_;
std::unique_ptr<MixtureRowWeights> row_weights_;
};
REGISTER_CALCULATOR(MotionAnalysisCalculator);
::mediapipe::Status MotionAnalysisCalculator::GetContract(
CalculatorContract* cc) {
if (cc->Inputs().HasTag("VIDEO")) {
cc->Inputs().Tag("VIDEO").Set<ImageFrame>();
}
// Optional input stream from frame selection calculator.
if (cc->Inputs().HasTag("SELECTION")) {
cc->Inputs().Tag("SELECTION").Set<FrameSelectionResult>();
}
RET_CHECK(cc->Inputs().HasTag("VIDEO") || cc->Inputs().HasTag("SELECTION"))
<< "Either VIDEO, SELECTION must be specified.";
if (cc->Outputs().HasTag("FLOW")) {
cc->Outputs().Tag("FLOW").Set<RegionFlowFeatureList>();
}
if (cc->Outputs().HasTag("CAMERA")) {
cc->Outputs().Tag("CAMERA").Set<CameraMotion>();
}
if (cc->Outputs().HasTag("SALIENCY")) {
cc->Outputs().Tag("SALIENCY").Set<SalientPointFrame>();
}
if (cc->Outputs().HasTag("VIZ")) {
cc->Outputs().Tag("VIZ").Set<ImageFrame>();
}
if (cc->Outputs().HasTag("DENSE_FG")) {
cc->Outputs().Tag("DENSE_FG").Set<ImageFrame>();
}
if (cc->Outputs().HasTag("VIDEO_OUT")) {
cc->Outputs().Tag("VIDEO_OUT").Set<ImageFrame>();
}
if (cc->Outputs().HasTag("GRAY_VIDEO_OUT")) {
// We only output grayscale video if we're actually performing full region-
// flow analysis on the video.
RET_CHECK(cc->Inputs().HasTag("VIDEO") &&
!cc->Inputs().HasTag("SELECTION"));
cc->Outputs().Tag("GRAY_VIDEO_OUT").Set<ImageFrame>();
}
if (cc->InputSidePackets().HasTag("CSV_FILE")) {
cc->InputSidePackets().Tag("CSV_FILE").Set<std::string>();
}
if (cc->InputSidePackets().HasTag("DOWNSAMPLE")) {
cc->InputSidePackets().Tag("DOWNSAMPLE").Set<float>();
}
if (cc->InputSidePackets().HasTag(kOptionsTag)) {
cc->InputSidePackets().Tag(kOptionsTag).Set<CalculatorOptions>();
}
return ::mediapipe::OkStatus();
}
::mediapipe::Status MotionAnalysisCalculator::Open(CalculatorContext* cc) {
options_ =
tool::RetrieveOptions(cc->Options<MotionAnalysisCalculatorOptions>(),
cc->InputSidePackets(), kOptionsTag);
video_input_ = cc->Inputs().HasTag("VIDEO");
selection_input_ = cc->Inputs().HasTag("SELECTION");
region_flow_feature_output_ = cc->Outputs().HasTag("FLOW");
camera_motion_output_ = cc->Outputs().HasTag("CAMERA");
saliency_output_ = cc->Outputs().HasTag("SALIENCY");
visualize_output_ = cc->Outputs().HasTag("VIZ");
dense_foreground_output_ = cc->Outputs().HasTag("DENSE_FG");
video_output_ = cc->Outputs().HasTag("VIDEO_OUT");
grayscale_output_ = cc->Outputs().HasTag("GRAY_VIDEO_OUT");
csv_file_input_ = cc->InputSidePackets().HasTag("CSV_FILE");
hybrid_meta_analysis_ = options_.meta_analysis() ==
MotionAnalysisCalculatorOptions::META_ANALYSIS_HYBRID;
if (video_output_) {
RET_CHECK(selection_input_) << "VIDEO_OUT requires SELECTION input";
}
if (selection_input_) {
switch (options_.selection_analysis()) {
case MotionAnalysisCalculatorOptions::NO_ANALYSIS_USE_SELECTION:
RET_CHECK(!visualize_output_)
<< "Visualization not supported for NO_ANALYSIS_USE_SELECTION";
RET_CHECK(!dense_foreground_output_)
<< "Dense foreground not supported for NO_ANALYSIS_USE_SELECTION";
RET_CHECK(!saliency_output_)
<< "Saliency output not supported for NO_ANALYSIS_USE_SELECTION";
break;
case MotionAnalysisCalculatorOptions::ANALYSIS_RECOMPUTE:
case MotionAnalysisCalculatorOptions::ANALYSIS_WITH_SEED:
RET_CHECK(video_input_) << "Need video input for feature tracking.";
break;
case MotionAnalysisCalculatorOptions::ANALYSIS_FROM_FEATURES:
// Nothing to add here.
break;
}
}
if (visualize_output_ || dense_foreground_output_ || video_output_) {
RET_CHECK(video_input_) << "Video input required.";
}
if (csv_file_input_) {
RET_CHECK(!selection_input_)
<< "Can not use selection input with csv input.";
if (!hybrid_meta_analysis_) {
RET_CHECK(!saliency_output_ && !visualize_output_ &&
!dense_foreground_output_ && !grayscale_output_)
<< "CSV file and meta input only supports flow and camera motion "
<< "output when using metadata only.";
}
}
if (csv_file_input_) {
// Read from file and parse.
const std::string filename =
cc->InputSidePackets().Tag("CSV_FILE").Get<std::string>();
std::string file_contents;
std::ifstream input_file(filename, std::ios::in);
input_file.seekg(0, std::ios::end);
const int file_length = input_file.tellg();
file_contents.resize(file_length);
input_file.seekg(0, std::ios::beg);
input_file.read(&file_contents[0], file_length);
input_file.close();
RET_CHECK(ParseModelCSV(file_contents, &meta_homographies_))
<< "Could not parse CSV file";
}
// Get video header from video or selection input if present.
const VideoHeader* video_header = nullptr;
if (video_input_ && !cc->Inputs().Tag("VIDEO").Header().IsEmpty()) {
video_header = &(cc->Inputs().Tag("VIDEO").Header().Get<VideoHeader>());
} else if (selection_input_ &&
!cc->Inputs().Tag("SELECTION").Header().IsEmpty()) {
video_header = &(cc->Inputs().Tag("SELECTION").Header().Get<VideoHeader>());
} else {
LOG(WARNING) << "No input video header found. Downstream calculators "
"expecting video headers are likely to fail.";
}
with_saliency_ = options_.analysis_options().compute_motion_saliency();
// Force computation of saliency if requested as output.
if (cc->Outputs().HasTag("SALIENCY")) {
with_saliency_ = true;
if (!options_.analysis_options().compute_motion_saliency()) {
LOG(WARNING) << "Enable saliency computation. Set "
<< "compute_motion_saliency to true to silence this "
<< "warning.";
options_.mutable_analysis_options()->set_compute_motion_saliency(true);
}
}
if (options_.bypass_mode()) {
cc->SetOffset(TimestampDiff(0));
}
if (cc->InputSidePackets().HasTag("DOWNSAMPLE")) {
options_.mutable_analysis_options()
->mutable_flow_options()
->set_downsample_factor(
cc->InputSidePackets().Tag("DOWNSAMPLE").Get<float>());
}
// If no video header is provided, just return and initialize on the first
// Process() call.
if (video_header == nullptr) {
return ::mediapipe::OkStatus();
}
////////////// EARLY RETURN; ONLY HEADER OUTPUT SHOULD GO HERE ///////////////
if (visualize_output_) {
cc->Outputs().Tag("VIZ").SetHeader(Adopt(new VideoHeader(*video_header)));
}
if (video_output_) {
cc->Outputs()
.Tag("VIDEO_OUT")
.SetHeader(Adopt(new VideoHeader(*video_header)));
}
if (cc->Outputs().HasTag("DENSE_FG")) {
std::unique_ptr<VideoHeader> foreground_header(
new VideoHeader(*video_header));
foreground_header->format = ImageFormat::GRAY8;
cc->Outputs().Tag("DENSE_FG").SetHeader(Adopt(foreground_header.release()));
}
if (cc->Outputs().HasTag("CAMERA")) {
cc->Outputs().Tag("CAMERA").SetHeader(
Adopt(new VideoHeader(*video_header)));
}
if (cc->Outputs().HasTag("SALIENCY")) {
cc->Outputs()
.Tag("SALIENCY")
.SetHeader(Adopt(new VideoHeader(*video_header)));
}
return ::mediapipe::OkStatus();
}
::mediapipe::Status MotionAnalysisCalculator::Process(CalculatorContext* cc) {
if (options_.bypass_mode()) {
return ::mediapipe::OkStatus();
}
InputStream* video_stream =
video_input_ ? &(cc->Inputs().Tag("VIDEO")) : nullptr;
InputStream* selection_stream =
selection_input_ ? &(cc->Inputs().Tag("SELECTION")) : nullptr;
// Checked on Open.
CHECK(video_stream || selection_stream);
// Lazy init.
if (frame_width_ < 0 || frame_height_ < 0) {
MP_RETURN_IF_ERROR(InitOnProcess(video_stream, selection_stream));
}
const Timestamp timestamp = cc->InputTimestamp();
if ((csv_file_input_) && !hybrid_meta_analysis_) {
if (camera_motion_output_) {
RET_CHECK(!meta_motions_.empty()) << "Insufficient metadata.";
CameraMotion output_motion = meta_motions_.front();
meta_motions_.pop_front();
output_motion.set_timestamp_usec(timestamp.Value());
cc->Outputs().Tag("CAMERA").Add(new CameraMotion(output_motion),
timestamp);
}
if (region_flow_feature_output_) {
RET_CHECK(!meta_features_.empty()) << "Insufficient frames in CSV file";
RegionFlowFeatureList output_features = meta_features_.front();
meta_features_.pop_front();
output_features.set_timestamp_usec(timestamp.Value());
cc->Outputs().Tag("FLOW").Add(new RegionFlowFeatureList(output_features),
timestamp);
}
++frame_idx_;
return ::mediapipe::OkStatus();
}
if (motion_analysis_ == nullptr) {
// We do not need MotionAnalysis when using just metadata.
motion_analysis_.reset(new MotionAnalysis(options_.analysis_options(),
frame_width_, frame_height_));
}
std::unique_ptr<FrameSelectionResult> frame_selection_result;
// Always use frame if selection is not activated.
bool use_frame = !selection_input_;
if (selection_input_) {
CHECK(selection_stream);
// Fill in timestamps we process.
if (!selection_stream->Value().IsEmpty()) {
ASSIGN_OR_RETURN(
frame_selection_result,
selection_stream->Value().ConsumeOrCopy<FrameSelectionResult>());
use_frame = true;
// Make sure both features and camera motion are present.
RET_CHECK(frame_selection_result->has_camera_motion() &&
frame_selection_result->has_features())
<< "Frame selection input error at: " << timestamp
<< " both camera motion and features need to be "
"present in FrameSelectionResult. "
<< frame_selection_result->has_camera_motion() << " , "
<< frame_selection_result->has_features();
}
}
if (selection_input_ && use_frame &&
options_.selection_analysis() ==
MotionAnalysisCalculatorOptions::NO_ANALYSIS_USE_SELECTION) {
// Output concatenated results, nothing to compute here.
if (camera_motion_output_) {
cc->Outputs().Tag("CAMERA").Add(
frame_selection_result->release_camera_motion(), timestamp);
}
if (region_flow_feature_output_) {
cc->Outputs().Tag("FLOW").Add(frame_selection_result->release_features(),
timestamp);
}
if (video_output_) {
cc->Outputs().Tag("VIDEO_OUT").AddPacket(video_stream->Value());
}
return ::mediapipe::OkStatus();
}
if (use_frame) {
if (!selection_input_) {
const cv::Mat input_view =
formats::MatView(&video_stream->Get<ImageFrame>());
if (hybrid_meta_analysis_) {
// Seed with meta homography.
RET_CHECK(hybrid_meta_offset_ < meta_motions_.size())
<< "Not enough metadata received for hybrid meta analysis";
Homography initial_transform =
meta_motions_[hybrid_meta_offset_].homography();
std::function<void(RegionFlowFeatureList*)> subtract_helper = std::bind(
&MotionAnalysisCalculator::SubtractMetaMotion, this,
meta_motions_[hybrid_meta_offset_], std::placeholders::_1);
// Keep original features before modification around.
motion_analysis_->AddFrameGeneric(
input_view, timestamp.Value(), initial_transform, nullptr, nullptr,
&subtract_helper, &meta_features_[hybrid_meta_offset_]);
++hybrid_meta_offset_;
} else {
motion_analysis_->AddFrame(input_view, timestamp.Value());
}
} else {
selected_motions_.push_back(frame_selection_result->camera_motion());
switch (options_.selection_analysis()) {
case MotionAnalysisCalculatorOptions::NO_ANALYSIS_USE_SELECTION:
return ::mediapipe::UnknownErrorBuilder(MEDIAPIPE_LOC)
<< "Should not reach this point!";
case MotionAnalysisCalculatorOptions::ANALYSIS_FROM_FEATURES:
motion_analysis_->AddFeatures(frame_selection_result->features());
break;
case MotionAnalysisCalculatorOptions::ANALYSIS_RECOMPUTE: {
const cv::Mat input_view =
formats::MatView(&video_stream->Get<ImageFrame>());
motion_analysis_->AddFrame(input_view, timestamp.Value());
break;
}
case MotionAnalysisCalculatorOptions::ANALYSIS_WITH_SEED: {
Homography homography;
CameraMotionToHomography(frame_selection_result->camera_motion(),
&homography);
const cv::Mat input_view =
formats::MatView(&video_stream->Get<ImageFrame>());
motion_analysis_->AddFrameGeneric(input_view, timestamp.Value(),
homography, &homography);
break;
}
}
}
timestamp_buffer_.push_back(timestamp);
++frame_idx_;
VLOG_EVERY_N(0, 100) << "Analyzed frame " << frame_idx_;
// Buffer input frames only if visualization is requested.
if (visualize_output_ || video_output_) {
packet_buffer_.push_back(video_stream->Value());
}
// If requested, output grayscale thumbnails
if (grayscale_output_) {
cv::Mat grayscale_mat = motion_analysis_->GetGrayscaleFrameFromResults();
std::unique_ptr<ImageFrame> grayscale_image(new ImageFrame(
ImageFormat::GRAY8, grayscale_mat.cols, grayscale_mat.rows));
cv::Mat image_frame_mat = formats::MatView(grayscale_image.get());
grayscale_mat.copyTo(image_frame_mat);
cc->Outputs()
.Tag("GRAY_VIDEO_OUT")
.Add(grayscale_image.release(), timestamp);
}
// Output other results, if we have any yet.
OutputMotionAnalyzedFrames(false, cc);
}
return ::mediapipe::OkStatus();
}
::mediapipe::Status MotionAnalysisCalculator::Close(CalculatorContext* cc) {
// Guard against empty videos.
if (motion_analysis_) {
OutputMotionAnalyzedFrames(true, cc);
}
if (csv_file_input_) {
if (!meta_motions_.empty()) {
LOG(ERROR) << "More motions than frames. Unexpected! Remainder: "
<< meta_motions_.size();
}
}
return ::mediapipe::OkStatus();
}
void MotionAnalysisCalculator::OutputMotionAnalyzedFrames(
bool flush, CalculatorContext* cc) {
std::vector<std::unique_ptr<RegionFlowFeatureList>> features;
std::vector<std::unique_ptr<CameraMotion>> camera_motions;
std::vector<std::unique_ptr<SalientPointFrame>> saliency;
const int buffer_size = timestamp_buffer_.size();
const int num_results = motion_analysis_->GetResults(
flush, &features, &camera_motions, with_saliency_ ? &saliency : nullptr);
CHECK_LE(num_results, buffer_size);
if (num_results == 0) {
return;
}
for (int k = 0; k < num_results; ++k) {
// Region flow features and camera motion for this frame.
auto& feature_list = features[k];
auto& camera_motion = camera_motions[k];
const Timestamp timestamp = timestamp_buffer_[k];
if (selection_input_ && options_.hybrid_selection_camera()) {
if (camera_motion->type() > selected_motions_.front().type()) {
// Composited type is more stable.
camera_motion->Swap(&selected_motions_.front());
}
selected_motions_.pop_front();
}
if (hybrid_meta_analysis_) {
AddMetaMotion(meta_motions_.front(), meta_features_.front(),
feature_list.get(), camera_motion.get());
meta_motions_.pop_front();
meta_features_.pop_front();
}
// Video frame for visualization.
std::unique_ptr<ImageFrame> visualization_frame;
cv::Mat visualization;
if (visualize_output_) {
// Initialize visualization frame with original frame.
visualization_frame.reset(new ImageFrame());
visualization_frame->CopyFrom(packet_buffer_[k].Get<ImageFrame>(), 16);
visualization = formats::MatView(visualization_frame.get());
motion_analysis_->RenderResults(
*feature_list, *camera_motion,
with_saliency_ ? saliency[k].get() : nullptr, &visualization);
cc->Outputs().Tag("VIZ").Add(visualization_frame.release(), timestamp);
}
// Output dense foreground mask.
if (dense_foreground_output_) {
std::unique_ptr<ImageFrame> foreground_frame(
new ImageFrame(ImageFormat::GRAY8, frame_width_, frame_height_));
cv::Mat foreground = formats::MatView(foreground_frame.get());
motion_analysis_->ComputeDenseForeground(*feature_list, *camera_motion,
&foreground);
cc->Outputs().Tag("DENSE_FG").Add(foreground_frame.release(), timestamp);
}
// Output flow features if requested.
if (region_flow_feature_output_) {
cc->Outputs().Tag("FLOW").Add(feature_list.release(), timestamp);
}
// Output camera motion.
if (camera_motion_output_) {
cc->Outputs().Tag("CAMERA").Add(camera_motion.release(), timestamp);
}
if (video_output_) {
cc->Outputs().Tag("VIDEO_OUT").AddPacket(packet_buffer_[k]);
}
// Output saliency.
if (saliency_output_) {
cc->Outputs().Tag("SALIENCY").Add(saliency[k].release(), timestamp);
}
}
if (hybrid_meta_analysis_) {
hybrid_meta_offset_ -= num_results;
CHECK_GE(hybrid_meta_offset_, 0);
}
timestamp_buffer_.erase(timestamp_buffer_.begin(),
timestamp_buffer_.begin() + num_results);
if (visualize_output_ || video_output_) {
packet_buffer_.erase(packet_buffer_.begin(),
packet_buffer_.begin() + num_results);
}
}
::mediapipe::Status MotionAnalysisCalculator::InitOnProcess(
InputStream* video_stream, InputStream* selection_stream) {
if (video_stream) {
frame_width_ = video_stream->Get<ImageFrame>().Width();
frame_height_ = video_stream->Get<ImageFrame>().Height();
// Ensure image options are set correctly.
auto* region_options =
options_.mutable_analysis_options()->mutable_flow_options();
// Use two possible formats to account for different channel orders.
RegionFlowComputationOptions::ImageFormat image_format;
RegionFlowComputationOptions::ImageFormat image_format2;
switch (video_stream->Get<ImageFrame>().Format()) {
case ImageFormat::GRAY8:
image_format = image_format2 =
RegionFlowComputationOptions::FORMAT_GRAYSCALE;
break;
case ImageFormat::SRGB:
image_format = RegionFlowComputationOptions::FORMAT_RGB;
image_format2 = RegionFlowComputationOptions::FORMAT_BGR;
break;
case ImageFormat::SRGBA:
image_format = RegionFlowComputationOptions::FORMAT_RGBA;
image_format2 = RegionFlowComputationOptions::FORMAT_BGRA;
break;
default:
RET_CHECK(false) << "Unsupported image format.";
}
if (region_options->image_format() != image_format &&
region_options->image_format() != image_format2) {
LOG(WARNING) << "Requested image format in RegionFlowComputation "
<< "does not match video stream format. Overriding.";
region_options->set_image_format(image_format);
}
// Account for downsampling mode INPUT_SIZE. In this case we are handed
// already downsampled frames but the resulting CameraMotion should
// be computed on higher resolution as specifed by the downsample scale.
if (region_options->downsample_mode() ==
RegionFlowComputationOptions::DOWNSAMPLE_TO_INPUT_SIZE) {
const float scale = region_options->downsample_factor();
frame_width_ = static_cast<int>(std::round(frame_width_ * scale));
frame_height_ = static_cast<int>(std::round(frame_height_ * scale));
}
} else if (selection_stream) {
const auto& camera_motion =
selection_stream->Get<FrameSelectionResult>().camera_motion();
frame_width_ = camera_motion.frame_width();
frame_height_ = camera_motion.frame_height();
} else {
LOG(FATAL) << "Either VIDEO or SELECTION stream need to be specified.";
}
// Filled by CSV file parsing.
if (!meta_homographies_.empty()) {
CHECK(csv_file_input_);
AppendCameraMotionsFromHomographies(meta_homographies_,
true, // append identity.
&meta_motions_, &meta_features_);
meta_homographies_.clear();
}
// Filter weights before using for hybrid mode.
if (hybrid_meta_analysis_) {
auto* motion_options =
options_.mutable_analysis_options()->mutable_motion_options();
motion_options->set_filter_initialized_irls_weights(true);
}
return ::mediapipe::OkStatus();
}
bool MotionAnalysisCalculator::ParseModelCSV(
const std::string& contents, std::deque<Homography>* homographies) {
std::vector<absl::string_view> values =
absl::StrSplit(contents, absl::ByAnyChar(",\n"));
// Trim off any empty lines.
while (values.back().empty()) {
values.pop_back();
}
// Convert to float.
std::vector<float> homog_values;
homog_values.reserve(values.size());
for (const auto& value : values) {
double value_64f;
if (!absl::SimpleAtod(value, &value_64f)) {
LOG(ERROR) << "Not a double, expected!";
return false;
}
homog_values.push_back(value_64f);
}
return HomographiesFromValues(homog_values, homographies);
}
bool MotionAnalysisCalculator::HomographiesFromValues(
const std::vector<float>& homog_values,
std::deque<Homography>* homographies) {
CHECK(homographies);
// Obvious constants are obvious :D
constexpr int kHomographyValues = 9;
if (homog_values.size() % kHomographyValues != 0) {
LOG(ERROR) << "Contents not a multiple of " << kHomographyValues;
return false;
}
for (int k = 0; k < homog_values.size(); k += kHomographyValues) {
std::vector<double> h_vals(kHomographyValues);
for (int l = 0; l < kHomographyValues; ++l) {
h_vals[l] = homog_values[k + l];
}
// Normalize last entry to 1.
if (h_vals[kHomographyValues - 1] == 0) {
LOG(ERROR) << "Degenerate homography, last entry is zero";
return false;
}
const double scale = 1.0f / h_vals[kHomographyValues - 1];
for (int l = 0; l < kHomographyValues; ++l) {
h_vals[l] *= scale;
}
Homography h = HomographyAdapter::FromDoublePointer(h_vals.data(), false);
homographies->push_back(h);
}
if (homographies->size() % options_.meta_models_per_frame() != 0) {
LOG(ERROR) << "Total homographies not a multiple of specified models "
<< "per frame.";
return false;
}
return true;
}
void MotionAnalysisCalculator::SubtractMetaMotion(
const CameraMotion& meta_motion, RegionFlowFeatureList* features) {
if (meta_motion.mixture_homography().model_size() > 0) {
CHECK(row_weights_ != nullptr);
RegionFlowFeatureListViaTransform(meta_motion.mixture_homography(),
features, -1.0f,
1.0f, // subtract transformed.
true, // replace feature loc.
row_weights_.get());
} else {
RegionFlowFeatureListViaTransform(meta_motion.homography(), features, -1.0f,
1.0f, // subtract transformed.
true); // replace feature loc.
}
// Clamp transformed features to domain and handle outliers.
const float domain_diam =
hypot(features->frame_width(), features->frame_height());
const float motion_mag = meta_motion.average_magnitude();
// Same irls fraction as used by MODEL_MIXTURE_HOMOGRAPHY scaling in
// MotionEstimation.
const float irls_fraction = options_.analysis_options()
.motion_options()
.irls_mixture_fraction_scale() *
options_.analysis_options()
.motion_options()
.irls_motion_magnitude_fraction();
float err_scale = std::max(1.0f, motion_mag * irls_fraction);
const float max_err =
options_.meta_outlier_domain_ratio() * domain_diam * err_scale;
const float max_err_sq = max_err * max_err;
for (auto& feature : *features->mutable_feature()) {
feature.set_x(
std::max(0.0f, std::min(features->frame_width() - 1.0f, feature.x())));
feature.set_y(
std::max(0.0f, std::min(features->frame_height() - 1.0f, feature.y())));
// Label anything with large residual motion an outlier.
if (FeatureFlow(feature).Norm2() > max_err_sq) {
feature.set_irls_weight(0.0f);
}
}
}
void MotionAnalysisCalculator::AddMetaMotion(
const CameraMotion& meta_motion, const RegionFlowFeatureList& meta_features,
RegionFlowFeatureList* features, CameraMotion* motion) {
// Restore old feature location.
CHECK_EQ(meta_features.feature_size(), features->feature_size());
for (int k = 0; k < meta_features.feature_size(); ++k) {
auto feature = features->mutable_feature(k);
const auto& meta_feature = meta_features.feature(k);
feature->set_x(meta_feature.x());
feature->set_y(meta_feature.y());
feature->set_dx(meta_feature.dx());
feature->set_dy(meta_feature.dy());
}
// Composite camera motion.
*motion = ComposeCameraMotion(*motion, meta_motion);
// Restore type from metadata, i.e. do not declare motions as invalid.
motion->set_type(meta_motion.type());
motion->set_match_frame(-1);
}
void MotionAnalysisCalculator::AppendCameraMotionsFromHomographies(
const std::deque<Homography>& homographies, bool append_identity,
std::deque<CameraMotion>* camera_motions,
std::deque<RegionFlowFeatureList>* features) {
CHECK(camera_motions);
CHECK(features);
CameraMotion identity;
identity.set_frame_width(frame_width_);
identity.set_frame_height(frame_height_);
*identity.mutable_translation() = TranslationModel();
*identity.mutable_linear_similarity() = LinearSimilarityModel();
*identity.mutable_homography() = Homography();
identity.set_type(CameraMotion::VALID);
identity.set_match_frame(0);
RegionFlowFeatureList empty_list;
empty_list.set_long_tracks(true);
empty_list.set_match_frame(-1);
empty_list.set_frame_width(frame_width_);
empty_list.set_frame_height(frame_height_);
if (append_identity) {
camera_motions->push_back(identity);
features->push_back(empty_list);
}
const int models_per_frame = options_.meta_models_per_frame();
CHECK_GT(models_per_frame, 0) << "At least one model per frame is needed";
CHECK_EQ(0, homographies.size() % models_per_frame);
const int num_frames = homographies.size() / models_per_frame;
// Heuristic sigma, similar to what we use for rolling shutter removal.
const float mixture_sigma = 1.0f / models_per_frame;
if (row_weights_ == nullptr) {
row_weights_.reset(new MixtureRowWeights(frame_height_,
frame_height_ / 10, // 10% margin
mixture_sigma * frame_height_,
1.0f, models_per_frame));
}
for (int f = 0; f < num_frames; ++f) {
MixtureHomography mix_homog;
const int model_start = f * models_per_frame;
for (int k = 0; k < models_per_frame; ++k) {
const Homography& homog = homographies[model_start + k];
*mix_homog.add_model() = ModelInvert(homog);
}
CameraMotion c = identity;
c.set_match_frame(-1);
if (mix_homog.model_size() > 1) {
*c.mutable_mixture_homography() = mix_homog;
c.set_mixture_row_sigma(mixture_sigma);
for (int k = 0; k < models_per_frame; ++k) {
c.add_mixture_inlier_coverage(1.0f);
}
*c.add_mixture_homography_spectrum() = mix_homog;
c.set_rolling_shutter_motion_index(0);
*c.mutable_homography() = ProjectViaFit<Homography>(
mix_homog, frame_width_, frame_height_, row_weights_.get());
} else {
// Guaranteed to exist because to check that models_per_frame > 0 above.
*c.mutable_homography() = mix_homog.model(0);
}
// Project remaining motions down.
*c.mutable_linear_similarity() = ProjectViaFit<LinearSimilarityModel>(
c.homography(), frame_width_, frame_height_);
*c.mutable_translation() = ProjectViaFit<TranslationModel>(
c.homography(), frame_width_, frame_height_);
c.set_average_magnitude(
std::hypot(c.translation().dx(), c.translation().dy()));
camera_motions->push_back(c);
features->push_back(empty_list);
}
}
} // namespace mediapipe
@@ -0,0 +1,111 @@
// 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.
syntax = "proto2";
package mediapipe;
import "mediapipe/framework/calculator.proto";
import "mediapipe/util/tracking/motion_analysis.proto";
// Next tag: 10
message MotionAnalysisCalculatorOptions {
extend CalculatorOptions {
optional MotionAnalysisCalculatorOptions ext = 270698255;
}
optional mediapipe.MotionAnalysisOptions analysis_options = 1;
// Determines how optional input SELECTION (if present) is used to compute
// the final camera motion.
enum SelectionAnalysis {
// Recompute camera motion for selected frame neighbors.
ANALYSIS_RECOMPUTE = 1;
// Use composited camera motion and region flow from SELECTION input. No
// tracking or re-computation is performed.
// Note that in this case only CAMERA, FLOW and VIDEO_OUT tags are
// supported as output.
NO_ANALYSIS_USE_SELECTION = 2;
// Recompute camera motion for selected frame neighbors using
// features supplied by SELECTION input. No feature tracking is performed.
ANALYSIS_FROM_FEATURES = 3;
// Recomputes camera motion for selected frame neighbors but seeds
// initial transform with camera motion from SELECTION input.
ANALYSIS_WITH_SEED = 4;
}
optional SelectionAnalysis selection_analysis = 4
[default = ANALYSIS_WITH_SEED];
// If activated when SELECTION input is activated, will replace the computed
// camera motion (for any of the ANALYSIS_* case above) with the one supplied
// by the frame selection, in case the frame selection one is more stable.
// For example, if recomputed camera motion is unstable but the one from
// the selection result is stable, will use the stable result instead.
optional bool hybrid_selection_camera = 5 [default = false];
// Determines how optional input META is used to compute the final camera
// motion.
enum MetaAnalysis {
// Uses metadata supplied motions as is.
META_ANALYSIS_USE_META = 1;
// Seeds visual tracking from metadata motions - estimates visual residual
// motion and combines with metadata.
META_ANALYSIS_HYBRID = 2;
}
optional MetaAnalysis meta_analysis = 8 [default = META_ANALYSIS_USE_META];
// Determines number of homography models per frame stored in the CSV file
// or the homography metadata in META.
// For values > 1, MixtureHomographies are created.
optional int32 meta_models_per_frame = 6 [default = 1];
// Used for META_ANALYSIS_HYBRID. Rejects features which flow deviates
// domain_ratio * image diagonal size from the ground truth metadata motion.
optional float meta_outlier_domain_ratio = 9 [default = 0.0015];
// If true, the MotionAnalysisCalculator will skip all processing and emit no
// packets on any output. This is useful for quickly creating different
// versions of a MediaPipe graph without changing its structure, assuming that
// downstream calculators can handle missing input packets.
// TODO: Remove this hack. See b/36485206 for more details.
optional bool bypass_mode = 7 [default = false];
}
// Taken from
// java/com/google/android/libraries/microvideo/proto/microvideo.proto to
// satisfy leakr requirements
// TODO: Remove and use above proto.
message HomographyData {
// For each frame, there are 12 homography matrices stored. Each matrix is
// 3x3 (9 elements). This field will contain 12 x 3 x 3 float values. The
// first row of the first homography matrix will be followed by the second row
// of the first homography matrix, followed by third row of first homography
// matrix, followed by the first row of the second homography matrix, etc.
repeated float motion_homography_data = 1 [packed = true];
// Vector containing histogram counts for individual patches in the frame.
repeated uint32 histogram_count_data = 2 [packed = true];
// The width of the frame at the time metadata was sampled.
optional int32 frame_width = 3;
// The height of the frame at the time metadata was sampled.
optional int32 frame_height = 4;
}
@@ -12,6 +12,8 @@
// See the License for the specific language governing permissions and // See the License for the specific language governing permissions and
// limitations under the License. // limitations under the License.
#include <stdlib.h>
#include "mediapipe/framework/calculator_framework.h" #include "mediapipe/framework/calculator_framework.h"
#include "mediapipe/framework/formats/image_format.pb.h" #include "mediapipe/framework/formats/image_format.pb.h"
#include "mediapipe/framework/formats/image_frame.h" #include "mediapipe/framework/formats/image_frame.h"
@@ -66,6 +68,22 @@ ImageFormat::Format GetImageFormat(int num_channels) {
// output_stream: "VIDEO:video_frames" // output_stream: "VIDEO:video_frames"
// output_stream: "VIDEO_PRESTREAM:video_header" // output_stream: "VIDEO_PRESTREAM:video_header"
// } // }
//
// OpenCV's VideoCapture doesn't decode audio tracks. If the audio tracks need
// to be saved, specify an output side packet with tag "SAVED_AUDIO_PATH".
// The calculator will call FFmpeg binary to save audio tracks as an aac file.
// If the audio tracks can't be extracted by FFmpeg, the output side packet
// will contain an empty std::string.
//
// Example config:
// node {
// calculator: "OpenCvVideoDecoderCalculator"
// input_side_packet: "INPUT_FILE_PATH:input_file_path"
// output_side_packet: "SAVED_AUDIO_PATH:audio_path"
// output_stream: "VIDEO:video_frames"
// output_stream: "VIDEO_PRESTREAM:video_header"
// }
//
class OpenCvVideoDecoderCalculator : public CalculatorBase { class OpenCvVideoDecoderCalculator : public CalculatorBase {
public: public:
static ::mediapipe::Status GetContract(CalculatorContract* cc) { static ::mediapipe::Status GetContract(CalculatorContract* cc) {
@@ -74,6 +92,9 @@ class OpenCvVideoDecoderCalculator : public CalculatorBase {
if (cc->Outputs().HasTag("VIDEO_PRESTREAM")) { if (cc->Outputs().HasTag("VIDEO_PRESTREAM")) {
cc->Outputs().Tag("VIDEO_PRESTREAM").Set<VideoHeader>(); cc->Outputs().Tag("VIDEO_PRESTREAM").Set<VideoHeader>();
} }
if (cc->OutputSidePackets().HasTag("SAVED_AUDIO_PATH")) {
cc->OutputSidePackets().Tag("SAVED_AUDIO_PATH").Set<std::string>();
}
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -127,6 +148,35 @@ class OpenCvVideoDecoderCalculator : public CalculatorBase {
} }
// Rewind to the very first frame. // Rewind to the very first frame.
cap_->set(cv::CAP_PROP_POS_AVI_RATIO, 0); cap_->set(cv::CAP_PROP_POS_AVI_RATIO, 0);
if (cc->OutputSidePackets().HasTag("SAVED_AUDIO_PATH")) {
#ifdef HAVE_FFMPEG
std::string saved_audio_path = std::tmpnam(nullptr);
std::string ffmpeg_command =
absl::StrCat("ffmpeg -nostats -loglevel 0 -i ", input_file_path,
" -vn -f adts ", saved_audio_path);
system(ffmpeg_command.c_str());
int status_code = system(absl::StrCat("ls ", saved_audio_path).c_str());
if (status_code == 0) {
cc->OutputSidePackets()
.Tag("SAVED_AUDIO_PATH")
.Set(MakePacket<std::string>(saved_audio_path));
} else {
LOG(WARNING) << "FFmpeg can't extract audio from " << input_file_path
<< " by executing the following command: "
<< ffmpeg_command;
cc->OutputSidePackets()
.Tag("SAVED_AUDIO_PATH")
.Set(MakePacket<std::string>(std::string()));
}
#else
return ::mediapipe::InvalidArgumentErrorBuilder(MEDIAPIPE_LOC)
<< "OpenCVVideoDecoderCalculator can't save the audio file "
"because FFmpeg is not installed. Please remove "
"output_side_packet: \"SAVED_AUDIO_PATH\" from the node "
"config.";
#endif
}
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -55,8 +55,12 @@ TEST(OpenCvVideoDecoderCalculatorTest, TestMp4Avc720pVideo) {
EXPECT_EQ(640, header.height); EXPECT_EQ(640, header.height);
EXPECT_FLOAT_EQ(6.0f, header.duration); EXPECT_FLOAT_EQ(6.0f, header.duration);
EXPECT_FLOAT_EQ(30.0f, header.frame_rate); EXPECT_FLOAT_EQ(30.0f, header.frame_rate);
EXPECT_EQ(180, runner.Outputs().Tag("VIDEO").packets.size()); // The number of the output packets should be 180.
for (int i = 0; i < 180; ++i) { // Some OpenCV version returns the first two frames with the same timestamp on
// macos and we might miss one frame here.
int num_of_packets = runner.Outputs().Tag("VIDEO").packets.size();
EXPECT_GE(num_of_packets, 179);
for (int i = 0; i < num_of_packets; ++i) {
Packet image_frame_packet = runner.Outputs().Tag("VIDEO").packets[i]; Packet image_frame_packet = runner.Outputs().Tag("VIDEO").packets[i];
cv::Mat output_mat = cv::Mat output_mat =
formats::MatView(&(image_frame_packet.Get<ImageFrame>())); formats::MatView(&(image_frame_packet.Get<ImageFrame>()));
@@ -141,8 +145,12 @@ TEST(OpenCvVideoDecoderCalculatorTest, TestMkvVp8Video) {
EXPECT_EQ(320, header.height); EXPECT_EQ(320, header.height);
EXPECT_FLOAT_EQ(6.0f, header.duration); EXPECT_FLOAT_EQ(6.0f, header.duration);
EXPECT_FLOAT_EQ(30.0f, header.frame_rate); EXPECT_FLOAT_EQ(30.0f, header.frame_rate);
EXPECT_EQ(180, runner.Outputs().Tag("VIDEO").packets.size()); // The number of the output packets should be 180.
for (int i = 0; i < 180; ++i) { // Some OpenCV version returns the first two frames with the same timestamp on
// macos and we might miss one frame here.
int num_of_packets = runner.Outputs().Tag("VIDEO").packets.size();
EXPECT_GE(num_of_packets, 179);
for (int i = 0; i < num_of_packets; ++i) {
Packet image_frame_packet = runner.Outputs().Tag("VIDEO").packets[i]; Packet image_frame_packet = runner.Outputs().Tag("VIDEO").packets[i];
cv::Mat output_mat = cv::Mat output_mat =
formats::MatView(&(image_frame_packet.Get<ImageFrame>())); formats::MatView(&(image_frame_packet.Get<ImageFrame>()));
@@ -12,6 +12,8 @@
// See the License for the specific language governing permissions and // See the License for the specific language governing permissions and
// limitations under the License. // limitations under the License.
#include <stdlib.h>
#include <memory> #include <memory>
#include <string> #include <string>
#include <vector> #include <vector>
@@ -39,8 +41,7 @@ namespace mediapipe {
// packet. Currently, the calculator only supports one video stream (in // packet. Currently, the calculator only supports one video stream (in
// mediapipe::ImageFrame). // mediapipe::ImageFrame).
// //
// Example config to generate the output video file: // Example config:
//
// node { // node {
// calculator: "OpenCvVideoEncoderCalculator" // calculator: "OpenCvVideoEncoderCalculator"
// input_stream: "VIDEO:video" // input_stream: "VIDEO:video"
@@ -53,6 +54,26 @@ namespace mediapipe {
// } // }
// } // }
// } // }
//
// OpenCV's VideoWriter doesn't encode audio. If an input side packet with tag
// "AUDIO_FILE_PATH" is specified, the calculator will call FFmpeg binary to
// attach the audio file to the video as the last step in Close().
//
// Example config:
// node {
// calculator: "OpenCvVideoEncoderCalculator"
// input_stream: "VIDEO:video"
// input_stream: "VIDEO_PRESTREAM:video_header"
// input_side_packet: "OUTPUT_FILE_PATH:output_file_path"
// input_side_packet: "AUDIO_FILE_PATH:audio_path"
// node_options {
// [type.googleapis.com/mediapipe.OpenCvVideoEncoderCalculatorOptions]: {
// codec: "avc1"
// video_format: "mp4"
// }
// }
// }
//
class OpenCvVideoEncoderCalculator : public CalculatorBase { class OpenCvVideoEncoderCalculator : public CalculatorBase {
public: public:
static ::mediapipe::Status GetContract(CalculatorContract* cc); static ::mediapipe::Status GetContract(CalculatorContract* cc);
@@ -77,6 +98,9 @@ class OpenCvVideoEncoderCalculator : public CalculatorBase {
} }
RET_CHECK(cc->InputSidePackets().HasTag("OUTPUT_FILE_PATH")); RET_CHECK(cc->InputSidePackets().HasTag("OUTPUT_FILE_PATH"));
cc->InputSidePackets().Tag("OUTPUT_FILE_PATH").Set<std::string>(); cc->InputSidePackets().Tag("OUTPUT_FILE_PATH").Set<std::string>();
if (cc->InputSidePackets().HasTag("AUDIO_FILE_PATH")) {
cc->InputSidePackets().Tag("AUDIO_FILE_PATH").Set<std::string>();
}
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }
@@ -155,6 +179,33 @@ class OpenCvVideoEncoderCalculator : public CalculatorBase {
if (writer_ && writer_->isOpened()) { if (writer_ && writer_->isOpened()) {
writer_->release(); writer_->release();
} }
if (cc->InputSidePackets().HasTag("AUDIO_FILE_PATH")) {
#ifdef HAVE_FFMPEG
const std::string& audio_file_path =
cc->InputSidePackets().Tag("AUDIO_FILE_PATH").Get<std::string>();
if (audio_file_path.empty()) {
LOG(WARNING) << "OpenCvVideoEncoderCalculator isn't able to attach the "
"audio tracks to the generated video because the audio "
"file path is not specified.";
} else {
// A temp output file is needed because FFmpeg can't do in-place editing.
const std::string temp_file_path = std::tmpnam(nullptr);
system(absl::StrCat("mv ", output_file_path_, " ", temp_file_path,
"&& ffmpeg -nostats -loglevel 0 -i ", temp_file_path,
" -i ", audio_file_path,
" -c copy -map 0:v:0 -map 1:a:0 ", output_file_path_,
"&& rm ", temp_file_path)
.c_str());
}
#else
return ::mediapipe::InvalidArgumentErrorBuilder(MEDIAPIPE_LOC)
<< "OpenCVVideoEncoderCalculator can't attach the audio tracks to "
"the video because FFmpeg is not installed. Please remove "
"input_side_packet: \"AUDIO_FILE_PATH\" from the node "
"config.";
#endif
}
return ::mediapipe::OkStatus(); return ::mediapipe::OkStatus();
} }

Some files were not shown because too many files have changed in this diff Show More