Skip to content

Commit f0fd1b1

Browse files
committed
audio: tensorflow: tune: export model blobs for topology and sof-ctl
Update sof_tflm_train.py to package the trained tflite model behind the SOF IPC4 ABI header (struct sof_abi_hdr) and export it to: - tools/topology/topology2/include/components/tflm/*.conf for embedding into ALSA topology v2 files. - tools/ctl/ipc4/tflm/*.txt for runtime application with sof-ctl. Signed-off-by: Seppo Ingalsuo <seppo.ingalsuo@linux.intel.com>
1 parent 4f35a49 commit f0fd1b1

4 files changed

Lines changed: 892 additions & 4 deletions

File tree

‎src/audio/tensorflow/tune/sof_tflm_train.py‎

Lines changed: 119 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -30,7 +30,9 @@
3030
from __future__ import annotations
3131

3232
import argparse
33+
import datetime
3334
import os
35+
import struct
3436
import sys
3537
from 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+
134154
def 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('\twords "\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+
186269
def 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

231314
def 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

‎tools/ctl/ipc4/tflm/banana_strawberry_orange.txt‎

Lines changed: 1 addition & 0 deletions
Large diffs are not rendered by default.

‎tools/topology/topology2/include/common/data.conf‎

Lines changed: 8 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -19,6 +19,14 @@ Class.Base."data" {
1919
type "string"
2020
}
2121

22+
DefineAttribute."shorts" {
23+
type "string"
24+
}
25+
26+
DefineAttribute."words" {
27+
type "string"
28+
}
29+
2230
attributes {
2331
!constructor [
2432
"name"

0 commit comments

Comments
 (0)