8#include "Chirale_TensorFlowLite.h"
14#include "tensorflow/lite/c/common.h"
15#include "tensorflow/lite/experimental/microfrontend/lib/frontend.h"
16#include "tensorflow/lite/experimental/microfrontend/lib/frontend_util.h"
17#include "tensorflow/lite/micro/all_ops_resolver.h"
18#include "tensorflow/lite/micro/kernels/micro_ops.h"
19#include "tensorflow/lite/micro/micro_interpreter.h"
20#include "tensorflow/lite/micro/micro_mutable_op_resolver.h"
21#include "tensorflow/lite/micro/system_setup.h"
22#include "tensorflow/lite/schema/schema_generated.h"
34class TfLiteAudioStreamBase;
35class TfLiteAbstractRecognizeCommands;
46 virtual int read(int16_t*data,
int len) = 0;
58 virtual bool write(
const int16_t sample) = 0;
70 const unsigned char*
model =
nullptr;
77 bool is_new_command) =
nullptr;
136 return kCategoryCount;
152 int kCategoryCount = 0;
153 const char** labels =
nullptr;
165 static int8_t
quantize(
float value,
float scale,
float zero_point){
166 if(scale==0.0&&zero_point==0)
return value;
167 return value / scale + zero_point;
170 static float dequantize(int8_t value,
float scale,
float zero_point){
171 if(scale==0.0&&zero_point==0)
return value;
172 return (value - zero_point) * scale;
176 float deq = (
static_cast<float>(value) - zero_point) * scale;
177 return clip(deq * new_range, new_range);
180 static float clip(
float value,
float range){
182 return value > range ? range : value;
184 return -value < -range ? -range : value;
198 virtual TfLiteStatus
getCommand(
const TfLiteTensor* latest_results,
const int32_t current_time_ms,
199 const char** found_command,uint8_t* score,
bool* is_new_command) = 0;
226 if (
cfg.labels ==
nullptr) {
227 LOGE(
"config.labels not defined");
234 virtual TfLiteStatus
getCommand(
const TfLiteTensor* latest_results,
236 const char** found_command,
238 bool* is_new_command)
override {
249 TfLiteStatus result =
validate(latest_results);
250 if (result!=kTfLiteOk){
253 return evaluate(found_command, score, is_new_command);
280 uint8_t top_score = std::numeric_limits<uint8_t>::min();
282 if (score[j]>top_score){
303 TfLiteStatus
evaluate(
const char** found_command, uint8_t* result_score,
bool* is_new_command) {
325 LOGE(
"Could not find max category")
330 *result_score = totals[maxIdx] / count[maxIdx];
331 *found_command =
cfg.labels[maxIdx];
338 *is_new_command =
true;
340 *is_new_command =
false;
343 LOGD(
"Category: %s, score: %d, is_new: %d",*found_command, *result_score, *is_new_command);
349 TfLiteStatus
validate(
const TfLiteTensor* latest_results) {
350 if ((latest_results->dims->size != 2) ||
351 (latest_results->dims->data[0] != 1) ||
354 "The results for recognition should contain %d "
355 "elements, but there are "
356 "%d in an %d-dimensional shape",
358 (
int)latest_results->dims->size);
362 if (latest_results->type != kTfLiteInt8) {
363 LOGE(
"The results for recognition should be int8 elements, but are %d",
364 (
int)latest_results->type);
370 LOGE(
"Results must be in increasing time order: timestamp %d < %d",
395 virtual size_t write(
const uint8_t* data,
size_t len)= 0;
431 LOGE(
"setup_recognizer");
437 if (init_status != kTfLiteOk) {
462 virtual bool write(int16_t sample) {
470 int8_t* feature_buffer =
addSlice();
506 virtual bool write1(
const int16_t sample) {
540 int audio_samples_size =
545 LOGE(
"audio_samples_size=%d != kMaxAudioSampleSize=%d",
553 int8_t* new_slice_data =
555 size_t num_samples_read = 0;
558 &num_samples_read) != kTfLiteOk) {
559 LOGE(
"Error generateMicroFeatures");
573 if (invoke_status != kTfLiteOk) {
574 LOGE(
"Invoke failed");
582 const char* found_command =
nullptr;
584 bool is_new_command =
false;
587 output,
current_time, &found_command, &score, &is_new_command);
588 if (process_status != kTfLiteOk) {
589 LOGE(
"TfLiteMicroSpeechRecognizeCommands::getCommand() failed");
608 Serial.println(
"------------");
630 LOGE(
"frontendPopulateState() failed");
637 int input_size, int8_t* output,
639 size_t* num_samples_read) {
641 const int16_t* frontend_input = input;
644 FrontendOutput frontend_output = FrontendProcessSamples(
648 if (output_size != frontend_output.size) {
649 LOGE(
"output_size=%d, frontend_output.size=%d", output_size,
650 frontend_output.size);
663 for (
size_t i = 0; i < frontend_output.size; ++i) {
677 constexpr int32_t value_scale = 256;
678 constexpr int32_t value_div =
679 static_cast<int32_t
>((25.6f * 26.0f) + 0.5f);
681 ((frontend_output.values[i] * value_scale) + (value_div / 2)) /
698 bool is_new_command) {
703 if (is_new_command) {
705 snprintf(buffer, 80,
"Result: %s, score: %d, is_new: %s", found_command,
706 score, is_new_command ?
"true" :
"false");
735 virtual int read(int16_t*data,
int sampleCount)
override {
737 float two_pi = 2 *
PI;
738 for (
int j=0; j<sampleCount; j+=
channels){
746 if(kTfLiteOk!= invoke_status){
747 LOGE(
"invoke_status not ok");
750 if(kTfLiteInt8 !=
output->type){
751 LOGE(
"Output type is not kTfLiteInt8");
761 LOGD(
"generate data for channels");
838 LOGI(
"AllocateTensors");
839 TfLiteStatus allocate_status =
p_interpreter->AllocateTensors();
840 if (allocate_status != kTfLiteOk) {
841 LOGE(
"AllocateTensors() failed");
853 LOGE(
"Bad input tensor parameters in model");
861 LOGE(
"p_tensor_buffer is null");
880 virtual size_t write(
const uint8_t* data,
size_t len)
override {
883 LOGE(
"cfg.output is null");
886 int16_t* samples = (int16_t*)data;
887 int16_t sample_count = len / 2;
888 for (
int j = 0; j < sample_count; j++) {
898 virtual size_t readBytes(uint8_t *data,
size_t len)
override {
901 return cfg.
reader->
read((int16_t*)data, (
int) len/
sizeof(int16_t)) *
sizeof(int16_t);
934 virtual bool setModel(
const unsigned char* model) {
936 p_model = tflite::GetModel(model);
937 if (
p_model->version() != TFLITE_SCHEMA_VERSION) {
939 "Model provided is schema version %d not equal "
940 "to supported version %d.",
941 p_model->version(), TFLITE_SCHEMA_VERSION);
965 tflite::AllOpsResolver resolver;
966 static tflite::MicroInterpreter static_interpreter{
971 static tflite::MicroMutableOpResolver<4> micro_op_resolver{};
972 if (micro_op_resolver.AddDepthwiseConv2D() != kTfLiteOk) {
975 if (micro_op_resolver.AddFullyConnected() != kTfLiteOk) {
978 if (micro_op_resolver.AddSoftmax() != kTfLiteOk) {
981 if (micro_op_resolver.AddReshape() != kTfLiteOk) {
985 static tflite::MicroInterpreter static_interpreter{
static HardwareSerial Serial
Definition Arduino.h:179
#define PI
Definition AudioEffectsSuite.h:27
#define LOGW(...)
Definition AudioLoggerIDF.h:29
#define TRACEI()
Definition AudioLoggerIDF.h:32
#define TRACED()
Definition AudioLoggerIDF.h:31
#define LOGI(...)
Definition AudioLoggerIDF.h:28
#define LOGD(...)
Definition AudioLoggerIDF.h:27
#define LOGE(...)
Definition AudioLoggerIDF.h:30
#define DEFAULT_BUFFER_SIZE
Definition avr.h:20