Compare commits

..
69 Commits
Author SHA1 Message Date
Sebastian SchmidtandCopybara-Service d392f8ad98 Ensure that -std=c++14/17 is the first argument passed to Glog
PiperOrigin-RevId: 552509553
2023-07-31 09:47:32 -07:00
Sebastian SchmidtandCopybara-Service 81cf7fa173 Updat WASM binaries for 0.10.3 release
PiperOrigin-RevId: 551975834
2023-07-28 16:17:30 -07:00
MediaPipe TeamandCopybara-Service 8e313b4b0c Fix typo in model maker requirements.txt
PiperOrigin-RevId: 551973577
2023-07-28 16:07:12 -07:00
Sebastian SchmidtandCopybara-Service b4bcfab4f5 Remove extra letter from text classifier API
PiperOrigin-RevId: 551942087
2023-07-28 13:56:56 -07:00
Sebastian SchmidtandCopybara-Service 8ab9185c1d Use C+++ 17 for Glog only on Windows
PiperOrigin-RevId: 551928369
2023-07-28 12:58:42 -07:00
MediaPipe TeamandCopybara-Service 3f7752561b No public description
PiperOrigin-RevId: 551914786
2023-07-28 12:05:38 -07:00
MediaPipe TeamandCopybara-Service 9edb059d9f No public description
PiperOrigin-RevId: 551868738
2023-07-28 09:09:39 -07:00
MediaPipe TeamandCopybara-Service 7db0c1944b Internal change
PiperOrigin-RevId: 551789915
2023-07-28 02:29:52 -07:00
MediaPipe TeamandCopybara-Service db9a72a5df Internal Changes
PiperOrigin-RevId: 551674542
2023-07-27 16:35:35 -07:00
MediaPipe TeamandCopybara-Service 5c007558f8 internal change.
PiperOrigin-RevId: 551645248
2023-07-27 14:45:18 -07:00
Sebastian SchmidtandCopybara-Service 4d5c6bd33a Internal
PiperOrigin-RevId: 551625147
2023-07-27 13:33:14 -07:00
Sebastian SchmidtandCopybara-Service fdea10d230 Add C Headers for Text Classifier
PiperOrigin-RevId: 551618765
2023-07-27 13:09:49 -07:00
MediaPipe TeamandCopybara-Service 5b31f1e3e9 Update glog to latest commit
PiperOrigin-RevId: 551601991
2023-07-27 12:07:20 -07:00
Sebastian SchmidtandCopybara-Service 7d9cb4ee67 No public description
PiperOrigin-RevId: 551586945
2023-07-27 11:17:54 -07:00
MediaPipe TeamandCopybara-Service f3f9e71ccb No public description
PiperOrigin-RevId: 551549511
2023-07-27 09:15:46 -07:00
MediaPipe TeamandCopybara-Service 6de275834d internal change.
PiperOrigin-RevId: 551366789
2023-07-26 17:59:06 -07:00
Sebastian SchmidtandCopybara-Service c9d79a0076 Rollback of "Fix duplicate condition error in :resource_util"
PiperOrigin-RevId: 551332734
2023-07-26 15:31:51 -07:00
Sebastian SchmidtandCopybara-Service dad46e1e90 Update glog to 0.6
PiperOrigin-RevId: 551330044
2023-07-26 15:20:34 -07:00
Sebastian SchmidtandCopybara-Service f156397e8f Fix Android build with any Protos
PiperOrigin-RevId: 551325541
2023-07-26 15:04:38 -07:00
MediaPipe TeamandCopybara-Service fa5c1b03d2 No public description
PiperOrigin-RevId: 551277242
2023-07-26 12:08:33 -07:00
Sebastian SchmidtandCopybara-Service 87b925795d Update glog to 0.6
PiperOrigin-RevId: 551269455
2023-07-26 11:41:24 -07:00
Sebastian SchmidtandCopybara-Service 750f498b14 Internal
PiperOrigin-RevId: 551247471
2023-07-26 10:33:19 -07:00
MediaPipe TeamandCopybara-Service 1f6851c577 C++ Image segmenter add output size parameters.
PiperOrigin-RevId: 550995124
2023-07-25 14:22:11 -07:00
MediaPipe TeamandCopybara-Service bd7888cc0c 1. Move evaluation onto GPU/TPU hardware if available.
2. Move desired_precision and desired_recall from evaluate to hyperparameters so recall@precision metrics will be reported for both training and evaluation. This also fixes a bug where recompiling the model with the previously initialized metric objects would not properly reset the metric states.
3. Remove redundant label_names from create_... class methods in text_classifier. This information is already provided by the datasets.
4. Change loss function to FocalLoss.
5. Re-enable text_classifier unit tests using ExBert
6. Add input names to avoid flaky auto-assigned input names.

PiperOrigin-RevId: 550992146
2023-07-25 14:12:26 -07:00
MediaPipe TeamandCopybara-Service 85c3fed70a Add class weights to core hyperparameters and classifier library.
PiperOrigin-RevId: 550962843
2023-07-25 12:29:46 -07:00
MediaPipe TeamandCopybara-Service 62538a9496 No public description
PiperOrigin-RevId: 550954023
2023-07-25 11:57:08 -07:00
MediaPipe TeamandCopybara-Service 113c9b30c2 No public description
PiperOrigin-RevId: 550616150
2023-07-24 11:09:31 -07:00
Copybara-Service 66cceddb5e Merge pull request #4639 from priankakariatyml:ios-image-segmenter-container-utils
PiperOrigin-RevId: 550608973
2023-07-24 10:46:56 -07:00
Prianka Liz Kariat 72c62f7d5d Added iOS Image Segmenter Header 2023-07-24 20:38:16 +05:30
MediaPipe TeamandCopybara-Service 25b01784de Fix documentation
PiperOrigin-RevId: 549968822
2023-07-21 09:33:48 -07:00
MediaPipe TeamandCopybara-Service 9af637b125 Java API add visibility and presence for landmarks.
PiperOrigin-RevId: 549709256
2023-07-20 12:42:01 -07:00
Copybara-Service 236a36e39a Merge pull request #4629 from priankakariatyml:ios-vision-library-fixes
PiperOrigin-RevId: 549659874
2023-07-20 09:52:11 -07:00
Prianka Liz Kariat 540f4f7fe6 Fixed swift name of iOS face landmarker delegate 2023-07-20 15:57:37 +05:30
Prianka Liz Kariat 3198ccf6a5 Added missing headers in ios vision framework build 2023-07-20 15:57:16 +05:30
Steven HicksonandCopybara-Service e47af74b15 Adding support for 2 things in tensors_to_image_calculator:
1) 1 channel support for conversion after inference.
2) multitask support by allowing for different tensor outputs.

PiperOrigin-RevId: 549412331
2023-07-19 13:41:46 -07:00
MediaPipe TeamandCopybara-Service 085840388b Move waitOnCpu and waitOnGpu out of the synchronized block, which can cause deadlock.
PiperOrigin-RevId: 549217916
2023-07-18 23:42:01 -07:00
MediaPipe TeamandCopybara-Service 4e72fcf0cb Replace CHECK with RET_CHECK in GetContract() implementation from six calculators.
PiperOrigin-RevId: 549158984
2023-07-18 17:38:44 -07:00
MediaPipe TeamandCopybara-Service 4c60fe7365 add pose landmarks connections in C++ API
PiperOrigin-RevId: 549108310
2023-07-18 14:21:02 -07:00
MediaPipe TeamandCopybara-Service 9b00582f21 add hand landmarks connections in C++ API.
PiperOrigin-RevId: 549108307
2023-07-18 14:16:27 -07:00
MediaPipe TeamandCopybara-Service cb915858fa Internal change
PiperOrigin-RevId: 549052451
2023-07-18 11:01:37 -07:00
Jiuqiang TangandCopybara-Service 0c01187cf5 Internal change
PiperOrigin-RevId: 548886447
2023-07-17 21:57:24 -07:00
MediaPipe TeamandCopybara-Service ef12ce8575 Internal change
PiperOrigin-RevId: 548821518
2023-07-17 15:56:05 -07:00
MediaPipe TeamandCopybara-Service f1f9f80cd9 Internal change
PiperOrigin-RevId: 548746432
2023-07-17 11:18:00 -07:00
MediaPipe TeamandCopybara-Service 17bc1a5ab5 Internal change
PiperOrigin-RevId: 548196034
2023-07-14 12:39:45 -07:00
MediaPipe TeamandCopybara-Service 2fae07375c Discard outdated packets earlier in MuxInputStreamHandler.
In our pipeline, a deadlock is detected because the packets in deselected
data streams get piled up. In the current implementation, those packets only get
removed in FillInputSet(), but we should also do that in GetNodeReadiness().

PiperOrigin-RevId: 548051369
2023-07-14 01:12:14 -07:00
MediaPipe TeamandCopybara-Service 723e91cec1 Generalize non-define registration with MEDIAPIPE_STATIC_REGISTRATOR_TEMPLATE
PiperOrigin-RevId: 547929982
2023-07-13 14:52:37 -07:00
MediaPipe TeamandCopybara-Service c2c67c20fa Internal change
PiperOrigin-RevId: 547924907
2023-07-13 14:37:40 -07:00
Sebastian SchmidtandCopybara-Service 327feb42d1 Support WASM asset loading for MediaPipe Task Web
PiperOrigin-RevId: 547882566
2023-07-13 12:26:59 -07:00
MediaPipe TeamandCopybara-Service 8b59567cb7 Add proto3 Any proto support for Java task api
PiperOrigin-RevId: 547836041
2023-07-13 10:10:17 -07:00
MediaPipe TeamandCopybara-Service e37bedd344 Fix Halide BUILD rules
PiperOrigin-RevId: 547755467
2023-07-13 04:47:34 -07:00
MediaPipe TeamandCopybara-Service 251c5421f6 Internal change
PiperOrigin-RevId: 547735699
2023-07-13 02:53:16 -07:00
MediaPipe TeamandCopybara-Service 450c933cb5 MEDIAPIPE_NODE/SUBGRAPH_IMPLEMENTATION to use common define for registration
PiperOrigin-RevId: 547669538
2023-07-12 20:10:15 -07:00
MediaPipe TeamandCopybara-Service cc2aa4f4cc InferenceCalculatorAdvancedGL save cache in Open().
PiperOrigin-RevId: 547652481
2023-07-12 18:09:51 -07:00
MediaPipe TeamandCopybara-Service a2cd3e7f95 Internal change
PiperOrigin-RevId: 547614484
2023-07-12 15:17:40 -07:00
MediaPipe TeamandCopybara-Service 37b68714b8 Internal change
PiperOrigin-RevId: 547424721
2023-07-12 01:32:51 -07:00
MediaPipe TeamandCopybara-Service 3e93cbc838 Internal change
PiperOrigin-RevId: 547404737
2023-07-12 00:04:40 -07:00
Yilei YangandCopybara-Service 917af2ce6b Internal change
PiperOrigin-RevId: 547346939
2023-07-11 17:52:07 -07:00
Sebastian SchmidtandCopybara-Service f2f49b9fc8 Add angle to BoundingBox
PiperOrigin-RevId: 547321781
2023-07-11 16:00:35 -07:00
MediaPipe TeamandCopybara-Service aabf61f28d Internal Change
PiperOrigin-RevId: 547299595
2023-07-11 14:35:18 -07:00
MediaPipe TeamandCopybara-Service 56bc019819 Model Maker allow core dataset library to handle datasets with unknown sizes.
PiperOrigin-RevId: 547268411
2023-07-11 12:47:37 -07:00
MediaPipe TeamandCopybara-Service 4788fddde9 Internal Change
PiperOrigin-RevId: 547265380
2023-07-11 12:34:32 -07:00
MediaPipe TeamandCopybara-Service e4ec4d2526 Internal change
PiperOrigin-RevId: 547258228
2023-07-11 12:05:58 -07:00
MediaPipe TeamandCopybara-Service bf6561ce91 add symmetric color style option
PiperOrigin-RevId: 547069284
2023-07-10 21:41:01 -07:00
MediaPipe TeamandCopybara-Service 0bde987a38 Removed internal dependency on OpenCV 3.x, migrating it to OpenCV 4.x
PiperOrigin-RevId: 546945166
2023-07-10 12:17:54 -07:00
Copybara-Service df3f4167ae Merge pull request #4600 from priankakariatyml:ios-orientation-fix
PiperOrigin-RevId: 546358930
2023-07-07 12:56:39 -07:00
Sebastian SchmidtandCopybara-Service 03bc9d64f2 Update glog to 0.6
PiperOrigin-RevId: 546349096
2023-07-07 12:22:13 -07:00
MediaPipe TeamandCopybara-Service d45b15ef84 Add face landmarks connections for C++.
PiperOrigin-RevId: 546345842
2023-07-07 12:08:09 -07:00
Prianka Liz Kariat cae10ea115 Updated documentation of MPImage 2023-07-07 22:03:15 +05:30
Prianka Liz Kariat 7556a3f1b4 Changed left and right image orientation angles to match iOS UIImageOrientation 2023-07-07 19:57:44 +05:30
140 changed files with 3419 additions and 960 deletions
+7 -7
View File
@@ -157,22 +157,22 @@ http_archive(
# 2020-08-21 # 2020-08-21
http_archive( http_archive(
name = "com_github_glog_glog", name = "com_github_glog_glog",
strip_prefix = "glog-0a2e5931bd5ff22fd3bf8999eb8ce776f159cda6", strip_prefix = "glog-3a0d4d22c5ae0b9a2216988411cfa6bf860cc372",
sha256 = "58c9b3b6aaa4dd8b836c0fd8f65d0f941441fb95e27212c5eeb9979cfd3592ab", sha256 = "170d08f80210b82d95563f4723a15095eff1aad1863000e8eeb569c96a98fefb",
urls = [ urls = [
"https://github.com/google/glog/archive/0a2e5931bd5ff22fd3bf8999eb8ce776f159cda6.zip", "https://github.com/google/glog/archive/3a0d4d22c5ae0b9a2216988411cfa6bf860cc372.zip",
], ],
) )
http_archive( http_archive(
name = "com_github_glog_glog_no_gflags", name = "com_github_glog_glog_no_gflags",
strip_prefix = "glog-0a2e5931bd5ff22fd3bf8999eb8ce776f159cda6", strip_prefix = "glog-3a0d4d22c5ae0b9a2216988411cfa6bf860cc372",
sha256 = "58c9b3b6aaa4dd8b836c0fd8f65d0f941441fb95e27212c5eeb9979cfd3592ab", sha256 = "170d08f80210b82d95563f4723a15095eff1aad1863000e8eeb569c96a98fefb",
build_file = "@//third_party:glog_no_gflags.BUILD", build_file = "@//third_party:glog_no_gflags.BUILD",
urls = [ urls = [
"https://github.com/google/glog/archive/0a2e5931bd5ff22fd3bf8999eb8ce776f159cda6.zip", "https://github.com/google/glog/archive/3a0d4d22c5ae0b9a2216988411cfa6bf860cc372.zip",
], ],
patches = [ patches = [
"@//third_party:com_github_glog_glog_9779e5ea6ef59562b030248947f787d1256132ae.diff", "@//third_party:com_github_glog_glog.diff",
], ],
patch_args = [ patch_args = [
"-p1", "-p1",
+93 -58
View File
@@ -68,30 +68,108 @@ config_setting(
visibility = ["//visibility:public"], visibility = ["//visibility:public"],
) )
# Note: this cannot just match "apple_platform_type": "macos" because that option # Generic MacOS.
# defaults to "macos" even when building on Linux! config_setting(
alias(
name = "macos", name = "macos",
actual = select({ constraint_values = [
":macos_i386": ":macos_i386", "@platforms//os:macos",
":macos_x86_64": ":macos_x86_64", ],
":macos_arm64": ":macos_arm64",
"//conditions:default": ":macos_i386", # Arbitrarily chosen from above.
}),
visibility = ["//visibility:public"], visibility = ["//visibility:public"],
) )
# Note: this also matches on crosstool_top so that it does not produce ambiguous # MacOS x86 64-bit.
# selectors when used together with "android". config_setting(
name = "macos_x86_64",
constraint_values = [
"@platforms//os:macos",
"@platforms//cpu:x86_64",
],
visibility = ["//visibility:public"],
)
# MacOS ARM64.
config_setting(
name = "macos_arm64",
constraint_values = [
"@platforms//os:macos",
"@platforms//cpu:arm64",
],
visibility = ["//visibility:public"],
)
# Generic iOS.
config_setting( config_setting(
name = "ios", name = "ios",
values = { constraint_values = [
"crosstool_top": "@bazel_tools//tools/cpp:toolchain", "@platforms//os:ios",
"apple_platform_type": "ios", ],
},
visibility = ["//visibility:public"], visibility = ["//visibility:public"],
) )
# iOS device ARM32.
config_setting(
name = "ios_armv7",
constraint_values = [
"@platforms//os:ios",
"@platforms//cpu:arm",
],
visibility = ["//visibility:public"],
)
# iOS device ARM64.
config_setting(
name = "ios_arm64",
constraint_values = [
"@platforms//os:ios",
"@platforms//cpu:arm64",
],
visibility = ["//visibility:public"],
)
# iOS device ARM64E.
config_setting(
name = "ios_arm64e",
constraint_values = [
"@platforms//os:ios",
"@platforms//cpu:arm64e",
],
visibility = ["//visibility:public"],
)
# iOS simulator x86 32-bit.
config_setting(
name = "ios_i386",
constraint_values = [
"@platforms//os:ios",
"@platforms//cpu:x86_32",
"@build_bazel_apple_support//constraints:simulator",
],
visibility = ["//visibility:public"],
)
# iOS simulator x86 64-bit.
config_setting(
name = "ios_x86_64",
constraint_values = [
"@platforms//os:ios",
"@platforms//cpu:x86_64",
"@build_bazel_apple_support//constraints:simulator",
],
visibility = ["//visibility:public"],
)
# iOS simulator ARM64.
config_setting(
name = "ios_sim_arm64",
constraint_values = [
"@platforms//os:ios",
"@platforms//cpu:arm64",
"@build_bazel_apple_support//constraints:simulator",
],
visibility = ["//visibility:public"],
)
# Generic Apple.
alias( alias(
name = "apple", name = "apple",
actual = select({ actual = select({
@@ -102,49 +180,6 @@ alias(
visibility = ["//visibility:public"], visibility = ["//visibility:public"],
) )
config_setting(
name = "macos_i386",
values = {
"apple_platform_type": "macos",
"cpu": "darwin",
},
visibility = ["//visibility:public"],
)
config_setting(
name = "macos_x86_64",
values = {
"apple_platform_type": "macos",
"cpu": "darwin_x86_64",
},
visibility = ["//visibility:public"],
)
config_setting(
name = "macos_arm64",
values = {
"apple_platform_type": "macos",
"cpu": "darwin_arm64",
},
visibility = ["//visibility:public"],
)
[
config_setting(
name = arch,
values = {"cpu": arch},
visibility = ["//visibility:public"],
)
for arch in [
"ios_i386",
"ios_x86_64",
"ios_armv7",
"ios_arm64",
"ios_arm64e",
"ios_sim_arm64",
]
]
config_setting( config_setting(
name = "windows", name = "windows",
values = {"cpu": "x64_windows"}, values = {"cpu": "x64_windows"},
-11
View File
@@ -381,17 +381,6 @@ cc_library(
alwayslink = 1, alwayslink = 1,
) )
cc_library(
name = "clip_detection_vector_size_calculator",
srcs = ["clip_detection_vector_size_calculator.cc"],
deps = [
":clip_vector_size_calculator",
"//mediapipe/framework:calculator_framework",
"//mediapipe/framework/formats:detection_cc_proto",
],
alwayslink = 1,
)
cc_test( cc_test(
name = "clip_vector_size_calculator_test", name = "clip_vector_size_calculator_test",
srcs = ["clip_vector_size_calculator_test.cc"], srcs = ["clip_vector_size_calculator_test.cc"],
@@ -1,26 +0,0 @@
// Copyright 2019 The MediaPipe Authors.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
#include <vector>
#include "mediapipe/calculators/core/clip_vector_size_calculator.h"
#include "mediapipe/framework/formats/detection.pb.h"
namespace mediapipe {
typedef ClipVectorSizeCalculator<::mediapipe::Detection>
ClipDetectionVectorSizeCalculator;
REGISTER_CALCULATOR(ClipDetectionVectorSizeCalculator);
} // namespace mediapipe
@@ -112,7 +112,7 @@ class BilateralFilterCalculator : public CalculatorBase {
REGISTER_CALCULATOR(BilateralFilterCalculator); REGISTER_CALCULATOR(BilateralFilterCalculator);
absl::Status BilateralFilterCalculator::GetContract(CalculatorContract* cc) { absl::Status BilateralFilterCalculator::GetContract(CalculatorContract* cc) {
CHECK_GE(cc->Inputs().NumEntries(), 1); RET_CHECK_GE(cc->Inputs().NumEntries(), 1);
if (cc->Inputs().HasTag(kInputFrameTag) && if (cc->Inputs().HasTag(kInputFrameTag) &&
cc->Inputs().HasTag(kInputFrameTagGpu)) { cc->Inputs().HasTag(kInputFrameTagGpu)) {
@@ -110,7 +110,7 @@ REGISTER_CALCULATOR(SegmentationSmoothingCalculator);
absl::Status SegmentationSmoothingCalculator::GetContract( absl::Status SegmentationSmoothingCalculator::GetContract(
CalculatorContract* cc) { CalculatorContract* cc) {
CHECK_GE(cc->Inputs().NumEntries(), 1); RET_CHECK_GE(cc->Inputs().NumEntries(), 1);
cc->Inputs().Tag(kCurrentMaskTag).Set<Image>(); cc->Inputs().Tag(kCurrentMaskTag).Set<Image>();
cc->Inputs().Tag(kPreviousMaskTag).Set<Image>(); cc->Inputs().Tag(kPreviousMaskTag).Set<Image>();
@@ -142,7 +142,7 @@ class SetAlphaCalculator : public CalculatorBase {
REGISTER_CALCULATOR(SetAlphaCalculator); REGISTER_CALCULATOR(SetAlphaCalculator);
absl::Status SetAlphaCalculator::GetContract(CalculatorContract* cc) { absl::Status SetAlphaCalculator::GetContract(CalculatorContract* cc) {
CHECK_GE(cc->Inputs().NumEntries(), 1); RET_CHECK_GE(cc->Inputs().NumEntries(), 1);
bool use_gpu = false; bool use_gpu = false;
@@ -69,6 +69,7 @@ class InferenceCalculatorGlAdvancedImpl
gpu_delegate_options); gpu_delegate_options);
absl::Status ReadGpuCaches(tflite::gpu::TFLiteGPURunner* gpu_runner) const; absl::Status ReadGpuCaches(tflite::gpu::TFLiteGPURunner* gpu_runner) const;
absl::Status SaveGpuCaches(tflite::gpu::TFLiteGPURunner* gpu_runner) const; absl::Status SaveGpuCaches(tflite::gpu::TFLiteGPURunner* gpu_runner) const;
bool UseSerializedModel() const { return use_serialized_model_; }
private: private:
bool use_kernel_caching_ = false; bool use_kernel_caching_ = false;
@@ -150,8 +151,6 @@ InferenceCalculatorGlAdvancedImpl::GpuInferenceRunner::Process(
} }
absl::Status InferenceCalculatorGlAdvancedImpl::GpuInferenceRunner::Close() { absl::Status InferenceCalculatorGlAdvancedImpl::GpuInferenceRunner::Close() {
MP_RETURN_IF_ERROR(
on_disk_cache_helper_.SaveGpuCaches(tflite_gpu_runner_.get()));
return gpu_helper_.RunInGlContext([this]() -> absl::Status { return gpu_helper_.RunInGlContext([this]() -> absl::Status {
tflite_gpu_runner_.reset(); tflite_gpu_runner_.reset();
return absl::OkStatus(); return absl::OkStatus();
@@ -226,9 +225,14 @@ InferenceCalculatorGlAdvancedImpl::GpuInferenceRunner::InitTFLiteGPURunner(
tflite_gpu_runner_->GetOutputShapes()[i].c}; tflite_gpu_runner_->GetOutputShapes()[i].c};
} }
if (on_disk_cache_helper_.UseSerializedModel()) {
tflite_gpu_runner_->ForceOpenCLInitFromSerializedModel();
}
MP_RETURN_IF_ERROR( MP_RETURN_IF_ERROR(
on_disk_cache_helper_.ReadGpuCaches(tflite_gpu_runner_.get())); on_disk_cache_helper_.ReadGpuCaches(tflite_gpu_runner_.get()));
return tflite_gpu_runner_->Build(); MP_RETURN_IF_ERROR(tflite_gpu_runner_->Build());
return on_disk_cache_helper_.SaveGpuCaches(tflite_gpu_runner_.get());
} }
#if defined(MEDIAPIPE_ANDROID) || defined(MEDIAPIPE_CHROMIUMOS) #if defined(MEDIAPIPE_ANDROID) || defined(MEDIAPIPE_CHROMIUMOS)
+7 -3
View File
@@ -406,8 +406,13 @@ cc_library(
alwayslink = 1, alwayslink = 1,
) )
# This dependency removed tensorflow_jellyfish_deps and xprofilez_with_server because they failed # This dependency removed the following 3 targets because they failed Boq conformance test:
# Boq conformance test. Weigh your use case to see if this will work for you. #
# tensorflow_jellyfish_deps
# jfprof_lib
# xprofilez_with_server
#
# If you need them plz consider tensorflow_inference_calculator_no_envelope_loader.
cc_library( cc_library(
name = "tensorflow_inference_calculator_for_boq", name = "tensorflow_inference_calculator_for_boq",
srcs = ["tensorflow_inference_calculator.cc"], srcs = ["tensorflow_inference_calculator.cc"],
@@ -927,7 +932,6 @@ cc_test(
"//mediapipe/framework:timestamp", "//mediapipe/framework:timestamp",
"//mediapipe/framework/formats:detection_cc_proto", "//mediapipe/framework/formats:detection_cc_proto",
"//mediapipe/framework/formats:image_frame", "//mediapipe/framework/formats:image_frame",
"//mediapipe/framework/formats:image_frame_opencv",
"//mediapipe/framework/formats:location", "//mediapipe/framework/formats:location",
"//mediapipe/framework/formats:location_opencv", "//mediapipe/framework/formats:location_opencv",
"//mediapipe/framework/port:gtest_main", "//mediapipe/framework/port:gtest_main",
@@ -164,8 +164,8 @@ class PackMediaSequenceCalculator : public CalculatorBase {
} }
} }
CHECK(cc->Outputs().HasTag(kSequenceExampleTag) || RET_CHECK(cc->Outputs().HasTag(kSequenceExampleTag) ||
cc->OutputSidePackets().HasTag(kSequenceExampleTag)) cc->OutputSidePackets().HasTag(kSequenceExampleTag))
<< "Neither the output stream nor the output side packet is set to " << "Neither the output stream nor the output side packet is set to "
"output the sequence example."; "output the sequence example.";
if (cc->Outputs().HasTag(kSequenceExampleTag)) { if (cc->Outputs().HasTag(kSequenceExampleTag)) {
@@ -23,7 +23,6 @@
#include "mediapipe/framework/calculator_runner.h" #include "mediapipe/framework/calculator_runner.h"
#include "mediapipe/framework/formats/detection.pb.h" #include "mediapipe/framework/formats/detection.pb.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/location.h" #include "mediapipe/framework/formats/location.h"
#include "mediapipe/framework/formats/location_opencv.h" #include "mediapipe/framework/formats/location_opencv.h"
#include "mediapipe/framework/port/gmock.h" #include "mediapipe/framework/port/gmock.h"
@@ -96,7 +95,8 @@ TEST_F(PackMediaSequenceCalculatorTest, PacksTwoImages) {
mpms::SetClipMediaId(test_video_id, input_sequence.get()); mpms::SetClipMediaId(test_video_id, input_sequence.get());
cv::Mat image(2, 3, CV_8UC3, cv::Scalar(0, 0, 255)); cv::Mat image(2, 3, CV_8UC3, cv::Scalar(0, 0, 255));
std::vector<uchar> bytes; std::vector<uchar> bytes;
ASSERT_TRUE(cv::imencode(".jpg", image, bytes, {80})); ASSERT_TRUE(
cv::imencode(".jpg", image, bytes, {cv::IMWRITE_HDR_COMPRESSION, 1}));
OpenCvImageEncoderCalculatorResults encoded_image; OpenCvImageEncoderCalculatorResults encoded_image;
encoded_image.set_encoded_image(bytes.data(), bytes.size()); encoded_image.set_encoded_image(bytes.data(), bytes.size());
encoded_image.set_width(2); encoded_image.set_width(2);
@@ -139,7 +139,8 @@ TEST_F(PackMediaSequenceCalculatorTest, PacksTwoPrefixedImages) {
mpms::SetClipMediaId(test_video_id, input_sequence.get()); mpms::SetClipMediaId(test_video_id, input_sequence.get());
cv::Mat image(2, 3, CV_8UC3, cv::Scalar(0, 0, 255)); cv::Mat image(2, 3, CV_8UC3, cv::Scalar(0, 0, 255));
std::vector<uchar> bytes; std::vector<uchar> bytes;
ASSERT_TRUE(cv::imencode(".jpg", image, bytes, {80})); ASSERT_TRUE(
cv::imencode(".jpg", image, bytes, {cv::IMWRITE_HDR_COMPRESSION, 1}));
OpenCvImageEncoderCalculatorResults encoded_image; OpenCvImageEncoderCalculatorResults encoded_image;
encoded_image.set_encoded_image(bytes.data(), bytes.size()); encoded_image.set_encoded_image(bytes.data(), bytes.size());
encoded_image.set_width(2); encoded_image.set_width(2);
@@ -378,7 +379,8 @@ TEST_F(PackMediaSequenceCalculatorTest, PacksAdditionalContext) {
Adopt(input_sequence.release()); Adopt(input_sequence.release());
cv::Mat image(2, 3, CV_8UC3, cv::Scalar(0, 0, 255)); cv::Mat image(2, 3, CV_8UC3, cv::Scalar(0, 0, 255));
std::vector<uchar> bytes; std::vector<uchar> bytes;
ASSERT_TRUE(cv::imencode(".jpg", image, bytes, {80})); ASSERT_TRUE(
cv::imencode(".jpg", image, bytes, {cv::IMWRITE_HDR_COMPRESSION, 1}));
OpenCvImageEncoderCalculatorResults encoded_image; OpenCvImageEncoderCalculatorResults encoded_image;
encoded_image.set_encoded_image(bytes.data(), bytes.size()); encoded_image.set_encoded_image(bytes.data(), bytes.size());
auto image_ptr = auto image_ptr =
@@ -410,7 +412,8 @@ TEST_F(PackMediaSequenceCalculatorTest, PacksTwoForwardFlowEncodeds) {
cv::Mat image(2, 3, CV_8UC3, cv::Scalar(0, 0, 255)); cv::Mat image(2, 3, CV_8UC3, cv::Scalar(0, 0, 255));
std::vector<uchar> bytes; std::vector<uchar> bytes;
ASSERT_TRUE(cv::imencode(".jpg", image, bytes, {80})); ASSERT_TRUE(
cv::imencode(".jpg", image, bytes, {cv::IMWRITE_HDR_COMPRESSION, 1}));
std::string test_flow_string(bytes.begin(), bytes.end()); std::string test_flow_string(bytes.begin(), bytes.end());
OpenCvImageEncoderCalculatorResults encoded_flow; OpenCvImageEncoderCalculatorResults encoded_flow;
encoded_flow.set_encoded_image(test_flow_string); encoded_flow.set_encoded_image(test_flow_string);
@@ -618,7 +621,8 @@ TEST_F(PackMediaSequenceCalculatorTest, PacksBBoxWithImages) {
} }
cv::Mat image(height, width, CV_8UC3, cv::Scalar(0, 0, 255)); cv::Mat image(height, width, CV_8UC3, cv::Scalar(0, 0, 255));
std::vector<uchar> bytes; std::vector<uchar> bytes;
ASSERT_TRUE(cv::imencode(".jpg", image, bytes, {80})); ASSERT_TRUE(
cv::imencode(".jpg", image, bytes, {cv::IMWRITE_HDR_COMPRESSION, 1}));
OpenCvImageEncoderCalculatorResults encoded_image; OpenCvImageEncoderCalculatorResults encoded_image;
encoded_image.set_encoded_image(bytes.data(), bytes.size()); encoded_image.set_encoded_image(bytes.data(), bytes.size());
encoded_image.set_width(width); encoded_image.set_width(width);
@@ -767,7 +771,8 @@ TEST_F(PackMediaSequenceCalculatorTest, MissingStreamOK) {
cv::Mat image(2, 3, CV_8UC3, cv::Scalar(0, 0, 255)); cv::Mat image(2, 3, CV_8UC3, cv::Scalar(0, 0, 255));
std::vector<uchar> bytes; std::vector<uchar> bytes;
ASSERT_TRUE(cv::imencode(".jpg", image, bytes, {80})); ASSERT_TRUE(
cv::imencode(".jpg", image, bytes, {cv::IMWRITE_HDR_COMPRESSION, 1}));
std::string test_flow_string(bytes.begin(), bytes.end()); std::string test_flow_string(bytes.begin(), bytes.end());
OpenCvImageEncoderCalculatorResults encoded_flow; OpenCvImageEncoderCalculatorResults encoded_flow;
encoded_flow.set_encoded_image(test_flow_string); encoded_flow.set_encoded_image(test_flow_string);
@@ -813,7 +818,8 @@ TEST_F(PackMediaSequenceCalculatorTest, MissingStreamNotOK) {
mpms::SetClipMediaId(test_video_id, input_sequence.get()); mpms::SetClipMediaId(test_video_id, input_sequence.get());
cv::Mat image(2, 3, CV_8UC3, cv::Scalar(0, 0, 255)); cv::Mat image(2, 3, CV_8UC3, cv::Scalar(0, 0, 255));
std::vector<uchar> bytes; std::vector<uchar> bytes;
ASSERT_TRUE(cv::imencode(".jpg", image, bytes, {80})); ASSERT_TRUE(
cv::imencode(".jpg", image, bytes, {cv::IMWRITE_HDR_COMPRESSION, 1}));
std::string test_flow_string(bytes.begin(), bytes.end()); std::string test_flow_string(bytes.begin(), bytes.end());
OpenCvImageEncoderCalculatorResults encoded_flow; OpenCvImageEncoderCalculatorResults encoded_flow;
encoded_flow.set_encoded_image(test_flow_string); encoded_flow.set_encoded_image(test_flow_string);
@@ -970,7 +976,8 @@ TEST_F(PackMediaSequenceCalculatorTest, TestReconcilingAnnotations) {
auto input_sequence = ::absl::make_unique<tf::SequenceExample>(); auto input_sequence = ::absl::make_unique<tf::SequenceExample>();
cv::Mat image(2, 3, CV_8UC3, cv::Scalar(0, 0, 255)); cv::Mat image(2, 3, CV_8UC3, cv::Scalar(0, 0, 255));
std::vector<uchar> bytes; std::vector<uchar> bytes;
ASSERT_TRUE(cv::imencode(".jpg", image, bytes, {80})); ASSERT_TRUE(
cv::imencode(".jpg", image, bytes, {cv::IMWRITE_HDR_COMPRESSION, 1}));
OpenCvImageEncoderCalculatorResults encoded_image; OpenCvImageEncoderCalculatorResults encoded_image;
encoded_image.set_encoded_image(bytes.data(), bytes.size()); encoded_image.set_encoded_image(bytes.data(), bytes.size());
encoded_image.set_width(2); encoded_image.set_width(2);
@@ -1021,7 +1028,8 @@ TEST_F(PackMediaSequenceCalculatorTest, TestOverwritingAndReconciling) {
auto input_sequence = ::absl::make_unique<tf::SequenceExample>(); auto input_sequence = ::absl::make_unique<tf::SequenceExample>();
cv::Mat image(2, 3, CV_8UC3, cv::Scalar(0, 0, 255)); cv::Mat image(2, 3, CV_8UC3, cv::Scalar(0, 0, 255));
std::vector<uchar> bytes; std::vector<uchar> bytes;
ASSERT_TRUE(cv::imencode(".jpg", image, bytes, {80})); ASSERT_TRUE(
cv::imencode(".jpg", image, bytes, {cv::IMWRITE_HDR_COMPRESSION, 1}));
OpenCvImageEncoderCalculatorResults encoded_image; OpenCvImageEncoderCalculatorResults encoded_image;
encoded_image.set_encoded_image(bytes.data(), bytes.size()); encoded_image.set_encoded_image(bytes.data(), bytes.size());
int height = 2; int height = 2;
@@ -172,7 +172,7 @@ class AnnotationOverlayCalculator : public CalculatorBase {
REGISTER_CALCULATOR(AnnotationOverlayCalculator); REGISTER_CALCULATOR(AnnotationOverlayCalculator);
absl::Status AnnotationOverlayCalculator::GetContract(CalculatorContract* cc) { absl::Status AnnotationOverlayCalculator::GetContract(CalculatorContract* cc) {
CHECK_GE(cc->Inputs().NumEntries(), 1); RET_CHECK_GE(cc->Inputs().NumEntries(), 1);
bool use_gpu = false; bool use_gpu = false;
@@ -189,13 +189,13 @@ absl::Status AnnotationOverlayCalculator::GetContract(CalculatorContract* cc) {
#if !MEDIAPIPE_DISABLE_GPU #if !MEDIAPIPE_DISABLE_GPU
if (cc->Inputs().HasTag(kGpuBufferTag)) { if (cc->Inputs().HasTag(kGpuBufferTag)) {
cc->Inputs().Tag(kGpuBufferTag).Set<mediapipe::GpuBuffer>(); cc->Inputs().Tag(kGpuBufferTag).Set<mediapipe::GpuBuffer>();
CHECK(cc->Outputs().HasTag(kGpuBufferTag)); RET_CHECK(cc->Outputs().HasTag(kGpuBufferTag));
use_gpu = true; use_gpu = true;
} }
#endif // !MEDIAPIPE_DISABLE_GPU #endif // !MEDIAPIPE_DISABLE_GPU
if (cc->Inputs().HasTag(kImageFrameTag)) { if (cc->Inputs().HasTag(kImageFrameTag)) {
cc->Inputs().Tag(kImageFrameTag).Set<ImageFrame>(); cc->Inputs().Tag(kImageFrameTag).Set<ImageFrame>();
CHECK(cc->Outputs().HasTag(kImageFrameTag)); RET_CHECK(cc->Outputs().HasTag(kImageFrameTag));
} }
// Data streams to render. // Data streams to render.
@@ -322,27 +322,30 @@ absl::Status LandmarksToRenderDataCalculator::Process(CalculatorContext* cc) {
options_.presence_threshold(), options_.connection_color(), thickness, options_.presence_threshold(), options_.connection_color(), thickness,
/*normalized=*/false, render_data.get()); /*normalized=*/false, render_data.get());
} }
for (int i = 0; i < landmarks.landmark_size(); ++i) { if (options_.render_landmarks()) {
const Landmark& landmark = landmarks.landmark(i); for (int i = 0; i < landmarks.landmark_size(); ++i) {
const Landmark& landmark = landmarks.landmark(i);
if (!IsLandmarkVisibleAndPresent<Landmark>( if (!IsLandmarkVisibleAndPresent<Landmark>(
landmark, options_.utilize_visibility(), landmark, options_.utilize_visibility(),
options_.visibility_threshold(), options_.utilize_presence(), options_.visibility_threshold(), options_.utilize_presence(),
options_.presence_threshold())) { options_.presence_threshold())) {
continue; continue;
} }
auto* landmark_data_render = AddPointRenderData( auto* landmark_data_render = AddPointRenderData(
options_.landmark_color(), thickness, render_data.get()); options_.landmark_color(), thickness, render_data.get());
if (visualize_depth) { if (visualize_depth) {
SetColorSizeValueFromZ(landmark.z(), z_min, z_max, landmark_data_render, SetColorSizeValueFromZ(landmark.z(), z_min, z_max,
options_.min_depth_circle_thickness(), landmark_data_render,
options_.max_depth_circle_thickness()); options_.min_depth_circle_thickness(),
options_.max_depth_circle_thickness());
}
auto* landmark_data = landmark_data_render->mutable_point();
landmark_data->set_normalized(false);
landmark_data->set_x(landmark.x());
landmark_data->set_y(landmark.y());
} }
auto* landmark_data = landmark_data_render->mutable_point();
landmark_data->set_normalized(false);
landmark_data->set_x(landmark.x());
landmark_data->set_y(landmark.y());
} }
} }
@@ -368,27 +371,30 @@ absl::Status LandmarksToRenderDataCalculator::Process(CalculatorContext* cc) {
options_.presence_threshold(), options_.connection_color(), thickness, options_.presence_threshold(), options_.connection_color(), thickness,
/*normalized=*/true, render_data.get()); /*normalized=*/true, render_data.get());
} }
for (int i = 0; i < landmarks.landmark_size(); ++i) { if (options_.render_landmarks()) {
const NormalizedLandmark& landmark = landmarks.landmark(i); for (int i = 0; i < landmarks.landmark_size(); ++i) {
const NormalizedLandmark& landmark = landmarks.landmark(i);
if (!IsLandmarkVisibleAndPresent<NormalizedLandmark>( if (!IsLandmarkVisibleAndPresent<NormalizedLandmark>(
landmark, options_.utilize_visibility(), landmark, options_.utilize_visibility(),
options_.visibility_threshold(), options_.utilize_presence(), options_.visibility_threshold(), options_.utilize_presence(),
options_.presence_threshold())) { options_.presence_threshold())) {
continue; continue;
} }
auto* landmark_data_render = AddPointRenderData( auto* landmark_data_render = AddPointRenderData(
options_.landmark_color(), thickness, render_data.get()); options_.landmark_color(), thickness, render_data.get());
if (visualize_depth) { if (visualize_depth) {
SetColorSizeValueFromZ(landmark.z(), z_min, z_max, landmark_data_render, SetColorSizeValueFromZ(landmark.z(), z_min, z_max,
options_.min_depth_circle_thickness(), landmark_data_render,
options_.max_depth_circle_thickness()); options_.min_depth_circle_thickness(),
options_.max_depth_circle_thickness());
}
auto* landmark_data = landmark_data_render->mutable_point();
landmark_data->set_normalized(true);
landmark_data->set_x(landmark.x());
landmark_data->set_y(landmark.y());
} }
auto* landmark_data = landmark_data_render->mutable_point();
landmark_data->set_normalized(true);
landmark_data->set_x(landmark.x());
landmark_data->set_y(landmark.y());
} }
} }
@@ -32,6 +32,10 @@ message LandmarksToRenderDataCalculatorOptions {
// Color of the landmarks. // Color of the landmarks.
optional Color landmark_color = 2; optional Color landmark_color = 2;
// Whether to render landmarks as points.
optional bool render_landmarks = 14 [default = true];
// Color of the connections. // Color of the connections.
optional Color connection_color = 3; optional Color connection_color = 3;
-4
View File
@@ -130,7 +130,6 @@ cc_library(
"//mediapipe/framework/formats:video_stream_header", "//mediapipe/framework/formats:video_stream_header",
"//mediapipe/framework/port:opencv_imgproc", "//mediapipe/framework/port:opencv_imgproc",
"//mediapipe/framework/port:opencv_video", "//mediapipe/framework/port:opencv_video",
"//mediapipe/framework/port:ret_check",
"//mediapipe/framework/port:status", "//mediapipe/framework/port:status",
"//mediapipe/framework/tool:status_util", "//mediapipe/framework/tool:status_util",
], ],
@@ -341,7 +340,6 @@ cc_test(
"//mediapipe/framework/port:opencv_core", "//mediapipe/framework/port:opencv_core",
"//mediapipe/framework/port:parse_text_proto", "//mediapipe/framework/port:parse_text_proto",
"//mediapipe/framework/tool:test_util", "//mediapipe/framework/tool:test_util",
"@com_google_absl//absl/flags:flag",
], ],
) )
@@ -367,7 +365,6 @@ cc_test(
"//mediapipe/framework/port:opencv_video", "//mediapipe/framework/port:opencv_video",
"//mediapipe/framework/port:parse_text_proto", "//mediapipe/framework/port:parse_text_proto",
"//mediapipe/framework/tool:test_util", "//mediapipe/framework/tool:test_util",
"@com_google_absl//absl/flags:flag",
], ],
) )
@@ -451,7 +448,6 @@ cc_test(
"//mediapipe/framework/tool:test_util", "//mediapipe/framework/tool:test_util",
"//mediapipe/util/tracking:box_tracker_cc_proto", "//mediapipe/util/tracking:box_tracker_cc_proto",
"//mediapipe/util/tracking:tracking_cc_proto", "//mediapipe/util/tracking:tracking_cc_proto",
"@com_google_absl//absl/flags:flag",
], ],
) )
@@ -1,6 +1,6 @@
distributionBase=GRADLE_USER_HOME distributionBase=GRADLE_USER_HOME
distributionPath=wrapper/dists distributionPath=wrapper/dists
distributionUrl=https\://services.gradle.org/distributions/gradle-7.6.1-bin.zip distributionUrl=https\://services.gradle.org/distributions/gradle-7.6.2-bin.zip
networkTimeout=10000 networkTimeout=10000
zipStoreBase=GRADLE_USER_HOME zipStoreBase=GRADLE_USER_HOME
zipStorePath=wrapper/dists zipStorePath=wrapper/dists
+3
View File
@@ -44,6 +44,9 @@ bzl_library(
"encode_binary_proto.bzl", "encode_binary_proto.bzl",
], ],
visibility = ["//visibility:public"], visibility = ["//visibility:public"],
deps = [
"@bazel_skylib//lib:paths",
],
) )
alias( alias(
+19 -73
View File
@@ -64,57 +64,13 @@ class CalculatorBaseFactoryFor<
namespace api2 { namespace api2 {
namespace internal { namespace internal {
// Defining a member of this type causes P to be ODR-used, which forces its MEDIAPIPE_STATIC_REGISTRATOR_TEMPLATE(
// instantiation if it's a static member of a template. NodeRegistrator, mediapipe::CalculatorBaseRegistry, T::kCalculatorName,
// Previously we depended on the pointer's value to determine whether the size absl::make_unique<mediapipe::internal::CalculatorBaseFactoryFor<T>>)
// of a character array is 0 or 1, forcing it to be instantiated so the
// compiler can determine the object's layout. But using it as a template
// argument is more compact.
template <auto* P>
struct ForceStaticInstantiation {
#ifdef _MSC_VER
// Just having it as the template argument does not count as a use for
// MSVC.
static constexpr bool Use() { return P != nullptr; }
char force_static[Use()];
#endif // _MSC_VER
};
// Helper template for forcing the definition of a static registration token. MEDIAPIPE_STATIC_REGISTRATOR_TEMPLATE(SubgraphRegistrator,
template <typename T> mediapipe::SubgraphRegistry,
struct NodeRegistrationStatic { T::kCalculatorName, absl::make_unique<T>)
static NoDestructor<mediapipe::RegistrationToken> registration;
static mediapipe::RegistrationToken Make() {
return mediapipe::CalculatorBaseRegistry::Register(
T::kCalculatorName,
absl::make_unique<mediapipe::internal::CalculatorBaseFactoryFor<T>>);
}
using RequireStatics = ForceStaticInstantiation<&registration>;
};
// Static members of template classes can be defined in the header.
template <typename T>
NoDestructor<mediapipe::RegistrationToken>
NodeRegistrationStatic<T>::registration(NodeRegistrationStatic<T>::Make());
template <typename T>
struct SubgraphRegistrationImpl {
static NoDestructor<mediapipe::RegistrationToken> registration;
static mediapipe::RegistrationToken Make() {
return mediapipe::SubgraphRegistry::Register(T::kCalculatorName,
absl::make_unique<T>);
}
using RequireStatics = ForceStaticInstantiation<&registration>;
};
template <typename T>
NoDestructor<mediapipe::RegistrationToken>
SubgraphRegistrationImpl<T>::registration(
SubgraphRegistrationImpl<T>::Make());
} // namespace internal } // namespace internal
@@ -127,14 +83,7 @@ template <class Impl = void>
class RegisteredNode; class RegisteredNode;
template <class Impl> template <class Impl>
class RegisteredNode : public Node { class RegisteredNode : public Node, private internal::NodeRegistrator<Impl> {};
private:
// The member below triggers instantiation of the registration static.
// Note that the constructor of calculator subclasses is only invoked through
// the registration token, and so we cannot simply use the static in the
// constructor.
typename internal::NodeRegistrationStatic<Impl>::RequireStatics register_;
};
// No-op version for backwards compatibility. // No-op version for backwards compatibility.
template <> template <>
@@ -216,30 +165,27 @@ class NodeImpl : public RegisteredNode<Impl>, public Intf {
// TODO: verify that the subgraph config fully implements the // TODO: verify that the subgraph config fully implements the
// declared interface. // declared interface.
template <class Intf, class Impl> template <class Intf, class Impl>
class SubgraphImpl : public Subgraph, public Intf { class SubgraphImpl : public Subgraph,
private: public Intf,
typename internal::SubgraphRegistrationImpl<Impl>::RequireStatics register_; private internal::SubgraphRegistrator<Impl> {};
};
// This macro is used to register a calculator that does not use automatic // This macro is used to register a calculator that does not use automatic
// registration. Deprecated. // registration. Deprecated.
#define MEDIAPIPE_NODE_IMPLEMENTATION(Impl) \ #define MEDIAPIPE_NODE_IMPLEMENTATION(Impl) \
static mediapipe::NoDestructor<mediapipe::RegistrationToken> \ MEDIAPIPE_REGISTER_FACTORY_FUNCTION_QUALIFIED( \
REGISTRY_STATIC_VAR(calculator_registration, \ mediapipe::CalculatorBaseRegistry, calculator_registration, \
__LINE__)(mediapipe::CalculatorBaseRegistry::Register( \ Impl::kCalculatorName, \
Impl::kCalculatorName, \ absl::make_unique<mediapipe::internal::CalculatorBaseFactoryFor<Impl>>)
absl::make_unique<mediapipe::internal::CalculatorBaseFactoryFor<Impl>>))
// This macro is used to register a non-split-contract calculator. Deprecated. // This macro is used to register a non-split-contract calculator. Deprecated.
#define MEDIAPIPE_REGISTER_NODE(name) REGISTER_CALCULATOR(name) #define MEDIAPIPE_REGISTER_NODE(name) REGISTER_CALCULATOR(name)
// This macro is used to define a subgraph that does not use automatic // This macro is used to define a subgraph that does not use automatic
// registration. Deprecated. // registration. Deprecated.
#define MEDIAPIPE_SUBGRAPH_IMPLEMENTATION(Impl) \ #define MEDIAPIPE_SUBGRAPH_IMPLEMENTATION(Impl) \
static mediapipe::NoDestructor<mediapipe::RegistrationToken> \ MEDIAPIPE_REGISTER_FACTORY_FUNCTION_QUALIFIED( \
REGISTRY_STATIC_VAR(subgraph_registration, \ mediapipe::SubgraphRegistry, subgraph_registration, \
__LINE__)(mediapipe::SubgraphRegistry::Register( \ Impl::kCalculatorName, absl::make_unique<Impl>)
Impl::kCalculatorName, absl::make_unique<Impl>))
} // namespace api2 } // namespace api2
} // namespace mediapipe } // namespace mediapipe
+82
View File
@@ -144,6 +144,23 @@ template <typename T>
struct WrapStatusOr<absl::StatusOr<T>> { struct WrapStatusOr<absl::StatusOr<T>> {
using type = absl::StatusOr<T>; using type = absl::StatusOr<T>;
}; };
// Defining a member of this type causes P to be ODR-used, which forces its
// instantiation if it's a static member of a template.
// Previously we depended on the pointer's value to determine whether the size
// of a character array is 0 or 1, forcing it to be instantiated so the
// compiler can determine the object's layout. But using it as a template
// argument is more compact.
template <auto* P>
struct ForceStaticInstantiation {
#ifdef _MSC_VER
// Just having it as the template argument does not count as a use for
// MSVC.
static constexpr bool Use() { return P != nullptr; }
char force_static[Use()];
#endif // _MSC_VER
};
} // namespace registration_internal } // namespace registration_internal
class NamespaceAllowlist { class NamespaceAllowlist {
@@ -396,11 +413,76 @@ class GlobalFactoryRegistry {
new mediapipe::RegistrationToken( \ new mediapipe::RegistrationToken( \
RegistryType::Register(#name, __VA_ARGS__)) RegistryType::Register(#name, __VA_ARGS__))
#define MEDIAPIPE_REGISTER_FACTORY_FUNCTION_QUALIFIED(RegistryType, var_name, \
name, ...) \
static auto* REGISTRY_STATIC_VAR(var_name, __LINE__) = \
new mediapipe::RegistrationToken( \
RegistryType::Register(name, __VA_ARGS__))
// TODO: migrate to the above.
#define REGISTER_FACTORY_FUNCTION_QUALIFIED(RegistryType, var_name, name, ...) \ #define REGISTER_FACTORY_FUNCTION_QUALIFIED(RegistryType, var_name, name, ...) \
static auto* REGISTRY_STATIC_VAR(var_name, __LINE__) = \ static auto* REGISTRY_STATIC_VAR(var_name, __LINE__) = \
new mediapipe::RegistrationToken( \ new mediapipe::RegistrationToken( \
RegistryType::Register(#name, __VA_ARGS__)) RegistryType::Register(#name, __VA_ARGS__))
// Defines a utility registrator class which can be used to automatically
// register factory functions.
//
// Example:
// === Defining a registry ================================================
//
// class Component {};
//
// using ComponentRegistry = GlobalFactoryRegistry<std::unique_ptr<Component>>;
//
// === Defining a registrator =============================================
//
// MEDIAPIPE_STATIC_REGISTRATOR_TEMPLATE(ComponentRegistrator,
// ComponentRegistry, T::kName,
// absl::make_unique<T>);
//
// === Defining and registering a new component. ==========================
//
// class MyComponent : public Component,
// private ComponentRegistrator<MyComponent> {
// public:
// static constexpr char kName[] = "MyComponent";
// ...
// };
//
// NOTE:
// - MyComponent is automatically registered in ComponentRegistry by
// "MyComponent" name.
// - Every component is require to provide its name (T::kName here.)
#define MEDIAPIPE_STATIC_REGISTRATOR_TEMPLATE(RegistratorName, RegistryType, \
name, ...) \
template <typename T> \
struct Internal##RegistratorName { \
static NoDestructor<mediapipe::RegistrationToken> registration; \
\
static mediapipe::RegistrationToken Make() { \
return RegistryType::Register(name, __VA_ARGS__); \
} \
\
using RequireStatics = \
registration_internal::ForceStaticInstantiation<&registration>; \
}; \
/* Static members of template classes can be defined in the header. */ \
template <typename T> \
NoDestructor<mediapipe::RegistrationToken> \
Internal##RegistratorName<T>::registration( \
Internal##RegistratorName<T>::Make()); \
\
template <typename T> \
class RegistratorName { \
private: \
/* The member below triggers instantiation of the registration static. */ \
/* Note that the constructor of calculator subclasses is only invoked */ \
/* through the registration token, and so we cannot simply use the */ \
/* static in theconstructor. */ \
typename Internal##RegistratorName<T>::RequireStatics register_; \
};
} // namespace mediapipe } // namespace mediapipe
#endif // MEDIAPIPE_DEPS_REGISTRATION_H_ #endif // MEDIAPIPE_DEPS_REGISTRATION_H_
+46 -31
View File
@@ -37,29 +37,33 @@ Args:
output: The desired name of the output file. Optional. output: The desired name of the output file. Optional.
""" """
load("@bazel_skylib//lib:paths.bzl", "paths")
PROTOC = "@com_google_protobuf//:protoc" PROTOC = "@com_google_protobuf//:protoc"
def _canonicalize_proto_path_oss(all_protos, genfile_path): def _canonicalize_proto_path_oss(f):
"""For the protos from external repository, canonicalize the proto path and the file name. if not f.root.path:
return struct(
proto_path = ".",
file_name = f.short_path,
)
Returns: # `f.path` looks like "<genfiles>/external/<repo>/(_virtual_imports/<library>/)?<file_name>"
Proto path list and proto source file list. repo_name, _, file_name = f.path[len(paths.join(f.root.path, "external") + "/"):].partition("/")
""" if file_name.startswith("_virtual_imports/"):
proto_paths = [] # This is a virtual import; move "_virtual_imports/<library>" from `repo_name` to `file_name`.
proto_file_names = [] repo_name = paths.join(repo_name, *file_name.split("/", 2)[:2])
for s in all_protos.to_list(): file_name = file_name.split("/", 2)[-1]
if s.path.startswith(genfile_path): return struct(
repo_name, _, file_name = s.path[len(genfile_path + "/external/"):].partition("/") proto_path = paths.join(f.root.path, "external", repo_name),
file_name = file_name,
)
# handle virtual imports def _map_root_path(f):
if file_name.startswith("_virtual_imports"): return _canonicalize_proto_path_oss(f).proto_path
repo_name = repo_name + "/" + "/".join(file_name.split("/", 2)[:2])
file_name = file_name.split("/", 2)[-1] def _map_short_path(f):
proto_paths.append(genfile_path + "/external/" + repo_name) return _canonicalize_proto_path_oss(f).file_name
proto_file_names.append(file_name)
else:
proto_file_names.append(s.path)
return ([" --proto_path=" + path for path in proto_paths], proto_file_names)
def _get_proto_provider(dep): def _get_proto_provider(dep):
"""Get the provider for protocol buffers from a dependnecy. """Get the provider for protocol buffers from a dependnecy.
@@ -90,24 +94,35 @@ def _encode_binary_proto_impl(ctx):
sibling = textpb, sibling = textpb,
) )
path_list, file_list = _canonicalize_proto_path_oss(all_protos, ctx.genfiles_dir.path) args = ctx.actions.args()
args.add(textpb)
args.add(binarypb)
args.add(ctx.executable._proto_compiler)
args.add(ctx.attr.message_type, format = "--encode=%s")
args.add("--proto_path=.")
args.add_all(
all_protos,
map_each = _map_root_path,
format_each = "--proto_path=%s",
uniquify = True,
)
args.add_all(
all_protos,
map_each = _map_short_path,
uniquify = True,
)
# Note: the combination of absolute_paths and proto_path, as well as the exact # Note: the combination of absolute_paths and proto_path, as well as the exact
# order of gendir before ., is needed for the proto compiler to resolve # order of gendir before ., is needed for the proto compiler to resolve
# import statements that reference proto files produced by a genrule. # import statements that reference proto files produced by a genrule.
ctx.actions.run_shell( ctx.actions.run_shell(
tools = all_protos.to_list() + [textpb, ctx.executable._proto_compiler], tools = depset(
outputs = [binarypb], direct = [textpb, ctx.executable._proto_compiler],
command = " ".join( transitive = [all_protos],
[
ctx.executable._proto_compiler.path,
"--encode=" + ctx.attr.message_type,
"--proto_path=" + ctx.genfiles_dir.path,
"--proto_path=" + ctx.bin_dir.path,
"--proto_path=.",
] + path_list + file_list +
["<", textpb.path, ">", binarypb.path],
), ),
outputs = [binarypb],
command = "${@:3} < $1 > $2",
arguments = [args],
mnemonic = "EncodeProto", mnemonic = "EncodeProto",
) )
+11 -2
View File
@@ -261,8 +261,8 @@ cc_library(
) )
cc_library( cc_library(
name = "opencv_highgui", name = "opencv_photo",
hdrs = ["opencv_highgui_inc.h"], hdrs = ["opencv_photo_inc.h"],
deps = [ deps = [
":opencv_core", ":opencv_core",
"//third_party:opencv", "//third_party:opencv",
@@ -297,6 +297,15 @@ cc_library(
], ],
) )
cc_library(
name = "opencv_highgui",
hdrs = ["opencv_highgui_inc.h"],
deps = [
":opencv_core",
"//third_party:opencv",
],
)
cc_library( cc_library(
name = "opencv_videoio", name = "opencv_videoio",
hdrs = ["opencv_videoio_inc.h"], hdrs = ["opencv_videoio_inc.h"],
@@ -1,4 +1,4 @@
// Copyright 2019 The MediaPipe Authors. // Copyright 2023 The MediaPipe Authors.
// //
// Licensed under the Apache License, Version 2.0 (the "License"); // Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License. // you may not use this file except in compliance with the License.
@@ -12,8 +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.
#ifndef MEDIAPIPE_PORT_OPENCV_HIGHGUI_INC_H_ #ifndef MEDIAPIPE_FRAMEWORK_PORT_OPENCV_HIGHGUI_INC_H_
#define MEDIAPIPE_PORT_OPENCV_HIGHGUI_INC_H_ #define MEDIAPIPE_FRAMEWORK_PORT_OPENCV_HIGHGUI_INC_H_
#include <opencv2/core/version.hpp> #include <opencv2/core/version.hpp>
@@ -25,4 +25,4 @@
#include <opencv2/highgui.hpp> #include <opencv2/highgui.hpp>
#endif #endif
#endif // MEDIAPIPE_PORT_OPENCV_HIGHGUI_INC_H_ #endif // MEDIAPIPE_FRAMEWORK_PORT_OPENCV_HIGHGUI_INC_H_
@@ -1,4 +1,4 @@
// Copyright 2019 The MediaPipe Authors. // Copyright 2022 The MediaPipe Authors.
// //
// Licensed under the Apache License, Version 2.0 (the "License"); // Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License. // you may not use this file except in compliance with the License.
+1 -1
View File
@@ -1,4 +1,4 @@
// Copyright 2019 The MediaPipe Authors. // Copyright 2022 The MediaPipe Authors.
// //
// Licensed under the Apache License, Version 2.0 (the "License"); // Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License. // you may not use this file except in compliance with the License.
@@ -48,6 +48,18 @@ class MuxInputStreamHandler : public InputStreamHandler {
: InputStreamHandler(std::move(tag_map), cc_manager, options, : InputStreamHandler(std::move(tag_map), cc_manager, options,
calculator_run_in_parallel) {} calculator_run_in_parallel) {}
private:
CollectionItemId GetControlStreamId() const {
return input_stream_managers_.EndId() - 1;
}
void RemoveOutdatedDataPackets(Timestamp timestamp) {
const CollectionItemId control_stream_id = GetControlStreamId();
for (CollectionItemId id = input_stream_managers_.BeginId();
id < control_stream_id; ++id) {
input_stream_managers_.Get(id)->ErasePacketsEarlierThan(timestamp);
}
}
protected: protected:
// In MuxInputStreamHandler, a node is "ready" if: // In MuxInputStreamHandler, a node is "ready" if:
// - the control stream is done (need to call Close() in this case), or // - the control stream is done (need to call Close() in this case), or
@@ -58,9 +70,15 @@ class MuxInputStreamHandler : public InputStreamHandler {
absl::MutexLock lock(&input_streams_mutex_); absl::MutexLock lock(&input_streams_mutex_);
const auto& control_stream = const auto& control_stream =
input_stream_managers_.Get(input_stream_managers_.EndId() - 1); input_stream_managers_.Get(GetControlStreamId());
bool empty; bool empty;
*min_stream_timestamp = control_stream->MinTimestampOrBound(&empty); *min_stream_timestamp = control_stream->MinTimestampOrBound(&empty);
// Data streams may contain some outdated packets which failed to be popped
// out during "FillInputSet". (This handler doesn't sync input streams,
// hence "FillInputSet" can be triggerred before every input stream is
// filled with packets corresponding to the same timestamp.)
RemoveOutdatedDataPackets(*min_stream_timestamp);
if (empty) { if (empty) {
if (*min_stream_timestamp == Timestamp::Done()) { if (*min_stream_timestamp == Timestamp::Done()) {
// Calculator is done if the control input stream is done. // Calculator is done if the control input stream is done.
@@ -78,11 +96,6 @@ class MuxInputStreamHandler : public InputStreamHandler {
const auto& data_stream = input_stream_managers_.Get( const auto& data_stream = input_stream_managers_.Get(
input_stream_managers_.BeginId() + control_value); input_stream_managers_.BeginId() + control_value);
// Data stream may contain some outdated packets which failed to be popped
// out during "FillInputSet". (This handler doesn't sync input streams,
// hence "FillInputSet" can be triggerred before every input stream is
// filled with packets corresponding to the same timestamp.)
data_stream->ErasePacketsEarlierThan(*min_stream_timestamp);
Timestamp stream_timestamp = data_stream->MinTimestampOrBound(&empty); Timestamp stream_timestamp = data_stream->MinTimestampOrBound(&empty);
if (empty) { if (empty) {
if (stream_timestamp <= *min_stream_timestamp) { if (stream_timestamp <= *min_stream_timestamp) {
@@ -111,8 +124,7 @@ class MuxInputStreamHandler : public InputStreamHandler {
CHECK(input_set); CHECK(input_set);
absl::MutexLock lock(&input_streams_mutex_); absl::MutexLock lock(&input_streams_mutex_);
const CollectionItemId control_stream_id = const CollectionItemId control_stream_id = GetControlStreamId();
input_stream_managers_.EndId() - 1;
auto& control_stream = input_stream_managers_.Get(control_stream_id); auto& control_stream = input_stream_managers_.Get(control_stream_id);
int num_packets_dropped = 0; int num_packets_dropped = 0;
bool stream_is_done = false; bool stream_is_done = false;
@@ -140,15 +152,8 @@ class MuxInputStreamHandler : public InputStreamHandler {
AddPacketToShard(&input_set->Get(data_stream_id), std::move(data_packet), AddPacketToShard(&input_set->Get(data_stream_id), std::move(data_packet),
stream_is_done); stream_is_done);
// Discard old packets on other streams. // Discard old packets on data streams.
// Note that control_stream_id is the last valid id. RemoveOutdatedDataPackets(input_timestamp.NextAllowedInStream());
auto next_timestamp = input_timestamp.NextAllowedInStream();
for (CollectionItemId id = input_stream_managers_.BeginId();
id < control_stream_id; ++id) {
if (id == data_stream_id) continue;
auto& other_stream = input_stream_managers_.Get(id);
other_stream->ErasePacketsEarlierThan(next_timestamp);
}
} }
private: private:
@@ -645,5 +645,41 @@ TEST(MuxInputStreamHandlerTest,
MP_ASSERT_OK(graph.WaitUntilDone()); MP_ASSERT_OK(graph.WaitUntilDone());
} }
TEST(MuxInputStreamHandlerTest, RemovesUnusedDataStreamPackets) {
CalculatorGraphConfig config =
mediapipe::ParseTextProtoOrDie<CalculatorGraphConfig>(R"pb(
input_stream: "input0"
input_stream: "input1"
input_stream: "select"
node {
calculator: "MuxCalculator"
input_stream: "INPUT:0:input0"
input_stream: "INPUT:1:input1"
input_stream: "SELECT:select"
output_stream: "OUTPUT:output"
input_stream_handler { input_stream_handler: "MuxInputStreamHandler" }
}
)pb");
config.set_max_queue_size(1);
config.set_report_deadlock(true);
CalculatorGraph graph;
MP_ASSERT_OK(graph.Initialize(config));
MP_ASSERT_OK(graph.StartRun({}));
MP_ASSERT_OK(graph.AddPacketToInputStream(
"select", MakePacket<int>(0).At(Timestamp(2))));
MP_ASSERT_OK(graph.AddPacketToInputStream(
"input0", MakePacket<int>(1000).At(Timestamp(2))));
MP_ASSERT_OK(graph.WaitUntilIdle());
// Add two delayed packets to the deselected input. They should be discarded
// instead of triggering the deadlock detection (max_queue_size = 1).
MP_ASSERT_OK(graph.AddPacketToInputStream(
"input1", MakePacket<int>(900).At(Timestamp(1))));
MP_ASSERT_OK(graph.AddPacketToInputStream(
"input1", MakePacket<int>(900).At(Timestamp(2))));
MP_ASSERT_OK(graph.WaitUntilIdle());
}
} // namespace } // namespace
} // namespace mediapipe } // namespace mediapipe
+2 -3
View File
@@ -109,9 +109,8 @@ absl::Status GlContext::CreateContext(
} }
MP_RETURN_IF_ERROR(status); MP_RETURN_IF_ERROR(status);
LOG(INFO) << "Successfully created a WebGL context with major version " VLOG(1) << "Successfully created a WebGL context with major version "
<< gl_major_version_ << " and handle " << context_; << gl_major_version_ << " and handle " << context_;
return absl::OkStatus(); return absl::OkStatus();
} }
+8 -1
View File
@@ -104,6 +104,7 @@ class GlScalerCalculator : public CalculatorBase {
bool vertical_flip_output_; bool vertical_flip_output_;
bool horizontal_flip_output_; bool horizontal_flip_output_;
FrameScaleMode scale_mode_ = FrameScaleMode::kStretch; FrameScaleMode scale_mode_ = FrameScaleMode::kStretch;
bool use_nearest_neighbor_interpolation_ = false;
}; };
REGISTER_CALCULATOR(GlScalerCalculator); REGISTER_CALCULATOR(GlScalerCalculator);
@@ -186,7 +187,8 @@ absl::Status GlScalerCalculator::Open(CalculatorContext* cc) {
scale_mode_ = scale_mode_ =
FrameScaleModeFromProto(options.scale_mode(), FrameScaleMode::kStretch); FrameScaleModeFromProto(options.scale_mode(), FrameScaleMode::kStretch);
} }
use_nearest_neighbor_interpolation_ =
options.use_nearest_neighbor_interpolation();
if (HasTagOrIndex(cc->InputSidePackets(), "OUTPUT_DIMENSIONS", 1)) { if (HasTagOrIndex(cc->InputSidePackets(), "OUTPUT_DIMENSIONS", 1)) {
const auto& dimensions = const auto& dimensions =
TagOrIndex(cc->InputSidePackets(), "OUTPUT_DIMENSIONS", 1) TagOrIndex(cc->InputSidePackets(), "OUTPUT_DIMENSIONS", 1)
@@ -297,6 +299,11 @@ absl::Status GlScalerCalculator::Process(CalculatorContext* cc) {
glBindTexture(src2.target(), src2.name()); glBindTexture(src2.target(), src2.name());
} }
if (use_nearest_neighbor_interpolation_) {
glTexParameteri(GL_TEXTURE_2D, GL_TEXTURE_MAG_FILTER, GL_NEAREST);
glTexParameteri(GL_TEXTURE_2D, GL_TEXTURE_MIN_FILTER, GL_NEAREST);
}
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_, horizontal_flip_output_, vertical_flip_output_, rotation_, horizontal_flip_output_, vertical_flip_output_,
+4 -1
View File
@@ -19,7 +19,7 @@ package mediapipe;
import "mediapipe/framework/calculator.proto"; import "mediapipe/framework/calculator.proto";
import "mediapipe/gpu/scale_mode.proto"; import "mediapipe/gpu/scale_mode.proto";
// Next id: 8. // Next id: 9.
message GlScalerCalculatorOptions { message GlScalerCalculatorOptions {
extend CalculatorOptions { extend CalculatorOptions {
optional GlScalerCalculatorOptions ext = 166373014; optional GlScalerCalculatorOptions ext = 166373014;
@@ -39,4 +39,7 @@ message GlScalerCalculatorOptions {
// Flip the output texture horizontally. This is applied after rotation. // Flip the output texture horizontally. This is applied after rotation.
optional bool flip_horizontal = 5; optional bool flip_horizontal = 5;
optional ScaleMode.Mode scale_mode = 6; optional ScaleMode.Mode scale_mode = 6;
// Whether to use nearest neighbor interpolation. Default to use linear
// interpolation.
optional bool use_nearest_neighbor_interpolation = 8 [default = false];
} }
+5
View File
@@ -100,6 +100,10 @@ const GlTextureInfo& GlTextureInfoForGpuBufferFormat(GpuBufferFormat format,
{GL_R8, GL_RED, GL_UNSIGNED_BYTE, 1}, {GL_R8, GL_RED, GL_UNSIGNED_BYTE, 1},
#endif // TARGET_OS_OSX #endif // TARGET_OS_OSX
}}, }},
{GpuBufferFormat::kOneComponent8Alpha,
{
{GL_ALPHA, GL_ALPHA, GL_UNSIGNED_BYTE, 1},
}},
{GpuBufferFormat::kOneComponent8Red, {GpuBufferFormat::kOneComponent8Red,
{ {
{GL_R8, GL_RED, GL_UNSIGNED_BYTE, 1}, {GL_R8, GL_RED, GL_UNSIGNED_BYTE, 1},
@@ -221,6 +225,7 @@ ImageFormat::Format ImageFormatForGpuBufferFormat(GpuBufferFormat format) {
case GpuBufferFormat::kRGBA32: case GpuBufferFormat::kRGBA32:
// TODO: this likely maps to ImageFormat::SRGBA // TODO: this likely maps to ImageFormat::SRGBA
case GpuBufferFormat::kGrayHalf16: case GpuBufferFormat::kGrayHalf16:
case GpuBufferFormat::kOneComponent8Alpha:
case GpuBufferFormat::kOneComponent8Red: case GpuBufferFormat::kOneComponent8Red:
case GpuBufferFormat::kTwoComponent8: case GpuBufferFormat::kTwoComponent8:
case GpuBufferFormat::kTwoComponentHalf16: case GpuBufferFormat::kTwoComponentHalf16:
+2
View File
@@ -43,6 +43,7 @@ enum class GpuBufferFormat : uint32_t {
kGrayFloat32 = MEDIAPIPE_FOURCC('L', '0', '0', 'f'), kGrayFloat32 = MEDIAPIPE_FOURCC('L', '0', '0', 'f'),
kGrayHalf16 = MEDIAPIPE_FOURCC('L', '0', '0', 'h'), kGrayHalf16 = MEDIAPIPE_FOURCC('L', '0', '0', 'h'),
kOneComponent8 = MEDIAPIPE_FOURCC('L', '0', '0', '8'), kOneComponent8 = MEDIAPIPE_FOURCC('L', '0', '0', '8'),
kOneComponent8Alpha = MEDIAPIPE_FOURCC('A', '0', '0', '8'),
kOneComponent8Red = MEDIAPIPE_FOURCC('R', '0', '0', '8'), kOneComponent8Red = MEDIAPIPE_FOURCC('R', '0', '0', '8'),
kTwoComponent8 = MEDIAPIPE_FOURCC('2', 'C', '0', '8'), kTwoComponent8 = MEDIAPIPE_FOURCC('2', 'C', '0', '8'),
kTwoComponentHalf16 = MEDIAPIPE_FOURCC('2', 'C', '0', 'h'), kTwoComponentHalf16 = MEDIAPIPE_FOURCC('2', 'C', '0', 'h'),
@@ -101,6 +102,7 @@ inline OSType CVPixelFormatForGpuBufferFormat(GpuBufferFormat format) {
return kCVPixelFormatType_OneComponent32Float; return kCVPixelFormatType_OneComponent32Float;
case GpuBufferFormat::kOneComponent8: case GpuBufferFormat::kOneComponent8:
return kCVPixelFormatType_OneComponent8; return kCVPixelFormatType_OneComponent8;
case GpuBufferFormat::kOneComponent8Alpha:
case GpuBufferFormat::kOneComponent8Red: case GpuBufferFormat::kOneComponent8Red:
return -1; return -1;
case GpuBufferFormat::kTwoComponent8: case GpuBufferFormat::kTwoComponent8:
@@ -78,17 +78,21 @@ public class AppTextureFrame implements TextureFrame {
* Use {@link waitUntilReleasedWithGpuSync} whenever possible. * Use {@link waitUntilReleasedWithGpuSync} whenever possible.
*/ */
public void waitUntilReleased() throws InterruptedException { public void waitUntilReleased() throws InterruptedException {
GlSyncToken tokenToRelease = null;
synchronized (this) { synchronized (this) {
while (inUse && releaseSyncToken == null) { while (inUse && releaseSyncToken == null) {
wait(); wait();
} }
if (releaseSyncToken != null) { if (releaseSyncToken != null) {
releaseSyncToken.waitOnCpu(); tokenToRelease = releaseSyncToken;
releaseSyncToken.release();
inUse = false; inUse = false;
releaseSyncToken = null; releaseSyncToken = null;
} }
} }
if (tokenToRelease != null) {
tokenToRelease.waitOnCpu();
tokenToRelease.release();
}
} }
/** /**
@@ -98,17 +102,21 @@ public class AppTextureFrame implements TextureFrame {
* TextureFrame. * TextureFrame.
*/ */
public void waitUntilReleasedWithGpuSync() throws InterruptedException { public void waitUntilReleasedWithGpuSync() throws InterruptedException {
GlSyncToken tokenToRelease = null;
synchronized (this) { synchronized (this) {
while (inUse && releaseSyncToken == null) { while (inUse && releaseSyncToken == null) {
wait(); wait();
} }
if (releaseSyncToken != null) { if (releaseSyncToken != null) {
releaseSyncToken.waitOnGpu(); tokenToRelease = releaseSyncToken;
releaseSyncToken.release();
inUse = false; inUse = false;
releaseSyncToken = null; releaseSyncToken = null;
} }
} }
if (tokenToRelease != null) {
tokenToRelease.waitOnGpu();
tokenToRelease.release();
}
} }
/** /**
@@ -239,7 +239,7 @@ public final class PacketGetter {
/** /**
* Assign the native image buffer array in given ByteBuffer array. It assumes given ByteBuffer * Assign the native image buffer array in given ByteBuffer array. It assumes given ByteBuffer
* array has the the same size of image list packet, and assumes the output buffer stores pixels * array has the same size of image list packet, and assumes the output buffer stores pixels
* contiguously. It returns false if this assumption does not hold. * contiguously. It returns false if this assumption does not hold.
* *
* <p>If deepCopy is true, it assumes the given buffersArray has allocated the required size of * <p>If deepCopy is true, it assumes the given buffersArray has allocated the required size of
+1
View File
@@ -24,6 +24,7 @@ package_group(
package_group( package_group(
name = "1p_client", name = "1p_client",
packages = [ packages = [
"//cloud/ml/applications/vision/model_garden/model_oss/mediapipe/...",
"//research/privacy/learning/fl_eval/pcvr/...", "//research/privacy/learning/fl_eval/pcvr/...",
], ],
) )
@@ -57,3 +57,14 @@ py_test(
srcs = ["classification_dataset_test.py"], srcs = ["classification_dataset_test.py"],
deps = [":classification_dataset"], deps = [":classification_dataset"],
) )
py_library(
name = "cache_files",
srcs = ["cache_files.py"],
)
py_test(
name = "cache_files_test",
srcs = ["cache_files_test.py"],
deps = [":cache_files"],
)
@@ -0,0 +1,112 @@
# Copyright 2023 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.
"""Common TFRecord cache files library."""
import dataclasses
import os
import tempfile
from typing import Any, Mapping, Sequence
import tensorflow as tf
import yaml
# Suffix of the meta data file name.
METADATA_FILE_SUFFIX = '_metadata.yaml'
@dataclasses.dataclass(frozen=True)
class TFRecordCacheFiles:
"""TFRecordCacheFiles dataclass to store and load cached TFRecord files.
Attributes:
cache_prefix_filename: The cache prefix filename. This is usually provided
as a hash of the original data source to avoid different data sources
resulting in the same cache file.
cache_dir: The cache directory to save TFRecord and metadata file. When
cache_dir is None, a temporary folder will be created and will not be
removed automatically after training which makes it can be used later.
num_shards: Number of shards for output tfrecord files.
"""
cache_prefix_filename: str = 'cache_prefix'
cache_dir: str = dataclasses.field(default_factory=tempfile.mkdtemp)
num_shards: int = 1
def __post_init__(self):
if not tf.io.gfile.exists(self.cache_dir):
tf.io.gfile.makedirs(self.cache_dir)
if not self.cache_prefix_filename:
raise ValueError('cache_prefix_filename cannot be empty.')
if self.num_shards <= 0:
raise ValueError(
f'num_shards must be greater than 0, got {self.num_shards}'
)
@property
def cache_prefix(self) -> str:
"""The cache prefix including the cache directory and the cache prefix filename."""
return os.path.join(self.cache_dir, self.cache_prefix_filename)
@property
def tfrecord_files(self) -> Sequence[str]:
"""The TFRecord files."""
tfrecord_files = [
self.cache_prefix + '-%05d-of-%05d.tfrecord' % (i, self.num_shards)
for i in range(self.num_shards)
]
return tfrecord_files
@property
def metadata_file(self) -> str:
"""The metadata file."""
return self.cache_prefix + METADATA_FILE_SUFFIX
def get_writers(self) -> Sequence[tf.io.TFRecordWriter]:
"""Gets an array of TFRecordWriter objects.
Note that these writers should each be closed using .close() when done.
Returns:
Array of TFRecordWriter objects
"""
return [tf.io.TFRecordWriter(path) for path in self.tfrecord_files]
def save_metadata(self, metadata):
"""Writes metadata to file.
Args:
metadata: A dictionary of metadata content to write. Exact format is
dependent on the specific dataset, but typically includes a 'size' and
'label_names' entry.
"""
with tf.io.gfile.GFile(self.metadata_file, 'w') as f:
yaml.dump(metadata, f)
def load_metadata(self) -> Mapping[Any, Any]:
"""Reads metadata from file.
Returns:
Dictionary object containing metadata
"""
if not tf.io.gfile.exists(self.metadata_file):
return {}
with tf.io.gfile.GFile(self.metadata_file, 'r') as f:
metadata = yaml.load(f, Loader=yaml.FullLoader)
return metadata
def is_cached(self) -> bool:
"""Checks whether this CacheFiles is already cached."""
all_cached_files = list(self.tfrecord_files) + [self.metadata_file]
return all(tf.io.gfile.exists(f) for f in all_cached_files)
@@ -0,0 +1,77 @@
# Copyright 2023 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.
import tensorflow as tf
from mediapipe.model_maker.python.core.data import cache_files
class CacheFilesTest(tf.test.TestCase):
def test_tfrecord_cache_files(self):
cf = cache_files.TFRecordCacheFiles(
cache_prefix_filename='tfrecord',
cache_dir='/tmp/cache_dir',
num_shards=2,
)
self.assertEqual(cf.cache_prefix, '/tmp/cache_dir/tfrecord')
self.assertEqual(
cf.metadata_file,
'/tmp/cache_dir/tfrecord' + cache_files.METADATA_FILE_SUFFIX,
)
expected_tfrecord_files = [
'/tmp/cache_dir/tfrecord-%05d-of-%05d.tfrecord' % (i, 2)
for i in range(2)
]
self.assertEqual(cf.tfrecord_files, expected_tfrecord_files)
# Writing TFRecord Files
self.assertFalse(cf.is_cached())
for tfrecord_file in cf.tfrecord_files:
self.assertFalse(tf.io.gfile.exists(tfrecord_file))
writers = cf.get_writers()
for writer in writers:
writer.close()
for tfrecord_file in cf.tfrecord_files:
self.assertTrue(tf.io.gfile.exists(tfrecord_file))
self.assertFalse(cf.is_cached())
# Writing Metadata Files
original_metadata = {'size': 10, 'label_names': ['label1', 'label2']}
cf.save_metadata(original_metadata)
self.assertTrue(cf.is_cached())
metadata = cf.load_metadata()
self.assertEqual(metadata, original_metadata)
def test_recordio_cache_files_error(self):
with self.assertRaisesRegex(
ValueError, 'cache_prefix_filename cannot be empty'
):
cache_files.TFRecordCacheFiles(
cache_prefix_filename='',
cache_dir='/tmp/cache_dir',
num_shards=2,
)
with self.assertRaisesRegex(
ValueError, 'num_shards must be greater than 0, got 0'
):
cache_files.TFRecordCacheFiles(
cache_prefix_filename='tfrecord',
cache_dir='/tmp/cache_dir',
num_shards=0,
)
if __name__ == '__main__':
tf.test.main()
@@ -13,7 +13,7 @@
# limitations under the License. # limitations under the License.
"""Common classification dataset library.""" """Common classification dataset library."""
from typing import List, Tuple from typing import List, Optional, Tuple
import tensorflow as tf import tensorflow as tf
@@ -23,8 +23,12 @@ from mediapipe.model_maker.python.core.data import dataset as ds
class ClassificationDataset(ds.Dataset): class ClassificationDataset(ds.Dataset):
"""Dataset Loader for classification models.""" """Dataset Loader for classification models."""
def __init__(self, dataset: tf.data.Dataset, size: int, def __init__(
label_names: List[str]): self,
dataset: tf.data.Dataset,
label_names: List[str],
size: Optional[int] = None,
):
super().__init__(dataset, size) super().__init__(dataset, size)
self._label_names = label_names self._label_names = label_names
@@ -36,9 +36,14 @@ class ClassificationDatasetTest(tf.test.TestCase):
value: A value variable stored by the mock dataset class for testing. value: A value variable stored by the mock dataset class for testing.
""" """
def __init__(self, dataset: tf.data.Dataset, size: int, def __init__(
label_names: List[str], value: Any): self,
super().__init__(dataset=dataset, size=size, label_names=label_names) dataset: tf.data.Dataset,
label_names: List[str],
value: Any,
size: int,
):
super().__init__(dataset=dataset, label_names=label_names, size=size)
self.value = value self.value = value
def split(self, fraction: float) -> Tuple[_DatasetT, _DatasetT]: def split(self, fraction: float) -> Tuple[_DatasetT, _DatasetT]:
@@ -52,7 +57,8 @@ class ClassificationDatasetTest(tf.test.TestCase):
# Create data loader from sample data. # Create data loader from sample data.
ds = tf.data.Dataset.from_tensor_slices([[0, 1], [1, 1], [0, 0], [1, 0]]) ds = tf.data.Dataset.from_tensor_slices([[0, 1], [1, 1], [0, 0], [1, 0]])
data = MagicClassificationDataset( data = MagicClassificationDataset(
dataset=ds, size=len(ds), label_names=label_names, value=magic_value) dataset=ds, label_names=label_names, value=magic_value, size=len(ds)
)
# Train/Test data split. # Train/Test data split.
fraction = .25 fraction = .25
@@ -56,15 +56,14 @@ class Dataset(object):
def size(self) -> Optional[int]: def size(self) -> Optional[int]:
"""Returns the size of the dataset. """Returns the size of the dataset.
Note that this function may return None becuase the exact size of the Same functionality as calling __len__. See the __len__ method definition for
dataset isn't a necessary parameter to create an instance of this class, more information.
and tf.data.Dataset donesn't support a function to get the length directly
since it's lazy-loaded and may be infinite. Raises:
In most cases, however, when an instance of this class is created by helper TypeError if self._size is not set and the cardinality of self._dataset
functions like 'from_folder', the size of the dataset will be preprocessed, is INFINITE_CARDINALITY or UNKNOWN_CARDINALITY.
and this function can return an int representing the size of the dataset.
""" """
return self._size return self.__len__()
def gen_tf_dataset( def gen_tf_dataset(
self, self,
@@ -116,8 +115,22 @@ class Dataset(object):
# here. # here.
return dataset return dataset
def __len__(self): def __len__(self) -> int:
"""Returns the number of element of the dataset.""" """Returns the number of element of the dataset.
If size is not set, this method will fallback to using the __len__ method
of the tf.data.Dataset in self._dataset. Calling __len__ on a
tf.data.Dataset instance may throw a TypeError because the dataset may
be lazy-loaded with an unknown size or have infinite size.
In most cases, however, when an instance of this class is created by helper
functions like 'from_folder', the size of the dataset will be preprocessed,
and the _size instance variable will be already set.
Raises:
TypeError if self._size is not set and the cardinality of self._dataset
is INFINITE_CARDINALITY or UNKNOWN_CARDINALITY.
"""
if self._size is not None: if self._size is not None:
return self._size return self._size
else: else:
@@ -152,15 +165,25 @@ class Dataset(object):
Returns: Returns:
The splitted two sub datasets. The splitted two sub datasets.
Raises:
ValueError: if the provided fraction is not between 0 and 1.
ValueError: if this dataset does not have a set size.
""" """
assert (fraction > 0 and fraction < 1) if not (fraction > 0 and fraction < 1):
raise ValueError(f'Fraction must be between 0 and 1. Got:{fraction}')
if not self._size:
raise ValueError(
'Dataset size unknown. Cannot split the dataset when '
'the size is unknown.'
)
dataset = self._dataset dataset = self._dataset
train_size = int(self._size * fraction) train_size = int(self._size * fraction)
trainset = self.__class__(dataset.take(train_size), train_size, *args) trainset = self.__class__(dataset.take(train_size), *args, size=train_size)
test_size = self._size - train_size test_size = self._size - train_size
testset = self.__class__(dataset.skip(train_size), test_size, *args) testset = self.__class__(dataset.skip(train_size), *args, size=test_size)
return trainset, testset return trainset, testset
@@ -15,7 +15,7 @@
import dataclasses import dataclasses
import tempfile import tempfile
from typing import Optional from typing import Mapping, Optional
import tensorflow as tf import tensorflow as tf
@@ -36,6 +36,8 @@ class BaseHParams:
steps_per_epoch: An optional integer indicate the number of training steps steps_per_epoch: An optional integer indicate the number of training steps
per epoch. If not set, the training pipeline calculates the default steps per epoch. If not set, the training pipeline calculates the default steps
per epoch as the training dataset size divided by batch size. per epoch as the training dataset size divided by batch size.
class_weights: An optional mapping of indices to weights for weighting the
loss function during training.
shuffle: True if the dataset is shuffled before training. shuffle: True if the dataset is shuffled before training.
export_dir: The location of the model checkpoint files. export_dir: The location of the model checkpoint files.
distribution_strategy: A string specifying which Distribution Strategy to distribution_strategy: A string specifying which Distribution Strategy to
@@ -57,6 +59,7 @@ class BaseHParams:
batch_size: int batch_size: int
epochs: int epochs: int
steps_per_epoch: Optional[int] = None steps_per_epoch: Optional[int] = None
class_weights: Optional[Mapping[int, float]] = None
# Dataset-related parameters # Dataset-related parameters
shuffle: bool = False shuffle: bool = False
@@ -110,7 +110,9 @@ class Classifier(custom_model.CustomModel):
# dataset is exhausted even if there are epochs remaining. # dataset is exhausted even if there are epochs remaining.
steps_per_epoch=None, steps_per_epoch=None,
validation_data=validation_dataset, validation_data=validation_dataset,
callbacks=self._callbacks) callbacks=self._callbacks,
class_weight=self._hparams.class_weights,
)
def evaluate(self, data: dataset.Dataset, batch_size: int = 32) -> Any: def evaluate(self, data: dataset.Dataset, batch_size: int = 32) -> Any:
"""Evaluates the classifier with the provided evaluation dataset. """Evaluates the classifier with the provided evaluation dataset.
@@ -59,7 +59,7 @@ class FocalLoss(tf.keras.losses.Loss):
""" """
def __init__(self, gamma, class_weight: Optional[Sequence[float]] = None): def __init__(self, gamma, class_weight: Optional[Sequence[float]] = None):
"""Constructor. """Initializes FocalLoss.
Args: Args:
gamma: Focal loss gamma, as described in class docs. gamma: Focal loss gamma, as described in class docs.
@@ -115,6 +115,51 @@ class FocalLoss(tf.keras.losses.Loss):
return tf.reduce_sum(losses) / batch_size return tf.reduce_sum(losses) / batch_size
class SparseFocalLoss(FocalLoss):
"""Sparse implementation of Focal Loss.
This is the same as FocalLoss, except the labels are expected to be class ids
instead of 1-hot encoded vectors. See FocalLoss class documentation defined
in this same file for more details.
Example usage:
>>> y_true = [1, 2]
>>> y_pred = [[0.05, 0.95, 0], [0.1, 0.8, 0.1]]
>>> gamma = 2
>>> focal_loss = SparseFocalLoss(gamma, 3)
>>> focal_loss(y_true, y_pred).numpy()
0.9326
>>> # Calling with 'sample_weight'.
>>> focal_loss(y_true, y_pred, sample_weight=tf.constant([0.3, 0.7])).numpy()
0.6528
"""
def __init__(
self, gamma, num_classes, class_weight: Optional[Sequence[float]] = None
):
"""Initializes SparseFocalLoss.
Args:
gamma: Focal loss gamma, as described in class docs.
num_classes: Number of classes.
class_weight: A weight to apply to the loss, one for each class. The
weight is applied for each input where the ground truth label matches.
"""
super().__init__(gamma, class_weight=class_weight)
self._num_classes = num_classes
def __call__(
self,
y_true: tf.Tensor,
y_pred: tf.Tensor,
sample_weight: Optional[tf.Tensor] = None,
) -> tf.Tensor:
y_true = tf.cast(tf.reshape(y_true, [-1]), tf.int32)
y_true_one_hot = tf.one_hot(y_true, self._num_classes)
return super().__call__(y_true_one_hot, y_pred, sample_weight=sample_weight)
@dataclasses.dataclass @dataclasses.dataclass
class PerceptualLossWeight: class PerceptualLossWeight:
"""The weight for each perceptual loss. """The weight for each perceptual loss.
@@ -101,6 +101,23 @@ class FocalLossTest(tf.test.TestCase, parameterized.TestCase):
self.assertNear(loss, expected_loss, 1e-4) self.assertNear(loss, expected_loss, 1e-4)
class SparseFocalLossTest(tf.test.TestCase):
def test_sparse_focal_loss_matches_focal_loss(self):
num_classes = 2
y_pred = tf.constant([[0.8, 0.2], [0.3, 0.7]])
y_true = tf.constant([1, 0])
y_true_one_hot = tf.one_hot(y_true, num_classes)
for gamma in [0.0, 0.5, 1.0]:
expected_loss_fn = loss_functions.FocalLoss(gamma=gamma)
loss_fn = loss_functions.SparseFocalLoss(
gamma=gamma, num_classes=num_classes
)
expected_loss = expected_loss_fn(y_true_one_hot, y_pred)
loss = loss_fn(y_true, y_pred)
self.assertNear(loss, expected_loss, 1e-4)
class MockPerceptualLoss(loss_functions.PerceptualLoss): class MockPerceptualLoss(loss_functions.PerceptualLoss):
"""A mock class with implementation of abstract methods for testing.""" """A mock class with implementation of abstract methods for testing."""
@@ -46,13 +46,17 @@ class BertModelSpec:
""" """
downloaded_files: file_util.DownloadedFiles downloaded_files: file_util.DownloadedFiles
hparams: hp.BaseHParams = hp.BaseHParams( hparams: hp.BaseHParams = dataclasses.field(
epochs=3, default_factory=lambda: hp.BaseHParams(
batch_size=32, epochs=3,
learning_rate=3e-5, batch_size=32,
distribution_strategy='mirrored') learning_rate=3e-5,
model_options: bert_model_options.BertModelOptions = ( distribution_strategy='mirrored',
bert_model_options.BertModelOptions()) )
)
model_options: bert_model_options.BertModelOptions = dataclasses.field(
default_factory=bert_model_options.BertModelOptions
)
do_lower_case: bool = True do_lower_case: bool = True
tflite_input_name: Dict[str, str] = dataclasses.field( tflite_input_name: Dict[str, str] = dataclasses.field(
default_factory=lambda: _DEFAULT_TFLITE_INPUT_NAME) default_factory=lambda: _DEFAULT_TFLITE_INPUT_NAME)
@@ -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.
# Placeholder for internal Python strict library and test compatibility macro. # Placeholder for internal Python strict binary and library compatibility macro.
# Placeholder for internal Python strict test compatibility macro. # Placeholder for internal Python strict test compatibility macro.
package(default_visibility = ["//mediapipe:__subpackages__"]) package(default_visibility = ["//mediapipe:__subpackages__"])
@@ -76,7 +76,10 @@ py_test(
py_library( py_library(
name = "dataset", name = "dataset",
srcs = ["dataset.py"], srcs = ["dataset.py"],
deps = ["//mediapipe/model_maker/python/core/data:classification_dataset"], deps = [
"//mediapipe/model_maker/python/core/data:cache_files",
"//mediapipe/model_maker/python/core/data:classification_dataset",
],
) )
py_test( py_test(
@@ -88,7 +91,10 @@ py_test(
py_library( py_library(
name = "preprocessor", name = "preprocessor",
srcs = ["preprocessor.py"], srcs = ["preprocessor.py"],
deps = [":dataset"], deps = [
":dataset",
"//mediapipe/model_maker/python/core/data:cache_files",
],
) )
py_test( py_test(
@@ -99,6 +105,7 @@ py_test(
":dataset", ":dataset",
":model_spec", ":model_spec",
":preprocessor", ":preprocessor",
"//mediapipe/model_maker/python/core/data:cache_files",
], ],
) )
@@ -124,6 +131,7 @@ py_library(
":text_classifier_options", ":text_classifier_options",
"//mediapipe/model_maker/python/core/data:dataset", "//mediapipe/model_maker/python/core/data:dataset",
"//mediapipe/model_maker/python/core/tasks:classifier", "//mediapipe/model_maker/python/core/tasks:classifier",
"//mediapipe/model_maker/python/core/utils:loss_functions",
"//mediapipe/model_maker/python/core/utils:metrics", "//mediapipe/model_maker/python/core/utils:metrics",
"//mediapipe/model_maker/python/core/utils:model_util", "//mediapipe/model_maker/python/core/utils:model_util",
"//mediapipe/model_maker/python/core/utils:quantization", "//mediapipe/model_maker/python/core/utils:quantization",
@@ -147,6 +155,7 @@ py_test(
], ],
deps = [ deps = [
":text_classifier_import", ":text_classifier_import",
"//mediapipe/model_maker/python/core/utils:loss_functions",
"//mediapipe/tasks/python/test:test_utils", "//mediapipe/tasks/python/test:test_utils",
], ],
) )
@@ -15,11 +15,15 @@
import csv import csv
import dataclasses import dataclasses
import hashlib
import os
import random import random
import tempfile
from typing import List, Optional, Sequence
from typing import Optional, Sequence
import tensorflow as tf import tensorflow as tf
from mediapipe.model_maker.python.core.data import cache_files as cache_files_lib
from mediapipe.model_maker.python.core.data import classification_dataset from mediapipe.model_maker.python.core.data import classification_dataset
@@ -46,21 +50,49 @@ class CSVParameters:
class Dataset(classification_dataset.ClassificationDataset): class Dataset(classification_dataset.ClassificationDataset):
"""Dataset library for text classifier.""" """Dataset library for text classifier."""
def __init__(
self,
dataset: tf.data.Dataset,
label_names: List[str],
tfrecord_cache_files: Optional[cache_files_lib.TFRecordCacheFiles] = None,
size: Optional[int] = None,
):
super().__init__(dataset, label_names, size)
if not tfrecord_cache_files:
tfrecord_cache_files = cache_files_lib.TFRecordCacheFiles(
cache_prefix_filename="tfrecord", num_shards=1
)
self.tfrecord_cache_files = tfrecord_cache_files
@classmethod @classmethod
def from_csv(cls, def from_csv(
filename: str, cls,
csv_params: CSVParameters, filename: str,
shuffle: bool = True) -> "Dataset": csv_params: CSVParameters,
shuffle: bool = True,
cache_dir: Optional[str] = None,
num_shards: int = 1,
) -> "Dataset":
"""Loads text with labels from a CSV file. """Loads text with labels from a CSV file.
Args: Args:
filename: Name of the CSV file. filename: Name of the CSV file.
csv_params: Parameters used for reading the CSV file. csv_params: Parameters used for reading the CSV file.
shuffle: If True, randomly shuffle the data. shuffle: If True, randomly shuffle the data.
cache_dir: Optional parameter to specify where to store the preprocessed
dataset. Only used for BERT models.
num_shards: Optional parameter for num shards of the preprocessed dataset.
Note that using more than 1 shard will reorder the dataset. Only used
for BERT models.
Returns: Returns:
Dataset containing (text, label) pairs and other related info. Dataset containing (text, label) pairs and other related info.
""" """
if cache_dir is None:
cache_dir = tempfile.mkdtemp()
# calculate hash for cache based off of files
hasher = hashlib.md5()
hasher.update(os.path.basename(filename).encode("utf-8"))
with tf.io.gfile.GFile(filename, "r") as f: with tf.io.gfile.GFile(filename, "r") as f:
reader = csv.DictReader( reader = csv.DictReader(
f, f,
@@ -69,6 +101,9 @@ class Dataset(classification_dataset.ClassificationDataset):
quotechar=csv_params.quotechar) quotechar=csv_params.quotechar)
lines = list(reader) lines = list(reader)
for line in lines:
hasher.update(str(line).encode("utf-8"))
if shuffle: if shuffle:
random.shuffle(lines) random.shuffle(lines)
@@ -81,8 +116,18 @@ class Dataset(classification_dataset.ClassificationDataset):
index_by_label[line[csv_params.label_column]] for line in lines index_by_label[line[csv_params.label_column]] for line in lines
] ]
label_index_ds = tf.data.Dataset.from_tensor_slices( label_index_ds = tf.data.Dataset.from_tensor_slices(
tf.cast(label_indices, tf.int64)) tf.cast(label_indices, tf.int64)
)
text_label_ds = tf.data.Dataset.zip((text_ds, label_index_ds)) text_label_ds = tf.data.Dataset.zip((text_ds, label_index_ds))
hasher.update(str(num_shards).encode("utf-8"))
cache_prefix_filename = hasher.hexdigest()
tfrecord_cache_files = cache_files_lib.TFRecordCacheFiles(
cache_prefix_filename, cache_dir, num_shards
)
return Dataset( return Dataset(
dataset=text_label_ds, size=len(texts), label_names=label_names) dataset=text_label_ds,
label_names=label_names,
tfrecord_cache_files=tfrecord_cache_files,
size=len(texts),
)
@@ -53,7 +53,7 @@ class DatasetTest(tf.test.TestCase):
def test_split(self): def test_split(self):
ds = tf.data.Dataset.from_tensor_slices(['good', 'bad', 'neutral', 'odd']) ds = tf.data.Dataset.from_tensor_slices(['good', 'bad', 'neutral', 'odd'])
data = dataset.Dataset(ds, 4, ['pos', 'neg']) data = dataset.Dataset(ds, ['pos', 'neg'], size=4)
train_data, test_data = data.split(0.5) train_data, test_data = data.split(0.5)
expected_train_data = [b'good', b'bad'] expected_train_data = [b'good', b'bad']
expected_test_data = [b'neutral', b'odd'] expected_test_data = [b'neutral', b'odd']
@@ -15,7 +15,7 @@
import dataclasses import dataclasses
import enum import enum
from typing import Union from typing import Sequence, Union
from mediapipe.model_maker.python.core import hyperparameters as hp from mediapipe.model_maker.python.core import hyperparameters as hp
@@ -39,16 +39,34 @@ class BertHParams(hp.BaseHParams):
Attributes: Attributes:
learning_rate: Learning rate to use for gradient descent training. learning_rate: Learning rate to use for gradient descent training.
batch_size: Batch size for training. end_learning_rate: End learning rate for linear decay. Defaults to 0.
epochs: Number of training iterations over the dataset. batch_size: Batch size for training. Defaults to 48.
optimizer: Optimizer to use for training. Only supported values are "adamw" epochs: Number of training iterations over the dataset. Defaults to 2.
and "lamb". optimizer: Optimizer to use for training. Supported values are defined in
BertOptimizer enum: ADAMW and LAMB.
weight_decay: Weight decay of the optimizer. Defaults to 0.01.
desired_precisions: If specified, adds a RecallAtPrecision metric per
desired_precisions[i] entry which tracks the recall given the constraint
on precision. Only supported for binary classification.
desired_recalls: If specified, adds a PrecisionAtRecall metric per
desired_recalls[i] entry which tracks the precision given the constraint
on recall. Only supported for binary classification.
gamma: Gamma parameter for focal loss. To use cross entropy loss, set this
value to 0. Defaults to 2.0.
""" """
learning_rate: float = 3e-5 learning_rate: float = 3e-5
end_learning_rate: float = 0.0
batch_size: int = 48 batch_size: int = 48
epochs: int = 2 epochs: int = 2
optimizer: BertOptimizer = BertOptimizer.ADAMW optimizer: BertOptimizer = BertOptimizer.ADAMW
weight_decay: float = 0.01
desired_precisions: Sequence[float] = dataclasses.field(default_factory=list)
desired_recalls: Sequence[float] = dataclasses.field(default_factory=list)
gamma: float = 2.0
HParams = Union[BertHParams, AverageWordEmbeddingHParams] HParams = Union[BertHParams, AverageWordEmbeddingHParams]
@@ -47,11 +47,14 @@ class AverageWordEmbeddingClassifierSpec:
""" """
# `learning_rate` is unused for the average word embedding model # `learning_rate` is unused for the average word embedding model
hparams: hp.AverageWordEmbeddingHParams = hp.AverageWordEmbeddingHParams( hparams: hp.AverageWordEmbeddingHParams = dataclasses.field(
epochs=10, batch_size=32, learning_rate=0 default_factory=lambda: hp.AverageWordEmbeddingHParams(
epochs=10, batch_size=32, learning_rate=0
)
)
model_options: mo.AverageWordEmbeddingModelOptions = dataclasses.field(
default_factory=mo.AverageWordEmbeddingModelOptions
) )
model_options: mo.AverageWordEmbeddingModelOptions = (
mo.AverageWordEmbeddingModelOptions())
name: str = 'AverageWordEmbedding' name: str = 'AverageWordEmbedding'
average_word_embedding_classifier_spec = functools.partial( average_word_embedding_classifier_spec = functools.partial(
@@ -66,7 +69,7 @@ class BertClassifierSpec(bert_model_spec.BertModelSpec):
inherited from the BertModelSpec. inherited from the BertModelSpec.
""" """
hparams: hp.BertHParams = hp.BertHParams() hparams: hp.BertHParams = dataclasses.field(default_factory=hp.BertHParams)
mobilebert_classifier_spec = functools.partial( mobilebert_classifier_spec = functools.partial(
@@ -76,11 +79,6 @@ mobilebert_classifier_spec = functools.partial(
epochs=3, batch_size=48, learning_rate=3e-5, distribution_strategy='off' epochs=3, batch_size=48, learning_rate=3e-5, distribution_strategy='off'
), ),
name='MobileBert', name='MobileBert',
tflite_input_name={
'ids': 'serving_default_input_1:0',
'segment_ids': 'serving_default_input_2:0',
'mask': 'serving_default_input_3:0',
},
) )
exbert_classifier_spec = functools.partial( exbert_classifier_spec = functools.partial(
@@ -90,11 +88,6 @@ exbert_classifier_spec = functools.partial(
epochs=3, batch_size=48, learning_rate=3e-5, distribution_strategy='off' epochs=3, batch_size=48, learning_rate=3e-5, distribution_strategy='off'
), ),
name='ExBert', name='ExBert',
tflite_input_name={
'ids': 'serving_default_input_1:0',
'segment_ids': 'serving_default_input_2:0',
'mask': 'serving_default_input_3:0',
},
) )
@@ -46,11 +46,13 @@ class ModelSpecTest(tf.test.TestCase):
self.assertTrue(os.path.exists(model_spec_obj.downloaded_files.get_path())) self.assertTrue(os.path.exists(model_spec_obj.downloaded_files.get_path()))
self.assertTrue(model_spec_obj.do_lower_case) self.assertTrue(model_spec_obj.do_lower_case)
self.assertEqual( self.assertEqual(
model_spec_obj.tflite_input_name, { model_spec_obj.tflite_input_name,
'ids': 'serving_default_input_1:0', {
'mask': 'serving_default_input_3:0', 'ids': 'serving_default_input_word_ids:0',
'segment_ids': 'serving_default_input_2:0' 'mask': 'serving_default_input_mask:0',
}) 'segment_ids': 'serving_default_input_type_ids:0',
},
)
self.assertEqual( self.assertEqual(
model_spec_obj.model_options, model_spec_obj.model_options,
classifier_model_options.BertModelOptions( classifier_model_options.BertModelOptions(
@@ -15,14 +15,15 @@
"""Preprocessors for text classification.""" """Preprocessors for text classification."""
import collections import collections
import hashlib
import os import os
import re import re
import tempfile
from typing import Mapping, Sequence, Tuple, Union from typing import Mapping, Sequence, Tuple, Union
import tensorflow as tf import tensorflow as tf
import tensorflow_hub import tensorflow_hub
from mediapipe.model_maker.python.core.data import cache_files as cache_files_lib
from mediapipe.model_maker.python.text.text_classifier import dataset as text_classifier_ds from mediapipe.model_maker.python.text.text_classifier import dataset as text_classifier_ds
from official.nlp.data import classifier_data_lib from official.nlp.data import classifier_data_lib
from official.nlp.tools import tokenization from official.nlp.tools import tokenization
@@ -75,19 +76,20 @@ def _decode_record(
return bert_features, example["label_ids"] return bert_features, example["label_ids"]
def _single_file_dataset( def _tfrecord_dataset(
input_file: str, name_to_features: Mapping[str, tf.io.FixedLenFeature] tfrecord_files: Sequence[str],
name_to_features: Mapping[str, tf.io.FixedLenFeature],
) -> tf.data.TFRecordDataset: ) -> tf.data.TFRecordDataset:
"""Creates a single-file dataset to be passed for BERT custom training. """Creates a single-file dataset to be passed for BERT custom training.
Args: Args:
input_file: Filepath for the dataset. tfrecord_files: Filepaths for the dataset.
name_to_features: Maps record keys to feature types. name_to_features: Maps record keys to feature types.
Returns: Returns:
Dataset containing BERT model input features and labels. Dataset containing BERT model input features and labels.
""" """
d = tf.data.TFRecordDataset(input_file) d = tf.data.TFRecordDataset(tfrecord_files)
d = d.map( d = d.map(
lambda record: _decode_record(record, name_to_features), lambda record: _decode_record(record, name_to_features),
num_parallel_calls=tf.data.AUTOTUNE) num_parallel_calls=tf.data.AUTOTUNE)
@@ -221,15 +223,23 @@ class BertClassifierPreprocessor:
seq_len: Length of the input sequence to the model. seq_len: Length of the input sequence to the model.
vocab_file: File containing the BERT vocab. vocab_file: File containing the BERT vocab.
tokenizer: BERT tokenizer. tokenizer: BERT tokenizer.
model_name: Name of the model provided by the model_spec. Used to associate
cached files with specific Bert model vocab.
""" """
def __init__(self, seq_len: int, do_lower_case: bool, uri: str): def __init__(
self, seq_len: int, do_lower_case: bool, uri: str, model_name: str
):
self._seq_len = seq_len self._seq_len = seq_len
# Vocab filepath is tied to the BERT module's URI. # Vocab filepath is tied to the BERT module's URI.
self._vocab_file = os.path.join( self._vocab_file = os.path.join(
tensorflow_hub.resolve(uri), "assets", "vocab.txt") tensorflow_hub.resolve(uri), "assets", "vocab.txt"
self._tokenizer = tokenization.FullTokenizer(self._vocab_file, )
do_lower_case) self._do_lower_case = do_lower_case
self._tokenizer = tokenization.FullTokenizer(
self._vocab_file, self._do_lower_case
)
self._model_name = model_name
def _get_name_to_features(self): def _get_name_to_features(self):
"""Gets the dictionary mapping record keys to feature types.""" """Gets the dictionary mapping record keys to feature types."""
@@ -244,8 +254,45 @@ class BertClassifierPreprocessor:
"""Returns the vocab file of the BertClassifierPreprocessor.""" """Returns the vocab file of the BertClassifierPreprocessor."""
return self._vocab_file return self._vocab_file
def _get_tfrecord_cache_files(
self, ds_cache_files
) -> cache_files_lib.TFRecordCacheFiles:
"""Helper to regenerate cache prefix filename using preprocessor info.
We need to update the dataset cache_prefix cache because the actual cached
dataset depends on the preprocessor parameters such as model_name, seq_len,
and do_lower_case in addition to the raw dataset parameters which is already
included in the ds_cache_files.cache_prefix_filename
Specifically, the new cache_prefix_filename used by the preprocessor will
be a hash generated from the following:
1. cache_prefix_filename of the initial raw dataset
2. model_name
3. seq_len
4. do_lower_case
Args:
ds_cache_files: TFRecordCacheFiles from the original raw dataset object
Returns:
A new TFRecordCacheFiles object which incorporates the preprocessor
parameters.
"""
hasher = hashlib.md5()
hasher.update(ds_cache_files.cache_prefix_filename.encode("utf-8"))
hasher.update(self._model_name.encode("utf-8"))
hasher.update(str(self._seq_len).encode("utf-8"))
hasher.update(str(self._do_lower_case).encode("utf-8"))
cache_prefix_filename = hasher.hexdigest()
return cache_files_lib.TFRecordCacheFiles(
cache_prefix_filename,
ds_cache_files.cache_dir,
ds_cache_files.num_shards,
)
def preprocess( def preprocess(
self, dataset: text_classifier_ds.Dataset) -> text_classifier_ds.Dataset: self, dataset: text_classifier_ds.Dataset
) -> text_classifier_ds.Dataset:
"""Preprocesses data into input for a BERT-based classifier. """Preprocesses data into input for a BERT-based classifier.
Args: Args:
@@ -254,32 +301,65 @@ class BertClassifierPreprocessor:
Returns: Returns:
Dataset containing (bert_features, label) data. Dataset containing (bert_features, label) data.
""" """
examples = [] ds_cache_files = dataset.tfrecord_cache_files
for index, (text, label) in enumerate(dataset.gen_tf_dataset()): # Get new tfrecord_cache_files by including preprocessor information.
_validate_text_and_label(text, label) tfrecord_cache_files = self._get_tfrecord_cache_files(ds_cache_files)
examples.append( if not tfrecord_cache_files.is_cached():
classifier_data_lib.InputExample( print(f"Writing new cache files to {tfrecord_cache_files.cache_prefix}")
guid=str(index), writers = tfrecord_cache_files.get_writers()
text_a=text.numpy()[0].decode("utf-8"), size = 0
text_b=None, for index, (text, label) in enumerate(dataset.gen_tf_dataset()):
# InputExample expects the label name rather than the int ID _validate_text_and_label(text, label)
label=dataset.label_names[label.numpy()[0]])) example = classifier_data_lib.InputExample(
guid=str(index),
text_a=text.numpy()[0].decode("utf-8"),
text_b=None,
# InputExample expects the label name rather than the int ID
# label=dataset.label_names[label.numpy()[0]])
label=label.numpy()[0],
)
feature = classifier_data_lib.convert_single_example(
index, example, None, self._seq_len, self._tokenizer
)
tfrecord_file = os.path.join(tempfile.mkdtemp(), "bert_features.tfrecord") def create_int_feature(values):
classifier_data_lib.file_based_convert_examples_to_features( f = tf.train.Feature(
examples=examples, int64_list=tf.train.Int64List(value=list(values))
label_list=dataset.label_names, )
max_seq_length=self._seq_len, return f
tokenizer=self._tokenizer,
output_file=tfrecord_file) features = collections.OrderedDict()
preprocessed_ds = _single_file_dataset(tfrecord_file, features["input_ids"] = create_int_feature(feature.input_ids)
self._get_name_to_features()) features["input_mask"] = create_int_feature(feature.input_mask)
features["segment_ids"] = create_int_feature(feature.segment_ids)
features["label_ids"] = create_int_feature([feature.label_id])
tf_example = tf.train.Example(
features=tf.train.Features(feature=features)
)
writers[index % len(writers)].write(tf_example.SerializeToString())
size = index + 1
for writer in writers:
writer.close()
metadata = {"size": size, "label_names": dataset.label_names}
tfrecord_cache_files.save_metadata(metadata)
else:
print(
f"Using existing cache files at {tfrecord_cache_files.cache_prefix}"
)
metadata = tfrecord_cache_files.load_metadata()
size = metadata["size"]
label_names = metadata["label_names"]
preprocessed_ds = _tfrecord_dataset(
tfrecord_cache_files.tfrecord_files, self._get_name_to_features()
)
return text_classifier_ds.Dataset( return text_classifier_ds.Dataset(
dataset=preprocessed_ds, dataset=preprocessed_ds,
size=dataset.size, size=size,
label_names=dataset.label_names) label_names=label_names,
tfrecord_cache_files=tfrecord_cache_files,
)
TextClassifierPreprocessor = ( TextClassifierPreprocessor = Union[
Union[BertClassifierPreprocessor, BertClassifierPreprocessor, AverageWordEmbeddingClassifierPreprocessor
AverageWordEmbeddingClassifierPreprocessor]) ]
@@ -13,14 +13,17 @@
# limitations under the License. # limitations under the License.
import csv import csv
import io
import os import os
import tempfile import tempfile
from unittest import mock as unittest_mock from unittest import mock as unittest_mock
import mock
import numpy as np import numpy as np
import numpy.testing as npt import numpy.testing as npt
import tensorflow as tf import tensorflow as tf
from mediapipe.model_maker.python.core.data import cache_files
from mediapipe.model_maker.python.text.text_classifier import dataset as text_classifier_ds from mediapipe.model_maker.python.text.text_classifier import dataset as text_classifier_ds
from mediapipe.model_maker.python.text.text_classifier import model_spec from mediapipe.model_maker.python.text.text_classifier import model_spec
from mediapipe.model_maker.python.text.text_classifier import preprocessor from mediapipe.model_maker.python.text.text_classifier import preprocessor
@@ -84,11 +87,12 @@ class PreprocessorTest(tf.test.TestCase):
csv_file = self._get_csv_file() csv_file = self._get_csv_file()
dataset = text_classifier_ds.Dataset.from_csv( dataset = text_classifier_ds.Dataset.from_csv(
filename=csv_file, csv_params=self.CSV_PARAMS_) filename=csv_file, csv_params=self.CSV_PARAMS_)
bert_spec = model_spec.SupportedModels.MOBILEBERT_CLASSIFIER.value() bert_spec = model_spec.SupportedModels.EXBERT_CLASSIFIER.value()
bert_preprocessor = preprocessor.BertClassifierPreprocessor( bert_preprocessor = preprocessor.BertClassifierPreprocessor(
seq_len=5, seq_len=5,
do_lower_case=bert_spec.do_lower_case, do_lower_case=bert_spec.do_lower_case,
uri=bert_spec.downloaded_files.get_path(), uri=bert_spec.downloaded_files.get_path(),
model_name=bert_spec.name,
) )
preprocessed_dataset = bert_preprocessor.preprocess(dataset) preprocessed_dataset = bert_preprocessor.preprocess(dataset)
labels = [] labels = []
@@ -97,18 +101,91 @@ class PreprocessorTest(tf.test.TestCase):
self.assertEqual(label.shape, [1]) self.assertEqual(label.shape, [1])
labels.append(label.numpy()[0]) labels.append(label.numpy()[0])
self.assertSameElements( self.assertSameElements(
features.keys(), ['input_word_ids', 'input_mask', 'input_type_ids']) features.keys(), ['input_word_ids', 'input_mask', 'input_type_ids']
)
for feature in features.values(): for feature in features.values():
self.assertEqual(feature.shape, [1, 5]) self.assertEqual(feature.shape, [1, 5])
input_masks.append(features['input_mask'].numpy()[0]) input_masks.append(features['input_mask'].numpy()[0])
npt.assert_array_equal(features['input_type_ids'].numpy()[0], npt.assert_array_equal(
[0, 0, 0, 0, 0]) features['input_type_ids'].numpy()[0], [0, 0, 0, 0, 0]
)
npt.assert_array_equal( npt.assert_array_equal(
np.stack(input_masks), np.array([[1, 1, 1, 1, 1], [1, 1, 1, 1, 0]])) np.stack(input_masks), np.array([[1, 1, 1, 1, 1], [1, 1, 1, 1, 0]])
)
self.assertEqual(labels, [1, 0]) self.assertEqual(labels, [1, 0])
def test_bert_preprocessor_cache(self):
csv_file = self._get_csv_file()
dataset = text_classifier_ds.Dataset.from_csv(
filename=csv_file,
csv_params=self.CSV_PARAMS_,
cache_dir=self.get_temp_dir(),
)
bert_spec = model_spec.SupportedModels.EXBERT_CLASSIFIER.value()
bert_preprocessor = preprocessor.BertClassifierPreprocessor(
seq_len=5,
do_lower_case=bert_spec.do_lower_case,
uri=bert_spec.downloaded_files.get_path(),
model_name=bert_spec.name,
)
ds_cache_files = dataset.tfrecord_cache_files
preprocessed_cache_files = bert_preprocessor._get_tfrecord_cache_files(
ds_cache_files
)
self.assertFalse(preprocessed_cache_files.is_cached())
preprocessed_dataset = bert_preprocessor.preprocess(dataset)
self.assertTrue(preprocessed_cache_files.is_cached())
self.assertEqual(
preprocessed_dataset.tfrecord_cache_files, preprocessed_cache_files
)
# The second time running preprocessor, it should load from cache directly
mock_stdout = io.StringIO()
with mock.patch('sys.stdout', mock_stdout):
_ = bert_preprocessor.preprocess(dataset)
self.assertEqual(
mock_stdout.getvalue(),
'Using existing cache files at'
f' {preprocessed_cache_files.cache_prefix}\n',
)
def _get_new_prefix(self, cf, bert_spec, seq_len, do_lower_case):
bert_preprocessor = preprocessor.BertClassifierPreprocessor(
seq_len=seq_len,
do_lower_case=do_lower_case,
uri=bert_spec.downloaded_files.get_path(),
model_name=bert_spec.name,
)
new_cf = bert_preprocessor._get_tfrecord_cache_files(cf)
return new_cf.cache_prefix_filename
def test_bert_get_tfrecord_cache_files(self):
# Test to ensure regenerated cache_files have different prefixes
all_cf_prefixes = set()
cf = cache_files.TFRecordCacheFiles(
cache_prefix_filename='cache_prefix',
cache_dir=self.get_temp_dir(),
num_shards=1,
)
exbert_spec = model_spec.SupportedModels.EXBERT_CLASSIFIER.value()
all_cf_prefixes.add(self._get_new_prefix(cf, exbert_spec, 5, True))
all_cf_prefixes.add(self._get_new_prefix(cf, exbert_spec, 10, True))
all_cf_prefixes.add(self._get_new_prefix(cf, exbert_spec, 5, False))
mobilebert_spec = model_spec.SupportedModels.MOBILEBERT_CLASSIFIER.value()
all_cf_prefixes.add(self._get_new_prefix(cf, mobilebert_spec, 5, True))
all_cf_prefixes.add(self._get_new_prefix(cf, mobilebert_spec, 10, True))
all_cf_prefixes.add(self._get_new_prefix(cf, mobilebert_spec, 5, False))
new_cf = cache_files.TFRecordCacheFiles(
cache_prefix_filename='new_cache_prefix',
cache_dir=self.get_temp_dir(),
num_shards=1,
)
all_cf_prefixes.add(self._get_new_prefix(new_cf, exbert_spec, 5, True))
# Each item of all_cf_prefixes should be unique, so 7 total.
self.assertLen(all_cf_prefixes, 7)
if __name__ == '__main__': if __name__ == '__main__':
# Load compressed models from tensorflow_hub # Load compressed models from tensorflow_hub
os.environ['TFHUB_MODEL_LOAD_FORMAT'] = 'COMPRESSED'
tf.test.main() tf.test.main()
@@ -16,8 +16,8 @@
} }
}, },
{ {
"name": "mask", "name": "segment_ids",
"description": "Mask with 1 for real tokens and 0 for padding tokens.", "description": "0 for the first sequence, 1 for the second sequence if exists.",
"content": { "content": {
"content_properties_type": "FeatureProperties", "content_properties_type": "FeatureProperties",
"content_properties": { "content_properties": {
@@ -27,8 +27,8 @@
} }
}, },
{ {
"name": "segment_ids", "name": "mask",
"description": "0 for the first sequence, 1 for the second sequence if exists.", "description": "Mask with 1 for real tokens and 0 for padding tokens.",
"content": { "content": {
"content_properties_type": "FeatureProperties", "content_properties_type": "FeatureProperties",
"content_properties": { "content_properties": {
@@ -24,6 +24,7 @@ import tensorflow_hub as hub
from mediapipe.model_maker.python.core.data import dataset as ds from mediapipe.model_maker.python.core.data import dataset as ds
from mediapipe.model_maker.python.core.tasks import classifier from mediapipe.model_maker.python.core.tasks import classifier
from mediapipe.model_maker.python.core.utils import loss_functions
from mediapipe.model_maker.python.core.utils import metrics from mediapipe.model_maker.python.core.utils import metrics
from mediapipe.model_maker.python.core.utils import model_util from mediapipe.model_maker.python.core.utils import model_util
from mediapipe.model_maker.python.core.utils import quantization from mediapipe.model_maker.python.core.utils import quantization
@@ -116,17 +117,14 @@ class TextClassifier(classifier.Classifier):
options.supported_model == ms.SupportedModels.MOBILEBERT_CLASSIFIER options.supported_model == ms.SupportedModels.MOBILEBERT_CLASSIFIER
or options.supported_model == ms.SupportedModels.EXBERT_CLASSIFIER or options.supported_model == ms.SupportedModels.EXBERT_CLASSIFIER
): ):
text_classifier = ( text_classifier = _BertClassifier.create_bert_classifier(
_BertClassifier.create_bert_classifier(train_data, validation_data, train_data, validation_data, options
options, )
train_data.label_names))
elif (options.supported_model == elif (options.supported_model ==
ms.SupportedModels.AVERAGE_WORD_EMBEDDING_CLASSIFIER): ms.SupportedModels.AVERAGE_WORD_EMBEDDING_CLASSIFIER):
text_classifier = ( text_classifier = _AverageWordEmbeddingClassifier.create_average_word_embedding_classifier(
_AverageWordEmbeddingClassifier train_data, validation_data, options
.create_average_word_embedding_classifier(train_data, validation_data, )
options,
train_data.label_names))
else: else:
raise ValueError(f"Unknown model {options.supported_model}") raise ValueError(f"Unknown model {options.supported_model}")
@@ -166,28 +164,8 @@ class TextClassifier(classifier.Classifier):
processed_data = self._text_preprocessor.preprocess(data) processed_data = self._text_preprocessor.preprocess(data)
dataset = processed_data.gen_tf_dataset(batch_size, is_training=False) dataset = processed_data.gen_tf_dataset(batch_size, is_training=False)
additional_metrics = [] with self._hparams.get_strategy().scope():
if desired_precisions and len(data.label_names) == 2: return self._model.evaluate(dataset)
for precision in desired_precisions:
additional_metrics.append(
metrics.BinarySparseRecallAtPrecision(
precision, name=f"recall_at_precision_{precision}"
)
)
if desired_recalls and len(data.label_names) == 2:
for recall in desired_recalls:
additional_metrics.append(
metrics.BinarySparsePrecisionAtRecall(
recall, name=f"precision_at_recall_{recall}"
)
)
metric_functions = self._metric_functions + additional_metrics
self._model.compile(
optimizer=self._optimizer,
loss=self._loss_function,
metrics=metric_functions,
)
return self._model.evaluate(dataset)
def export_model( def export_model(
self, self,
@@ -255,16 +233,17 @@ class _AverageWordEmbeddingClassifier(TextClassifier):
@classmethod @classmethod
def create_average_word_embedding_classifier( def create_average_word_embedding_classifier(
cls, train_data: text_ds.Dataset, validation_data: text_ds.Dataset, cls,
train_data: text_ds.Dataset,
validation_data: text_ds.Dataset,
options: text_classifier_options.TextClassifierOptions, options: text_classifier_options.TextClassifierOptions,
label_names: Sequence[str]) -> "_AverageWordEmbeddingClassifier": ) -> "_AverageWordEmbeddingClassifier":
"""Creates, trains, and returns an Average Word Embedding classifier. """Creates, trains, and returns an Average Word Embedding classifier.
Args: Args:
train_data: Training data. train_data: Training data.
validation_data: Validation data. validation_data: Validation data.
options: Options for creating and training the text classifier. options: Options for creating and training the text classifier.
label_names: Label names used in the data.
Returns: Returns:
An Average Word Embedding classifier. An Average Word Embedding classifier.
@@ -370,28 +349,25 @@ class _BertClassifier(TextClassifier):
self._callbacks = model_util.get_default_callbacks(self._hparams.export_dir) self._callbacks = model_util.get_default_callbacks(self._hparams.export_dir)
self._model_options = model_options self._model_options = model_options
with self._hparams.get_strategy().scope(): with self._hparams.get_strategy().scope():
self._loss_function = tf.keras.losses.SparseCategoricalCrossentropy() self._loss_function = loss_functions.SparseFocalLoss(
self._metric_functions = [ self._hparams.gamma, self._num_classes
tf.keras.metrics.SparseCategoricalAccuracy( )
"test_accuracy", dtype=tf.float32 self._metric_functions = self._create_metrics()
), self._text_preprocessor: preprocessor.BertClassifierPreprocessor = None
metrics.SparsePrecision(name="precision", dtype=tf.float32),
metrics.SparseRecall(name="recall", dtype=tf.float32),
]
self._text_preprocessor: preprocessor.BertClassifierPreprocessor = None
@classmethod @classmethod
def create_bert_classifier( def create_bert_classifier(
cls, train_data: text_ds.Dataset, validation_data: text_ds.Dataset, cls,
train_data: text_ds.Dataset,
validation_data: text_ds.Dataset,
options: text_classifier_options.TextClassifierOptions, options: text_classifier_options.TextClassifierOptions,
label_names: Sequence[str]) -> "_BertClassifier": ) -> "_BertClassifier":
"""Creates, trains, and returns a BERT-based classifier. """Creates, trains, and returns a BERT-based classifier.
Args: Args:
train_data: Training data. train_data: Training data.
validation_data: Validation data. validation_data: Validation data.
options: Options for creating and training the text classifier. options: Options for creating and training the text classifier.
label_names: Label names used in the data.
Returns: Returns:
A BERT-based classifier. A BERT-based classifier.
@@ -435,9 +411,59 @@ class _BertClassifier(TextClassifier):
seq_len=self._model_options.seq_len, seq_len=self._model_options.seq_len,
do_lower_case=self._model_spec.do_lower_case, do_lower_case=self._model_spec.do_lower_case,
uri=self._model_spec.downloaded_files.get_path(), uri=self._model_spec.downloaded_files.get_path(),
model_name=self._model_spec.name,
) )
return (self._text_preprocessor.preprocess(train_data), return (
self._text_preprocessor.preprocess(validation_data)) self._text_preprocessor.preprocess(train_data),
self._text_preprocessor.preprocess(validation_data),
)
def _create_metrics(self):
"""Creates metrics for training and evaluation.
The default metrics are accuracy, precision, and recall.
For binary classification tasks only (num_classes=2):
Users can configure PrecisionAtRecall and RecallAtPrecision metrics using
the desired_presisions and desired_recalls fields in BertHParams.
Returns:
A list of tf.keras.Metric subclasses which can be used with model.compile
"""
metric_functions = [
tf.keras.metrics.SparseCategoricalAccuracy(
"accuracy", dtype=tf.float32
),
metrics.SparsePrecision(name="precision", dtype=tf.float32),
metrics.SparseRecall(name="recall", dtype=tf.float32),
]
if self._num_classes == 2:
if self._hparams.desired_precisions:
for desired_precision in self._hparams.desired_precisions:
metric_functions.append(
metrics.BinarySparseRecallAtPrecision(
desired_precision,
name=f"recall_at_precision_{desired_precision}",
num_thresholds=1000,
)
)
if self._hparams.desired_recalls:
for desired_recall in self._hparams.desired_recalls:
metric_functions.append(
metrics.BinarySparseRecallAtPrecision(
desired_recall,
name=f"precision_at_recall_{desired_recall}",
num_thresholds=1000,
)
)
else:
if self._hparams.desired_precisions or self._hparams.desired_recalls:
raise ValueError(
"desired_recalls and desired_precisions parameters are binary"
" metrics and not supported for num_classes > 2. Found"
f" num_classes: {self._num_classes}"
)
return metric_functions
def _create_model(self): def _create_model(self):
"""Creates a BERT-based classifier model. """Creates a BERT-based classifier model.
@@ -447,11 +473,20 @@ class _BertClassifier(TextClassifier):
""" """
encoder_inputs = dict( encoder_inputs = dict(
input_word_ids=tf.keras.layers.Input( input_word_ids=tf.keras.layers.Input(
shape=(self._model_options.seq_len,), dtype=tf.int32), shape=(self._model_options.seq_len,),
dtype=tf.int32,
name="input_word_ids",
),
input_mask=tf.keras.layers.Input( input_mask=tf.keras.layers.Input(
shape=(self._model_options.seq_len,), dtype=tf.int32), shape=(self._model_options.seq_len,),
dtype=tf.int32,
name="input_mask",
),
input_type_ids=tf.keras.layers.Input( input_type_ids=tf.keras.layers.Input(
shape=(self._model_options.seq_len,), dtype=tf.int32), shape=(self._model_options.seq_len,),
dtype=tf.int32,
name="input_type_ids",
),
) )
encoder = hub.KerasLayer( encoder = hub.KerasLayer(
self._model_spec.downloaded_files.get_path(), self._model_spec.downloaded_files.get_path(),
@@ -493,16 +528,21 @@ class _BertClassifier(TextClassifier):
lr_schedule = tf.keras.optimizers.schedules.PolynomialDecay( lr_schedule = tf.keras.optimizers.schedules.PolynomialDecay(
initial_learning_rate=initial_lr, initial_learning_rate=initial_lr,
decay_steps=total_steps, decay_steps=total_steps,
end_learning_rate=0.0, end_learning_rate=self._hparams.end_learning_rate,
power=1.0) power=1.0,
)
if warmup_steps: if warmup_steps:
lr_schedule = model_util.WarmUp( lr_schedule = model_util.WarmUp(
initial_learning_rate=initial_lr, initial_learning_rate=initial_lr,
decay_schedule_fn=lr_schedule, decay_schedule_fn=lr_schedule,
warmup_steps=warmup_steps) warmup_steps=warmup_steps,
)
if self._hparams.optimizer == hp.BertOptimizer.ADAMW: if self._hparams.optimizer == hp.BertOptimizer.ADAMW:
self._optimizer = tf.keras.optimizers.experimental.AdamW( self._optimizer = tf.keras.optimizers.experimental.AdamW(
lr_schedule, weight_decay=0.01, epsilon=1e-6, global_clipnorm=1.0 lr_schedule,
weight_decay=self._hparams.weight_decay,
epsilon=1e-6,
global_clipnorm=1.0,
) )
self._optimizer.exclude_from_weight_decay( self._optimizer.exclude_from_weight_decay(
var_names=["LayerNorm", "layer_norm", "bias"] var_names=["LayerNorm", "layer_norm", "bias"]
@@ -510,7 +550,7 @@ class _BertClassifier(TextClassifier):
elif self._hparams.optimizer == hp.BertOptimizer.LAMB: elif self._hparams.optimizer == hp.BertOptimizer.LAMB:
self._optimizer = tfa_optimizers.LAMB( self._optimizer = tfa_optimizers.LAMB(
lr_schedule, lr_schedule,
weight_decay_rate=0.01, weight_decay_rate=self._hparams.weight_decay,
epsilon=1e-6, epsilon=1e-6,
exclude_from_weight_decay=["LayerNorm", "layer_norm", "bias"], exclude_from_weight_decay=["LayerNorm", "layer_norm", "bias"],
global_clipnorm=1.0, global_clipnorm=1.0,
@@ -84,8 +84,8 @@ def run(data_dir,
options) options)
# Gets evaluation results. # Gets evaluation results.
_, acc = model.evaluate(validation_data) metrics = model.evaluate(validation_data)
print('Eval accuracy: %f' % acc) print('Eval accuracy: %f' % metrics[1])
model.export_model(quantization_config=quantization_config) model.export_model(quantization_config=quantization_config)
model.export_labels(export_dir=options.hparams.export_dir) model.export_labels(export_dir=options.hparams.export_dir)
@@ -16,17 +16,17 @@ import csv
import filecmp import filecmp
import os import os
import tempfile import tempfile
import unittest
from unittest import mock as unittest_mock from unittest import mock as unittest_mock
from absl.testing import parameterized
import tensorflow as tf import tensorflow as tf
from mediapipe.model_maker.python.core.utils import loss_functions
from mediapipe.model_maker.python.text import text_classifier from mediapipe.model_maker.python.text import text_classifier
from mediapipe.tasks.python.test import test_utils from mediapipe.tasks.python.test import test_utils
@unittest.skip('b/275624089') class TextClassifierTest(tf.test.TestCase, parameterized.TestCase):
class TextClassifierTest(tf.test.TestCase):
_AVERAGE_WORD_EMBEDDING_JSON_FILE = ( _AVERAGE_WORD_EMBEDDING_JSON_FILE = (
test_utils.get_test_data_path('average_word_embedding_metadata.json')) test_utils.get_test_data_path('average_word_embedding_metadata.json'))
@@ -78,8 +78,8 @@ class TextClassifierTest(tf.test.TestCase):
text_classifier.TextClassifier.create(train_data, validation_data, text_classifier.TextClassifier.create(train_data, validation_data,
options)) options))
_, accuracy = average_word_embedding_classifier.evaluate(validation_data) metrics = average_word_embedding_classifier.evaluate(validation_data)
self.assertGreaterEqual(accuracy, 0.0) self.assertGreaterEqual(metrics[1], 0.0) # metrics[1] is accuracy
# Test export_model # Test export_model
average_word_embedding_classifier.export_model() average_word_embedding_classifier.export_model()
@@ -98,12 +98,25 @@ class TextClassifierTest(tf.test.TestCase):
filecmp.cmp( filecmp.cmp(
output_metadata_file, output_metadata_file,
self._AVERAGE_WORD_EMBEDDING_JSON_FILE, self._AVERAGE_WORD_EMBEDDING_JSON_FILE,
shallow=False)) shallow=False,
)
)
def test_create_and_train_bert(self): @parameterized.named_parameters(
# Skipping mobilebert b/c OSS test timeout/flakiness: b/275624089
# dict(
# testcase_name='mobilebert',
# supported_model=text_classifier.SupportedModels.MOBILEBERT_CLASSIFIER,
# ),
dict(
testcase_name='exbert',
supported_model=text_classifier.SupportedModels.EXBERT_CLASSIFIER,
),
)
def test_create_and_train_bert(self, supported_model):
train_data, validation_data = self._get_data() train_data, validation_data = self._get_data()
options = text_classifier.TextClassifierOptions( options = text_classifier.TextClassifierOptions(
supported_model=text_classifier.SupportedModels.MOBILEBERT_CLASSIFIER, supported_model=supported_model,
model_options=text_classifier.BertModelOptions( model_options=text_classifier.BertModelOptions(
do_fine_tuning=False, seq_len=2 do_fine_tuning=False, seq_len=2
), ),
@@ -117,8 +130,8 @@ class TextClassifierTest(tf.test.TestCase):
bert_classifier = text_classifier.TextClassifier.create( bert_classifier = text_classifier.TextClassifier.create(
train_data, validation_data, options) train_data, validation_data, options)
_, accuracy = bert_classifier.evaluate(validation_data) metrics = bert_classifier.evaluate(validation_data)
self.assertGreaterEqual(accuracy, 0.0) self.assertGreaterEqual(metrics[1], 0.0) # metrics[1] is accuracy
# Test export_model # Test export_model
bert_classifier.export_model() bert_classifier.export_model()
@@ -142,45 +155,93 @@ class TextClassifierTest(tf.test.TestCase):
) )
def test_label_mismatch(self): def test_label_mismatch(self):
options = ( options = text_classifier.TextClassifierOptions(
text_classifier.TextClassifierOptions( supported_model=(text_classifier.SupportedModels.EXBERT_CLASSIFIER)
supported_model=( )
text_classifier.SupportedModels.MOBILEBERT_CLASSIFIER)))
train_tf_dataset = tf.data.Dataset.from_tensor_slices([[0]]) train_tf_dataset = tf.data.Dataset.from_tensor_slices([[0]])
train_data = text_classifier.Dataset(train_tf_dataset, 1, ['foo']) train_data = text_classifier.Dataset(train_tf_dataset, ['foo'], 1)
validation_tf_dataset = tf.data.Dataset.from_tensor_slices([[0]]) validation_tf_dataset = tf.data.Dataset.from_tensor_slices([[0]])
validation_data = text_classifier.Dataset(validation_tf_dataset, 1, ['bar']) validation_data = text_classifier.Dataset(validation_tf_dataset, ['bar'], 1)
with self.assertRaisesRegex( with self.assertRaisesRegex(
ValueError, ValueError,
'Training data label names .* not equal to validation data label names' 'Training data label names .* not equal to validation data label names',
): ):
text_classifier.TextClassifier.create(train_data, validation_data, text_classifier.TextClassifier.create(
options) train_data, validation_data, options
)
def test_options_mismatch(self): def test_options_mismatch(self):
train_data, validation_data = self._get_data() train_data, validation_data = self._get_data()
avg_options = ( avg_options = text_classifier.TextClassifierOptions(
text_classifier.TextClassifierOptions( supported_model=(text_classifier.SupportedModels.EXBERT_CLASSIFIER),
supported_model=( model_options=text_classifier.AverageWordEmbeddingModelOptions(),
text_classifier.SupportedModels.MOBILEBERT_CLASSIFIER), )
model_options=text_classifier.AverageWordEmbeddingModelOptions())) with self.assertRaisesWithLiteralMatch(
with self.assertRaisesRegex( ValueError,
ValueError, 'Expected AVERAGE_WORD_EMBEDDING_CLASSIFIER, got' 'Expected AVERAGE_WORD_EMBEDDING_CLASSIFIER, got'
' SupportedModels.MOBILEBERT_CLASSIFIER'): ' SupportedModels.EXBERT_CLASSIFIER',
text_classifier.TextClassifier.create(train_data, validation_data, ):
avg_options) text_classifier.TextClassifier.create(
train_data, validation_data, avg_options
)
bert_options = ( bert_options = text_classifier.TextClassifierOptions(
text_classifier.TextClassifierOptions( supported_model=(
supported_model=(text_classifier.SupportedModels text_classifier.SupportedModels.AVERAGE_WORD_EMBEDDING_CLASSIFIER
.AVERAGE_WORD_EMBEDDING_CLASSIFIER), ),
model_options=text_classifier.BertModelOptions())) model_options=text_classifier.BertModelOptions(),
with self.assertRaisesRegex( )
ValueError, 'Expected MOBILEBERT_CLASSIFIER, got' with self.assertRaisesWithLiteralMatch(
' SupportedModels.AVERAGE_WORD_EMBEDDING_CLASSIFIER'): ValueError,
text_classifier.TextClassifier.create(train_data, validation_data, 'Expected a Bert Classifier(MobileBERT or EXBERT), got'
bert_options) ' SupportedModels.AVERAGE_WORD_EMBEDDING_CLASSIFIER',
):
text_classifier.TextClassifier.create(
train_data, validation_data, bert_options
)
def test_bert_loss_and_metrics_creation(self):
train_data, validation_data = self._get_data()
supported_model = text_classifier.SupportedModels.EXBERT_CLASSIFIER
hparams = text_classifier.BertHParams(
desired_recalls=[0.2],
desired_precisions=[0.9],
epochs=1,
batch_size=1,
learning_rate=3e-5,
distribution_strategy='off',
gamma=3.5,
)
options = text_classifier.TextClassifierOptions(
supported_model=supported_model, hparams=hparams
)
bert_classifier = text_classifier.TextClassifier.create(
train_data, validation_data, options
)
loss_fn = bert_classifier._loss_function
self.assertIsInstance(loss_fn, loss_functions.SparseFocalLoss)
self.assertEqual(loss_fn._gamma, 3.5)
self.assertEqual(loss_fn._num_classes, 2)
metric_names = [m.name for m in bert_classifier._metric_functions]
expected_metric_names = [
'accuracy',
'recall',
'precision',
'precision_at_recall_0.2',
'recall_at_precision_0.9',
]
self.assertCountEqual(metric_names, expected_metric_names)
# Non-binary data
tf_dataset = tf.data.Dataset.from_tensor_slices([[0]])
data = text_classifier.Dataset(tf_dataset, ['foo', 'bar', 'baz'], 1)
with self.assertRaisesWithLiteralMatch(
ValueError,
'desired_recalls and desired_precisions parameters are binary metrics'
' and not supported for num_classes > 2. Found num_classes: 3',
):
text_classifier.TextClassifier.create(data, data, options)
if __name__ == '__main__': if __name__ == '__main__':
@@ -115,5 +115,7 @@ class Dataset(classification_dataset.ClassificationDataset):
', '.join(label_names), ', '.join(label_names),
) )
return Dataset( return Dataset(
dataset=image_label_ds, size=all_image_size, label_names=label_names dataset=image_label_ds,
label_names=label_names,
size=all_image_size,
) )
@@ -13,7 +13,7 @@
# limitations under the License. # limitations under the License.
# Placeholder for internal Python strict test compatibility macro. # Placeholder for internal Python strict test compatibility macro.
# Placeholder for internal Python strict library and test compatibility macro. # Placeholder for internal Python strict binary and library compatibility macro.
licenses(["notice"]) licenses(["notice"])
@@ -249,5 +249,6 @@ class Dataset(classification_dataset.ClassificationDataset):
len(valid_hand_data), len(label_names), ','.join(label_names))) len(valid_hand_data), len(label_names), ','.join(label_names)))
return Dataset( return Dataset(
dataset=hand_embedding_label_ds, dataset=hand_embedding_label_ds,
label_names=label_names,
size=len(valid_hand_data), size=len(valid_hand_data),
label_names=label_names) )
@@ -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.
# Placeholder for internal Python strict library and test compatibility macro. # Placeholder for internal Python strict binary and library compatibility macro.
# Placeholder for internal Python library rule. # Placeholder for internal Python library rule.
licenses(["notice"]) licenses(["notice"])
@@ -15,28 +15,12 @@
import os import os
import random import random
from typing import List, Optional
import tensorflow as tf import tensorflow as tf
import tensorflow_datasets as tfds
from mediapipe.model_maker.python.core.data import classification_dataset from mediapipe.model_maker.python.core.data import classification_dataset
from mediapipe.model_maker.python.vision.core import image_utils from mediapipe.model_maker.python.vision.core import image_utils
def _create_data(
name: str, data: tf.data.Dataset, info: tfds.core.DatasetInfo,
label_names: List[str]
) -> Optional[classification_dataset.ClassificationDataset]:
"""Creates a Dataset object from tfds data."""
if name not in data:
return None
data = data[name]
data = data.map(lambda a: (a['image'], a['label']))
size = info.splits[name].num_examples
return Dataset(data, size, label_names)
class Dataset(classification_dataset.ClassificationDataset): class Dataset(classification_dataset.ClassificationDataset):
"""Dataset library for image classifier.""" """Dataset library for image classifier."""
@@ -99,4 +83,5 @@ class Dataset(classification_dataset.ClassificationDataset):
'Load image with size: %d, num_label: %d, labels: %s.', all_image_size, 'Load image with size: %d, num_label: %d, labels: %s.', all_image_size,
all_label_size, ', '.join(label_names)) all_label_size, ', '.join(label_names))
return Dataset( return Dataset(
dataset=image_label_ds, size=all_image_size, label_names=label_names) dataset=image_label_ds, label_names=label_names, size=all_image_size
)
@@ -41,7 +41,7 @@ class DatasetTest(tf.test.TestCase):
def test_split(self): def test_split(self):
ds = tf.data.Dataset.from_tensor_slices([[0, 1], [1, 1], [0, 0], [1, 0]]) ds = tf.data.Dataset.from_tensor_slices([[0, 1], [1, 1], [0, 0], [1, 0]])
data = dataset.Dataset(dataset=ds, size=4, label_names=['pos', 'neg']) data = dataset.Dataset(dataset=ds, label_names=['pos', 'neg'], size=4)
train_data, test_data = data.split(fraction=0.5) train_data, test_data = data.split(fraction=0.5)
self.assertLen(train_data, 2) self.assertLen(train_data, 2)
@@ -52,8 +52,9 @@ class ImageClassifierTest(tf.test.TestCase, parameterized.TestCase):
ds = tf.data.Dataset.from_generator( ds = tf.data.Dataset.from_generator(
self._gen, (tf.uint8, tf.int64), (tf.TensorShape( self._gen, (tf.uint8, tf.int64), (tf.TensorShape(
[self.IMAGE_SIZE, self.IMAGE_SIZE, 3]), tf.TensorShape([]))) [self.IMAGE_SIZE, self.IMAGE_SIZE, 3]), tf.TensorShape([])))
data = image_classifier.Dataset(ds, self.IMAGES_PER_CLASS * 3, data = image_classifier.Dataset(
['cyan', 'magenta', 'yellow']) ds, ['cyan', 'magenta', 'yellow'], self.IMAGES_PER_CLASS * 3
)
return data return data
def setUp(self): def setUp(self):
@@ -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.
# Placeholder for internal Python strict library and test compatibility macro. # Placeholder for internal Python strict binary and library compatibility macro.
# Placeholder for internal Python strict test compatibility macro. # Placeholder for internal Python strict test compatibility macro.
licenses(["notice"]) licenses(["notice"])
@@ -54,6 +54,7 @@ py_library(
srcs = ["dataset.py"], srcs = ["dataset.py"],
deps = [ deps = [
":dataset_util", ":dataset_util",
"//mediapipe/model_maker/python/core/data:cache_files",
"//mediapipe/model_maker/python/core/data:classification_dataset", "//mediapipe/model_maker/python/core/data:classification_dataset",
], ],
) )
@@ -73,6 +74,7 @@ py_test(
py_library( py_library(
name = "dataset_util", name = "dataset_util",
srcs = ["dataset_util.py"], srcs = ["dataset_util.py"],
deps = ["//mediapipe/model_maker/python/core/data:cache_files"],
) )
py_test( py_test(
@@ -16,8 +16,8 @@
from typing import Optional from typing import Optional
import tensorflow as tf import tensorflow as tf
import yaml
from mediapipe.model_maker.python.core.data import cache_files
from mediapipe.model_maker.python.core.data import classification_dataset from mediapipe.model_maker.python.core.data import classification_dataset
from mediapipe.model_maker.python.vision.object_detector import dataset_util from mediapipe.model_maker.python.vision.object_detector import dataset_util
from official.vision.dataloaders import tf_example_decoder from official.vision.dataloaders import tf_example_decoder
@@ -76,14 +76,16 @@ class Dataset(classification_dataset.ClassificationDataset):
ValueError: If the label_name for id 0 is set to something other than ValueError: If the label_name for id 0 is set to something other than
the 'background' class. the 'background' class.
""" """
cache_files = dataset_util.get_cache_files_coco(data_dir, cache_dir) tfrecord_cache_files = dataset_util.get_cache_files_coco(
if not dataset_util.is_cached(cache_files): data_dir, cache_dir
)
if not tfrecord_cache_files.is_cached():
label_map = dataset_util.get_label_map_coco(data_dir) label_map = dataset_util.get_label_map_coco(data_dir)
cache_writer = dataset_util.COCOCacheFilesWriter( cache_writer = dataset_util.COCOCacheFilesWriter(
label_map=label_map, max_num_images=max_num_images label_map=label_map, max_num_images=max_num_images
) )
cache_writer.write_files(cache_files, data_dir) cache_writer.write_files(tfrecord_cache_files, data_dir)
return cls.from_cache(cache_files.cache_prefix) return cls.from_cache(tfrecord_cache_files)
@classmethod @classmethod
def from_pascal_voc_folder( def from_pascal_voc_folder(
@@ -134,47 +136,48 @@ class Dataset(classification_dataset.ClassificationDataset):
Raises: Raises:
ValueError: if the input data directory is empty. ValueError: if the input data directory is empty.
""" """
cache_files = dataset_util.get_cache_files_pascal_voc(data_dir, cache_dir) tfrecord_cache_files = dataset_util.get_cache_files_pascal_voc(
if not dataset_util.is_cached(cache_files): data_dir, cache_dir
)
if not tfrecord_cache_files.is_cached():
label_map = dataset_util.get_label_map_pascal_voc(data_dir) label_map = dataset_util.get_label_map_pascal_voc(data_dir)
cache_writer = dataset_util.PascalVocCacheFilesWriter( cache_writer = dataset_util.PascalVocCacheFilesWriter(
label_map=label_map, max_num_images=max_num_images label_map=label_map, max_num_images=max_num_images
) )
cache_writer.write_files(cache_files, data_dir) cache_writer.write_files(tfrecord_cache_files, data_dir)
return cls.from_cache(cache_files.cache_prefix) return cls.from_cache(tfrecord_cache_files)
@classmethod @classmethod
def from_cache(cls, cache_prefix: str) -> 'Dataset': def from_cache(
cls, tfrecord_cache_files: cache_files.TFRecordCacheFiles
) -> 'Dataset':
"""Loads the TFRecord data from cache. """Loads the TFRecord data from cache.
Args: Args:
cache_prefix: The cache prefix including the cache directory and the cache tfrecord_cache_files: The TFRecordCacheFiles object containing the already
prefix filename, e.g: '/tmp/cache/train'. cached TFRecord and metadata files.
Returns: Returns:
ObjectDetectorDataset object. ObjectDetectorDataset object.
Raises:
ValueError if tfrecord_cache_files are not already cached.
""" """
# Get TFRecord Files if not tfrecord_cache_files.is_cached():
tfrecord_file_pattern = cache_prefix + '*.tfrecord' raise ValueError(
matched_files = tf.io.gfile.glob(tfrecord_file_pattern) 'Cache files must be already cached to use the from_cache method.'
if not matched_files: )
raise ValueError('TFRecord files are empty.')
# Load meta_data. metadata = tfrecord_cache_files.load_metadata()
meta_data_file = cache_prefix + dataset_util.META_DATA_FILE_SUFFIX
if not tf.io.gfile.exists(meta_data_file):
raise ValueError("Metadata file %s doesn't exist." % meta_data_file)
with tf.io.gfile.GFile(meta_data_file, 'r') as f:
meta_data = yaml.load(f, Loader=yaml.FullLoader)
dataset = tf.data.TFRecordDataset(matched_files) dataset = tf.data.TFRecordDataset(tfrecord_cache_files.tfrecord_files)
decoder = tf_example_decoder.TfExampleDecoder(regenerate_source_id=False) decoder = tf_example_decoder.TfExampleDecoder(regenerate_source_id=False)
dataset = dataset.map(decoder.decode, num_parallel_calls=tf.data.AUTOTUNE) dataset = dataset.map(decoder.decode, num_parallel_calls=tf.data.AUTOTUNE)
label_map = meta_data['label_map'] label_map = metadata['label_map']
label_names = [label_map[k] for k in sorted(label_map.keys())] label_names = [label_map[k] for k in sorted(label_map.keys())]
return Dataset( return Dataset(
dataset=dataset, size=meta_data['size'], label_names=label_names dataset=dataset, label_names=label_names, size=metadata['size']
) )
@@ -15,25 +15,20 @@
import abc import abc
import collections import collections
import dataclasses
import hashlib import hashlib
import json import json
import math import math
import os import os
import tempfile import tempfile
from typing import Any, Dict, List, Mapping, Optional, Sequence from typing import Any, Dict, List, Mapping, Optional
import xml.etree.ElementTree as ET import xml.etree.ElementTree as ET
import tensorflow as tf import tensorflow as tf
import yaml
from mediapipe.model_maker.python.core.data import cache_files
from official.vision.data import tfrecord_lib from official.vision.data import tfrecord_lib
# Suffix of the meta data file name.
META_DATA_FILE_SUFFIX = '_meta_data.yaml'
def _xml_get(node: ET.Element, name: str) -> ET.Element: def _xml_get(node: ET.Element, name: str) -> ET.Element:
"""Gets a named child from an XML Element node. """Gets a named child from an XML Element node.
@@ -71,18 +66,9 @@ def _get_dir_basename(data_dir: str) -> str:
return os.path.basename(os.path.abspath(data_dir)) return os.path.basename(os.path.abspath(data_dir))
@dataclasses.dataclass(frozen=True)
class CacheFiles:
"""Cache files for object detection."""
cache_prefix: str
tfrecord_files: Sequence[str]
meta_data_file: str
def _get_cache_files( def _get_cache_files(
cache_dir: Optional[str], cache_prefix_filename: str, num_shards: int = 10 cache_dir: Optional[str], cache_prefix_filename: str, num_shards: int = 10
) -> CacheFiles: ) -> cache_files.TFRecordCacheFiles:
"""Creates an object of CacheFiles class. """Creates an object of CacheFiles class.
Args: Args:
@@ -96,28 +82,16 @@ def _get_cache_files(
An object of CacheFiles class. An object of CacheFiles class.
""" """
cache_dir = _get_cache_dir_or_create(cache_dir) cache_dir = _get_cache_dir_or_create(cache_dir)
# The cache prefix including the cache directory and the cache prefix return cache_files.TFRecordCacheFiles(
# filename, e.g: '/tmp/cache/train'. cache_prefix_filename=cache_prefix_filename,
cache_prefix = os.path.join(cache_dir, cache_prefix_filename) cache_dir=cache_dir,
tf.compat.v1.logging.info( num_shards=num_shards,
'Cache will be stored in %s with prefix filename %s. Cache_prefix is %s'
% (cache_dir, cache_prefix_filename, cache_prefix)
)
# Cached files including the TFRecord files and the meta data file.
tfrecord_files = [
cache_prefix + '-%05d-of-%05d.tfrecord' % (i, num_shards)
for i in range(num_shards)
]
meta_data_file = cache_prefix + META_DATA_FILE_SUFFIX
return CacheFiles(
cache_prefix=cache_prefix,
tfrecord_files=tuple(tfrecord_files),
meta_data_file=meta_data_file,
) )
def get_cache_files_coco(data_dir: str, cache_dir: str) -> CacheFiles: def get_cache_files_coco(
data_dir: str, cache_dir: str
) -> cache_files.TFRecordCacheFiles:
"""Creates an object of CacheFiles class using a COCO formatted dataset. """Creates an object of CacheFiles class using a COCO formatted dataset.
Args: Args:
@@ -152,7 +126,9 @@ def get_cache_files_coco(data_dir: str, cache_dir: str) -> CacheFiles:
return _get_cache_files(cache_dir, cache_prefix_filename, num_shards) return _get_cache_files(cache_dir, cache_prefix_filename, num_shards)
def get_cache_files_pascal_voc(data_dir: str, cache_dir: str) -> CacheFiles: def get_cache_files_pascal_voc(
data_dir: str, cache_dir: str
) -> cache_files.TFRecordCacheFiles:
"""Gets an object of CacheFiles using a PASCAL VOC formatted dataset. """Gets an object of CacheFiles using a PASCAL VOC formatted dataset.
Args: Args:
@@ -181,14 +157,6 @@ def get_cache_files_pascal_voc(data_dir: str, cache_dir: str) -> CacheFiles:
return _get_cache_files(cache_dir, cache_prefix_filename, num_shards) return _get_cache_files(cache_dir, cache_prefix_filename, num_shards)
def is_cached(cache_files: CacheFiles) -> bool:
"""Checks whether cache files are already cached."""
all_cached_files = list(cache_files.tfrecord_files) + [
cache_files.meta_data_file
]
return all(tf.io.gfile.exists(path) for path in all_cached_files)
class CacheFilesWriter(abc.ABC): class CacheFilesWriter(abc.ABC):
"""CacheFilesWriter class to write the cached files.""" """CacheFilesWriter class to write the cached files."""
@@ -208,19 +176,22 @@ class CacheFilesWriter(abc.ABC):
self.label_map = label_map self.label_map = label_map
self.max_num_images = max_num_images self.max_num_images = max_num_images
def write_files(self, cache_files: CacheFiles, *args, **kwargs) -> None: def write_files(
"""Writes TFRecord and meta_data files. self,
tfrecord_cache_files: cache_files.TFRecordCacheFiles,
*args,
**kwargs,
) -> None:
"""Writes TFRecord and metadata files.
Args: Args:
cache_files: CacheFiles object including a list of TFRecord files and the tfrecord_cache_files: TFRecordCacheFiles object including a list of
meta data yaml file to save the meta_data including data size and TFRecord files and the meta data yaml file to save the metadata
label_map. including data size and label_map.
*args: Non-keyword of parameters used in the `_get_example` method. *args: Non-keyword of parameters used in the `_get_example` method.
**kwargs: Keyword parameters used in the `_get_example` method. **kwargs: Keyword parameters used in the `_get_example` method.
""" """
writers = [ writers = tfrecord_cache_files.get_writers()
tf.io.TFRecordWriter(path) for path in cache_files.tfrecord_files
]
# Writes tf.Example into TFRecord files. # Writes tf.Example into TFRecord files.
size = 0 size = 0
@@ -235,10 +206,9 @@ class CacheFilesWriter(abc.ABC):
for writer in writers: for writer in writers:
writer.close() writer.close()
# Writes meta_data into meta_data_file. # Writes metadata into metadata_file.
meta_data = {'size': size, 'label_map': self.label_map} metadata = {'size': size, 'label_map': self.label_map}
with tf.io.gfile.GFile(cache_files.meta_data_file, 'w') as f: tfrecord_cache_files.save_metadata(metadata)
yaml.dump(meta_data, f)
@abc.abstractmethod @abc.abstractmethod
def _get_example(self, *args, **kwargs): def _get_example(self, *args, **kwargs):
@@ -19,7 +19,6 @@ import shutil
from unittest import mock as unittest_mock from unittest import mock as unittest_mock
import tensorflow as tf import tensorflow as tf
import yaml
from mediapipe.model_maker.python.vision.core import test_utils from mediapipe.model_maker.python.vision.core import test_utils
from mediapipe.model_maker.python.vision.object_detector import dataset_util from mediapipe.model_maker.python.vision.object_detector import dataset_util
@@ -30,13 +29,10 @@ class DatasetUtilTest(tf.test.TestCase):
def _assert_cache_files_equal(self, cf1, cf2): def _assert_cache_files_equal(self, cf1, cf2):
self.assertEqual(cf1.cache_prefix, cf2.cache_prefix) self.assertEqual(cf1.cache_prefix, cf2.cache_prefix)
self.assertCountEqual(cf1.tfrecord_files, cf2.tfrecord_files) self.assertEqual(cf1.num_shards, cf2.num_shards)
self.assertEqual(cf1.meta_data_file, cf2.meta_data_file)
def _assert_cache_files_not_equal(self, cf1, cf2): def _assert_cache_files_not_equal(self, cf1, cf2):
self.assertNotEqual(cf1.cache_prefix, cf2.cache_prefix) self.assertNotEqual(cf1.cache_prefix, cf2.cache_prefix)
self.assertNotEqual(cf1.tfrecord_files, cf2.tfrecord_files)
self.assertNotEqual(cf1.meta_data_file, cf2.meta_data_file)
def _get_cache_files_and_assert_neq_fn(self, cache_files_fn): def _get_cache_files_and_assert_neq_fn(self, cache_files_fn):
def get_cache_files_and_assert_neq(cf, data_dir, cache_dir): def get_cache_files_and_assert_neq(cf, data_dir, cache_dir):
@@ -57,7 +53,7 @@ class DatasetUtilTest(tf.test.TestCase):
self.assertEqual( self.assertEqual(
cache_files.tfrecord_files[0], '/tmp/train-00000-of-00001.tfrecord' cache_files.tfrecord_files[0], '/tmp/train-00000-of-00001.tfrecord'
) )
self.assertEqual(cache_files.meta_data_file, '/tmp/train_meta_data.yaml') self.assertEqual(cache_files.metadata_file, '/tmp/train_metadata.yaml')
def test_matching_get_cache_files_coco(self): def test_matching_get_cache_files_coco(self):
cache_dir = self.create_tempdir() cache_dir = self.create_tempdir()
@@ -118,7 +114,7 @@ class DatasetUtilTest(tf.test.TestCase):
self.assertEqual( self.assertEqual(
cache_files.tfrecord_files[0], '/tmp/train-00000-of-00001.tfrecord' cache_files.tfrecord_files[0], '/tmp/train-00000-of-00001.tfrecord'
) )
self.assertEqual(cache_files.meta_data_file, '/tmp/train_meta_data.yaml') self.assertEqual(cache_files.metadata_file, '/tmp/train_metadata.yaml')
def test_matching_get_cache_files_pascal_voc(self): def test_matching_get_cache_files_pascal_voc(self):
cache_dir = self.create_tempdir() cache_dir = self.create_tempdir()
@@ -173,13 +169,13 @@ class DatasetUtilTest(tf.test.TestCase):
cache_files = dataset_util.get_cache_files_coco( cache_files = dataset_util.get_cache_files_coco(
tasks_test_utils.get_test_data_path('coco_data'), cache_dir=tempdir tasks_test_utils.get_test_data_path('coco_data'), cache_dir=tempdir
) )
self.assertFalse(dataset_util.is_cached(cache_files)) self.assertFalse(cache_files.is_cached())
with open(cache_files.tfrecord_files[0], 'w') as f: with open(cache_files.tfrecord_files[0], 'w') as f:
f.write('test') f.write('test')
self.assertFalse(dataset_util.is_cached(cache_files)) self.assertFalse(cache_files.is_cached())
with open(cache_files.meta_data_file, 'w') as f: with open(cache_files.metadata_file, 'w') as f:
f.write('test') f.write('test')
self.assertTrue(dataset_util.is_cached(cache_files)) self.assertTrue(cache_files.is_cached())
def test_get_label_map_coco(self): def test_get_label_map_coco(self):
coco_dir = tasks_test_utils.get_test_data_path('coco_data') coco_dir = tasks_test_utils.get_test_data_path('coco_data')
@@ -203,13 +199,11 @@ class DatasetUtilTest(tf.test.TestCase):
self.assertTrue(os.path.isfile(cache_files.tfrecord_files[0])) self.assertTrue(os.path.isfile(cache_files.tfrecord_files[0]))
self.assertGreater(os.path.getsize(cache_files.tfrecord_files[0]), 0) self.assertGreater(os.path.getsize(cache_files.tfrecord_files[0]), 0)
# Checks the meta_data file # Checks the metadata file
self.assertTrue(os.path.isfile(cache_files.meta_data_file)) self.assertTrue(os.path.isfile(cache_files.metadata_file))
self.assertGreater(os.path.getsize(cache_files.meta_data_file), 0) self.assertGreater(os.path.getsize(cache_files.metadata_file), 0)
with tf.io.gfile.GFile(cache_files.meta_data_file, 'r') as f: metadata_dict = cache_files.load_metadata()
meta_data_dict = yaml.load(f, Loader=yaml.FullLoader) self.assertEqual(metadata_dict['size'], expected_size)
# Size is 3 because some examples are skipped for having poor bboxes
self.assertEqual(meta_data_dict['size'], expected_size)
def test_coco_cache_files_writer(self): def test_coco_cache_files_writer(self):
tempdir = self.create_tempdir() tempdir = self.create_tempdir()
+1 -1
View File
@@ -5,4 +5,4 @@ opencv-python
tensorflow>=2.10 tensorflow>=2.10
tensorflow-datasets tensorflow-datasets
tensorflow-hub tensorflow-hub
tf-models-official==2.11.6 tf-models-official>=2.13.1
+2 -2
View File
@@ -13,17 +13,17 @@
# limitations under the License. # limitations under the License.
"""MediaPipe solution drawing utils.""" """MediaPipe solution drawing utils."""
import dataclasses
import math import math
from typing import List, Mapping, Optional, Tuple, Union from typing import List, Mapping, Optional, Tuple, Union
import cv2 import cv2
import dataclasses
import matplotlib.pyplot as plt import matplotlib.pyplot as plt
import numpy as np import numpy as np
from mediapipe.framework.formats import detection_pb2 from mediapipe.framework.formats import detection_pb2
from mediapipe.framework.formats import location_data_pb2
from mediapipe.framework.formats import landmark_pb2 from mediapipe.framework.formats import landmark_pb2
from mediapipe.framework.formats import location_data_pb2
_PRESENCE_THRESHOLD = 0.5 _PRESENCE_THRESHOLD = 0.5
_VISIBILITY_THRESHOLD = 0.5 _VISIBILITY_THRESHOLD = 0.5
@@ -20,7 +20,6 @@ import cv2
import numpy as np import numpy as np
from google.protobuf import text_format from google.protobuf import text_format
from mediapipe.framework.formats import detection_pb2 from mediapipe.framework.formats import detection_pb2
from mediapipe.framework.formats import landmark_pb2 from mediapipe.framework.formats import landmark_pb2
from mediapipe.python.solutions import drawing_utils from mediapipe.python.solutions import drawing_utils
@@ -0,0 +1,29 @@
# TODO: describe this package.
# Copyright 2022 The MediaPipe Authors.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
package(default_visibility = ["//mediapipe/tasks:internal"])
licenses(["notice"])
cc_library(
name = "category",
hdrs = ["category.h"],
)
cc_library(
name = "classification_result",
hdrs = ["classification_result.h"],
)
@@ -0,0 +1,42 @@
/* Copyright 2023 The MediaPipe Authors.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#ifndef MEDIAPIPE_TASKS_C_COMPONENTS_CONTAINERS_CATEGORY_H_
#define MEDIAPIPE_TASKS_C_COMPONENTS_CONTAINERS_CATEGORY_H_
// Defines a single classification result.
//
// The label maps packed into the TFLite Model Metadata [1] are used to populate
// the 'category_name' and 'display_name' fields.
//
// [1]: https://www.tensorflow.org/lite/convert/metadata
struct Category {
// The index of the category in the classification model output.
int index;
// The score for this category, e.g. (but not necessarily) a probability in
// [0,1].
float score;
// The optional ID for the category, read from the label map packed in the
// TFLite Model Metadata if present. Not necessarily human-readable.
char* category_name;
// The optional human-readable name for the category, read from the label map
// packed in the TFLite Model Metadata if present.
char* display_name;
};
#endif // MEDIAPIPE_TASKS_C_COMPONENTS_CONTAINERS_CATEGORY_H_
@@ -0,0 +1,60 @@
/* Copyright 2023 The MediaPipe Authors.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#ifndef MEDIAPIPE_TASKS_C_COMPONENTS_CONTAINERS_CLASSIFICATION_RESULT_H_
#define MEDIAPIPE_TASKS_C_COMPONENTS_CONTAINERS_CLASSIFICATION_RESULT_H_
#include <stdbool.h>
#include <stdint.h>
// Defines classification results for a given classifier head.
struct Classifications {
// The array of predicted categories, usually sorted by descending scores,
// e.g. from high to low probability.
struct Category* categories;
// The number of elements in the categories array.
uint32_t categories_count;
// The index of the classifier head (i.e. output tensor) these categories
// refer to. This is useful for multi-head models.
int head_index;
// The optional name of the classifier head, as provided in the TFLite Model
// Metadata [1] if present. This is useful for multi-head models.
//
// [1]: https://www.tensorflow.org/lite/convert/metadata
char* head_name;
};
// Defines classification results of a model.
struct ClassificationResult {
// The classification results for each head of the model.
struct Classifications* classifications;
// The number of classifications in the classifications array.
uint32_t classifications_count;
// The optional timestamp (in milliseconds) of the start of the chunk of data
// corresponding to these results.
//
// This is only used for classification on time series (e.g. audio
// classification). In these use cases, the amount of data to process might
// exceed the maximum size that the model can process: to solve this, the
// input data is split into multiple chunks starting at different timestamps.
int64_t timestamp_ms;
// Specifies whether the timestamp contains a valid value.
bool has_timestamp_ms;
};
#endif // MEDIAPIPE_TASKS_C_COMPONENTS_CONTAINERS_CLASSIFICATION_RESULT_H_
@@ -0,0 +1,22 @@
# Copyright 2023 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.
package(default_visibility = ["//mediapipe/tasks:internal"])
licenses(["notice"])
cc_library(
name = "classifier_options",
hdrs = ["classifier_options.h"],
)
@@ -0,0 +1,51 @@
/* Copyright 2023 The MediaPipe Authors.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#ifndef MEDIAPIPE_TASKS_C_COMPONENTS_PROCESSORS_CLASSIFIER_OPTIONS_H_
#define MEDIAPIPE_TASKS_C_COMPONENTS_PROCESSORS_CLASSIFIER_OPTIONS_H_
#include <stdint.h>
// Classifier options for MediaPipe C classification Tasks.
struct ClassifierOptions {
// The locale to use for display names specified through the TFLite Model
// Metadata, if any. Defaults to English.
char* display_names_locale;
// The maximum number of top-scored classification results to return. If < 0,
// all available results will be returned. If 0, an invalid argument error is
// returned.
int max_results;
// Score threshold to override the one provided in the model metadata (if
// any). Results below this value are rejected.
float score_threshold;
// The allowlist of category names. If non-empty, detection results whose
// category name is not in this set will be filtered out. Duplicate or unknown
// category names are ignored. Mutually exclusive with category_denylist.
char** category_allowlist;
// The number of elements in the category allowlist.
uint32_t category_allowlist_count;
// The denylist of category names. If non-empty, detection results whose
// category name is in this set will be filtered out. Duplicate or unknown
// category names are ignored. Mutually exclusive with category_allowlist.
char** category_denylist = {};
// The number of elements in the category denylist.
uint32_t category_denylist_count;
};
#endif // MEDIAPIPE_TASKS_C_COMPONENTS_PROCESSORS_CLASSIFIER_OPTIONS_H_
+22
View File
@@ -0,0 +1,22 @@
# Copyright 2023 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.
package(default_visibility = ["//mediapipe/tasks:internal"])
licenses(["notice"])
cc_library(
name = "base_options",
hdrs = ["base_options.h"],
)
+28
View File
@@ -0,0 +1,28 @@
/* Copyright 2023 The MediaPipe Authors.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#ifndef MEDIAPIPE_TASKS_C_CORE_BASE_OPTIONS_H_
#define MEDIAPIPE_TASKS_C_CORE_BASE_OPTIONS_H_
// Base options for MediaPipe C Tasks.
struct BaseOptions {
// The model asset file contents as a string.
char* model_asset_buffer;
// The path to the model asset to open and mmap in memory.
char* model_asset_path;
};
#endif // MEDIAPIPE_TASKS_C_CORE_BASE_OPTIONS_H_
@@ -0,0 +1,28 @@
# Copyright 2023 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.
package(default_visibility = ["//mediapipe/tasks:internal"])
licenses(["notice"])
cc_library(
name = "text_classifier",
hdrs = ["text_classifier.h"],
visibility = ["//visibility:public"],
deps = [
"//mediapipe/tasks/c/components/containers:classification_result",
"//mediapipe/tasks/c/components/processors:classifier_options",
"//mediapipe/tasks/c/core:base_options",
],
)
@@ -0,0 +1,46 @@
/* Copyright 2023 The MediaPipe Authors.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#ifndef MEDIAPIPE_TASKS_C_TEXT_TEXT_CLASSIFIER_TEXT_CLASSIFIER_H_
#define MEDIAPIPE_TASKS_C_TEXT_TEXT_CLASSIFIER_TEXT_CLASSIFIER_H_
#include "mediapipe/tasks/c/components/containers/classification_result.h"
#include "mediapipe/tasks/c/components/processors/classifier_options.h"
#include "mediapipe/tasks/c/core/base_options.h"
typedef ClassificationResult TextClassifierResult;
// The options for configuring a MediaPipe text classifier task.
struct TextClassifierOptions {
// Base options for configuring MediaPipe Tasks, such as specifying the model
// file with metadata, accelerator options, op resolver, etc.
struct BaseOptions base_options;
// Options for configuring the classifier behavior, such as score threshold,
// number of results, etc.
struct ClassifierOptions classifier_options;
};
// Creates a TextClassifier from the provided `options`.
void* text_classifier_create(struct TextClassifierOptions options);
// Performs classification on the input `text`.
TextClassifierResult text_classifier_classify(void* classifier,
char* utf8_text);
// Shuts down the TextClassifier when all the work is done. Frees all memory.
void text_classifier_close(void* classifier);
#endif // MEDIAPIPE_TASKS_C_TEXT_TEXT_CLASSIFIER_TEXT_CLASSIFIER_H_
@@ -217,3 +217,8 @@ cc_library(
], ],
alwayslink = 1, alwayslink = 1,
) )
cc_library(
name = "face_landmarks_connections",
hdrs = ["face_landmarks_connections.h"],
)
@@ -0,0 +1,651 @@
/* Copyright 2023 The MediaPipe Authors.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#ifndef MEDIAPIPE_TASKS_CC_VISION_FACE_LANDMARKER_FACE_LANDMARKS_CONNECTIONS_H_
#define MEDIAPIPE_TASKS_CC_VISION_FACE_LANDMARKER_FACE_LANDMARKS_CONNECTIONS_H_
#include <array>
namespace mediapipe {
namespace tasks {
namespace vision {
namespace face_landmarker {
struct FaceLandmarksConnections {
static constexpr std::array<std::array<int, 2>, 40> kFaceLandmarksLips{
{{61, 146}, {146, 91}, {91, 181}, {181, 84}, {84, 17}, {17, 314},
{314, 405}, {405, 321}, {321, 375}, {375, 291}, {61, 185}, {185, 40},
{40, 39}, {39, 37}, {37, 0}, {0, 267}, {267, 269}, {269, 270},
{270, 409}, {409, 291}, {78, 95}, {95, 88}, {88, 178}, {178, 87},
{87, 14}, {14, 317}, {317, 402}, {402, 318}, {318, 324}, {324, 308},
{78, 191}, {191, 80}, {80, 81}, {81, 82}, {82, 13}, {13, 312},
{312, 311}, {311, 310}, {310, 415}, {415, 308}}};
static constexpr std::array<std::array<int, 2>, 16> kFaceLandmarksLeftEye{
{{263, 249},
{249, 390},
{390, 373},
{373, 374},
{374, 380},
{380, 381},
{381, 382},
{382, 362},
{263, 466},
{466, 388},
{388, 387},
{387, 386},
{386, 385},
{385, 384},
{384, 398},
{398, 362}}};
static constexpr std::array<std::array<int, 2>, 8> kFaceLandmarksLeftEyeBrow{
{{276, 283},
{283, 282},
{282, 295},
{295, 285},
{300, 293},
{293, 334},
{334, 296},
{296, 336}}};
static constexpr std::array<std::array<int, 2>, 4> kFaceLandmarksLeftIris{
{{474, 475}, {475, 476}, {476, 477}, {477, 474}}};
static constexpr std::array<std::array<int, 2>, 16> kFaceLandmarksRightEye{
{{33, 7},
{7, 163},
{163, 144},
{144, 145},
{145, 153},
{153, 154},
{154, 155},
{155, 133},
{33, 246},
{246, 161},
{161, 160},
{160, 159},
{159, 158},
{158, 157},
{157, 173},
{173, 133}}};
static constexpr std::array<std::array<int, 2>, 8> kFaceLandmarksRightEyeBrow{
{{46, 53},
{53, 52},
{52, 65},
{65, 55},
{70, 63},
{63, 105},
{105, 66},
{66, 107}}};
static constexpr std::array<std::array<int, 2>, 4> kFaceLandmarksRightIris{
{{469, 470}, {470, 471}, {471, 472}, {472, 469}}};
static constexpr std::array<std::array<int, 2>, 36> kFaceLandmarksFaceOval{
{{10, 338}, {338, 297}, {297, 332}, {332, 284}, {284, 251}, {251, 389},
{389, 356}, {356, 454}, {454, 323}, {323, 361}, {361, 288}, {288, 397},
{397, 365}, {365, 379}, {379, 378}, {378, 400}, {400, 377}, {377, 152},
{152, 148}, {148, 176}, {176, 149}, {149, 150}, {150, 136}, {136, 172},
{172, 58}, {58, 132}, {132, 93}, {93, 234}, {234, 127}, {127, 162},
{162, 21}, {21, 54}, {54, 103}, {103, 67}, {67, 109}, {109, 10}}};
// Lips + Left Eye + Left Eye Brows + Right Eye + Right Eye Brows + Face Oval.
static constexpr std::array<std::array<int, 2>, 132> kFaceLandmarksConnectors{
{{61, 146}, {146, 91}, {91, 181}, {181, 84}, {84, 17}, {17, 314},
{314, 405}, {405, 321}, {321, 375}, {375, 291}, {61, 185}, {185, 40},
{40, 39}, {39, 37}, {37, 0}, {0, 267}, {267, 269}, {269, 270},
{270, 409}, {409, 291}, {78, 95}, {95, 88}, {88, 178}, {178, 87},
{87, 14}, {14, 317}, {317, 402}, {402, 318}, {318, 324}, {324, 308},
{78, 191}, {191, 80}, {80, 81}, {81, 82}, {82, 13}, {13, 312},
{312, 311}, {311, 310}, {310, 415}, {415, 30}, {263, 249}, {249, 390},
{390, 373}, {373, 374}, {374, 380}, {380, 381}, {381, 382}, {382, 362},
{263, 466}, {466, 388}, {388, 387}, {387, 386}, {386, 385}, {385, 384},
{384, 398}, {398, 362}, {276, 283}, {283, 282}, {282, 295}, {295, 285},
{300, 293}, {293, 334}, {334, 296}, {296, 336}, {33, 7}, {7, 163},
{163, 144}, {144, 145}, {145, 153}, {153, 154}, {154, 155}, {155, 133},
{33, 246}, {246, 161}, {161, 160}, {160, 159}, {159, 158}, {158, 157},
{157, 173}, {173, 13}, {46, 53}, {53, 52}, {52, 65}, {65, 55},
{70, 63}, {63, 105}, {105, 66}, {66, 107}, {10, 338}, {338, 297},
{297, 332}, {332, 284}, {284, 251}, {251, 389}, {389, 356}, {356, 454},
{454, 323}, {323, 361}, {361, 288}, {288, 397}, {397, 365}, {365, 379},
{379, 378}, {378, 400}, {400, 377}, {377, 152}, {152, 148}, {148, 176},
{176, 149}, {149, 150}, {150, 136}, {136, 172}, {172, 58}, {58, 132},
{132, 93}, {93, 234}, {234, 127}, {127, 162}, {162, 21}, {21, 54},
{54, 103}, {103, 67}, {67, 109}, {109, 10}}};
static constexpr std::array<std::array<int, 2>, 2556>
kFaceLandmarksTesselation{
{{127, 34}, {34, 139}, {139, 127}, {11, 0}, {0, 37},
{37, 11}, {232, 231}, {231, 120}, {120, 232}, {72, 37},
{37, 39}, {39, 72}, {128, 121}, {121, 47}, {47, 128},
{232, 121}, {121, 128}, {128, 232}, {104, 69}, {69, 67},
{67, 104}, {175, 171}, {171, 148}, {148, 175}, {118, 50},
{50, 101}, {101, 118}, {73, 39}, {39, 40}, {40, 73},
{9, 151}, {151, 108}, {108, 9}, {48, 115}, {115, 131},
{131, 48}, {194, 204}, {204, 211}, {211, 194}, {74, 40},
{40, 185}, {185, 74}, {80, 42}, {42, 183}, {183, 80},
{40, 92}, {92, 186}, {186, 40}, {230, 229}, {229, 118},
{118, 230}, {202, 212}, {212, 214}, {214, 202}, {83, 18},
{18, 17}, {17, 83}, {76, 61}, {61, 146}, {146, 76},
{160, 29}, {29, 30}, {30, 160}, {56, 157}, {157, 173},
{173, 56}, {106, 204}, {204, 194}, {194, 106}, {135, 214},
{214, 192}, {192, 135}, {203, 165}, {165, 98}, {98, 203},
{21, 71}, {71, 68}, {68, 21}, {51, 45}, {45, 4},
{4, 51}, {144, 24}, {24, 23}, {23, 144}, {77, 146},
{146, 91}, {91, 77}, {205, 50}, {50, 187}, {187, 205},
{201, 200}, {200, 18}, {18, 201}, {91, 106}, {106, 182},
{182, 91}, {90, 91}, {91, 181}, {181, 90}, {85, 84},
{84, 17}, {17, 85}, {206, 203}, {203, 36}, {36, 206},
{148, 171}, {171, 140}, {140, 148}, {92, 40}, {40, 39},
{39, 92}, {193, 189}, {189, 244}, {244, 193}, {159, 158},
{158, 28}, {28, 159}, {247, 246}, {246, 161}, {161, 247},
{236, 3}, {3, 196}, {196, 236}, {54, 68}, {68, 104},
{104, 54}, {193, 168}, {168, 8}, {8, 193}, {117, 228},
{228, 31}, {31, 117}, {189, 193}, {193, 55}, {55, 189},
{98, 97}, {97, 99}, {99, 98}, {126, 47}, {47, 100},
{100, 126}, {166, 79}, {79, 218}, {218, 166}, {155, 154},
{154, 26}, {26, 155}, {209, 49}, {49, 131}, {131, 209},
{135, 136}, {136, 150}, {150, 135}, {47, 126}, {126, 217},
{217, 47}, {223, 52}, {52, 53}, {53, 223}, {45, 51},
{51, 134}, {134, 45}, {211, 170}, {170, 140}, {140, 211},
{67, 69}, {69, 108}, {108, 67}, {43, 106}, {106, 91},
{91, 43}, {230, 119}, {119, 120}, {120, 230}, {226, 130},
{130, 247}, {247, 226}, {63, 53}, {53, 52}, {52, 63},
{238, 20}, {20, 242}, {242, 238}, {46, 70}, {70, 156},
{156, 46}, {78, 62}, {62, 96}, {96, 78}, {46, 53},
{53, 63}, {63, 46}, {143, 34}, {34, 227}, {227, 143},
{123, 117}, {117, 111}, {111, 123}, {44, 125}, {125, 19},
{19, 44}, {236, 134}, {134, 51}, {51, 236}, {216, 206},
{206, 205}, {205, 216}, {154, 153}, {153, 22}, {22, 154},
{39, 37}, {37, 167}, {167, 39}, {200, 201}, {201, 208},
{208, 200}, {36, 142}, {142, 100}, {100, 36}, {57, 212},
{212, 202}, {202, 57}, {20, 60}, {60, 99}, {99, 20},
{28, 158}, {158, 157}, {157, 28}, {35, 226}, {226, 113},
{113, 35}, {160, 159}, {159, 27}, {27, 160}, {204, 202},
{202, 210}, {210, 204}, {113, 225}, {225, 46}, {46, 113},
{43, 202}, {202, 204}, {204, 43}, {62, 76}, {76, 77},
{77, 62}, {137, 123}, {123, 116}, {116, 137}, {41, 38},
{38, 72}, {72, 41}, {203, 129}, {129, 142}, {142, 203},
{64, 98}, {98, 240}, {240, 64}, {49, 102}, {102, 64},
{64, 49}, {41, 73}, {73, 74}, {74, 41}, {212, 216},
{216, 207}, {207, 212}, {42, 74}, {74, 184}, {184, 42},
{169, 170}, {170, 211}, {211, 169}, {170, 149}, {149, 176},
{176, 170}, {105, 66}, {66, 69}, {69, 105}, {122, 6},
{6, 168}, {168, 122}, {123, 147}, {147, 187}, {187, 123},
{96, 77}, {77, 90}, {90, 96}, {65, 55}, {55, 107},
{107, 65}, {89, 90}, {90, 180}, {180, 89}, {101, 100},
{100, 120}, {120, 101}, {63, 105}, {105, 104}, {104, 63},
{93, 137}, {137, 227}, {227, 93}, {15, 86}, {86, 85},
{85, 15}, {129, 102}, {102, 49}, {49, 129}, {14, 87},
{87, 86}, {86, 14}, {55, 8}, {8, 9}, {9, 55},
{100, 47}, {47, 121}, {121, 100}, {145, 23}, {23, 22},
{22, 145}, {88, 89}, {89, 179}, {179, 88}, {6, 122},
{122, 196}, {196, 6}, {88, 95}, {95, 96}, {96, 88},
{138, 172}, {172, 136}, {136, 138}, {215, 58}, {58, 172},
{172, 215}, {115, 48}, {48, 219}, {219, 115}, {42, 80},
{80, 81}, {81, 42}, {195, 3}, {3, 51}, {51, 195},
{43, 146}, {146, 61}, {61, 43}, {171, 175}, {175, 199},
{199, 171}, {81, 82}, {82, 38}, {38, 81}, {53, 46},
{46, 225}, {225, 53}, {144, 163}, {163, 110}, {110, 144},
{52, 65}, {65, 66}, {66, 52}, {229, 228}, {228, 117},
{117, 229}, {34, 127}, {127, 234}, {234, 34}, {107, 108},
{108, 69}, {69, 107}, {109, 108}, {108, 151}, {151, 109},
{48, 64}, {64, 235}, {235, 48}, {62, 78}, {78, 191},
{191, 62}, {129, 209}, {209, 126}, {126, 129}, {111, 35},
{35, 143}, {143, 111}, {117, 123}, {123, 50}, {50, 117},
{222, 65}, {65, 52}, {52, 222}, {19, 125}, {125, 141},
{141, 19}, {221, 55}, {55, 65}, {65, 221}, {3, 195},
{195, 197}, {197, 3}, {25, 7}, {7, 33}, {33, 25},
{220, 237}, {237, 44}, {44, 220}, {70, 71}, {71, 139},
{139, 70}, {122, 193}, {193, 245}, {245, 122}, {247, 130},
{130, 33}, {33, 247}, {71, 21}, {21, 162}, {162, 71},
{170, 169}, {169, 150}, {150, 170}, {188, 174}, {174, 196},
{196, 188}, {216, 186}, {186, 92}, {92, 216}, {2, 97},
{97, 167}, {167, 2}, {141, 125}, {125, 241}, {241, 141},
{164, 167}, {167, 37}, {37, 164}, {72, 38}, {38, 12},
{12, 72}, {38, 82}, {82, 13}, {13, 38}, {63, 68},
{68, 71}, {71, 63}, {226, 35}, {35, 111}, {111, 226},
{101, 50}, {50, 205}, {205, 101}, {206, 92}, {92, 165},
{165, 206}, {209, 198}, {198, 217}, {217, 209}, {165, 167},
{167, 97}, {97, 165}, {220, 115}, {115, 218}, {218, 220},
{133, 112}, {112, 243}, {243, 133}, {239, 238}, {238, 241},
{241, 239}, {214, 135}, {135, 169}, {169, 214}, {190, 173},
{173, 133}, {133, 190}, {171, 208}, {208, 32}, {32, 171},
{125, 44}, {44, 237}, {237, 125}, {86, 87}, {87, 178},
{178, 86}, {85, 86}, {86, 179}, {179, 85}, {84, 85},
{85, 180}, {180, 84}, {83, 84}, {84, 181}, {181, 83},
{201, 83}, {83, 182}, {182, 201}, {137, 93}, {93, 132},
{132, 137}, {76, 62}, {62, 183}, {183, 76}, {61, 76},
{76, 184}, {184, 61}, {57, 61}, {61, 185}, {185, 57},
{212, 57}, {57, 186}, {186, 212}, {214, 207}, {207, 187},
{187, 214}, {34, 143}, {143, 156}, {156, 34}, {79, 239},
{239, 237}, {237, 79}, {123, 137}, {137, 177}, {177, 123},
{44, 1}, {1, 4}, {4, 44}, {201, 194}, {194, 32},
{32, 201}, {64, 102}, {102, 129}, {129, 64}, {213, 215},
{215, 138}, {138, 213}, {59, 166}, {166, 219}, {219, 59},
{242, 99}, {99, 97}, {97, 242}, {2, 94}, {94, 141},
{141, 2}, {75, 59}, {59, 235}, {235, 75}, {24, 110},
{110, 228}, {228, 24}, {25, 130}, {130, 226}, {226, 25},
{23, 24}, {24, 229}, {229, 23}, {22, 23}, {23, 230},
{230, 22}, {26, 22}, {22, 231}, {231, 26}, {112, 26},
{26, 232}, {232, 112}, {189, 190}, {190, 243}, {243, 189},
{221, 56}, {56, 190}, {190, 221}, {28, 56}, {56, 221},
{221, 28}, {27, 28}, {28, 222}, {222, 27}, {29, 27},
{27, 223}, {223, 29}, {30, 29}, {29, 224}, {224, 30},
{247, 30}, {30, 225}, {225, 247}, {238, 79}, {79, 20},
{20, 238}, {166, 59}, {59, 75}, {75, 166}, {60, 75},
{75, 240}, {240, 60}, {147, 177}, {177, 215}, {215, 147},
{20, 79}, {79, 166}, {166, 20}, {187, 147}, {147, 213},
{213, 187}, {112, 233}, {233, 244}, {244, 112}, {233, 128},
{128, 245}, {245, 233}, {128, 114}, {114, 188}, {188, 128},
{114, 217}, {217, 174}, {174, 114}, {131, 115}, {115, 220},
{220, 131}, {217, 198}, {198, 236}, {236, 217}, {198, 131},
{131, 134}, {134, 198}, {177, 132}, {132, 58}, {58, 177},
{143, 35}, {35, 124}, {124, 143}, {110, 163}, {163, 7},
{7, 110}, {228, 110}, {110, 25}, {25, 228}, {356, 389},
{389, 368}, {368, 356}, {11, 302}, {302, 267}, {267, 11},
{452, 350}, {350, 349}, {349, 452}, {302, 303}, {303, 269},
{269, 302}, {357, 343}, {343, 277}, {277, 357}, {452, 453},
{453, 357}, {357, 452}, {333, 332}, {332, 297}, {297, 333},
{175, 152}, {152, 377}, {377, 175}, {347, 348}, {348, 330},
{330, 347}, {303, 304}, {304, 270}, {270, 303}, {9, 336},
{336, 337}, {337, 9}, {278, 279}, {279, 360}, {360, 278},
{418, 262}, {262, 431}, {431, 418}, {304, 408}, {408, 409},
{409, 304}, {310, 415}, {415, 407}, {407, 310}, {270, 409},
{409, 410}, {410, 270}, {450, 348}, {348, 347}, {347, 450},
{422, 430}, {430, 434}, {434, 422}, {313, 314}, {314, 17},
{17, 313}, {306, 307}, {307, 375}, {375, 306}, {387, 388},
{388, 260}, {260, 387}, {286, 414}, {414, 398}, {398, 286},
{335, 406}, {406, 418}, {418, 335}, {364, 367}, {367, 416},
{416, 364}, {423, 358}, {358, 327}, {327, 423}, {251, 284},
{284, 298}, {298, 251}, {281, 5}, {5, 4}, {4, 281},
{373, 374}, {374, 253}, {253, 373}, {307, 320}, {320, 321},
{321, 307}, {425, 427}, {427, 411}, {411, 425}, {421, 313},
{313, 18}, {18, 421}, {321, 405}, {405, 406}, {406, 321},
{320, 404}, {404, 405}, {405, 320}, {315, 16}, {16, 17},
{17, 315}, {426, 425}, {425, 266}, {266, 426}, {377, 400},
{400, 369}, {369, 377}, {322, 391}, {391, 269}, {269, 322},
{417, 465}, {465, 464}, {464, 417}, {386, 257}, {257, 258},
{258, 386}, {466, 260}, {260, 388}, {388, 466}, {456, 399},
{399, 419}, {419, 456}, {284, 332}, {332, 333}, {333, 284},
{417, 285}, {285, 8}, {8, 417}, {346, 340}, {340, 261},
{261, 346}, {413, 441}, {441, 285}, {285, 413}, {327, 460},
{460, 328}, {328, 327}, {355, 371}, {371, 329}, {329, 355},
{392, 439}, {439, 438}, {438, 392}, {382, 341}, {341, 256},
{256, 382}, {429, 420}, {420, 360}, {360, 429}, {364, 394},
{394, 379}, {379, 364}, {277, 343}, {343, 437}, {437, 277},
{443, 444}, {444, 283}, {283, 443}, {275, 440}, {440, 363},
{363, 275}, {431, 262}, {262, 369}, {369, 431}, {297, 338},
{338, 337}, {337, 297}, {273, 375}, {375, 321}, {321, 273},
{450, 451}, {451, 349}, {349, 450}, {446, 342}, {342, 467},
{467, 446}, {293, 334}, {334, 282}, {282, 293}, {458, 461},
{461, 462}, {462, 458}, {276, 353}, {353, 383}, {383, 276},
{308, 324}, {324, 325}, {325, 308}, {276, 300}, {300, 293},
{293, 276}, {372, 345}, {345, 447}, {447, 372}, {352, 345},
{345, 340}, {340, 352}, {274, 1}, {1, 19}, {19, 274},
{456, 248}, {248, 281}, {281, 456}, {436, 427}, {427, 425},
{425, 436}, {381, 256}, {256, 252}, {252, 381}, {269, 391},
{391, 393}, {393, 269}, {200, 199}, {199, 428}, {428, 200},
{266, 330}, {330, 329}, {329, 266}, {287, 273}, {273, 422},
{422, 287}, {250, 462}, {462, 328}, {328, 250}, {258, 286},
{286, 384}, {384, 258}, {265, 353}, {353, 342}, {342, 265},
{387, 259}, {259, 257}, {257, 387}, {424, 431}, {431, 430},
{430, 424}, {342, 353}, {353, 276}, {276, 342}, {273, 335},
{335, 424}, {424, 273}, {292, 325}, {325, 307}, {307, 292},
{366, 447}, {447, 345}, {345, 366}, {271, 303}, {303, 302},
{302, 271}, {423, 266}, {266, 371}, {371, 423}, {294, 455},
{455, 460}, {460, 294}, {279, 278}, {278, 294}, {294, 279},
{271, 272}, {272, 304}, {304, 271}, {432, 434}, {434, 427},
{427, 432}, {272, 407}, {407, 408}, {408, 272}, {394, 430},
{430, 431}, {431, 394}, {395, 369}, {369, 400}, {400, 395},
{334, 333}, {333, 299}, {299, 334}, {351, 417}, {417, 168},
{168, 351}, {352, 280}, {280, 411}, {411, 352}, {325, 319},
{319, 320}, {320, 325}, {295, 296}, {296, 336}, {336, 295},
{319, 403}, {403, 404}, {404, 319}, {330, 348}, {348, 349},
{349, 330}, {293, 298}, {298, 333}, {333, 293}, {323, 454},
{454, 447}, {447, 323}, {15, 16}, {16, 315}, {315, 15},
{358, 429}, {429, 279}, {279, 358}, {14, 15}, {15, 316},
{316, 14}, {285, 336}, {336, 9}, {9, 285}, {329, 349},
{349, 350}, {350, 329}, {374, 380}, {380, 252}, {252, 374},
{318, 402}, {402, 403}, {403, 318}, {6, 197}, {197, 419},
{419, 6}, {318, 319}, {319, 325}, {325, 318}, {367, 364},
{364, 365}, {365, 367}, {435, 367}, {367, 397}, {397, 435},
{344, 438}, {438, 439}, {439, 344}, {272, 271}, {271, 311},
{311, 272}, {195, 5}, {5, 281}, {281, 195}, {273, 287},
{287, 291}, {291, 273}, {396, 428}, {428, 199}, {199, 396},
{311, 271}, {271, 268}, {268, 311}, {283, 444}, {444, 445},
{445, 283}, {373, 254}, {254, 339}, {339, 373}, {282, 334},
{334, 296}, {296, 282}, {449, 347}, {347, 346}, {346, 449},
{264, 447}, {447, 454}, {454, 264}, {336, 296}, {296, 299},
{299, 336}, {338, 10}, {10, 151}, {151, 338}, {278, 439},
{439, 455}, {455, 278}, {292, 407}, {407, 415}, {415, 292},
{358, 371}, {371, 355}, {355, 358}, {340, 345}, {345, 372},
{372, 340}, {346, 347}, {347, 280}, {280, 346}, {442, 443},
{443, 282}, {282, 442}, {19, 94}, {94, 370}, {370, 19},
{441, 442}, {442, 295}, {295, 441}, {248, 419}, {419, 197},
{197, 248}, {263, 255}, {255, 359}, {359, 263}, {440, 275},
{275, 274}, {274, 440}, {300, 383}, {383, 368}, {368, 300},
{351, 412}, {412, 465}, {465, 351}, {263, 467}, {467, 466},
{466, 263}, {301, 368}, {368, 389}, {389, 301}, {395, 378},
{378, 379}, {379, 395}, {412, 351}, {351, 419}, {419, 412},
{436, 426}, {426, 322}, {322, 436}, {2, 164}, {164, 393},
{393, 2}, {370, 462}, {462, 461}, {461, 370}, {164, 0},
{0, 267}, {267, 164}, {302, 11}, {11, 12}, {12, 302},
{268, 12}, {12, 13}, {13, 268}, {293, 300}, {300, 301},
{301, 293}, {446, 261}, {261, 340}, {340, 446}, {330, 266},
{266, 425}, {425, 330}, {426, 423}, {423, 391}, {391, 426},
{429, 355}, {355, 437}, {437, 429}, {391, 327}, {327, 326},
{326, 391}, {440, 457}, {457, 438}, {438, 440}, {341, 382},
{382, 362}, {362, 341}, {459, 457}, {457, 461}, {461, 459},
{434, 430}, {430, 394}, {394, 434}, {414, 463}, {463, 362},
{362, 414}, {396, 369}, {369, 262}, {262, 396}, {354, 461},
{461, 457}, {457, 354}, {316, 403}, {403, 402}, {402, 316},
{315, 404}, {404, 403}, {403, 315}, {314, 405}, {405, 404},
{404, 314}, {313, 406}, {406, 405}, {405, 313}, {421, 418},
{418, 406}, {406, 421}, {366, 401}, {401, 361}, {361, 366},
{306, 408}, {408, 407}, {407, 306}, {291, 409}, {409, 408},
{408, 291}, {287, 410}, {410, 409}, {409, 287}, {432, 436},
{436, 410}, {410, 432}, {434, 416}, {416, 411}, {411, 434},
{264, 368}, {368, 383}, {383, 264}, {309, 438}, {438, 457},
{457, 309}, {352, 376}, {376, 401}, {401, 352}, {274, 275},
{275, 4}, {4, 274}, {421, 428}, {428, 262}, {262, 421},
{294, 327}, {327, 358}, {358, 294}, {433, 416}, {416, 367},
{367, 433}, {289, 455}, {455, 439}, {439, 289}, {462, 370},
{370, 326}, {326, 462}, {2, 326}, {326, 370}, {370, 2},
{305, 460}, {460, 455}, {455, 305}, {254, 449}, {449, 448},
{448, 254}, {255, 261}, {261, 446}, {446, 255}, {253, 450},
{450, 449}, {449, 253}, {252, 451}, {451, 450}, {450, 252},
{256, 452}, {452, 451}, {451, 256}, {341, 453}, {453, 452},
{452, 341}, {413, 464}, {464, 463}, {463, 413}, {441, 413},
{413, 414}, {414, 441}, {258, 442}, {442, 441}, {441, 258},
{257, 443}, {443, 442}, {442, 257}, {259, 444}, {444, 443},
{443, 259}, {260, 445}, {445, 444}, {444, 260}, {467, 342},
{342, 445}, {445, 467}, {459, 458}, {458, 250}, {250, 459},
{289, 392}, {392, 290}, {290, 289}, {290, 328}, {328, 460},
{460, 290}, {376, 433}, {433, 435}, {435, 376}, {250, 290},
{290, 392}, {392, 250}, {411, 416}, {416, 433}, {433, 411},
{341, 463}, {463, 464}, {464, 341}, {453, 464}, {464, 465},
{465, 453}, {357, 465}, {465, 412}, {412, 357}, {343, 412},
{412, 399}, {399, 343}, {360, 363}, {363, 440}, {440, 360},
{437, 399}, {399, 456}, {456, 437}, {420, 456}, {456, 363},
{363, 420}, {401, 435}, {435, 288}, {288, 401}, {372, 383},
{383, 353}, {353, 372}, {339, 255}, {255, 249}, {249, 339},
{448, 261}, {261, 255}, {255, 448}, {133, 243}, {243, 190},
{190, 133}, {133, 155}, {155, 112}, {112, 133}, {33, 246},
{246, 247}, {247, 33}, {33, 130}, {130, 25}, {25, 33},
{398, 384}, {384, 286}, {286, 398}, {362, 398}, {398, 414},
{414, 362}, {362, 463}, {463, 341}, {341, 362}, {263, 359},
{359, 467}, {467, 263}, {263, 249}, {249, 255}, {255, 263},
{466, 467}, {467, 260}, {260, 466}, {75, 60}, {60, 166},
{166, 75}, {238, 239}, {239, 79}, {79, 238}, {162, 127},
{127, 139}, {139, 162}, {72, 11}, {11, 37}, {37, 72},
{121, 232}, {232, 120}, {120, 121}, {73, 72}, {72, 39},
{39, 73}, {114, 128}, {128, 47}, {47, 114}, {233, 232},
{232, 128}, {128, 233}, {103, 104}, {104, 67}, {67, 103},
{152, 175}, {175, 148}, {148, 152}, {119, 118}, {118, 101},
{101, 119}, {74, 73}, {73, 40}, {40, 74}, {107, 9},
{9, 108}, {108, 107}, {49, 48}, {48, 131}, {131, 49},
{32, 194}, {194, 211}, {211, 32}, {184, 74}, {74, 185},
{185, 184}, {191, 80}, {80, 183}, {183, 191}, {185, 40},
{40, 186}, {186, 185}, {119, 230}, {230, 118}, {118, 119},
{210, 202}, {202, 214}, {214, 210}, {84, 83}, {83, 17},
{17, 84}, {77, 76}, {76, 146}, {146, 77}, {161, 160},
{160, 30}, {30, 161}, {190, 56}, {56, 173}, {173, 190},
{182, 106}, {106, 194}, {194, 182}, {138, 135}, {135, 192},
{192, 138}, {129, 203}, {203, 98}, {98, 129}, {54, 21},
{21, 68}, {68, 54}, {5, 51}, {51, 4}, {4, 5},
{145, 144}, {144, 23}, {23, 145}, {90, 77}, {77, 91},
{91, 90}, {207, 205}, {205, 187}, {187, 207}, {83, 201},
{201, 18}, {18, 83}, {181, 91}, {91, 182}, {182, 181},
{180, 90}, {90, 181}, {181, 180}, {16, 85}, {85, 17},
{17, 16}, {205, 206}, {206, 36}, {36, 205}, {176, 148},
{148, 140}, {140, 176}, {165, 92}, {92, 39}, {39, 165},
{245, 193}, {193, 244}, {244, 245}, {27, 159}, {159, 28},
{28, 27}, {30, 247}, {247, 161}, {161, 30}, {174, 236},
{236, 196}, {196, 174}, {103, 54}, {54, 104}, {104, 103},
{55, 193}, {193, 8}, {8, 55}, {111, 117}, {117, 31},
{31, 111}, {221, 189}, {189, 55}, {55, 221}, {240, 98},
{98, 99}, {99, 240}, {142, 126}, {126, 100}, {100, 142},
{219, 166}, {166, 218}, {218, 219}, {112, 155}, {155, 26},
{26, 112}, {198, 209}, {209, 131}, {131, 198}, {169, 135},
{135, 150}, {150, 169}, {114, 47}, {47, 217}, {217, 114},
{224, 223}, {223, 53}, {53, 224}, {220, 45}, {45, 134},
{134, 220}, {32, 211}, {211, 140}, {140, 32}, {109, 67},
{67, 108}, {108, 109}, {146, 43}, {43, 91}, {91, 146},
{231, 230}, {230, 120}, {120, 231}, {113, 226}, {226, 247},
{247, 113}, {105, 63}, {63, 52}, {52, 105}, {241, 238},
{238, 242}, {242, 241}, {124, 46}, {46, 156}, {156, 124},
{95, 78}, {78, 96}, {96, 95}, {70, 46}, {46, 63},
{63, 70}, {116, 143}, {143, 227}, {227, 116}, {116, 123},
{123, 111}, {111, 116}, {1, 44}, {44, 19}, {19, 1},
{3, 236}, {236, 51}, {51, 3}, {207, 216}, {216, 205},
{205, 207}, {26, 154}, {154, 22}, {22, 26}, {165, 39},
{39, 167}, {167, 165}, {199, 200}, {200, 208}, {208, 199},
{101, 36}, {36, 100}, {100, 101}, {43, 57}, {57, 202},
{202, 43}, {242, 20}, {20, 99}, {99, 242}, {56, 28},
{28, 157}, {157, 56}, {124, 35}, {35, 113}, {113, 124},
{29, 160}, {160, 27}, {27, 29}, {211, 204}, {204, 210},
{210, 211}, {124, 113}, {113, 46}, {46, 124}, {106, 43},
{43, 204}, {204, 106}, {96, 62}, {62, 77}, {77, 96},
{227, 137}, {137, 116}, {116, 227}, {73, 41}, {41, 72},
{72, 73}, {36, 203}, {203, 142}, {142, 36}, {235, 64},
{64, 240}, {240, 235}, {48, 49}, {49, 64}, {64, 48},
{42, 41}, {41, 74}, {74, 42}, {214, 212}, {212, 207},
{207, 214}, {183, 42}, {42, 184}, {184, 183}, {210, 169},
{169, 211}, {211, 210}, {140, 170}, {170, 176}, {176, 140},
{104, 105}, {105, 69}, {69, 104}, {193, 122}, {122, 168},
{168, 193}, {50, 123}, {123, 187}, {187, 50}, {89, 96},
{96, 90}, {90, 89}, {66, 65}, {65, 107}, {107, 66},
{179, 89}, {89, 180}, {180, 179}, {119, 101}, {101, 120},
{120, 119}, {68, 63}, {63, 104}, {104, 68}, {234, 93},
{93, 227}, {227, 234}, {16, 15}, {15, 85}, {85, 16},
{209, 129}, {129, 49}, {49, 209}, {15, 14}, {14, 86},
{86, 15}, {107, 55}, {55, 9}, {9, 107}, {120, 100},
{100, 121}, {121, 120}, {153, 145}, {145, 22}, {22, 153},
{178, 88}, {88, 179}, {179, 178}, {197, 6}, {6, 196},
{196, 197}, {89, 88}, {88, 96}, {96, 89}, {135, 138},
{138, 136}, {136, 135}, {138, 215}, {215, 172}, {172, 138},
{218, 115}, {115, 219}, {219, 218}, {41, 42}, {42, 81},
{81, 41}, {5, 195}, {195, 51}, {51, 5}, {57, 43},
{43, 61}, {61, 57}, {208, 171}, {171, 199}, {199, 208},
{41, 81}, {81, 38}, {38, 41}, {224, 53}, {53, 225},
{225, 224}, {24, 144}, {144, 110}, {110, 24}, {105, 52},
{52, 66}, {66, 105}, {118, 229}, {229, 117}, {117, 118},
{227, 34}, {34, 234}, {234, 227}, {66, 107}, {107, 69},
{69, 66}, {10, 109}, {109, 151}, {151, 10}, {219, 48},
{48, 235}, {235, 219}, {183, 62}, {62, 191}, {191, 183},
{142, 129}, {129, 126}, {126, 142}, {116, 111}, {111, 143},
{143, 116}, {118, 117}, {117, 50}, {50, 118}, {223, 222},
{222, 52}, {52, 223}, {94, 19}, {19, 141}, {141, 94},
{222, 221}, {221, 65}, {65, 222}, {196, 3}, {3, 197},
{197, 196}, {45, 220}, {220, 44}, {44, 45}, {156, 70},
{70, 139}, {139, 156}, {188, 122}, {122, 245}, {245, 188},
{139, 71}, {71, 162}, {162, 139}, {149, 170}, {170, 150},
{150, 149}, {122, 188}, {188, 196}, {196, 122}, {206, 216},
{216, 92}, {92, 206}, {164, 2}, {2, 167}, {167, 164},
{242, 141}, {141, 241}, {241, 242}, {0, 164}, {164, 37},
{37, 0}, {11, 72}, {72, 12}, {12, 11}, {12, 38},
{38, 13}, {13, 12}, {70, 63}, {63, 71}, {71, 70},
{31, 226}, {226, 111}, {111, 31}, {36, 101}, {101, 205},
{205, 36}, {203, 206}, {206, 165}, {165, 203}, {126, 209},
{209, 217}, {217, 126}, {98, 165}, {165, 97}, {97, 98},
{237, 220}, {220, 218}, {218, 237}, {237, 239}, {239, 241},
{241, 237}, {210, 214}, {214, 169}, {169, 210}, {140, 171},
{171, 32}, {32, 140}, {241, 125}, {125, 237}, {237, 241},
{179, 86}, {86, 178}, {178, 179}, {180, 85}, {85, 179},
{179, 180}, {181, 84}, {84, 180}, {180, 181}, {182, 83},
{83, 181}, {181, 182}, {194, 201}, {201, 182}, {182, 194},
{177, 137}, {137, 132}, {132, 177}, {184, 76}, {76, 183},
{183, 184}, {185, 61}, {61, 184}, {184, 185}, {186, 57},
{57, 185}, {185, 186}, {216, 212}, {212, 186}, {186, 216},
{192, 214}, {214, 187}, {187, 192}, {139, 34}, {34, 156},
{156, 139}, {218, 79}, {79, 237}, {237, 218}, {147, 123},
{123, 177}, {177, 147}, {45, 44}, {44, 4}, {4, 45},
{208, 201}, {201, 32}, {32, 208}, {98, 64}, {64, 129},
{129, 98}, {192, 213}, {213, 138}, {138, 192}, {235, 59},
{59, 219}, {219, 235}, {141, 242}, {242, 97}, {97, 141},
{97, 2}, {2, 141}, {141, 97}, {240, 75}, {75, 235},
{235, 240}, {229, 24}, {24, 228}, {228, 229}, {31, 25},
{25, 226}, {226, 31}, {230, 23}, {23, 229}, {229, 230},
{231, 22}, {22, 230}, {230, 231}, {232, 26}, {26, 231},
{231, 232}, {233, 112}, {112, 232}, {232, 233}, {244, 189},
{189, 243}, {243, 244}, {189, 221}, {221, 190}, {190, 189},
{222, 28}, {28, 221}, {221, 222}, {223, 27}, {27, 222},
{222, 223}, {224, 29}, {29, 223}, {223, 224}, {225, 30},
{30, 224}, {224, 225}, {113, 247}, {247, 225}, {225, 113},
{99, 60}, {60, 240}, {240, 99}, {213, 147}, {147, 215},
{215, 213}, {60, 20}, {20, 166}, {166, 60}, {192, 187},
{187, 213}, {213, 192}, {243, 112}, {112, 244}, {244, 243},
{244, 233}, {233, 245}, {245, 244}, {245, 128}, {128, 188},
{188, 245}, {188, 114}, {114, 174}, {174, 188}, {134, 131},
{131, 220}, {220, 134}, {174, 217}, {217, 236}, {236, 174},
{236, 198}, {198, 134}, {134, 236}, {215, 177}, {177, 58},
{58, 215}, {156, 143}, {143, 124}, {124, 156}, {25, 110},
{110, 7}, {7, 25}, {31, 228}, {228, 25}, {25, 31},
{264, 356}, {356, 368}, {368, 264}, {0, 11}, {11, 267},
{267, 0}, {451, 452}, {452, 349}, {349, 451}, {267, 302},
{302, 269}, {269, 267}, {350, 357}, {357, 277}, {277, 350},
{350, 452}, {452, 357}, {357, 350}, {299, 333}, {333, 297},
{297, 299}, {396, 175}, {175, 377}, {377, 396}, {280, 347},
{347, 330}, {330, 280}, {269, 303}, {303, 270}, {270, 269},
{151, 9}, {9, 337}, {337, 151}, {344, 278}, {278, 360},
{360, 344}, {424, 418}, {418, 431}, {431, 424}, {270, 304},
{304, 409}, {409, 270}, {272, 310}, {310, 407}, {407, 272},
{322, 270}, {270, 410}, {410, 322}, {449, 450}, {450, 347},
{347, 449}, {432, 422}, {422, 434}, {434, 432}, {18, 313},
{313, 17}, {17, 18}, {291, 306}, {306, 375}, {375, 291},
{259, 387}, {387, 260}, {260, 259}, {424, 335}, {335, 418},
{418, 424}, {434, 364}, {364, 416}, {416, 434}, {391, 423},
{423, 327}, {327, 391}, {301, 251}, {251, 298}, {298, 301},
{275, 281}, {281, 4}, {4, 275}, {254, 373}, {373, 253},
{253, 254}, {375, 307}, {307, 321}, {321, 375}, {280, 425},
{425, 411}, {411, 280}, {200, 421}, {421, 18}, {18, 200},
{335, 321}, {321, 406}, {406, 335}, {321, 320}, {320, 405},
{405, 321}, {314, 315}, {315, 17}, {17, 314}, {423, 426},
{426, 266}, {266, 423}, {396, 377}, {377, 369}, {369, 396},
{270, 322}, {322, 269}, {269, 270}, {413, 417}, {417, 464},
{464, 413}, {385, 386}, {386, 258}, {258, 385}, {248, 456},
{456, 419}, {419, 248}, {298, 284}, {284, 333}, {333, 298},
{168, 417}, {417, 8}, {8, 168}, {448, 346}, {346, 261},
{261, 448}, {417, 413}, {413, 285}, {285, 417}, {326, 327},
{327, 328}, {328, 326}, {277, 355}, {355, 329}, {329, 277},
{309, 392}, {392, 438}, {438, 309}, {381, 382}, {382, 256},
{256, 381}, {279, 429}, {429, 360}, {360, 279}, {365, 364},
{364, 379}, {379, 365}, {355, 277}, {277, 437}, {437, 355},
{282, 443}, {443, 283}, {283, 282}, {281, 275}, {275, 363},
{363, 281}, {395, 431}, {431, 369}, {369, 395}, {299, 297},
{297, 337}, {337, 299}, {335, 273}, {273, 321}, {321, 335},
{348, 450}, {450, 349}, {349, 348}, {359, 446}, {446, 467},
{467, 359}, {283, 293}, {293, 282}, {282, 283}, {250, 458},
{458, 462}, {462, 250}, {300, 276}, {276, 383}, {383, 300},
{292, 308}, {308, 325}, {325, 292}, {283, 276}, {276, 293},
{293, 283}, {264, 372}, {372, 447}, {447, 264}, {346, 352},
{352, 340}, {340, 346}, {354, 274}, {274, 19}, {19, 354},
{363, 456}, {456, 281}, {281, 363}, {426, 436}, {436, 425},
{425, 426}, {380, 381}, {381, 252}, {252, 380}, {267, 269},
{269, 393}, {393, 267}, {421, 200}, {200, 428}, {428, 421},
{371, 266}, {266, 329}, {329, 371}, {432, 287}, {287, 422},
{422, 432}, {290, 250}, {250, 328}, {328, 290}, {385, 258},
{258, 384}, {384, 385}, {446, 265}, {265, 342}, {342, 446},
{386, 387}, {387, 257}, {257, 386}, {422, 424}, {424, 430},
{430, 422}, {445, 342}, {342, 276}, {276, 445}, {422, 273},
{273, 424}, {424, 422}, {306, 292}, {292, 307}, {307, 306},
{352, 366}, {366, 345}, {345, 352}, {268, 271}, {271, 302},
{302, 268}, {358, 423}, {423, 371}, {371, 358}, {327, 294},
{294, 460}, {460, 327}, {331, 279}, {279, 294}, {294, 331},
{303, 271}, {271, 304}, {304, 303}, {436, 432}, {432, 427},
{427, 436}, {304, 272}, {272, 408}, {408, 304}, {395, 394},
{394, 431}, {431, 395}, {378, 395}, {395, 400}, {400, 378},
{296, 334}, {334, 299}, {299, 296}, {6, 351}, {351, 168},
{168, 6}, {376, 352}, {352, 411}, {411, 376}, {307, 325},
{325, 320}, {320, 307}, {285, 295}, {295, 336}, {336, 285},
{320, 319}, {319, 404}, {404, 320}, {329, 330}, {330, 349},
{349, 329}, {334, 293}, {293, 333}, {333, 334}, {366, 323},
{323, 447}, {447, 366}, {316, 15}, {15, 315}, {315, 316},
{331, 358}, {358, 279}, {279, 331}, {317, 14}, {14, 316},
{316, 317}, {8, 285}, {285, 9}, {9, 8}, {277, 329},
{329, 350}, {350, 277}, {253, 374}, {374, 252}, {252, 253},
{319, 318}, {318, 403}, {403, 319}, {351, 6}, {6, 419},
{419, 351}, {324, 318}, {318, 325}, {325, 324}, {397, 367},
{367, 365}, {365, 397}, {288, 435}, {435, 397}, {397, 288},
{278, 344}, {344, 439}, {439, 278}, {310, 272}, {272, 311},
{311, 310}, {248, 195}, {195, 281}, {281, 248}, {375, 273},
{273, 291}, {291, 375}, {175, 396}, {396, 199}, {199, 175},
{312, 311}, {311, 268}, {268, 312}, {276, 283}, {283, 445},
{445, 276}, {390, 373}, {373, 339}, {339, 390}, {295, 282},
{282, 296}, {296, 295}, {448, 449}, {449, 346}, {346, 448},
{356, 264}, {264, 454}, {454, 356}, {337, 336}, {336, 299},
{299, 337}, {337, 338}, {338, 151}, {151, 337}, {294, 278},
{278, 455}, {455, 294}, {308, 292}, {292, 415}, {415, 308},
{429, 358}, {358, 355}, {355, 429}, {265, 340}, {340, 372},
{372, 265}, {352, 346}, {346, 280}, {280, 352}, {295, 442},
{442, 282}, {282, 295}, {354, 19}, {19, 370}, {370, 354},
{285, 441}, {441, 295}, {295, 285}, {195, 248}, {248, 197},
{197, 195}, {457, 440}, {440, 274}, {274, 457}, {301, 300},
{300, 368}, {368, 301}, {417, 351}, {351, 465}, {465, 417},
{251, 301}, {301, 389}, {389, 251}, {394, 395}, {395, 379},
{379, 394}, {399, 412}, {412, 419}, {419, 399}, {410, 436},
{436, 322}, {322, 410}, {326, 2}, {2, 393}, {393, 326},
{354, 370}, {370, 461}, {461, 354}, {393, 164}, {164, 267},
{267, 393}, {268, 302}, {302, 12}, {12, 268}, {312, 268},
{268, 13}, {13, 312}, {298, 293}, {293, 301}, {301, 298},
{265, 446}, {446, 340}, {340, 265}, {280, 330}, {330, 425},
{425, 280}, {322, 426}, {426, 391}, {391, 322}, {420, 429},
{429, 437}, {437, 420}, {393, 391}, {391, 326}, {326, 393},
{344, 440}, {440, 438}, {438, 344}, {458, 459}, {459, 461},
{461, 458}, {364, 434}, {434, 394}, {394, 364}, {428, 396},
{396, 262}, {262, 428}, {274, 354}, {354, 457}, {457, 274},
{317, 316}, {316, 402}, {402, 317}, {316, 315}, {315, 403},
{403, 316}, {315, 314}, {314, 404}, {404, 315}, {314, 313},
{313, 405}, {405, 314}, {313, 421}, {421, 406}, {406, 313},
{323, 366}, {366, 361}, {361, 323}, {292, 306}, {306, 407},
{407, 292}, {306, 291}, {291, 408}, {408, 306}, {291, 287},
{287, 409}, {409, 291}, {287, 432}, {432, 410}, {410, 287},
{427, 434}, {434, 411}, {411, 427}, {372, 264}, {264, 383},
{383, 372}, {459, 309}, {309, 457}, {457, 459}, {366, 352},
{352, 401}, {401, 366}, {1, 274}, {274, 4}, {4, 1},
{418, 421}, {421, 262}, {262, 418}, {331, 294}, {294, 358},
{358, 331}, {435, 433}, {433, 367}, {367, 435}, {392, 289},
{289, 439}, {439, 392}, {328, 462}, {462, 326}, {326, 328},
{94, 2}, {2, 370}, {370, 94}, {289, 305}, {305, 455},
{455, 289}, {339, 254}, {254, 448}, {448, 339}, {359, 255},
{255, 446}, {446, 359}, {254, 253}, {253, 449}, {449, 254},
{253, 252}, {252, 450}, {450, 253}, {252, 256}, {256, 451},
{451, 252}, {256, 341}, {341, 452}, {452, 256}, {414, 413},
{413, 463}, {463, 414}, {286, 441}, {441, 414}, {414, 286},
{286, 258}, {258, 441}, {441, 286}, {258, 257}, {257, 442},
{442, 258}, {257, 259}, {259, 443}, {443, 257}, {259, 260},
{260, 444}, {444, 259}, {260, 467}, {467, 445}, {445, 260},
{309, 459}, {459, 250}, {250, 309}, {305, 289}, {289, 290},
{290, 305}, {305, 290}, {290, 460}, {460, 305}, {401, 376},
{376, 435}, {435, 401}, {309, 250}, {250, 392}, {392, 309},
{376, 411}, {411, 433}, {433, 376}, {453, 341}, {341, 464},
{464, 453}, {357, 453}, {453, 465}, {465, 357}, {343, 357},
{357, 412}, {412, 343}, {437, 343}, {343, 399}, {399, 437},
{344, 360}, {360, 440}, {440, 344}, {420, 437}, {437, 456},
{456, 420}, {360, 420}, {420, 363}, {363, 360}, {361, 401},
{401, 288}, {288, 361}, {265, 372}, {372, 353}, {353, 265},
{390, 339}, {339, 249}, {249, 390}, {339, 448}, {448, 255},
{255, 339}}};
};
} // namespace face_landmarker
} // namespace vision
} // namespace tasks
} // namespace mediapipe
#endif // MEDIAPIPE_TASKS_CC_VISION_FACE_LANDMARKER_FACE_LANDMARKS_CONNECTIONS_H_
@@ -111,6 +111,7 @@ class TensorsToImageCalculator : public Node {
private: private:
TensorsToImageCalculatorOptions options_; TensorsToImageCalculatorOptions options_;
absl::Status CpuProcess(CalculatorContext* cc); absl::Status CpuProcess(CalculatorContext* cc);
int tensor_position_;
#if !MEDIAPIPE_DISABLE_GPU #if !MEDIAPIPE_DISABLE_GPU
#if MEDIAPIPE_METAL_ENABLED #if MEDIAPIPE_METAL_ENABLED
@@ -166,6 +167,7 @@ absl::Status TensorsToImageCalculator::Open(CalculatorContext* cc) {
<< "Must specify either `input_tensor_float_range` or " << "Must specify either `input_tensor_float_range` or "
"`input_tensor_uint_range` in the calculator options"; "`input_tensor_uint_range` in the calculator options";
} }
tensor_position_ = options_.tensor_position();
return absl::OkStatus(); return absl::OkStatus();
} }
@@ -202,17 +204,23 @@ absl::Status TensorsToImageCalculator::CpuProcess(CalculatorContext* cc) {
return absl::OkStatus(); return absl::OkStatus();
} }
const auto& input_tensors = kInputTensors(cc).Get(); const auto& input_tensors = kInputTensors(cc).Get();
RET_CHECK_EQ(input_tensors.size(), 1) RET_CHECK_GT(input_tensors.size(), tensor_position_)
<< "Expect 1 input tensor, but have " << input_tensors.size(); << "Expect input tensor at position " << tensor_position_
<< ", but have tensors of size " << input_tensors.size();
const auto& input_tensor = input_tensors[0]; const auto& input_tensor = input_tensors[tensor_position_];
const int tensor_in_height = input_tensor.shape().dims[1]; const int tensor_in_height = input_tensor.shape().dims[1];
const int tensor_in_width = input_tensor.shape().dims[2]; const int tensor_in_width = input_tensor.shape().dims[2];
const int tensor_in_channels = input_tensor.shape().dims[3]; const int tensor_in_channels = input_tensor.shape().dims[3];
RET_CHECK_EQ(tensor_in_channels, 3); RET_CHECK(tensor_in_channels == 3 || tensor_in_channels == 1);
auto output_frame = std::make_shared<ImageFrame>( auto format = mediapipe::ImageFormat::SRGB;
mediapipe::ImageFormat::SRGB, tensor_in_width, tensor_in_height); if (tensor_in_channels == 1) {
format = mediapipe::ImageFormat::GRAY8;
}
auto output_frame =
std::make_shared<ImageFrame>(format, tensor_in_width, tensor_in_height);
cv::Mat output_matview = mediapipe::formats::MatView(output_frame.get()); cv::Mat output_matview = mediapipe::formats::MatView(output_frame.get());
constexpr float kOutputImageRangeMin = 0.0f; constexpr float kOutputImageRangeMin = 0.0f;
@@ -227,8 +235,9 @@ absl::Status TensorsToImageCalculator::CpuProcess(CalculatorContext* cc) {
GetValueRangeTransformation( GetValueRangeTransformation(
input_range.min(), input_range.max(), input_range.min(), input_range.max(),
kOutputImageRangeMin, kOutputImageRangeMax)); kOutputImageRangeMin, kOutputImageRangeMax));
tensor_matview.convertTo(output_matview, CV_8UC3, transform.scale, tensor_matview.convertTo(output_matview,
transform.offset); CV_MAKETYPE(CV_8U, tensor_in_channels),
transform.scale, transform.offset);
} else if (input_tensor.element_type() == Tensor::ElementType::kUInt8) { } else if (input_tensor.element_type() == Tensor::ElementType::kUInt8) {
cv::Mat tensor_matview( cv::Mat tensor_matview(
cv::Size(tensor_in_width, tensor_in_height), cv::Size(tensor_in_width, tensor_in_height),
@@ -239,8 +248,9 @@ absl::Status TensorsToImageCalculator::CpuProcess(CalculatorContext* cc) {
GetValueRangeTransformation( GetValueRangeTransformation(
input_range.min(), input_range.max(), input_range.min(), input_range.max(),
kOutputImageRangeMin, kOutputImageRangeMax)); kOutputImageRangeMin, kOutputImageRangeMax));
tensor_matview.convertTo(output_matview, CV_8UC3, transform.scale, tensor_matview.convertTo(output_matview,
transform.offset); CV_MAKETYPE(CV_8U, tensor_in_channels),
transform.scale, transform.offset);
} else { } else {
return absl::InvalidArgumentError( return absl::InvalidArgumentError(
absl::Substitute("Type of tensor must be kFloat32 or kUInt8, got: $0", absl::Substitute("Type of tensor must be kFloat32 or kUInt8, got: $0",
@@ -264,10 +274,14 @@ absl::Status TensorsToImageCalculator::MetalProcess(CalculatorContext* cc) {
return absl::OkStatus(); return absl::OkStatus();
} }
const auto& input_tensors = kInputTensors(cc).Get(); const auto& input_tensors = kInputTensors(cc).Get();
RET_CHECK_EQ(input_tensors.size(), 1) RET_CHECK_GT(input_tensors.size(), tensor_position_)
<< "Expect 1 input tensor, but have " << input_tensors.size(); << "Expect input tensor at position " << tensor_position_
const int tensor_width = input_tensors[0].shape().dims[2]; << ", but have tensors of size " << input_tensors.size();
const int tensor_height = input_tensors[0].shape().dims[1]; const int tensor_width = input_tensors[tensor_position_].shape().dims[2];
const int tensor_height = input_tensors[tensor_position_].shape().dims[1];
const int tensor_channels = input_tensors[tensor_position_].shape().dims[3];
// TODO: Add 1 channel support.
RET_CHECK(tensor_channels == 3);
// TODO: Fix unused variable // TODO: Fix unused variable
[[maybe_unused]] id<MTLDevice> device = gpu_helper_.mtlDevice; [[maybe_unused]] id<MTLDevice> device = gpu_helper_.mtlDevice;
@@ -277,8 +291,8 @@ absl::Status TensorsToImageCalculator::MetalProcess(CalculatorContext* cc) {
[command_buffer computeCommandEncoder]; [command_buffer computeCommandEncoder];
[compute_encoder setComputePipelineState:to_buffer_program_]; [compute_encoder setComputePipelineState:to_buffer_program_];
auto input_view = auto input_view = mediapipe::MtlBufferView::GetReadView(
mediapipe::MtlBufferView::GetReadView(input_tensors[0], command_buffer); input_tensors[tensor_position_], command_buffer);
[compute_encoder setBuffer:input_view.buffer() offset:0 atIndex:0]; [compute_encoder setBuffer:input_view.buffer() offset:0 atIndex:0];
mediapipe::GpuBuffer output = mediapipe::GpuBuffer output =
@@ -355,7 +369,7 @@ absl::Status TensorsToImageCalculator::GlSetup(CalculatorContext* cc) {
absl::StrCat(tflite::gpu::gl::GetShaderHeader(workgroup_size_), R"( absl::StrCat(tflite::gpu::gl::GetShaderHeader(workgroup_size_), R"(
precision highp float; precision highp float;
layout(rgba8, binding = 0) writeonly uniform highp image2D output_texture; layout(rgba8, binding = 0) writeonly uniform highp image2D output_texture;
uniform ivec2 out_size; uniform ivec3 out_size;
)"); )");
const std::string shader_body = R"( const std::string shader_body = R"(
@@ -366,10 +380,11 @@ absl::Status TensorsToImageCalculator::GlSetup(CalculatorContext* cc) {
void main() { void main() {
int out_width = out_size.x; int out_width = out_size.x;
int out_height = out_size.y; int out_height = out_size.y;
int out_channels = out_size.z;
ivec2 gid = ivec2(gl_GlobalInvocationID.xy); ivec2 gid = ivec2(gl_GlobalInvocationID.xy);
if (gid.x >= out_width || gid.y >= out_height) { return; } if (gid.x >= out_width || gid.y >= out_height) { return; }
int linear_index = 3 * (gid.y * out_width + gid.x); int linear_index = out_channels * (gid.y * out_width + gid.x);
#ifdef FLIP_Y_COORD #ifdef FLIP_Y_COORD
int y_coord = out_height - gid.y - 1; int y_coord = out_height - gid.y - 1;
@@ -377,8 +392,14 @@ absl::Status TensorsToImageCalculator::GlSetup(CalculatorContext* cc) {
int y_coord = gid.y; int y_coord = gid.y;
#endif // defined(FLIP_Y_COORD) #endif // defined(FLIP_Y_COORD)
vec4 out_value;
ivec2 out_coordinate = ivec2(gid.x, y_coord); ivec2 out_coordinate = ivec2(gid.x, y_coord);
vec4 out_value = vec4(input_data.elements[linear_index], input_data.elements[linear_index + 1], input_data.elements[linear_index + 2], 1.0); if (out_channels == 3) {
out_value = vec4(input_data.elements[linear_index], input_data.elements[linear_index + 1], input_data.elements[linear_index + 2], 1.0);
} else {
float in_value = input_data.elements[linear_index];
out_value = vec4(in_value, in_value, in_value, 1.0);
}
imageStore(output_texture, out_coordinate, out_value); imageStore(output_texture, out_coordinate, out_value);
})"; })";
@@ -438,10 +459,15 @@ absl::Status TensorsToImageCalculator::GlProcess(CalculatorContext* cc) {
return absl::OkStatus(); return absl::OkStatus();
} }
const auto& input_tensors = kInputTensors(cc).Get(); const auto& input_tensors = kInputTensors(cc).Get();
RET_CHECK_EQ(input_tensors.size(), 1) RET_CHECK_GT(input_tensors.size(), tensor_position_)
<< "Expect 1 input tensor, but have " << input_tensors.size(); << "Expect input tensor at position " << tensor_position_
const int tensor_width = input_tensors[0].shape().dims[2]; << ", but have tensors of size " << input_tensors.size();
const int tensor_height = input_tensors[0].shape().dims[1];
const auto& input_tensor = input_tensors[tensor_position_];
const int tensor_width = input_tensor.shape().dims[2];
const int tensor_height = input_tensor.shape().dims[1];
const int tensor_in_channels = input_tensor.shape().dims[3];
RET_CHECK(tensor_in_channels == 3 || tensor_in_channels == 1);
#if MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_31 #if MEDIAPIPE_OPENGL_ES_VERSION >= MEDIAPIPE_OPENGL_ES_31
@@ -454,7 +480,7 @@ absl::Status TensorsToImageCalculator::GlProcess(CalculatorContext* cc) {
glBindImageTexture(output_index, out_texture->id(), 0, GL_FALSE, 0, glBindImageTexture(output_index, out_texture->id(), 0, GL_FALSE, 0,
GL_WRITE_ONLY, GL_RGBA8); GL_WRITE_ONLY, GL_RGBA8);
auto read_view = input_tensors[0].GetOpenGlBufferReadView(); auto read_view = input_tensor.GetOpenGlBufferReadView();
glBindBufferBase(GL_SHADER_STORAGE_BUFFER, 2, read_view.name()); glBindBufferBase(GL_SHADER_STORAGE_BUFFER, 2, read_view.name());
const tflite::gpu::uint3 workload = {tensor_width, tensor_height, 1}; const tflite::gpu::uint3 workload = {tensor_width, tensor_height, 1};
@@ -462,8 +488,8 @@ absl::Status TensorsToImageCalculator::GlProcess(CalculatorContext* cc) {
tflite::gpu::DivideRoundUp(workload, workgroup_size_); tflite::gpu::DivideRoundUp(workload, workgroup_size_);
glUseProgram(gl_compute_program_->id()); glUseProgram(gl_compute_program_->id());
glUniform2i(glGetUniformLocation(gl_compute_program_->id(), "out_size"), glUniform3i(glGetUniformLocation(gl_compute_program_->id(), "out_size"),
tensor_width, tensor_height); tensor_width, tensor_height, tensor_in_channels);
MP_RETURN_IF_ERROR(gl_compute_program_->Dispatch(workgroups)); MP_RETURN_IF_ERROR(gl_compute_program_->Dispatch(workgroups));
@@ -481,8 +507,8 @@ absl::Status TensorsToImageCalculator::GlProcess(CalculatorContext* cc) {
#else #else
if (!input_tensors[0].ready_as_opengl_texture_2d()) { if (!input_tensor.ready_as_opengl_texture_2d()) {
(void)input_tensors[0].GetCpuReadView(); (void)input_tensor.GetCpuReadView();
} }
auto output_texture = auto output_texture =
@@ -490,7 +516,7 @@ absl::Status TensorsToImageCalculator::GlProcess(CalculatorContext* cc) {
gl_helper_.BindFramebuffer(output_texture); // GL_TEXTURE0 gl_helper_.BindFramebuffer(output_texture); // GL_TEXTURE0
glActiveTexture(GL_TEXTURE1); glActiveTexture(GL_TEXTURE1);
glBindTexture(GL_TEXTURE_2D, glBindTexture(GL_TEXTURE_2D,
input_tensors[0].GetOpenGlTexture2dReadView().name()); input_tensor.GetOpenGlTexture2dReadView().name());
MP_RETURN_IF_ERROR(gl_renderer_->GlRender( MP_RETURN_IF_ERROR(gl_renderer_->GlRender(
tensor_width, tensor_height, output_texture.width(), tensor_width, tensor_height, output_texture.width(),
@@ -48,4 +48,8 @@ message TensorsToImageCalculatorOptions {
FloatRange input_tensor_float_range = 2; FloatRange input_tensor_float_range = 2;
UIntRange input_tensor_uint_range = 3; UIntRange input_tensor_uint_range = 3;
} }
// Determines which output tensor to slice when there are multiple output
// tensors available (e.g. network has multiple heads)
optional int32 tensor_position = 4 [default = 0];
} }
@@ -153,6 +153,11 @@ cc_library(
alwayslink = 1, alwayslink = 1,
) )
cc_library(
name = "hand_landmarks_connections",
hdrs = ["hand_landmarks_connections.h"],
)
# TODO: open source hand joints graph # TODO: open source hand joints graph
cc_library( cc_library(
@@ -0,0 +1,54 @@
/* Copyright 2023 The MediaPipe Authors.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#ifndef MEDIAPIPE_TASKS_CC_VISION_HAND_LANDMARKER_HAND_LANDMARKS_CONNECTIONS_H_
#define MEDIAPIPE_TASKS_CC_VISION_HAND_LANDMARKER_HAND_LANDMARKS_CONNECTIONS_H_
#include <array>
namespace mediapipe {
namespace tasks {
namespace vision {
namespace hand_landmarker {
static constexpr std::array<std::array<int, 2>, 6> kHandPalmConnections{
{{0, 1}, {0, 5}, {9, 13}, {13, 17}, {5, 9}, {0, 17}}};
static constexpr std::array<std::array<int, 2>, 3> kHandThumbConnections{
{{1, 2}, {2, 3}, {3, 4}}};
static constexpr std::array<std::array<int, 2>, 3> kHandIndexFingerConnections{
{{5, 6}, {6, 7}, {7, 8}}};
static constexpr std::array<std::array<int, 2>, 3> kHandMiddleFingerConnections{
{{9, 10}, {10, 11}, {11, 12}}};
static constexpr std::array<std::array<int, 2>, 3> kHandRingFingerConnections{
{{13, 14}, {14, 15}, {15, 16}}};
static constexpr std::array<std::array<int, 2>, 3> kHandPinkyFingerConnections{
{{17, 18}, {18, 19}, {19, 20}}};
static constexpr std::array<std::array<int, 2>, 21> kHandConnections{
{{0, 1}, {0, 5}, {9, 13}, {13, 17}, {5, 9}, {0, 17}, {1, 2},
{2, 3}, {3, 4}, {5, 6}, {6, 7}, {7, 8}, {9, 10}, {10, 11},
{11, 12}, {13, 14}, {14, 15}, {15, 16}, {17, 18}, {18, 19}, {19, 20}}};
} // namespace hand_landmarker
} // namespace vision
} // namespace tasks
} // namespace mediapipe
#endif // MEDIAPIPE_TASKS_CC_VISION_HAND_LANDMARKER_HAND_LANDMARKS_CONNECTIONS_H_
@@ -16,6 +16,7 @@ limitations under the License.
#include "mediapipe/tasks/cc/vision/image_segmenter/image_segmenter.h" #include "mediapipe/tasks/cc/vision/image_segmenter/image_segmenter.h"
#include <optional> #include <optional>
#include <utility>
#include "absl/strings/str_format.h" #include "absl/strings/str_format.h"
#include "mediapipe/framework/api2/builder.h" #include "mediapipe/framework/api2/builder.h"
@@ -41,6 +42,8 @@ constexpr char kConfidenceMasksTag[] = "CONFIDENCE_MASKS";
constexpr char kConfidenceMasksStreamName[] = "confidence_masks"; constexpr char kConfidenceMasksStreamName[] = "confidence_masks";
constexpr char kCategoryMaskTag[] = "CATEGORY_MASK"; constexpr char kCategoryMaskTag[] = "CATEGORY_MASK";
constexpr char kCategoryMaskStreamName[] = "category_mask"; constexpr char kCategoryMaskStreamName[] = "category_mask";
constexpr char kOutputSizeTag[] = "OUTPUT_SIZE";
constexpr char kOutputSizeStreamName[] = "output_size";
constexpr char kImageInStreamName[] = "image_in"; constexpr char kImageInStreamName[] = "image_in";
constexpr char kImageOutStreamName[] = "image_out"; constexpr char kImageOutStreamName[] = "image_out";
constexpr char kImageTag[] = "IMAGE"; constexpr char kImageTag[] = "IMAGE";
@@ -70,6 +73,7 @@ CalculatorGraphConfig CreateGraphConfig(
options.get()); options.get());
graph.In(kImageTag).SetName(kImageInStreamName); graph.In(kImageTag).SetName(kImageInStreamName);
graph.In(kNormRectTag).SetName(kNormRectStreamName); graph.In(kNormRectTag).SetName(kNormRectStreamName);
graph.In(kOutputSizeTag).SetName(kOutputSizeStreamName);
if (output_confidence_masks) { if (output_confidence_masks) {
task_subgraph.Out(kConfidenceMasksTag) task_subgraph.Out(kConfidenceMasksTag)
.SetName(kConfidenceMasksStreamName) >> .SetName(kConfidenceMasksStreamName) >>
@@ -85,10 +89,12 @@ CalculatorGraphConfig CreateGraphConfig(
graph.Out(kImageTag); graph.Out(kImageTag);
if (enable_flow_limiting) { if (enable_flow_limiting) {
return tasks::core::AddFlowLimiterCalculator( return tasks::core::AddFlowLimiterCalculator(
graph, task_subgraph, {kImageTag, kNormRectTag}, kConfidenceMasksTag); graph, task_subgraph, {kImageTag, kNormRectTag, kOutputSizeTag},
kConfidenceMasksTag);
} }
graph.In(kImageTag) >> task_subgraph.In(kImageTag); graph.In(kImageTag) >> task_subgraph.In(kImageTag);
graph.In(kNormRectTag) >> task_subgraph.In(kNormRectTag); graph.In(kNormRectTag) >> task_subgraph.In(kNormRectTag);
graph.In(kOutputSizeTag) >> task_subgraph.In(kOutputSizeTag);
return graph.GetConfig(); return graph.GetConfig();
} }
@@ -211,6 +217,13 @@ absl::StatusOr<std::unique_ptr<ImageSegmenter>> ImageSegmenter::Create(
absl::StatusOr<ImageSegmenterResult> ImageSegmenter::Segment( absl::StatusOr<ImageSegmenterResult> ImageSegmenter::Segment(
mediapipe::Image image, mediapipe::Image image,
std::optional<core::ImageProcessingOptions> image_processing_options) { std::optional<core::ImageProcessingOptions> image_processing_options) {
return Segment(image, image.width(), image.height(),
std::move(image_processing_options));
}
absl::StatusOr<ImageSegmenterResult> ImageSegmenter::Segment(
mediapipe::Image image, int output_width, int output_height,
std::optional<core::ImageProcessingOptions> image_processing_options) {
if (image.UsesGpu()) { if (image.UsesGpu()) {
return CreateStatusWithPayload( return CreateStatusWithPayload(
absl::StatusCode::kInvalidArgument, absl::StatusCode::kInvalidArgument,
@@ -225,7 +238,10 @@ absl::StatusOr<ImageSegmenterResult> ImageSegmenter::Segment(
ProcessImageData( ProcessImageData(
{{kImageInStreamName, mediapipe::MakePacket<Image>(std::move(image))}, {{kImageInStreamName, mediapipe::MakePacket<Image>(std::move(image))},
{kNormRectStreamName, {kNormRectStreamName,
MakePacket<NormalizedRect>(std::move(norm_rect))}})); MakePacket<NormalizedRect>(std::move(norm_rect))},
{kOutputSizeStreamName,
MakePacket<std::pair<int, int>>(
std::make_pair(output_width, output_height))}}));
std::optional<std::vector<Image>> confidence_masks; std::optional<std::vector<Image>> confidence_masks;
if (output_confidence_masks_) { if (output_confidence_masks_) {
confidence_masks = confidence_masks =
@@ -243,6 +259,14 @@ absl::StatusOr<ImageSegmenterResult> ImageSegmenter::Segment(
absl::StatusOr<ImageSegmenterResult> ImageSegmenter::SegmentForVideo( absl::StatusOr<ImageSegmenterResult> ImageSegmenter::SegmentForVideo(
mediapipe::Image image, int64_t timestamp_ms, mediapipe::Image image, int64_t timestamp_ms,
std::optional<core::ImageProcessingOptions> image_processing_options) { std::optional<core::ImageProcessingOptions> image_processing_options) {
return SegmentForVideo(image, image.width(), image.height(), timestamp_ms,
image_processing_options);
}
absl::StatusOr<ImageSegmenterResult> ImageSegmenter::SegmentForVideo(
mediapipe::Image image, int output_width, int output_height,
int64_t timestamp_ms,
std::optional<core::ImageProcessingOptions> image_processing_options) {
if (image.UsesGpu()) { if (image.UsesGpu()) {
return CreateStatusWithPayload( return CreateStatusWithPayload(
absl::StatusCode::kInvalidArgument, absl::StatusCode::kInvalidArgument,
@@ -260,6 +284,10 @@ absl::StatusOr<ImageSegmenterResult> ImageSegmenter::SegmentForVideo(
.At(Timestamp(timestamp_ms * kMicroSecondsPerMilliSecond))}, .At(Timestamp(timestamp_ms * kMicroSecondsPerMilliSecond))},
{kNormRectStreamName, {kNormRectStreamName,
MakePacket<NormalizedRect>(std::move(norm_rect)) MakePacket<NormalizedRect>(std::move(norm_rect))
.At(Timestamp(timestamp_ms * kMicroSecondsPerMilliSecond))},
{kOutputSizeStreamName,
MakePacket<std::pair<int, int>>(
std::make_pair(output_width, output_height))
.At(Timestamp(timestamp_ms * kMicroSecondsPerMilliSecond))}})); .At(Timestamp(timestamp_ms * kMicroSecondsPerMilliSecond))}}));
std::optional<std::vector<Image>> confidence_masks; std::optional<std::vector<Image>> confidence_masks;
if (output_confidence_masks_) { if (output_confidence_masks_) {
@@ -278,6 +306,13 @@ absl::StatusOr<ImageSegmenterResult> ImageSegmenter::SegmentForVideo(
absl::Status ImageSegmenter::SegmentAsync( absl::Status ImageSegmenter::SegmentAsync(
Image image, int64_t timestamp_ms, Image image, int64_t timestamp_ms,
std::optional<core::ImageProcessingOptions> image_processing_options) { std::optional<core::ImageProcessingOptions> image_processing_options) {
return SegmentAsync(image, image.width(), image.height(), timestamp_ms,
image_processing_options);
}
absl::Status ImageSegmenter::SegmentAsync(
Image image, int output_width, int output_height, int64_t timestamp_ms,
std::optional<core::ImageProcessingOptions> image_processing_options) {
if (image.UsesGpu()) { if (image.UsesGpu()) {
return CreateStatusWithPayload( return CreateStatusWithPayload(
absl::StatusCode::kInvalidArgument, absl::StatusCode::kInvalidArgument,
@@ -293,6 +328,10 @@ absl::Status ImageSegmenter::SegmentAsync(
.At(Timestamp(timestamp_ms * kMicroSecondsPerMilliSecond))}, .At(Timestamp(timestamp_ms * kMicroSecondsPerMilliSecond))},
{kNormRectStreamName, {kNormRectStreamName,
MakePacket<NormalizedRect>(std::move(norm_rect)) MakePacket<NormalizedRect>(std::move(norm_rect))
.At(Timestamp(timestamp_ms * kMicroSecondsPerMilliSecond))},
{kOutputSizeStreamName,
MakePacket<std::pair<int, int>>(
std::make_pair(output_width, output_height))
.At(Timestamp(timestamp_ms * kMicroSecondsPerMilliSecond))}}); .At(Timestamp(timestamp_ms * kMicroSecondsPerMilliSecond))}});
} }
@@ -102,17 +102,36 @@ class ImageSegmenter : tasks::vision::core::BaseVisionTaskApi {
// //
// The image can be of any size with format RGB or RGBA. // The image can be of any size with format RGB or RGBA.
// //
// The output size is the same as the input image size.
//
// The optional 'image_processing_options' parameter can be used to specify // The optional 'image_processing_options' parameter can be used to specify
// the rotation to apply to the image before performing segmentation, by // the rotation to apply to the image before performing segmentation, by
// setting its 'rotation_degrees' field. Note that specifying a // setting its 'rotation_degrees' field. Note that specifying a
// region-of-interest using the 'region_of_interest' field is NOT supported // region-of-interest using the 'region_of_interest' field is NOT supported
// and will result in an invalid argument error being returned. // and will result in an invalid argument error being returned.
absl::StatusOr<ImageSegmenterResult> Segment( absl::StatusOr<ImageSegmenterResult> Segment(
mediapipe::Image image, mediapipe::Image image,
std::optional<core::ImageProcessingOptions> image_processing_options = std::optional<core::ImageProcessingOptions> image_processing_options =
std::nullopt); std::nullopt);
// Performs image segmentation on the provided single image.
// Only use this method when the ImageSegmenter is created with the image
// running mode.
//
// The image can be of any size with format RGB or RGBA.
//
// The output width and height specify the size of the resulted mask.
//
// The optional 'image_processing_options' parameter can be used to specify
// the rotation to apply to the image before performing segmentation, by
// setting its 'rotation_degrees' field. Note that specifying a
// region-of-interest using the 'region_of_interest' field is NOT supported
// and will result in an invalid argument error being returned.
absl::StatusOr<ImageSegmenterResult> Segment(
mediapipe::Image image, int output_width, int output_height,
std::optional<core::ImageProcessingOptions> image_processing_options =
std::nullopt);
// Performs image segmentation on the provided video frame. // Performs image segmentation on the provided video frame.
// Only use this method when the ImageSegmenter is created with the video // Only use this method when the ImageSegmenter is created with the video
// running mode. // running mode.
@@ -121,16 +140,39 @@ class ImageSegmenter : tasks::vision::core::BaseVisionTaskApi {
// provide the video frame's timestamp (in milliseconds). The input timestamps // provide the video frame's timestamp (in milliseconds). The input timestamps
// must be monotonically increasing. // must be monotonically increasing.
// //
// The optional 'image_processing_options' parameter can be used to specify // The output size is the same as the input image size.
// the rotation to apply to the image before performing segmentation, by //
// setting its 'rotation_degrees' field. Note that specifying a // The optional 'image_processing_options' parameter can be used
// region-of-interest using the 'region_of_interest' field is NOT supported // to specify the rotation to apply to the image before performing
// segmentation, by setting its 'rotation_degrees' field. Note that specifying
// a region-of-interest using the 'region_of_interest' field is NOT supported
// and will result in an invalid argument error being returned. // and will result in an invalid argument error being returned.
absl::StatusOr<ImageSegmenterResult> SegmentForVideo( absl::StatusOr<ImageSegmenterResult> SegmentForVideo(
mediapipe::Image image, int64_t timestamp_ms, mediapipe::Image image, int64_t timestamp_ms,
std::optional<core::ImageProcessingOptions> image_processing_options = std::optional<core::ImageProcessingOptions> image_processing_options =
std::nullopt); std::nullopt);
// Performs image segmentation on the provided video frame.
// Only use this method when the ImageSegmenter is created with the video
// running mode.
//
// The image can be of any size with format RGB or RGBA. It's required to
// provide the video frame's timestamp (in milliseconds). The input timestamps
// must be monotonically increasing.
//
// The output width and height specify the size of the resulted mask.
//
// The optional 'image_processing_options' parameter can be used
// to specify the rotation to apply to the image before performing
// segmentation, by setting its 'rotation_degrees' field. Note that specifying
// a region-of-interest using the 'region_of_interest' field is NOT supported
// and will result in an invalid argument error being returned.
absl::StatusOr<ImageSegmenterResult> SegmentForVideo(
mediapipe::Image image, int output_width, int output_height,
int64_t timestamp_ms,
std::optional<core::ImageProcessingOptions> image_processing_options =
std::nullopt);
// Sends live image data to perform image segmentation, and the results will // Sends live image data to perform image segmentation, and the results will
// be available via the "result_callback" provided in the // be available via the "result_callback" provided in the
// ImageSegmenterOptions. Only use this method when the ImageSegmenter is // ImageSegmenterOptions. Only use this method when the ImageSegmenter is
@@ -141,6 +183,8 @@ class ImageSegmenter : tasks::vision::core::BaseVisionTaskApi {
// sent to the image segmenter. The input timestamps must be monotonically // sent to the image segmenter. The input timestamps must be monotonically
// increasing. // increasing.
// //
// The output size is the same as the input image size.
//
// The optional 'image_processing_options' parameter can be used to specify // The optional 'image_processing_options' parameter can be used to specify
// the rotation to apply to the image before performing segmentation, by // the rotation to apply to the image before performing segmentation, by
// setting its 'rotation_degrees' field. Note that specifying a // setting its 'rotation_degrees' field. Note that specifying a
@@ -158,6 +202,36 @@ class ImageSegmenter : tasks::vision::core::BaseVisionTaskApi {
std::optional<core::ImageProcessingOptions> std::optional<core::ImageProcessingOptions>
image_processing_options = std::nullopt); image_processing_options = std::nullopt);
// Sends live image data to perform image segmentation, and the results will
// be available via the "result_callback" provided in the
// ImageSegmenterOptions. Only use this method when the ImageSegmenter is
// created with the live stream running mode.
//
// The image can be of any size with format RGB or RGBA. It's required to
// provide a timestamp (in milliseconds) to indicate when the input image is
// sent to the image segmenter. The input timestamps must be monotonically
// increasing.
//
// The output width and height specify the size of the resulted mask.
//
// The optional 'image_processing_options' parameter can be used to specify
// the rotation to apply to the image before performing segmentation, by
// setting its 'rotation_degrees' field. Note that specifying a
// region-of-interest using the 'region_of_interest' field is NOT supported
// and will result in an invalid argument error being returned.
//
// The "result_callback" prvoides
// - An ImageSegmenterResult.
// - The const reference to the corresponding input image that the image
// segmentation runs on. Note that the const reference to the image will
// no longer be valid when the callback returns. To access the image data
// outside of the callback, callers need to make a copy of the image.
// - The input timestamp in milliseconds.
absl::Status SegmentAsync(mediapipe::Image image, int output_width,
int output_height, int64_t timestamp_ms,
std::optional<core::ImageProcessingOptions>
image_processing_options = std::nullopt);
// Shuts down the ImageSegmenter when all works are done. // Shuts down the ImageSegmenter when all works are done.
absl::Status Close() { return runner_->Close(); } absl::Status Close() { return runner_->Close(); }
@@ -82,6 +82,7 @@ constexpr char kImageGpuTag[] = "IMAGE_GPU";
constexpr char kNormRectTag[] = "NORM_RECT"; constexpr char kNormRectTag[] = "NORM_RECT";
constexpr char kTensorsTag[] = "TENSORS"; constexpr char kTensorsTag[] = "TENSORS";
constexpr char kOutputSizeTag[] = "OUTPUT_SIZE"; constexpr char kOutputSizeTag[] = "OUTPUT_SIZE";
constexpr char kSizeTag[] = "SIZE";
constexpr char kQualityScoresTag[] = "QUALITY_SCORES"; constexpr char kQualityScoresTag[] = "QUALITY_SCORES";
constexpr char kSegmentationMetadataName[] = "SEGMENTER_METADATA"; constexpr char kSegmentationMetadataName[] = "SEGMENTER_METADATA";
@@ -356,6 +357,9 @@ absl::StatusOr<ImageAndTensorsOnDevice> ConvertImageToTensors(
// Describes image rotation and region of image to perform detection // Describes image rotation and region of image to perform detection
// on. // on.
// @Optional: rect covering the whole image is used if not specified. // @Optional: rect covering the whole image is used if not specified.
// OUTPUT_SIZE - std::pair<int, int> @Optional
// The output size of the mask, in width and height. If not specified, the
// output size of the input image is used.
// //
// Outputs: // Outputs:
// CONFIDENCE_MASK - mediapipe::Image @Multiple // CONFIDENCE_MASK - mediapipe::Image @Multiple
@@ -400,11 +404,16 @@ class ImageSegmenterGraph : public core::ModelTaskGraph {
if (!options.segmenter_options().has_output_type()) { if (!options.segmenter_options().has_output_type()) {
MP_RETURN_IF_ERROR(SanityCheck(sc)); MP_RETURN_IF_ERROR(SanityCheck(sc));
} }
std::optional<Source<std::pair<int, int>>> output_size;
if (HasInput(sc->OriginalNode(), kOutputSizeTag)) {
output_size = graph.In(kOutputSizeTag).Cast<std::pair<int, int>>();
}
ASSIGN_OR_RETURN( ASSIGN_OR_RETURN(
auto output_streams, auto output_streams,
BuildSegmentationTask( BuildSegmentationTask(
options, *model_resources, graph[Input<Image>(kImageTag)], options, *model_resources, graph[Input<Image>(kImageTag)],
graph[Input<NormalizedRect>::Optional(kNormRectTag)], graph)); graph[Input<NormalizedRect>::Optional(kNormRectTag)], output_size,
graph));
// TODO: remove deprecated output type support. // TODO: remove deprecated output type support.
if (options.segmenter_options().has_output_type()) { if (options.segmenter_options().has_output_type()) {
@@ -469,7 +478,8 @@ class ImageSegmenterGraph : public core::ModelTaskGraph {
absl::StatusOr<ImageSegmenterOutputs> BuildSegmentationTask( absl::StatusOr<ImageSegmenterOutputs> BuildSegmentationTask(
const ImageSegmenterGraphOptions& task_options, const ImageSegmenterGraphOptions& task_options,
const core::ModelResources& model_resources, Source<Image> image_in, const core::ModelResources& model_resources, Source<Image> image_in,
Source<NormalizedRect> norm_rect_in, Graph& graph) { Source<NormalizedRect> norm_rect_in,
std::optional<Source<std::pair<int, int>>> output_size, Graph& graph) {
MP_RETURN_IF_ERROR(SanityCheckOptions(task_options)); MP_RETURN_IF_ERROR(SanityCheckOptions(task_options));
// Adds preprocessing calculators and connects them to the graph input image // Adds preprocessing calculators and connects them to the graph input image
@@ -514,10 +524,14 @@ class ImageSegmenterGraph : public core::ModelTaskGraph {
image_and_tensors.tensors >> inference.In(kTensorsTag); image_and_tensors.tensors >> inference.In(kTensorsTag);
inference.Out(kTensorsTag) >> tensor_to_images.In(kTensorsTag); inference.Out(kTensorsTag) >> tensor_to_images.In(kTensorsTag);
// Adds image property calculator for output size. if (output_size.has_value()) {
auto& image_properties = graph.AddNode("ImagePropertiesCalculator"); *output_size >> tensor_to_images.In(kOutputSizeTag);
image_in >> image_properties.In("IMAGE"); } else {
image_properties.Out("SIZE") >> tensor_to_images.In(kOutputSizeTag); // Adds image property calculator for output size.
auto& image_properties = graph.AddNode("ImagePropertiesCalculator");
image_in >> image_properties.In(kImageTag);
image_properties.Out(kSizeTag) >> tensor_to_images.In(kOutputSizeTag);
}
// Exports multiple segmented masks. // Exports multiple segmented masks.
// TODO: remove deprecated output type support. // TODO: remove deprecated output type support.
@@ -155,3 +155,8 @@ cc_library(
"//mediapipe/tasks/cc/components/containers:landmark", "//mediapipe/tasks/cc/components/containers:landmark",
], ],
) )
cc_library(
name = "pose_landmarks_connections",
hdrs = ["pose_landmarks_connections.h"],
)
@@ -0,0 +1,39 @@
/* Copyright 2023 The MediaPipe Authors.
Licensed under the Apache License, Version 2.0 (the "License");
you may not use this file except in compliance with the License.
You may obtain a copy of the License at
http://www.apache.org/licenses/LICENSE-2.0
Unless required by applicable law or agreed to in writing, software
distributed under the License is distributed on an "AS IS" BASIS,
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
See the License for the specific language governing permissions and
limitations under the License.
==============================================================================*/
#ifndef MEDIAPIPE_TASKS_CC_VISION_POSE_LANDMARKER_POSE_LANDMARKS_CONNECTIONS_H_
#define MEDIAPIPE_TASKS_CC_VISION_POSE_LANDMARKER_POSE_LANDMARKS_CONNECTIONS_H_
#include <array>
namespace mediapipe {
namespace tasks {
namespace vision {
namespace pose_landmarker {
static constexpr std::array<std::array<int, 2>, 34> kPoseLandmarksConnections{{
{1, 2}, {0, 1}, {2, 3}, {3, 7}, {0, 4}, {4, 5}, {5, 6},
{6, 8}, {9, 10}, {11, 12}, {11, 13}, {13, 15}, {15, 17}, {15, 19},
{15, 21}, {17, 19}, {12, 14}, {14, 16}, {16, 18}, {16, 20}, {16, 22},
{18, 20}, {11, 23}, {12, 24}, {23, 24}, {23, 25}, {24, 26}, {25, 27},
{26, 28}, {27, 29}, {28, 30}, {29, 31}, {30, 32}, {27, 31},
}};
} // namespace pose_landmarker
} // namespace vision
} // namespace tasks
} // namespace mediapipe
#endif // MEDIAPIPE_TASKS_CC_VISION_POSE_LANDMARKER_POSE_LANDMARKS_CONNECTIONS_H_
+4
View File
@@ -66,7 +66,9 @@ strip_api_include_path_prefix(
"//mediapipe/tasks/ios/components/containers:sources/MPPClassificationResult.h", "//mediapipe/tasks/ios/components/containers:sources/MPPClassificationResult.h",
"//mediapipe/tasks/ios/components/containers:sources/MPPEmbedding.h", "//mediapipe/tasks/ios/components/containers:sources/MPPEmbedding.h",
"//mediapipe/tasks/ios/components/containers:sources/MPPEmbeddingResult.h", "//mediapipe/tasks/ios/components/containers:sources/MPPEmbeddingResult.h",
"//mediapipe/tasks/ios/components/containers:sources/MPPConnection.h",
"//mediapipe/tasks/ios/components/containers:sources/MPPDetection.h", "//mediapipe/tasks/ios/components/containers:sources/MPPDetection.h",
"//mediapipe/tasks/ios/components/containers:sources/MPPLandmark.h",
"//mediapipe/tasks/ios/core:sources/MPPBaseOptions.h", "//mediapipe/tasks/ios/core:sources/MPPBaseOptions.h",
"//mediapipe/tasks/ios/core:sources/MPPTaskOptions.h", "//mediapipe/tasks/ios/core:sources/MPPTaskOptions.h",
"//mediapipe/tasks/ios/core:sources/MPPTaskResult.h", "//mediapipe/tasks/ios/core:sources/MPPTaskResult.h",
@@ -160,6 +162,8 @@ apple_static_xcframework(
":MPPCategory.h", ":MPPCategory.h",
":MPPClassificationResult.h", ":MPPClassificationResult.h",
":MPPDetection.h", ":MPPDetection.h",
":MPPLandmark.h",
":MPPConnection.h",
":MPPCommon.h", ":MPPCommon.h",
":MPPTaskOptions.h", ":MPPTaskOptions.h",
":MPPTaskResult.h", ":MPPTaskResult.h",
@@ -25,7 +25,7 @@
static NSDictionary *const kPortraitImage = static NSDictionary *const kPortraitImage =
@{@"name" : @"portrait", @"type" : @"jpg", @"orientation" : @(UIImageOrientationUp)}; @{@"name" : @"portrait", @"type" : @"jpg", @"orientation" : @(UIImageOrientationUp)};
static NSDictionary *const kPortraitRotatedImage = static NSDictionary *const kPortraitRotatedImage =
@{@"name" : @"portrait_rotated", @"type" : @"jpg", @"orientation" : @(UIImageOrientationRight)}; @{@"name" : @"portrait_rotated", @"type" : @"jpg", @"orientation" : @(UIImageOrientationLeft)};
static NSDictionary *const kCatImage = @{@"name" : @"cat", @"type" : @"jpg"}; static NSDictionary *const kCatImage = @{@"name" : @"cat", @"type" : @"jpg"};
static NSString *const kShortRangeBlazeFaceModel = @"face_detection_short_range"; static NSString *const kShortRangeBlazeFaceModel = @"face_detection_short_range";
static NSArray<NSArray *> *const kPortraitExpectedKeypoints = @[ static NSArray<NSArray *> *const kPortraitExpectedKeypoints = @[
@@ -343,7 +343,7 @@ static NSString *const kLiveStreamTestsDictExpectationKey = @"expectation";
MPPGestureRecognizer *gestureRecognizer = MPPGestureRecognizer *gestureRecognizer =
[self createGestureRecognizerWithOptionsSucceeds:gestureRecognizerOptions]; [self createGestureRecognizerWithOptionsSucceeds:gestureRecognizerOptions];
MPPImage *mppImage = [self imageWithFileInfo:kPointingUpRotatedImage MPPImage *mppImage = [self imageWithFileInfo:kPointingUpRotatedImage
orientation:UIImageOrientationRight]; orientation:UIImageOrientationLeft];
MPPGestureRecognizerResult *gestureRecognizerResult = [gestureRecognizer recognizeImage:mppImage MPPGestureRecognizerResult *gestureRecognizerResult = [gestureRecognizer recognizeImage:mppImage
error:nil]; error:nil];
@@ -402,7 +402,7 @@ static NSString *const kLiveStreamTestsDictExpectationKey = @"expectation";
]; ];
MPPImage *image = [self imageWithFileInfo:kBurgerRotatedImage MPPImage *image = [self imageWithFileInfo:kBurgerRotatedImage
orientation:UIImageOrientationRight]; orientation:UIImageOrientationLeft];
[self assertResultsOfClassifyImage:image [self assertResultsOfClassifyImage:image
usingImageClassifier:imageClassifier usingImageClassifier:imageClassifier
@@ -425,7 +425,7 @@ static NSString *const kLiveStreamTestsDictExpectationKey = @"expectation";
displayName:nil] ]; displayName:nil] ];
MPPImage *image = [self imageWithFileInfo:kMultiObjectsRotatedImage MPPImage *image = [self imageWithFileInfo:kMultiObjectsRotatedImage
orientation:UIImageOrientationRight]; orientation:UIImageOrientationLeft];
// roi around folding chair // roi around folding chair
MPPImageClassifierResult *imageClassifierResult = MPPImageClassifierResult *imageClassifierResult =
@@ -438,7 +438,7 @@ static NSString *const kLiveStreamTestsDictExpectationKey = @"expectation";
[[MPPObjectDetectorResult alloc] initWithDetections:detections timestampInMilliseconds:0]; [[MPPObjectDetectorResult alloc] initWithDetections:detections timestampInMilliseconds:0];
MPPImage *image = [self imageWithFileInfo:kCatsAndDogsRotatedImage MPPImage *image = [self imageWithFileInfo:kCatsAndDogsRotatedImage
orientation:UIImageOrientationRight]; orientation:UIImageOrientationLeft];
[self assertResultsOfDetectInImage:image [self assertResultsOfDetectInImage:image
usingObjectDetector:objectDetector usingObjectDetector:objectDetector
@@ -62,10 +62,10 @@ NS_SWIFT_NAME(MPImage)
/** /**
* Initializes an `MPPImage` object with the given `UIImage`. * Initializes an `MPPImage` object with the given `UIImage`.
* The orientation of the newly created `MPPImage` will be `UIImageOrientationUp`. * The orientation of the newly created `MPPImage` will be equal to the `imageOrientation` of
* Hence, if this image is used as input for any MediaPipe vision tasks, inference will be * `UIImage` and when sent to the vision tasks for inference, rotation will be applied accordingly.
* performed on the it without any rotation. To create an `MPPImage` with a different orientation, * To create an `MPPImage` with an orientation different from its `imageOrientation`, please use
* please use `[MPPImage initWithImage:orientation:error:]`. * `[MPPImage initWithImage:orientation:error:]`.
* *
* @param image The image to use as the source. Its `CGImage` property must not be `NULL`. * @param image The image to use as the source. Its `CGImage` property must not be `NULL`.
* @param error An optional error parameter populated when there is an error in initializing the * @param error An optional error parameter populated when there is an error in initializing the
@@ -77,14 +77,19 @@ NS_SWIFT_NAME(MPImage)
- (nullable instancetype)initWithUIImage:(UIImage *)image error:(NSError **)error; - (nullable instancetype)initWithUIImage:(UIImage *)image error:(NSError **)error;
/** /**
* Initializes an `MPPImage` object with the given `UIImabe` and orientation. * Initializes an `MPPImage` object with the given `UIImage` and orientation. The given orientation
* will be used to calculate the rotation to be applied to the `UIImage` before inference is
* performed on it by the vision tasks. The `imageOrientation` stored in the `UIImage` is ignored
* when `MPImage` objects created by this method are sent to the vision tasks for inference. Use
* `[MPPImage initWithImage:orientation:error:]` to initialize images with the `imageOrientation` of
* `UIImage`.
* *
* If the newly created `MPPImage` is used as input for any MediaPipe vision tasks, inference * If the newly created `MPPImage` is used as input for any MediaPipe vision tasks, inference
* will be performed on a copy of the image rotated according to the orientation. * will be performed on a copy of the image rotated according to the orientation.
* *
* @param image The image to use as the source. Its `CGImage` property must not be `NULL`. * @param image The image to use as the source. Its `CGImage` property must not be `NULL`.
* @param orientation The display orientation of the image. This will be stored in the property * @param orientation The display orientation of the image. This will be stored in the property
* `orientation`. `MPPImage`. * `orientation` `MPPImage` and will override the `imageOrientation` of the passed in `UIImage`.
* @param error An optional error parameter populated when there is an error in initializing the * @param error An optional error parameter populated when there is an error in initializing the
* `MPPImage`. * `MPPImage`.
* *
@@ -30,13 +30,13 @@ using ::mediapipe::tasks::core::PacketsCallback;
} // namespace } // namespace
/** Rotation degrees for a 90 degree rotation to the right. */ /** Rotation degrees for a 90 degree rotation to the right. */
static const NSInteger kMPPOrientationDegreesRight = -90; static const NSInteger kMPPOrientationDegreesRight = -270;
/** Rotation degrees for a 180 degree rotation. */ /** Rotation degrees for a 180 degree rotation. */
static const NSInteger kMPPOrientationDegreesDown = -180; static const NSInteger kMPPOrientationDegreesDown = -180;
/** Rotation degrees for a 90 degree rotation to the left. */ /** Rotation degrees for a 90 degree rotation to the left. */
static const NSInteger kMPPOrientationDegreesLeft = -270; static const NSInteger kMPPOrientationDegreesLeft = -90;
static NSString *const kTaskPrefix = @"com.mediapipe.tasks.vision"; static NSString *const kTaskPrefix = @"com.mediapipe.tasks.vision";
@@ -30,7 +30,7 @@ NS_ASSUME_NONNULL_BEGIN
* The delegate of `MPPFaceLandmarker` must adopt `MPPFaceLandmarkerLiveStreamDelegate` protocol. * The delegate of `MPPFaceLandmarker` must adopt `MPPFaceLandmarkerLiveStreamDelegate` protocol.
* The methods in this protocol are optional. * The methods in this protocol are optional.
*/ */
NS_SWIFT_NAME(FaceDetectorLiveStreamDelegate) NS_SWIFT_NAME(FaceLandmarkerLiveStreamDelegate)
@protocol MPPFaceLandmarkerLiveStreamDelegate <NSObject> @protocol MPPFaceLandmarkerLiveStreamDelegate <NSObject>
/** /**

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