@@ -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). */
164177static 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
227240static bool g_tflm_initialized ;
228241static int g_tflm_instance_count ;
229242static 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