Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
3b6d3c4058 | ||
|
|
252a5713c7 | ||
|
|
de4fbc10e6 | ||
|
|
d144e564d8 | ||
|
|
dd02df1dbe | ||
|
|
66b377c825 | ||
|
|
bf5185f122 | ||
|
|
a2823541e6 | ||
|
|
ae6be10afe | ||
|
|
38ee2603a7 | ||
|
|
86b3283b2f | ||
|
|
7d470a1335 | ||
|
|
d16cc3be5b | ||
|
|
137867d088 | ||
|
|
446d7cf6b6 | ||
|
|
90f72bd851 | ||
|
|
4285aeddfc | ||
|
|
37287925b0 | ||
|
|
48bcbb115f |
@@ -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
@@ -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" && \
|
||||||
|
|||||||
@@ -1,7 +1,7 @@
|
|||||||

|

|
||||||
=======================================================================
|
=======================================================================
|
||||||
|
|
||||||
[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.
|
||||||
|
|
||||||

|

|
||||||
|
|
||||||
@@ -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)
|
||||||
|
|
||||||

|
|
||||||

|

|
||||||
|

|
||||||
|

|
||||||

|

|
||||||

|

|
||||||
|
|
||||||
## 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).
|
||||||
|
|||||||
@@ -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"
|
||||||
],
|
],
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
@@ -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 |
@@ -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>();
|
||||||
|
|||||||
+1
-1
@@ -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>();
|
||||||
|
|||||||
+1
-1
@@ -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)));
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
@@ -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
Reference in New Issue
Block a user