Depends on TFLite shim header.

PiperOrigin-RevId: 508491302
This commit is contained in:
MediaPipe Team
2023-02-09 15:29:47 -08:00
committed by Copybara-Service
parent 99fc975f49
commit fd764dae0a
9 changed files with 60 additions and 41 deletions
+6 -2
View File
@@ -13,6 +13,8 @@
# limitations under the License.
#
load("@org_tensorflow//tensorflow/lite/core/shims:cc_library_with_tflite.bzl", "cc_library_with_tflite")
licenses(["notice"])
package(default_visibility = [
@@ -110,10 +112,13 @@ cc_library(
],
)
cc_library(
cc_library_with_tflite(
name = "tflite_model_loader",
srcs = ["tflite_model_loader.cc"],
hdrs = ["tflite_model_loader.h"],
tflite_deps = [
"@org_tensorflow//tensorflow/lite/core/shims:framework_stable",
],
visibility = ["//visibility:public"],
deps = [
"//mediapipe/framework/api2:packet",
@@ -121,6 +126,5 @@ cc_library(
"//mediapipe/framework/port:status",
"//mediapipe/framework/port:statusor",
"//mediapipe/util:resource_util",
"@org_tensorflow//tensorflow/lite:framework",
],
)
+5 -3
View File
@@ -19,6 +19,8 @@
namespace mediapipe {
using FlatBufferModel = ::tflite_shims::FlatBufferModel;
absl::StatusOr<api2::Packet<TfLiteModelPtr>> TfLiteModelLoader::LoadFromPath(
const std::string& path) {
std::string model_path = path;
@@ -36,12 +38,12 @@ absl::StatusOr<api2::Packet<TfLiteModelPtr>> TfLiteModelLoader::LoadFromPath(
mediapipe::GetResourceContents(resolved_path, &model_blob));
}
auto model = tflite::FlatBufferModel::VerifyAndBuildFromBuffer(
model_blob.data(), model_blob.size());
auto model = FlatBufferModel::VerifyAndBuildFromBuffer(model_blob.data(),
model_blob.size());
RET_CHECK(model) << "Failed to load model from path " << model_path;
return api2::MakePacket<TfLiteModelPtr>(
model.release(),
[model_blob = std::move(model_blob)](tflite::FlatBufferModel* model) {
[model_blob = std::move(model_blob)](FlatBufferModel* model) {
// It's required that model_blob is deleted only after
// model is deleted, hence capturing model_blob.
delete model;
+7 -3
View File
@@ -15,16 +15,20 @@
#ifndef MEDIAPIPE_UTIL_TFLITE_TFLITE_MODEL_LOADER_H_
#define MEDIAPIPE_UTIL_TFLITE_TFLITE_MODEL_LOADER_H_
#include <functional>
#include <memory>
#include <string>
#include "mediapipe/framework/api2/packet.h"
#include "mediapipe/framework/port/status.h"
#include "mediapipe/framework/port/statusor.h"
#include "tensorflow/lite/model.h"
#include "tensorflow/lite/core/shims/cc/model.h"
namespace mediapipe {
// Represents a TfLite model as a FlatBuffer.
using TfLiteModelPtr =
std::unique_ptr<tflite::FlatBufferModel,
std::function<void(tflite::FlatBufferModel*)>>;
std::unique_ptr<tflite_shims::FlatBufferModel,
std::function<void(tflite_shims::FlatBufferModel*)>>;
class TfLiteModelLoader {
public: