3030from __future__ import annotations
3131
3232import argparse
33+ import datetime
3334import os
35+ import struct
3436import sys
3537from pathlib import Path
3638
@@ -131,6 +133,24 @@ def convert_to_tflite_int8(
131133 return converter .convert ()
132134
133135
136+ def set_tflite_model_description (tflite_bytes : bytes , description : str ) -> bytes :
137+ """Embed comma-separated labels into TFLite FlatBuffer Model.description."""
138+ try :
139+ from tensorflow .lite .python import schema_py_generated as schema_fb
140+ import flatbuffers
141+
142+ model_obj = schema_fb .ModelT .InitFromObj (
143+ schema_fb .Model .GetRootAsModel (tflite_bytes , 0 )
144+ )
145+ model_obj .description = description
146+ builder = flatbuffers .Builder (len (tflite_bytes ) + len (description ) + 256 )
147+ builder .Finish (model_obj .Pack (builder ), "TFL3" )
148+ return bytes (builder .Output ())
149+ except Exception as exc :
150+ print (f" [WARN] could not embed description in tflite FlatBuffer: { exc } " )
151+ return tflite_bytes
152+
153+
134154def emit_c_array (
135155 tflite_bytes : bytes ,
136156 out_cc : Path ,
@@ -183,6 +203,69 @@ def emit_labels_header(labels: list[str], out_h: Path) -> None:
183203 f .write (f"#endif // { guard } \n " )
184204
185205
206+ # SOF IPC4 ABI definitions for binary control blobs
207+ SOF_IPC4_ABI_MAGIC = 0x34464F53 # 'SOF4' in little-endian
208+ SOF_ABI_VERSION = (3 << 24 ) | (29 << 12 ) | 1 # 3.29.1 (0x0301d001)
209+ SOF_CTRL_CMD_BINARY = 3
210+
211+
212+ def build_ipc4_abi_blob (tflite_bytes : bytes , param_id : int = 0 ) -> bytes :
213+ """Pack tflite model bytes behind a standard 32-byte struct sof_abi_hdr."""
214+ size = len (tflite_bytes )
215+ abi_header = struct .pack ("<IIIIIIII" , SOF_IPC4_ABI_MAGIC , param_id , size , SOF_ABI_VERSION , 0 , 0 , 0 , 0 )
216+ pad = (4 - (size % 4 )) % 4
217+ return abi_header + tflite_bytes + b"\x00 " * pad
218+
219+
220+ def emit_topology2_conf (
221+ tflite_bytes : bytes ,
222+ out_conf : Path ,
223+ model_name : str ,
224+ labels : list [str ],
225+ ) -> None :
226+ """Emit Topology2 configuration blob in .conf format for inclusion to topology."""
227+ blob8 = build_ipc4_abi_blob (tflite_bytes )
228+ words = struct .unpack (f"<{ len (blob8 ) // 4 } I" , blob8 )
229+ out_conf .parent .mkdir (parents = True , exist_ok = True )
230+ today = datetime .date .today ().strftime ("%d-%b-%Y" )
231+ labels_str = ", " .join (labels )
232+ data_name = f"tflm_config_{ model_name } "
233+ with open (out_conf , "w" ) as f :
234+ f .write (f"# Exported TFLM Model Control Words { today } \n " )
235+ f .write (f"# Model: { model_name } (classes: { labels_str } )\n " )
236+ f .write (f"# Auto-generated by sof_tflm_train.py — do not edit.\n " )
237+ f .write (f'Object.Base.data."{ data_name } " {{\n ' )
238+ f .write ('\t words "\n ' )
239+ lines = []
240+ for i in range (0 , len (words ), 8 ):
241+ chunk = words [i : i + 8 ]
242+ lines .append ("," .join (f"0x{ w :08x} " for w in chunk ))
243+ for idx , line in enumerate (lines ):
244+ if idx == len (lines ) - 1 :
245+ f .write (f"\t \t { line } \" \n " )
246+ else :
247+ f .write (f"\t \t { line } ,\n " )
248+ f .write ("}\n " )
249+
250+
251+ def emit_sofctl_ipc4_txt (
252+ tflite_bytes : bytes ,
253+ out_txt : Path ,
254+ ) -> None :
255+ """Emit sof-ctl IPC4 text blob in comma-separated uint32 CSV format."""
256+ size = len (tflite_bytes )
257+ tlv_cmd = SOF_CTRL_CMD_BINARY
258+ tlv_size = 32 + size
259+ pad = (4 - (size % 4 )) % 4
260+ payload_padded = tflite_bytes + b"\x00 " * pad
261+ words = [tlv_cmd , tlv_size , SOF_IPC4_ABI_MAGIC , 0 , size , SOF_ABI_VERSION , 0 , 0 , 0 , 0 ]
262+ payload_words = list (struct .unpack (f"<{ len (payload_padded ) // 4 } I" , payload_padded ))
263+ all_words = words + payload_words
264+ out_txt .parent .mkdir (parents = True , exist_ok = True )
265+ with open (out_txt , "w" ) as f :
266+ f .write ("," .join (str (w ) for w in all_words ) + "\n " )
267+
268+
186269def parse_args () -> argparse .Namespace :
187270 ap = argparse .ArgumentParser (description = __doc__ )
188271 ap .add_argument ("--feat-root" , required = True )
@@ -230,11 +313,12 @@ def parse_args() -> argparse.Namespace:
230313
231314def main () -> int :
232315 args = parse_args ()
316+ reserved = {"silence" , "unknown" }
317+ keywords = [lbl for lbl in args .labels if lbl not in reserved ]
318+
233319 if not args .name :
234- # Derive a sensible default: first keyword label past silence/unknown.
235- reserved = {"silence" , "unknown" }
236- keywords = [lbl for lbl in args .labels if lbl not in reserved ]
237- args .name = keywords [0 ] if keywords else "wov_model"
320+ # Default name contains all trained keywords joined with underscore
321+ args .name = "_" .join (keywords ) if keywords else "wov_model"
238322 print (f">>> --name not set, using { args .name !r} " )
239323 tf .keras .utils .set_random_seed (args .seed )
240324
@@ -284,6 +368,7 @@ def main() -> int:
284368
285369 print (">>> Converting to int8 tflite" )
286370 tflite_bytes = convert_to_tflite_int8 (model , X_tr , y_tr , args .rep_samples )
371+ tflite_bytes = set_tflite_model_description (tflite_bytes , "," .join (args .labels ))
287372
288373 out_dir = Path (args .out_dir )
289374 out_dir .mkdir (parents = True , exist_ok = True )
@@ -309,6 +394,36 @@ def main() -> int:
309394 labels_path .write_text ("\n " .join (args .labels ) + "\n " )
310395 print (f" wrote { labels_path } (archive copy of the label list)" )
311396
397+ # Export Topology2 config blob (.conf) and sof-ctl IPC4 text blob (.txt)
398+ out_conf = out_dir / f"{ args .name } .conf"
399+ emit_topology2_conf (tflite_bytes , out_conf , args .name , args .labels )
400+ print (f" wrote { out_conf } " )
401+
402+ out_txt = out_dir / f"{ args .name } .txt"
403+ emit_sofctl_ipc4_txt (tflite_bytes , out_txt )
404+ print (f" wrote { out_txt } " )
405+
406+ # Also export directly to SOF tools directories if present in workspace/repo
407+ script_dir = Path (__file__ ).resolve ().parent
408+ repo_root = None
409+ for parent in script_dir .parents :
410+ if (parent / "tools/topology/topology2" ).is_dir ():
411+ repo_root = parent
412+ break
413+ if repo_root is None and "SOF_WORKSPACE" in os .environ :
414+ candidate = Path (os .environ ["SOF_WORKSPACE" ]) / "sof"
415+ if (candidate / "tools/topology/topology2" ).is_dir ():
416+ repo_root = candidate
417+
418+ if repo_root is not None :
419+ tplg_tflm_dir = repo_root / "tools/topology/topology2/include/components/tflm"
420+ emit_topology2_conf (tflite_bytes , tplg_tflm_dir / f"{ args .name } .conf" , args .name , args .labels )
421+ print (f" exported { tplg_tflm_dir / f'{ args .name } .conf' } " )
422+
423+ ctl_tflm_dir = repo_root / "tools/ctl/ipc4/tflm"
424+ emit_sofctl_ipc4_txt (tflite_bytes , ctl_tflm_dir / f"{ args .name } .txt" )
425+ print (f" exported { ctl_tflm_dir / f'{ args .name } .txt' } " )
426+
312427 return 0
313428
314429
0 commit comments