blob: 9f595f0c08b30cdc3677d1854ccd8cf0bdd42d9e [file] [edit]
From 951dda10343553f8e7fbb291b7b4d6a3b1fe9d21 Mon Sep 17 00:00:00 2001
From: TFLite Patcher <tflite@localhost>
Date: Fri, 18 Oct 2024 18:04:20 +0800
Subject: [PATCH] Support Convolution2DTransposeBias custom op
PATCH_NAME=Convolution2DTransposeBias
---
.../delegates/gpu/cl/testing/performance_profiling.cc | 8 ++++++++
tensorflow/lite/tools/benchmark/benchmark_tflite_model.cc | 6 +++++-
.../tools/evaluation/stages/tflite_inference_stage.cc | 2 ++
3 files changed, 15 insertions(+), 1 deletion(-)
diff --git a/tensorflow/lite/delegates/gpu/cl/testing/performance_profiling.cc b/tensorflow/lite/delegates/gpu/cl/testing/performance_profiling.cc
index 67c6ed69..3a470b04 100644
--- a/tensorflow/lite/delegates/gpu/cl/testing/performance_profiling.cc
+++ b/tensorflow/lite/delegates/gpu/cl/testing/performance_profiling.cc
@@ -47,6 +47,8 @@ absl::Status RunPredefinedLayoutSample(const std::string& model_name) {
auto flatbuffer = tflite::FlatBufferModel::BuildFromFile(model_name.c_str());
GraphFloat32 graph_cl;
ops::builtin::BuiltinOpResolver op_resolver;
+ TfLiteRegistration reg = { nullptr, nullptr, nullptr, nullptr };
+ op_resolver.AddCustom("Convolution2DTransposeBias", &reg);
RETURN_IF_ERROR(BuildFromFlatBuffer(*flatbuffer, op_resolver, &graph_cl,
/*allow_quant_ops=*/true));
@@ -91,6 +93,8 @@ absl::Status RunExternalImmutableSample(const std::string& model_name) {
auto flatbuffer = tflite::FlatBufferModel::BuildFromFile(model_name.c_str());
GraphFloat32 graph_cl;
ops::builtin::BuiltinOpResolver op_resolver;
+ TfLiteRegistration reg = { nullptr, nullptr, nullptr, nullptr };
+ op_resolver.AddCustom("Convolution2DTransposeBias", &reg);
RETURN_IF_ERROR(BuildFromFlatBuffer(*flatbuffer, op_resolver, &graph_cl,
/*allow_quant_ops*/ true));
@@ -143,6 +147,8 @@ absl::Status RunSerializedTest(const std::string& model_name) {
auto flatbuffer = tflite::FlatBufferModel::BuildFromFile(model_name.c_str());
GraphFloat32 graph_cl;
ops::builtin::BuiltinOpResolver op_resolver;
+ TfLiteRegistration reg = { nullptr, nullptr, nullptr, nullptr };
+ op_resolver.AddCustom("Convolution2DTransposeBias", &reg);
RETURN_IF_ERROR(BuildFromFlatBuffer(*flatbuffer, op_resolver, &graph_cl,
/*allow_quant_ops*/ true));
@@ -275,6 +281,8 @@ absl::Status RunModelSample(const std::string& model_name) {
auto flatbuffer = tflite::FlatBufferModel::BuildFromFile(model_name.c_str());
GraphFloat32 graph_cl;
ops::builtin::BuiltinOpResolver op_resolver;
+ TfLiteRegistration reg = { nullptr, nullptr, nullptr, nullptr };
+ op_resolver.AddCustom("Convolution2DTransposeBias", &reg);
RETURN_IF_ERROR(BuildFromFlatBuffer(*flatbuffer, op_resolver, &graph_cl,
/*allow_quant_ops*/ true));
diff --git a/tensorflow/lite/tools/benchmark/benchmark_tflite_model.cc b/tensorflow/lite/tools/benchmark/benchmark_tflite_model.cc
index 37e488f0..750858c2 100644
--- a/tensorflow/lite/tools/benchmark/benchmark_tflite_model.cc
+++ b/tensorflow/lite/tools/benchmark/benchmark_tflite_model.cc
@@ -69,7 +69,11 @@ void RegisterSelectedOps(::tflite::MutableOpResolver* resolver);
// library with another definition of this function (presumably to actually
// register custom ops), that version will be used instead.
void ABSL_ATTRIBUTE_WEAK
-RegisterSelectedOps(::tflite::MutableOpResolver* resolver) {}
+RegisterSelectedOps(::tflite::MutableOpResolver* resolver) {
+ static TfLiteRegistration reg = { nullptr, nullptr, nullptr, nullptr };
+ resolver->AddCustom("Convolution2DTransposeBias",
+ &reg);
+}
namespace tflite {
namespace benchmark {
diff --git a/tensorflow/lite/tools/evaluation/stages/tflite_inference_stage.cc b/tensorflow/lite/tools/evaluation/stages/tflite_inference_stage.cc
index 9d6f7f9b..700bf461 100644
--- a/tensorflow/lite/tools/evaluation/stages/tflite_inference_stage.cc
+++ b/tensorflow/lite/tools/evaluation/stages/tflite_inference_stage.cc
@@ -146,6 +146,8 @@ TfLiteStatus TfliteInferenceStage::Init(
resolver_ = std::make_unique<
ops::builtin::BuiltinOpResolverWithoutDefaultDelegates>();
}
+ TfLiteRegistration reg = { nullptr, nullptr, nullptr, nullptr };
+ resolver_->AddCustom("Convolution2DTransposeBias", &reg);
RegisterSelectedOps(resolver_.get());
InterpreterBuilder(*model_, *resolver_)(&interpreter_);
if (!interpreter_) {