Skip to content
Open
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
67 changes: 62 additions & 5 deletions cpp/HybridTfliteModule.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -2,11 +2,18 @@
#include "TfliteHelpers.hpp"

#include <memory>
#include <utility>
#include <vector>

#if defined(ANDROID)
#include <tflite/c/c_api.h>
#include <tflite/delegates/gpu/delegate.h>
#include <tflite/delegates/nnapi/nnapi_delegate_c_api.h>
#elif defined(__APPLE__)
#include <TensorFlowLiteC/TensorFlowLiteC.h>
#if FAST_TFLITE_ENABLE_CORE_ML
#include <TensorFlowLiteCCoreML/TensorFlowLiteCCoreML.h>
#endif
#else
#error "Invalid Platform!"
#endif
Expand All @@ -32,6 +39,37 @@ TfLiteDelegate* getDelegate(TensorflowModelDelegate delegateType) {
"\"!");
}

/**
* TFLite's C API does not transfer delegate ownership to the interpreter: the
* caller must keep a delegate alive for the interpreter's lifetime and free it
* afterwards with the delegate's own delete function.
*/
struct DelegateDeleter {
TensorflowModelDelegate delegateType;

void operator()(TfLiteDelegate* delegate) const {
switch (delegateType) {
#if defined(__APPLE__) && FAST_TFLITE_ENABLE_CORE_ML
case TensorflowModelDelegate::CORE_ML:
TfLiteCoreMlDelegateDelete(delegate);
return;
#endif
#if defined(ANDROID)
case TensorflowModelDelegate::ANDROID_GPU:
TfLiteGpuDelegateV2Delete(delegate);
return;
case TensorflowModelDelegate::NNAPI:
TfLiteNnapiDelegateDelete(delegate);
return;
#endif
default:
// getDelegate() throws for every other type on this platform.
return;
}
}
};
using OwnedDelegate = std::unique_ptr<TfLiteDelegate, DelegateDeleter>;

std::shared_ptr<HybridTfliteModelSpec>
HybridTfliteModule::createModel(const std::shared_ptr<ArrayBuffer>& modelData,
const std::vector<TensorflowModelDelegate>& delegates) {
Expand All @@ -50,20 +88,39 @@ HybridTfliteModule::createModel(const std::shared_ptr<ArrayBuffer>& modelData,

// Add all hardware accelerated delegates (e.g. GPU, NPU, ...)
// if any. The default CPU delegate will always be available.
std::vector<TensorflowModelDelegate> effectiveDelegates;
std::vector<OwnedDelegate> ownedDelegates;
effectiveDelegates.reserve(delegates.size());
ownedDelegates.reserve(delegates.size());
for (const TensorflowModelDelegate& delegateType : delegates) {
TfLiteDelegate* delegate = getDelegate(delegateType);
TfLiteInterpreterOptionsAddDelegate(options.get(), delegate);
OwnedDelegate delegate(getDelegate(delegateType), DelegateDeleter{delegateType});
if (delegate == nullptr) {
// e.g. CoreML on devices without a Neural Engine — fall back to CPU
// instead of registering a null delegate with the interpreter.
continue;
}
TfLiteInterpreterOptionsAddDelegate(options.get(), delegate.get());
effectiveDelegates.push_back(delegateType);
ownedDelegates.push_back(std::move(delegate));
}

TfLiteInterpreter* rawInterpreter = TfLiteInterpreterCreate(model.get(), options.get());
if (rawInterpreter == nullptr) {
// `ownedDelegates` frees the delegates on unwind.
throw std::runtime_error("Failed to create TFLite interpreter!");
}
// The delegates travel with the interpreter and are freed right after it,
// so they can never be deleted while the interpreter still uses them.
const std::shared_ptr<TfLiteInterpreter> interpreter(
rawInterpreter, [modelData](TfLiteInterpreter* value) { TfLiteInterpreterDelete(value); });
rawInterpreter,
[modelData, ownedDelegates = std::move(ownedDelegates)](TfLiteInterpreter* value) mutable {
TfLiteInterpreterDelete(value);
ownedDelegates.clear();
});

// Wrap in HybridTfliteModel — stores shared_ptr<ArrayBuffer> to keep model data bytes alive
return std::make_shared<HybridTfliteModel>(interpreter, modelData, delegates);
// Wrap in HybridTfliteModel — stores shared_ptr<ArrayBuffer> to keep model data bytes alive.
// Only the delegates that were actually registered are reported via `getDelegates()`.
return std::make_shared<HybridTfliteModel>(interpreter, modelData, effectiveDelegates);
}

} // namespace margelo::nitro::tflite