Depends on TFLite shim header.
PiperOrigin-RevId: 508491302
This commit is contained in:
committed by
Copybara-Service
parent
99fc975f49
commit
fd764dae0a
@@ -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",
|
||||
],
|
||||
)
|
||||
|
||||
@@ -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;
|
||||
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user