Current section
Files
Jump to
Current section
Files
c_src/tflite/interpreter.cpp
#include <string>
#include <erl_nif.h>
#include "../nif_utils.hpp"
#include "../erlang_nif_resource.h"
#include "../helper.h"
#include "tensorflow/lite/interpreter.h"
#include "tensorflow/lite/kernels/register.h"
#include "tensorflow/lite/model.h"
#include "interpreter.h"
#include "tflitetensor.h"
#include "status.h"
ERL_NIF_TERM interpreter_new(ErlNifEnv *env, int argc, const ERL_NIF_TERM argv[]) {
NifResInterpreter * res = nullptr;
ERL_NIF_TERM ret;
if (!(res = NifResInterpreter::allocate_resource(env, ret))) {
return ret;
}
res->val = new tflite::Interpreter();
ret = enif_make_resource(env, res);
enif_release_resource(res);
return erlang::nif::ok(env, ret);
}
ERL_NIF_TERM interpreter_set_inputs(ErlNifEnv *env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
ERL_NIF_TERM self_nif = argv[0];
ERL_NIF_TERM inputs_nif = argv[1];
NifResInterpreter * self_res;
std::vector<int> inputs;
ERL_NIF_TERM ret;
if (!(self_res = NifResInterpreter::get_resource(env, self_nif, ret))) {
return ret;
}
if (!erlang::nif::get_list(env, inputs_nif, inputs)) {
return erlang::nif::error(env, "expecting `inputs` to be a list of non-negative integers");
}
TfLiteStatus status = self_res->val->SetInputs(inputs);
return tflite_status_to_erl_term(env, status);
}
ERL_NIF_TERM interpreter_set_outputs(ErlNifEnv *env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
ERL_NIF_TERM self_nif = argv[0];
ERL_NIF_TERM outputs_nif = argv[1];
NifResInterpreter * self_res;
std::vector<int> outputs;
ERL_NIF_TERM ret;
if (!(self_res = NifResInterpreter::get_resource(env, self_nif, ret))) {
return ret;
}
if (!erlang::nif::get_list(env, outputs_nif, outputs)) {
return erlang::nif::error(env, "expecting `outputs` to be a list of non-negative integers");
}
TfLiteStatus status = self_res->val->SetOutputs(outputs);
return tflite_status_to_erl_term(env, status);
}
ERL_NIF_TERM interpreter_set_variables(ErlNifEnv *env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
ERL_NIF_TERM self_nif = argv[0];
ERL_NIF_TERM variables_nif = argv[1];
NifResInterpreter * self_res;
std::vector<int> variables;
ERL_NIF_TERM ret;
if (!(self_res = NifResInterpreter::get_resource(env, self_nif, ret))) {
return ret;
}
if (!erlang::nif::get_list(env, variables_nif, variables)) {
return erlang::nif::error(env, "expecting `variables` to be a list of non-negative integers");
}
TfLiteStatus status = self_res->val->SetVariables(variables);
return tflite_status_to_erl_term(env, status);
}
ERL_NIF_TERM interpreter_inputs(ErlNifEnv *env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
ERL_NIF_TERM self_nif = argv[0];
NifResInterpreter *self_res;
ERL_NIF_TERM ret;
if (!(self_res = NifResInterpreter::get_resource(env, self_nif, ret))) {
return ret;
}
const std::vector<int>& inputs = self_res->val->inputs();
size_t cnt = inputs.size();
if (cnt > 0) {
ERL_NIF_TERM * arr = (ERL_NIF_TERM *)enif_alloc(sizeof(ERL_NIF_TERM) * cnt);
if (!arr) {
return erlang::nif::error(env, "enif_alloc failed");
}
for (size_t i = 0; i < cnt; i++) {
arr[i] = enif_make_int(env, inputs[i]);
}
ret = enif_make_list_from_array(env, arr, (unsigned)cnt);
enif_free((void *)arr);
} else {
// Returns an empty list if cnt is 0.
ret = enif_make_list(env, 0, nullptr);
}
return erlang::nif::ok(env, ret);
}
ERL_NIF_TERM interpreter_get_input_name(ErlNifEnv *env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
ERL_NIF_TERM self_nif = argv[0];
ERL_NIF_TERM index_nif = argv[1];
int index;
NifResInterpreter *self_res;
ERL_NIF_TERM ret;
if (!(self_res = NifResInterpreter::get_resource(env, self_nif, ret))) {
return ret;
}
if (!enif_get_int(env, index_nif, &index)) {
return erlang::nif::error(env, "expecting index to be an integer");
}
const auto& inputs = self_res->val->inputs();
if (inputs.size() <= index || index < 0) {
return erlang::nif::error(env, "index out of bound");
}
const char * name = self_res->val->GetInputName(index);
if (name == nullptr) {
return erlang::nif::error(env, "cannot get tensor's name");
}
return erlang::nif::ok(env, erlang::nif::make_binary(env, name));
}
ERL_NIF_TERM interpreter_outputs(ErlNifEnv *env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
ERL_NIF_TERM self_nif = argv[0];
NifResInterpreter *self_res;
ERL_NIF_TERM ret;
if (!(self_res = NifResInterpreter::get_resource(env, self_nif, ret))) {
return ret;
}
const std::vector<int>& outputs = self_res->val->outputs();
if (erlang::nif::make(env, outputs, ret)) {
return erlang::nif::error(env, "enif_alloc failed");
}
return erlang::nif::ok(env, ret);
}
ERL_NIF_TERM interpreter_variables(ErlNifEnv *env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
ERL_NIF_TERM self_nif = argv[0];
NifResInterpreter *self_res;
ERL_NIF_TERM ret;
if (!(self_res = NifResInterpreter::get_resource(env, self_nif, ret))) {
return ret;
}
const std::vector<int>& variables = self_res->val->variables();
if (erlang::nif::make(env, variables, ret)) {
return erlang::nif::error(env, "enif_alloc failed");
}
return erlang::nif::ok(env, ret);
}
ERL_NIF_TERM interpreter_get_output_name(ErlNifEnv *env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
ERL_NIF_TERM self_nif = argv[0];
ERL_NIF_TERM index_nif = argv[1];
int index;
NifResInterpreter *self_res;
ERL_NIF_TERM ret;
if (!(self_res = NifResInterpreter::get_resource(env, self_nif, ret))) {
return ret;
}
if (!enif_get_int(env, index_nif, &index)) {
return erlang::nif::error(env, "expecting index to be an integer");
}
const auto& outputs = self_res->val->outputs();
if (outputs.size() <= index || index < 0) {
return erlang::nif::error(env, "index out of bound");
}
const char * name = self_res->val->GetOutputName(index);
if (name == nullptr) {
return erlang::nif::error(env, "cannot get tensor's name");
}
return erlang::nif::ok(env, erlang::nif::make_binary(env, name));
}
ERL_NIF_TERM interpreter_tensors_size(ErlNifEnv *env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
ERL_NIF_TERM self_nif = argv[0];
NifResInterpreter *self_res;
ERL_NIF_TERM ret;
if (!(self_res = NifResInterpreter::get_resource(env, self_nif, ret))) {
return ret;
}
return erlang::nif::ok(env, enif_make_uint64(env, self_res->val->tensors_size()));
}
ERL_NIF_TERM interpreter_nodes_size(ErlNifEnv *env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
ERL_NIF_TERM self_nif = argv[0];
NifResInterpreter *self_res;
ERL_NIF_TERM ret;
if (!(self_res = NifResInterpreter::get_resource(env, self_nif, ret))) {
return ret;
}
return erlang::nif::ok(env, enif_make_uint64(env, self_res->val->nodes_size()));
}
ERL_NIF_TERM interpreter_execution_plan(ErlNifEnv *env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
ERL_NIF_TERM self_nif = argv[0];
NifResInterpreter *self_res;
ERL_NIF_TERM ret;
if (!(self_res = NifResInterpreter::get_resource(env, self_nif, ret))) {
return ret;
}
const std::vector<int>& execution_plan = self_res->val->execution_plan();
if (erlang::nif::make(env, execution_plan, ret)) {
return erlang::nif::error(env, "enif_alloc failed");
}
return erlang::nif::ok(env, ret);
}
ERL_NIF_TERM interpreter_tensor(ErlNifEnv *env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
ERL_NIF_TERM self_nif = argv[0];
ERL_NIF_TERM index_nif = argv[1];
int index;
NifResInterpreter *self_res;
ERL_NIF_TERM ret;
if (!(self_res = NifResInterpreter::get_resource(env, self_nif, ret))) {
return ret;
}
if (!enif_get_int(env, index_nif, &index)) {
return erlang::nif::error(env, "expecting index to be an integer");
}
const size_t num_tensors = self_res->val->tensors_size();
if (num_tensors <= index || index < 0) {
return erlang::nif::error(env, "index out of bound");
}
NifResTfLiteTensor * tensor_res = nullptr;
auto cached_tensor_res = self_res->tensors->find(index);
if (cached_tensor_res != self_res->tensors->end()) {
tensor_res = cached_tensor_res->second;
} else {
if (!(tensor_res = NifResTfLiteTensor::allocate_resource(env, ret))) {
return ret;
}
tensor_res->val = self_res->val->tensor(index);
tensor_res->borrowed = true;
}
ERL_NIF_TERM tensor_type;
if (!_tflitetensor_type(env, tensor_res->val, tensor_type)) {
tensor_type = erlang::nif::atom(env, "unknown");
}
ERL_NIF_TERM tensor_shape;
if (!_tflitetensor_shape(env, tensor_res->val, tensor_shape)) {
return erlang::nif::error(env, "cannot allocate memory for tensor shape");
}
ERL_NIF_TERM tensor_shape_signature;
if (!_tflitetensor_shape_signature(env, tensor_res->val, tensor_shape_signature)) {
return erlang::nif::error(env, "cannot allocate memory for tensor shape signature");
}
ERL_NIF_TERM tensor_name;
if (!_tflitetensor_name(env, tensor_res->val, tensor_name)) {
return erlang::nif::error(env, "cannot allocate memory for tensor name");
}
ERL_NIF_TERM tensor_quantization_params;
if (!_tflitetensor_quantization_params(env, tensor_res->val, tensor_quantization_params)) {
return erlang::nif::error(env, "cannot allocate memory for tensor quantization params");
}
ERL_NIF_TERM tensor_sparsity_params;
if (!_tflitetensor_sparsity_params(env, tensor_res->val, tensor_sparsity_params)) {
return erlang::nif::error(env, "cannot allocate memory for tensor sparsity params");
}
ERL_NIF_TERM tensor_reference = enif_make_resource(env, tensor_res);
if (cached_tensor_res == self_res->tensors->end()) {
enif_keep_resource(tensor_res);
(*self_res->tensors)[index] = tensor_res;
}
return erlang::nif::ok(env, enif_make_tuple8(
env,
tensor_name,
index_nif,
tensor_shape,
tensor_shape_signature,
tensor_type,
tensor_quantization_params,
tensor_sparsity_params,
tensor_reference
));
}
ERL_NIF_TERM interpreter_signature_keys(ErlNifEnv *env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
ERL_NIF_TERM self_nif = argv[0];
NifResInterpreter *self_res;
ERL_NIF_TERM ret;
if (!(self_res = NifResInterpreter::get_resource(env, self_nif, ret))) {
return ret;
}
const std::vector<const std::string*> signature_keys = self_res->val->signature_keys();
if (erlang::nif::make(env, signature_keys, ret)) {
return erlang::nif::error(env, "enif_alloc failed");
}
return erlang::nif::ok(env, ret);
}
ERL_NIF_TERM interpreter_input_tensor(ErlNifEnv *env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 3) return enif_make_badarg(env);
ERL_NIF_TERM self_nif = argv[0];
ERL_NIF_TERM index_nif = argv[1];
ERL_NIF_TERM data_nif = argv[2];
int index;
ErlNifBinary data;
NifResInterpreter *self_res;
ERL_NIF_TERM ret;
if (!(self_res = NifResInterpreter::get_resource(env, self_nif, ret))) {
return ret;
}
if (!enif_get_int(env, index_nif, &index)) {
return erlang::nif::error(env, "expecting index to be an integer");
}
if (!enif_inspect_binary(env, data_nif, &data)) {
return erlang::nif::error(env, "cannot get input data");
}
const auto& inputs = self_res->val->inputs();
if (inputs.size() <= index || index < 0) {
return erlang::nif::error(env, "index out of bound");
}
auto input_tensor = self_res->val->input_tensor(index);
if (input_tensor->data.data == nullptr) {
return erlang::nif::error(env, "tensor is not allocated yet? Please call TFLiteBEAM.Interpreter.allocate_tensors first");
}
size_t maximum_bytes = input_tensor->bytes;
if (data.size < maximum_bytes) {
maximum_bytes = data.size;
}
memcpy(input_tensor->data.data, data.data, maximum_bytes);
return erlang::nif::ok(env);
}
ERL_NIF_TERM interpreter_output_tensor(ErlNifEnv *env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
ERL_NIF_TERM self_nif = argv[0];
ERL_NIF_TERM index_nif = argv[1];
int index;
NifResInterpreter *self_res;
ERL_NIF_TERM ret;
if (!(self_res = NifResInterpreter::get_resource(env, self_nif, ret))) {
return ret;
}
if (!enif_get_int(env, index_nif, &index)) {
return erlang::nif::error(env, "expecting index to be an integer");
}
const auto& outputs = self_res->val->outputs();
if (outputs.size() <= index || index < 0) {
return erlang::nif::error(env, "index out of bound");
}
auto t = self_res->val->output_tensor(index);
ErlNifBinary tensor_data;
size_t tensor_size = t->bytes;
if (!enif_alloc_binary(tensor_size, &tensor_data)) {
return erlang::nif::error(env, "cannot allocate enough memory for the tensor");
}
memcpy(tensor_data.data, t->data.data, tensor_size);
return erlang::nif::ok(env, enif_make_binary(env, &tensor_data));
}
ERL_NIF_TERM interpreter_allocate_tensors(ErlNifEnv *env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
ERL_NIF_TERM self_nif = argv[0];
NifResInterpreter * self_res;
ERL_NIF_TERM ret;
if (!(self_res = NifResInterpreter::get_resource(env, self_nif, ret))) {
return ret;
}
switch (self_res->val->AllocateTensors()) {
case kTfLiteOk:
return erlang::nif::atom(env, "ok");
case kTfLiteError:
return erlang::nif::error(env, "General runtime error");
case kTfLiteDelegateError:
return erlang::nif::error(env, "TfLiteDelegate");
case kTfLiteApplicationError:
return erlang::nif::error(env, "Application");
case kTfLiteDelegateDataNotFound:
return erlang::nif::error(env, "DelegateDataNotFound");
case kTfLiteDelegateDataWriteError:
return erlang::nif::error(env, "DelegateDataWriteError");
case kTfLiteDelegateDataReadError:
return erlang::nif::error(env, "DelegateDataReadError");
default:
return erlang::nif::error(env, "unknown error");
}
}
ERL_NIF_TERM interpreter_get_signature_defs(ErlNifEnv *env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
ERL_NIF_TERM self_nif = argv[0];
NifResInterpreter * self_res;
ERL_NIF_TERM ret;
if (!(self_res = NifResInterpreter::get_resource(env, self_nif, ret))) {
return ret;
}
auto interpreter_ = self_res->val;
ERL_NIF_TERM result;
size_t num_items = interpreter_->signature_keys().size();
if (num_items == 0) {
return erlang::nif::ok(env, erlang::nif::atom(env, "nil"));
}
ERL_NIF_TERM * keys = (ERL_NIF_TERM *)enif_alloc(sizeof(ERL_NIF_TERM) * num_items);
if (keys == nullptr) {
return erlang::nif::error(env, "out of memory");
}
ERL_NIF_TERM * vals = (ERL_NIF_TERM *)enif_alloc(sizeof(ERL_NIF_TERM) * num_items);
if (vals == nullptr) {
enif_free(keys);
return erlang::nif::error(env, "out of memory");
}
size_t sig_key_index = 0;
ERL_NIF_TERM signature_def_keys[2];
signature_def_keys[0] = erlang::nif::atom(env, "inputs");
signature_def_keys[1] = erlang::nif::atom(env, "outputs");
for (const auto& sig_key : interpreter_->signature_keys()) {
ERL_NIF_TERM signature_def_vals[2];
const auto& signature_def_inputs = interpreter_->signature_inputs(sig_key->c_str());
const auto& signature_def_outputs = interpreter_->signature_outputs(sig_key->c_str());
size_t inputs_items = signature_def_inputs.size();
ERL_NIF_TERM * inputs_keys = (ERL_NIF_TERM *)enif_alloc(sizeof(ERL_NIF_TERM) * inputs_items);
if (inputs_keys == nullptr) {
enif_free(keys);
enif_free(vals);
return erlang::nif::error(env, "out of memory");
}
ERL_NIF_TERM * inputs_vals = (ERL_NIF_TERM *)enif_alloc(sizeof(ERL_NIF_TERM) * inputs_items);
if (inputs_keys == nullptr) {
enif_free(keys);
enif_free(vals);
enif_free(inputs_keys);
return erlang::nif::error(env, "out of memory");
}
size_t input_item_index = 0;
for (const auto& input : signature_def_inputs) {
if (input.first.length() > 0) {
inputs_keys[input_item_index] = erlang::nif::atom(env, input.first.c_str());
inputs_vals[input_item_index] = erlang::nif::make(env, (long)input.second);
input_item_index++;
}
}
if (!enif_make_map_from_arrays(env, inputs_keys, inputs_vals, input_item_index, &signature_def_vals[0])) {
enif_free(keys);
enif_free(vals);
enif_free(inputs_keys);
enif_free(inputs_vals);
return erlang::nif::error(env, "duplicate keys found in signature_def_inputs");
}
enif_free(inputs_keys);
enif_free(inputs_vals);
size_t outputs_items = signature_def_inputs.size();
ERL_NIF_TERM * outputs_keys = (ERL_NIF_TERM *)enif_alloc(sizeof(ERL_NIF_TERM) * outputs_items);
if (outputs_keys == nullptr) {
enif_free(keys);
enif_free(vals);
return erlang::nif::error(env, "out of memory");
}
ERL_NIF_TERM * outputs_vals = (ERL_NIF_TERM *)enif_alloc(sizeof(ERL_NIF_TERM) * outputs_items);
if (outputs_keys == nullptr) {
enif_free(keys);
enif_free(vals);
enif_free(outputs_keys);
return erlang::nif::error(env, "out of memory");
}
size_t output_item_index = 0;
for (const auto& output : signature_def_outputs) {
if (output.first.length()) {
outputs_keys[output_item_index] = erlang::nif::atom(env, output.first.c_str());
outputs_vals[output_item_index] = erlang::nif::make(env, (long)output.second);
output_item_index++;
}
}
if (!enif_make_map_from_arrays(env, outputs_keys, outputs_vals, output_item_index, &signature_def_vals[1])) {
enif_free(keys);
enif_free(vals);
enif_free(outputs_keys);
enif_free(outputs_vals);
return erlang::nif::error(env, "duplicate keys found in signature_def_outputs");
}
enif_free(outputs_keys);
enif_free(outputs_vals);
keys[sig_key_index] = erlang::nif::atom(env, sig_key->c_str());
enif_make_map_from_arrays(env, signature_def_keys, signature_def_vals, 2, &vals[sig_key_index]);
sig_key_index++;
}
enif_make_map_from_arrays(env, keys, vals, num_items, &result);
enif_free(keys);
enif_free(vals);
return erlang::nif::ok(env, result);
}
ERL_NIF_TERM interpreter_invoke(ErlNifEnv *env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 1) return enif_make_badarg(env);
ERL_NIF_TERM self_nif = argv[0];
NifResInterpreter *self_res;
ERL_NIF_TERM ret;
if (!(self_res = NifResInterpreter::get_resource(env, self_nif, ret))) {
return ret;
}
return tflite_status_to_erl_term(env, self_res->val->Invoke());
}
ERL_NIF_TERM interpreter_set_num_threads(ErlNifEnv *env, int argc, const ERL_NIF_TERM argv[]) {
if (argc != 2) return enif_make_badarg(env);
ERL_NIF_TERM self_nif = argv[0];
ERL_NIF_TERM num_threads_nif = argv[1];
int num_threads = 1;
NifResInterpreter * self_res;
ERL_NIF_TERM ret;
if (!(self_res = NifResInterpreter::get_resource(env, self_nif, ret))) {
return ret;
}
if (!enif_get_int(env, num_threads_nif, &num_threads) || num_threads < 1) {
return erlang::nif::error(env, "expecting num_threads to be an positive integer");
}
auto status = self_res->val->SetNumThreads(num_threads);
return tflite_status_to_erl_term(env, status);
}