Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
156 changes: 144 additions & 12 deletions src/audio/tensorflow/README.md
Original file line number Diff line number Diff line change
@@ -1,22 +1,154 @@
# TensorFlow Lite Micro (TFLM) Architecture
# TensorFlow Lite Micro (TFLM) & Wake-on-Voice (WoV) Architecture

This directory acts as the bridge for running ML models.
This directory provides the TensorFlow Lite for Microcontrollers (TFLM) classification module (`TFLMCLY`) for Sound Open Firmware (SOF), including integration with MFCC feature extraction, mtrace logging, IPC host notifications, and Key Phrase Buffer (KPB) Wake-on-Voice (WoV) trigger infrastructure.

---

## Overview

Integrates TensorFlow Lite for Microcontrollers into the SOF audio pipeline. Evaluates pre-trained neural network topologies inline with the audio stream for tasks like wake-word, noise cancellation, or sound classification.
The TFLM module evaluates pre-trained micro speech neural network models inline within the SOF audio processing graph. It receives pre-processed audio feature tensors (e.g. 40-bin mel spectrograms from the MFCC component), runs model inference, logs keyword detections to `mtrace`, issues IPC4 notifications to the host audio driver, and signals the KPB module to drain buffered pre-keyword audio for Wake-on-Voice.

---

## Architecture & Data Flow

## Architecture Diagram
### Dual-Path Wake-on-Voice (WoV) Architecture

To allow continuous keyword evaluation without streaming audio to the host until a keyword is detected, the pipeline separates real-time keyword detection from host PCM draining via KPB:

```mermaid
graph LR
Feat[Audio Features] --> TFLM[TFLM Inference Engine]
SubGraph[FlatBuffer Model] -.-> TFLM
TFLM --> Out[Inference Labels/Scores]
graph TD
DAI[HDA Mic DAI] --> Gain[Gain Component]
Gain --> KPB[KPB Buffer Module]

subgraph "Real-Time Detection Path (KPB Pin 1)"
KPB -- Live Audio Stream --> SRC[SRC: 48kHz -> 16kHz]
SRC --> MFCC[MFCC Feature Extractor]
MFCC -- 40-bin Mel Tensors --> TFLM[TFLM Classifier: tflmcly]
TFLM --> VSink[Virtual Sink: virtual.tflm_sink]
end

subgraph "Host Draining Path (KPB Pin 2)"
KPB -- History Draining Stream --> Host[Host Copier: PCM Capture]
end

TFLM -- "1. Log to mtrace (comp_info)" --> MTrace[mtrace / SOF Trace Log]
TFLM -- "2. IPC4 Host Notification" --> IPC[Host Audio Driver]
TFLM -- "3. KPB_EVENT_BEGIN_DRAINING (notifier_event)" --> KPB
```

---

## Pipeline Execution & Event Flow

1. **Continuous Real-Time Listening**:
- Live microphone audio is captured by the HDA DAI and passed into `KPB` (`kpb.2.1`).
- KPB Output Pin 1 continuously streams audio to `SRC` (resampling 48kHz $\to$ 16kHz), `MFCC` (generating 40-bin `int8_t` mel spectrogram features), and `TFLM` (`tflmcly.1.1`).
- KPB stores the raw PCM audio continuously in its circular history buffer (e.g. 2100ms – 3000ms history).

2. **Inference & Keyword Detection**:
- `tflm_process()` feeds 1960-byte feature tensors ($40 \text{ features} \times 49 \text{ windows}$) into the Micro Speech TFLM interpreter (`[1, 49, 40]` `int8_t` input tensor).
- Upon `TF_ProcessClassify()`, predictions are evaluated across categories:
- `0`: `silence`
- `1`: `unknown`
- `2`: `yes` (Keyword)
- `3`: `no` (Keyword)

3. **Wake-on-Voice Trigger & Notification**:
When a keyword (`yes` or `no`) is detected with $\ge 0.70$ confidence:
- **mtrace Logging**: Logs detection with keyword label and confidence score via `comp_info()`:
`"TFLM keyword detected: yes (confidence 0.852)"`
- **Host IPC Notification**: Sends an IPC4 module notification (`SOF_IPC4_MODULE_NOTIFICATION`) to inform the host driver.
- **KPB Draining Signal**: Fires a system notification event (`NOTIFIER_ID_KPB_CLIENT_EVT`) with `KPB_EVENT_BEGIN_DRAINING`.
- **Host PCM Draining**: KPB opens Output Pin 2 to `host-copier`, draining pre-keyword history buffer audio followed by live mic audio to the host capture stream.

---

## Topology v2 Integration & Usage

### 1. Component Widget (`include/components/tflm.conf`)

Defines `Class.Widget."tflmcly"`:
- **UUID**: `42:c6:1d:c5:e1:a2:df:48:a4:90:e2:74:8c:b6:36:3e` (`c51dc642-a2e1-48df-a490e2748cb6363e`)
- **Type**: `effect`

### 2. Detection Pipeline Template (`include/pipelines/cavs/host-gateway-src-mfcc-tflm-capture.conf`)

Instantiates the real-time detection graph:
```conf
Object.Widget {
virtual."1" { name "virtual.tflm_sink" }
src."1" { ... }
mfcc."1" { ... }
tflmcly."1" { ... }
}
```

## Configuration and Scripts
### 3. Top-Level Topology Configuration (`sof-hda-tflm.conf`)

Instantiates the complete HDA Mic WoV topology with dual-path KPB routing:
```conf
Object.Base.route [
# DAI -> Gain -> KPB
{ source "dai-copier.HDA.Analog.capture"; sink "gain.2.1" }
{ source "gain.2.1"; sink "kpb.2.1" }

# KPB Pin 1 -> Real-time Detection Path
{ source "kpb.2.1"; sink "src.1.1" }

# KPB Pin 2 -> Host WoV Draining Path
{ source "kpb.2.1"; sink "host-copier.0.capture" }
]
```

---

## Building and Testing Topologies

### Pre-processing and Compiling with `alsatplg`

Build `sof-hda-tflm.tplg` using the Topology v2 pre-processor:

```bash
ALSA_CONFIG_DIR=tools/topology/topology2 \
tools/bin/alsatplg \
-I tools/topology/topology2/ \
-p -c tools/topology/topology2/sof-hda-tflm.conf \
-o build/sof-hda-tflm.tplg
```

### Inspecting Decoded Topology Graphs

Decode and verify the compiled `.tplg` binary:

```bash
tools/bin/alsatplg -d build/sof-hda-tflm.tplg -o decoded.txt
grep -A 20 "SectionGraph" decoded.txt
```

Expected graph output:
```text
SectionGraph {
set0 {
gain.2.1 <- dai-copier.HDA.Analog.capture
kpb.2.1 <- gain.2.1
src.1.1 <- kpb.2.1 (Real-Time Detection Path)
host-copier.0.capture <- kpb.2.1 (Host WoV Draining Path)
}
set1 {
mfcc.1.1 <- src.1.1
tflmcly.1.1 <- mfcc.1.1
virtual.tflm_sink <- tflmcly.1.1 (Detection Sink Termination)
}
}
```

---

## Source Files

- **Kconfig**: Enforces requirements for C++17 support and core framework staging logic (`COMP_TENSORFLOW`).
- **CMakeLists.txt**: An intricate build specification linking the Tensilica neural network library block computations (`nn_hifi_lib`) and the TensorFlow Lite micro core engine (`tflm_lib`). Also hooks the `tflm-classify.c` SOF adapter via compiler flags explicitly enforcing memory, precision, and XTENSA optimizations.
- **tflmcly.toml**: Topology definition for the specific TFLM Classifier implementation binding the engine against the UUID `UUIDREG_STR_TFLMCLY`.
- **[tflm-classify.c](file:///home/lrg/work/sof-ptl/sof/src/audio/tensorflow/tflm-classify.c)**: SOF module adapter implementation for TFLM.
- **[speech.cc](file:///home/lrg/work/sof-ptl/sof/src/audio/tensorflow/speech.cc)** / **[speech.h](file:///home/lrg/work/sof-ptl/sof/src/audio/tensorflow/speech.h)**: TFLM C++ API bridge & micro speech tensor wrapper.
- **[tflm.conf](file:///home/lrg/work/sof-ptl/sof/tools/topology/topology2/include/components/tflm.conf)**: Topology v2 widget class definition.
- **[host-gateway-src-mfcc-tflm-capture.conf](file:///home/lrg/work/sof-ptl/sof/tools/topology/topology2/include/pipelines/cavs/host-gateway-src-mfcc-tflm-capture.conf)**: Detection pipeline template.
- **[sof-hda-tflm.conf](file:///home/lrg/work/sof-ptl/sof/tools/topology/topology2/sof-hda-tflm.conf)**: Top-level WoV topology configuration.
134 changes: 122 additions & 12 deletions src/audio/tensorflow/tflm-classify.c
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,13 @@
#include <stddef.h>
#include <stdint.h>

#include <ipc4/base-config.h>
#include <ipc4/header.h>
#include <ipc4/module.h>
#include <ipc4/notification.h>
#include <sof/audio/kpb.h>
#include <sof/lib/notifier.h>

#include "speech.h"

SOF_DEFINE_REG_UUID(tflmcly);
Expand All @@ -44,8 +51,87 @@ static const char * const prediction[] = TFLM_CATEGORY_DATA;
struct tflm_comp_data {
struct comp_data_blob_handler *model_handler;
struct tf_classify tfc;
struct ipc_msg *msg;
struct kpb_event_data event_data;
struct kpb_client client_data;
};

static void tflm_notify_kpb(struct processing_module *mod)
{
struct tflm_comp_data *cd = module_get_private_data(mod);
struct comp_dev *dev = mod->dev;

comp_info(dev, "TFLM keyword trigger -> notifying KPB to begin draining");

cd->client_data.r_ptr = NULL;
cd->client_data.sink = NULL;
cd->client_data.id = 0;
cd->event_data.event_id = KPB_EVENT_BEGIN_DRAINING;
cd->event_data.client_data = &cd->client_data;

notifier_event(dev, NOTIFIER_ID_KPB_CLIENT_EVT,
NOTIFIER_TARGET_CORE_ALL_MASK, &cd->event_data,
sizeof(cd->event_data));
}

static int tflm_ipc_notification_init(struct processing_module *mod)
{
struct tflm_comp_data *cd = module_get_private_data(mod);
struct ipc_msg msg_proto;
struct comp_dev *dev = mod->dev;
struct comp_ipc_config *ipc_config = &dev->ipc_config;
union ipc4_notification_header *primary =
(union ipc4_notification_header *)&msg_proto.header;
struct sof_ipc4_notify_module_data *msg_module_data;
struct sof_ipc4_control_msg_payload *msg_payload;

memset_s(&msg_proto, sizeof(msg_proto), 0, sizeof(msg_proto));
primary->r.notif_type = SOF_IPC4_MODULE_NOTIFICATION;
primary->r.type = SOF_IPC4_GLB_NOTIFICATION;
primary->r.rsp = SOF_IPC4_MESSAGE_DIR_MSG_REQUEST;
primary->r.msg_tgt = SOF_IPC4_MESSAGE_TARGET_FW_GEN_MSG;
cd->msg = ipc_msg_w_ext_init(msg_proto.header, msg_proto.extension,
sizeof(struct sof_ipc4_notify_module_data) +
sizeof(struct sof_ipc4_control_msg_payload) +
sizeof(struct sof_ipc4_ctrl_value_chan));
if (!cd->msg) {
comp_err(dev, "Failed to initialize TFLM notification");
return -ENOMEM;
}

msg_module_data = (struct sof_ipc4_notify_module_data *)cd->msg->tx_data;
msg_module_data->instance_id = IPC4_INST_ID(ipc_config->id);
msg_module_data->module_id = IPC4_MOD_ID(ipc_config->id);
msg_module_data->event_id = SOF_IPC4_NOTIFY_MODULE_EVENTID_ALSA_MAGIC_VAL |
SOF_IPC4_SWITCH_CONTROL_PARAM_ID;
msg_module_data->event_data_size = sizeof(struct sof_ipc4_control_msg_payload) +
sizeof(struct sof_ipc4_ctrl_value_chan);

msg_payload = (struct sof_ipc4_control_msg_payload *)msg_module_data->event_data;
msg_payload->id = 0;
msg_payload->num_elems = 1;
msg_payload->chanv[0].channel = 0;

comp_dbg(dev, "TFLM notification init: instance_id = 0x%08x, module_id = 0x%08x",
msg_module_data->instance_id, msg_module_data->module_id);
return 0;
}

static void tflm_send_keyword_notification(struct processing_module *mod, uint32_t category_idx)
{
struct tflm_comp_data *cd = module_get_private_data(mod);
struct sof_ipc4_notify_module_data *msg_module_data;
struct sof_ipc4_control_msg_payload *msg_payload;

if (!cd->msg)
return;

msg_module_data = (struct sof_ipc4_notify_module_data *)cd->msg->tx_data;
msg_payload = (struct sof_ipc4_control_msg_payload *)msg_module_data->event_data;
msg_payload->chanv[0].value = category_idx;
ipc_msg_send(cd->msg, NULL, false);
}

__cold static int tflm_init(struct processing_module *mod)
{
struct module_data *md = &mod->priv;
Expand Down Expand Up @@ -97,11 +183,19 @@ __cold static int tflm_init(struct processing_module *mod)
goto fail;
}

ret = tflm_ipc_notification_init(mod);
if (ret < 0) {
comp_err(dev, "failed to init notification");
goto fail;
}

return ret;

fail:
/* Passing NULL pointer to free functions is Ok */
mod_data_blob_handler_free(mod, cd->model_handler);
if (cd->msg)
ipc_msg_free(cd->msg);
mod_free(mod, cd);
return ret;
}
Expand All @@ -112,6 +206,8 @@ __cold static int tflm_free(struct processing_module *mod)

assert_can_be_cold();

if (cd->msg)
ipc_msg_free(cd->msg);
mod_data_blob_handler_free(mod, cd->model_handler);
mod_free(mod, cd);
return 0;
Expand Down Expand Up @@ -180,24 +276,23 @@ static int tflm_process(struct processing_module *mod,
{
struct tflm_comp_data *cd = module_get_private_data(mod);
struct comp_dev *dev = mod->dev;
size_t frame_bytes = source_get_frame_bytes(sources[0]);
int features = source_get_data_frames_available(sources[0]);
int features_available = source_get_data_frames_available(sources[0]);
const void *data_ptr, *buf_start;
size_t buf_size;
int ret;
int ret = 0;

comp_dbg(dev, "entry");

/* Window size is TFLM_FEATURE_ELEM_COUNT and we increment
* by TFLM_FEATURE_SIZE until buffer empty.
/* Window size is TFLM_FEATURE_ELEM_COUNT (1960 bytes int8_t) and we increment
* by TFLM_FEATURE_SIZE (40 bytes int8_t) until buffer empty.
*/
while (features >= TFLM_FEATURE_ELEM_COUNT) {
ret = source_get_data(sources[0], TFLM_FEATURE_ELEM_COUNT * frame_bytes,
while (features_available >= TFLM_FEATURE_ELEM_COUNT) {
ret = source_get_data(sources[0], TFLM_FEATURE_ELEM_COUNT,
&data_ptr, &buf_start, &buf_size);
if (ret)
return ret;

cd->tfc.audio_features = data_ptr;
cd->tfc.audio_features = (int8_t *)data_ptr;
cd->tfc.audio_data_size = TFLM_FEATURE_ELEM_COUNT;
ret = TF_ProcessClassify(&cd->tfc);
if (!ret) {
Expand All @@ -207,15 +302,30 @@ static int tflm_process(struct processing_module *mod,
return ret;
}

/* debug - dump the output */
/* Evaluate predictions */
int max_idx = 0;
float max_score = cd->tfc.predictions[0];

for (int i = 0; i < cd->tfc.categories; i++) {
comp_dbg(dev, "tf: predictions %1.3f %s",
cd->tfc.predictions[i], prediction[i]);
if (cd->tfc.predictions[i] > max_score) {
max_score = cd->tfc.predictions[i];
max_idx = i;
}
}

/* Check if a keyword ("yes" or "no", category indices 2 and 3) is detected */
if (max_idx >= 2 && max_score >= 0.70f) {
comp_info(dev, "TFLM keyword detected: %s (confidence %1.3f)",
prediction[max_idx], (double)max_score);
tflm_send_keyword_notification(mod, max_idx);
tflm_notify_kpb(mod);
}

/* advance by one stride */
source_release_data(sources[0], TFLM_FEATURE_SIZE * frame_bytes);
features = source_get_data_frames_available(sources[0]);
/* advance by one 20ms stride (40 int8_t features) */
source_release_data(sources[0], TFLM_FEATURE_SIZE);
features_available = source_get_data_frames_available(sources[0]);
}

return ret;
Expand Down
Loading
Loading