Skip to content

Commit 4f35a49

Browse files
committed
audio: tensorflow: support loading model from bytes control
Add Kconfig option COMP_TENSORFLOW_MODEL_FROM_CONTROL to allow loading the TFLM wake-word model flatbuffer at prepare time from a configuration data blob via binary kcontrol (bytes control) instead of using the static C array compiled into the firmware binary. When this option is enabled, the firmware omits linking sof_tflm_quantized_model_data.cc and retrieves the model from the registered blob handler. Signed-off-by: Seppo Ingalsuo <seppo.ingalsuo@linux.intel.com>
1 parent 1aee10e commit 4f35a49

5 files changed

Lines changed: 153 additions & 22 deletions

File tree

‎src/audio/tensorflow/CMakeLists.txt‎

Lines changed: 6 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -306,10 +306,15 @@ add_library(tflm_lib STATIC
306306
${TFLM_PATH}/tensorflow/lite/micro/memory_planner/linear_memory_planner.cc
307307
${TFLM_PATH}/tensorflow/lite/micro/micro_utils.cc
308308
${TFLM_PATH}/tensorflow/lite/micro/micro_interpreter.cc
309-
sof_tflm_quantized_model_data.cc
310309
speech.cc
311310
)
312311

312+
if(CONFIG_COMP_TENSORFLOW_MODEL_FROM_CONTROL)
313+
target_compile_definitions(tflm_lib PRIVATE CONFIG_COMP_TENSORFLOW_MODEL_FROM_CONTROL=1)
314+
else()
315+
target_sources(tflm_lib PRIVATE sof_tflm_quantized_model_data.cc)
316+
endif()
317+
313318
target_include_directories(tflm_lib PRIVATE
314319
${TFLM_PATH}
315320
${FLATBUFFERS_PATH}/include

‎src/audio/tensorflow/Kconfig‎

Lines changed: 9 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -31,4 +31,13 @@ config COMP_TENSORFLOW_DEBUG_TRACE
3131
or per-node timing. The keyword-detected line and the shutdown
3232
summary remain enabled regardless of this option.
3333

34+
config COMP_TENSORFLOW_MODEL_FROM_CONTROL
35+
bool "Load TFLM model from bytes control instead of built-in header"
36+
default n
37+
help
38+
When enabled, the TFLM wake-word model is loaded from a runtime
39+
or topology configuration data blob via binary kcontrol (bytes control)
40+
instead of using the static g_sof_tflm_quantized_model_data embedded
41+
in the firmware image.
42+
3443
endif # COMP_TENSORFLOW

‎src/audio/tensorflow/speech.cc‎

Lines changed: 57 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,9 @@
1414
#include "tensorflow/lite/micro/testing/micro_test.h"
1515
#include "speech.h"
1616

17+
#if !CONFIG_COMP_TENSORFLOW_MODEL_FROM_CONTROL
1718
#include "sof_tflm_quantized_model_data.h"
19+
#endif
1820

1921
// The following values are derived from values used during model training.
2022
// If you change the way you preprocess the input, update all these constants.
@@ -122,27 +124,75 @@ static int Init_Interpreter(struct tf_classify *tfc)
122124
return -EINVAL;
123125
}
124126

125-
// Accept any output rank as long as total element count == categories.
127+
// Accept any output rank as long as 1 <= total element count <= TFLM_MAX_CATEGORY_COUNT.
126128
int output_elems = 1;
127129
for (int i = 0; i < output->dims->size; i++)
128130
output_elems *= output->dims->data[i];
129-
if (tfc->categories != output_elems) {
130-
tfc->error = "output shape != categories";
131-
MicroPrintf("TFLM: output rank=%d elems=%d categories=%d",
132-
output->dims->size, output_elems, tfc->categories);
131+
if (output_elems <= 0 || output_elems > TFLM_MAX_CATEGORY_COUNT) {
132+
tfc->error = "output shape out of range";
133+
MicroPrintf("TFLM: output rank=%d elems=%d max=%d",
134+
output->dims->size, output_elems, TFLM_MAX_CATEGORY_COUNT);
133135
return -EINVAL;
134136
}
135137

138+
tfc->categories = output_elems;
139+
140+
// Populate category names: first check if model description has comma-separated labels
141+
bool have_labels = false;
142+
if (model->description() && model->description()->c_str()) {
143+
const char *desc = model->description()->c_str();
144+
if (desc[0] != '\0') {
145+
int cat_idx = 0;
146+
int char_idx = 0;
147+
for (int i = 0; desc[i] != '\0' && cat_idx < tfc->categories; i++) {
148+
if (desc[i] == ',') {
149+
tfc->category_names[cat_idx][char_idx] = '\0';
150+
cat_idx++;
151+
char_idx = 0;
152+
} else if (char_idx < TFLM_MAX_LABEL_LEN - 1) {
153+
tfc->category_names[cat_idx][char_idx++] = desc[i];
154+
}
155+
}
156+
if (cat_idx < tfc->categories) {
157+
tfc->category_names[cat_idx][char_idx] = '\0';
158+
cat_idx++;
159+
}
160+
if (cat_idx == tfc->categories)
161+
have_labels = true;
162+
}
163+
}
164+
165+
// Fallback to static labels if description was absent or didn't match category count
166+
if (!have_labels) {
167+
static const char * const default_labels[] = TFLM_CATEGORY_DATA;
168+
for (int i = 0; i < tfc->categories; i++) {
169+
if (i < (int)(sizeof(default_labels) / sizeof(default_labels[0]))) {
170+
strncpy(tfc->category_names[i], default_labels[i], TFLM_MAX_LABEL_LEN - 1);
171+
} else {
172+
snprintf(tfc->category_names[i], TFLM_MAX_LABEL_LEN, "class_%d", i);
173+
}
174+
tfc->category_names[i][TFLM_MAX_LABEL_LEN - 1] = '\0';
175+
}
176+
}
177+
136178
return 0;
137179
}
138180

139181
int TF_SetModel(struct tf_classify *tfc, unsigned char *model_tflite)
140182
{
141-
// ignore passed in model today until we can load via binary kcontrol
183+
#if !CONFIG_COMP_TENSORFLOW_MODEL_FROM_CONTROL
184+
if (!model_tflite)
185+
model_tflite = const_cast<unsigned char *>(g_sof_tflm_quantized_model_data);
186+
#endif
187+
188+
if (!model_tflite) {
189+
tfc->error = "no model provided";
190+
return -EINVAL;
191+
}
142192

143193
// Map the model into a usable data structure. This doesn't involve any
144194
// copying or parsing, it's a very lightweight operation.
145-
model = tflite::GetModel(g_sof_tflm_quantized_model_data);
195+
model = tflite::GetModel(model_tflite);
146196
if (model->version() != TFLITE_SCHEMA_VERSION) {
147197
tfc->error = "failed to load model";
148198
return -EINVAL;

‎src/audio/tensorflow/speech.h‎

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -20,14 +20,17 @@
2020
#define TFLM_FEATURE_ELEM_COUNT (TFLM_FEATURE_SIZE * TFLM_FEATURE_COUNT)
2121
#define TFLM_FEATURE_STRIDE_MS 20
2222
#define TFLM_FEATURE_DURATION_MS 30
23+
#define TFLM_MAX_CATEGORY_COUNT 16
24+
#define TFLM_MAX_LABEL_LEN 16
2325

2426
struct tf_classify {
2527
int8_t *audio_features;
2628
size_t audio_data_size;
2729
int categories;
30+
char category_names[TFLM_MAX_CATEGORY_COUNT][TFLM_MAX_LABEL_LEN];
2831
const char *error;
29-
float predictions[TFLM_CATEGORY_COUNT];
30-
int8_t raw_output[TFLM_CATEGORY_COUNT];
32+
float predictions[TFLM_MAX_CATEGORY_COUNT];
33+
int8_t raw_output[TFLM_MAX_CATEGORY_COUNT];
3134
int op_count;
3235
uint32_t node_cycles[10];
3336
int node_codes[10];

‎src/audio/tensorflow/tflm-classify.c‎

Lines changed: 76 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -154,11 +154,24 @@ struct tflm_comp_data {
154154
/* Persistent AGC gain applied to every Q9.23 mel value (see AGC defines). */
155155
int32_t agc_gain_q23;
156156
/* Per-instance shutdown-summary counters. */
157-
uint32_t category_totals[TFLM_CATEGORY_COUNT];
157+
uint32_t category_totals[TFLM_MAX_CATEGORY_COUNT];
158158
uint32_t total_inferences;
159159
uint32_t kpb_trigger_events;
160160
} __attribute__((aligned(8)));
161161

162+
static const char *tflm_category_name(const struct tflm_comp_data *cd, int idx)
163+
{
164+
if (cd && idx >= 0 && idx < cd->tfc.categories && cd->tfc.category_names[idx][0] != '\0')
165+
return cd->tfc.category_names[idx];
166+
if (idx >= 0 && idx < (int)ARRAY_SIZE(prediction))
167+
return prediction[idx];
168+
if (idx == 0)
169+
return "silence";
170+
if (idx == 1)
171+
return "unknown";
172+
return "keyword";
173+
}
174+
162175
#if CONFIG_AMS
163176
/* Key-phrase detected AMS message UUID (matches KPB consumer). */
164177
static const ams_uuid_t tflm_ams_kpd_msg_uuid = AMS_KPD_MSG_UUID;
@@ -227,25 +240,30 @@ static __maybe_unused void tflm_send_keyword_notification(struct processing_modu
227240
static bool g_tflm_initialized;
228241
static int g_tflm_instance_count;
229242
static int g_tflm_shared_categories;
243+
static char g_tflm_shared_category_names[TFLM_MAX_CATEGORY_COUNT][TFLM_MAX_LABEL_LEN];
230244

231245
__cold static void tflm_log_summary_at_shutdown(struct processing_module *mod)
232246
{
233247
struct tflm_comp_data *cd = mod ? module_get_private_data(mod) : NULL;
234248
struct comp_dev *dev = mod ? mod->dev : NULL;
235249
char summary_buf[256];
236250
size_t off;
251+
int num_cats = cd ? cd->tfc.categories : g_tflm_shared_categories;
237252
int i;
238253

239254
if (!cd)
240255
return;
241256

257+
if (num_cats <= 0 || num_cats > TFLM_MAX_CATEGORY_COUNT)
258+
num_cats = TFLM_CATEGORY_COUNT;
259+
242260
off = snprintk(summary_buf, sizeof(summary_buf),
243261
"[TFLM STREAM SHUTDOWN SUMMARY] Total Inferences=%u | Keyword Events:",
244262
cd->total_inferences);
245-
for (i = 0; i < TFLM_CATEGORY_COUNT && off < sizeof(summary_buf); i++) {
263+
for (i = 0; i < num_cats && off < sizeof(summary_buf); i++) {
246264
off += snprintk(summary_buf + off, sizeof(summary_buf) - off,
247265
"%s %s=%u", i ? "," : "",
248-
prediction[i], cd->category_totals[i]);
266+
tflm_category_name(cd, i), cd->category_totals[i]);
249267
}
250268
if (off < sizeof(summary_buf))
251269
snprintk(summary_buf + off, sizeof(summary_buf) - off,
@@ -276,6 +294,13 @@ __cold static int tflm_init(struct processing_module *mod)
276294
}
277295

278296
md->private = cd;
297+
cd->model_handler = mod_data_blob_handler_new(mod);
298+
if (!cd->model_handler) {
299+
printk("[TFLM INIT] FAILED: mod_data_blob_handler_new failed\n");
300+
rfree(cd);
301+
return -ENOMEM;
302+
}
303+
279304
cd->tfc.categories = TFLM_CATEGORY_COUNT;
280305
cd->drain_req_ms = TFLM_KPB_DRAIN_REQ_MS;
281306
cd->agc_gain_q23 = TFLM_AGC_GAIN_TARGET_Q23;
@@ -299,6 +324,7 @@ __cold static int tflm_free(struct processing_module *mod)
299324
g_tflm_instance_count = 0;
300325
g_tflm_initialized = false;
301326
}
327+
mod_data_blob_handler_free(mod, cd->model_handler);
302328
rfree(cd);
303329
return 0;
304330
}
@@ -308,7 +334,26 @@ __cold static int tflm_set_config(struct processing_module *mod, uint32_t param_
308334
const uint8_t *fragment, size_t fragment_size, uint8_t *response,
309335
size_t response_size)
310336
{
311-
return 0;
337+
struct tflm_comp_data *cd = module_get_private_data(mod);
338+
339+
if (mod->dev->state != COMP_STATE_INIT && mod->dev->state != COMP_STATE_READY) {
340+
comp_warn(mod->dev, "tflm_set_config(): model update ignored while not idle (state %d)",
341+
mod->dev->state);
342+
return 0;
343+
}
344+
345+
return comp_data_blob_set(cd->model_handler, pos, data_offset_size,
346+
fragment, fragment_size);
347+
}
348+
349+
__cold static int tflm_get_config(struct processing_module *mod,
350+
uint32_t config_id, uint32_t *data_offset_size,
351+
uint8_t *fragment, size_t fragment_size)
352+
{
353+
struct sof_ipc_ctrl_data *cdata = (struct sof_ipc_ctrl_data *)fragment;
354+
struct tflm_comp_data *cd = module_get_private_data(mod);
355+
356+
return comp_data_blob_get_cmd(cd->model_handler, cdata, fragment_size);
312357
}
313358

314359
/*
@@ -579,12 +624,12 @@ static int tflm_process(struct processing_module *mod,
579624
char raw_buf[160];
580625
int off = snprintk(raw_buf, sizeof(raw_buf),
581626
"[DBG raw_output] ret=%d", ret);
582-
for (int i = 0; i < TFLM_CATEGORY_COUNT &&
627+
for (int i = 0; i < cd->tfc.categories &&
583628
off < (int)sizeof(raw_buf); i++)
584629
off += snprintk(raw_buf + off,
585630
sizeof(raw_buf) - off,
586631
" %s=%d",
587-
prediction[i],
632+
tflm_category_name(cd, i),
588633
cd->tfc.raw_output[i]);
589634
sof_ut_log(raw_buf);
590635
}
@@ -608,7 +653,7 @@ static int tflm_process(struct processing_module *mod,
608653
}
609654

610655
cd->total_inferences++;
611-
if (max_idx >= 0 && max_idx < TFLM_CATEGORY_COUNT)
656+
if (max_idx >= 0 && max_idx < TFLM_MAX_CATEGORY_COUNT)
612657
cd->category_totals[max_idx]++;
613658

614659
#if CONFIG_COMP_TENSORFLOW_DEBUG_TRACE
@@ -618,14 +663,14 @@ static int tflm_process(struct processing_module *mod,
618663
if (max_pct_dbg < 0) max_pct_dbg = 0;
619664
int off = snprintk(result_buf, sizeof(result_buf),
620665
"TFLM top prediction: %s confidence=%d pct (inferences=%u):",
621-
prediction[max_idx], max_pct_dbg,
666+
tflm_category_name(cd, max_idx), max_pct_dbg,
622667
cd->total_inferences);
623-
for (int i = 0; i < TFLM_CATEGORY_COUNT &&
668+
for (int i = 0; i < cd->tfc.categories &&
624669
off < (int)sizeof(result_buf); i++)
625670
off += snprintk(result_buf + off,
626671
sizeof(result_buf) - off,
627672
" %s=%u",
628-
prediction[i],
673+
tflm_category_name(cd, i),
629674
cd->category_totals[i]);
630675
sof_ut_log(result_buf);
631676
}
@@ -640,7 +685,7 @@ static int tflm_process(struct processing_module *mod,
640685
if (max_pct < 0) max_pct = 0;
641686
snprintk(kw_buf, sizeof(kw_buf),
642687
"TFLM KEYWORD DETECTED: %s confidence=%d pct",
643-
prediction[max_idx], max_pct);
688+
tflm_category_name(cd, max_idx), max_pct);
644689
sof_ut_log(kw_buf);
645690

646691
cd->kpb_trigger_events++;
@@ -680,11 +725,27 @@ static int tflm_prepare(struct processing_module *mod,
680725
if (g_tflm_initialized) {
681726
/* Shared TFLM engine already up. */
682727
cd->tfc.categories = g_tflm_shared_categories;
728+
memcpy(cd->tfc.category_names, g_tflm_shared_category_names,
729+
sizeof(g_tflm_shared_category_names));
683730
printk("[TFLM PREPARE] shared engine already initialized; attach instance\n");
684731
goto post_init;
685732
}
686733

687-
int ret = TF_SetModel(&cd->tfc, NULL);
734+
unsigned char *model_ptr = NULL;
735+
736+
#if CONFIG_COMP_TENSORFLOW_MODEL_FROM_CONTROL
737+
size_t blob_size;
738+
739+
model_ptr = comp_get_data_blob(cd->model_handler, &blob_size, NULL);
740+
if (!model_ptr || !blob_size) {
741+
printk("[TFLM PREPARE] FAILED: model blob not set from control\n");
742+
comp_err(mod->dev, "TFLM: model blob not set from control");
743+
return -EINVAL;
744+
}
745+
printk("[TFLM PREPARE] loaded model from control blob, size=%zu\n", blob_size);
746+
#endif
747+
748+
int ret = TF_SetModel(&cd->tfc, model_ptr);
688749
if (ret < 0) {
689750
printk("[TFLM PREPARE] FAILED: TF_SetModel returned %d\n", ret);
690751
return ret;
@@ -698,6 +759,8 @@ static int tflm_prepare(struct processing_module *mod,
698759
}
699760

700761
g_tflm_shared_categories = cd->tfc.categories;
762+
memcpy(g_tflm_shared_category_names, cd->tfc.category_names,
763+
sizeof(g_tflm_shared_category_names));
701764
g_tflm_initialized = true;
702765
printk("[TFLM PREPARE] TFLM model & ops initialized successfully!\n");
703766
printk("[TFLM PREPARE] arena_used=%zu / capacity=%zu bytes\n",
@@ -747,6 +810,7 @@ static const struct module_interface tflmcly_interface = {
747810
.prepare = tflm_prepare,
748811
.process = tflm_process,
749812
.set_configuration = tflm_set_config,
813+
.get_configuration = tflm_get_config,
750814
.reset = tflm_reset,
751815
.free = tflm_free
752816
};

0 commit comments

Comments
 (0)