2023-02-17 22:56:15 -05:00
{
"cells": [
{
"cell_type": "markdown",
"id": "825fe381",
"metadata": {},
"source": [
"# Introduction\n",
"\n",
"This notebook demonstrates the process of training a new openWakeWord model, using synthetic speech generated with open-source TTS models, and negative data representing music, noise, and speech. While the process here is complete, only small samples of datasets are utilized so that a new model can be trained on CPUs. In practice, much larger volumes of data (both positive and negitive examples) is needed to produce robust models. See the [documentation](https://github.com/dscripka/openWakeWord/tree/main/docs/models) for the pre-trained openWakeWord models for more information about how these models were trained."
]
},
{
"cell_type": "markdown",
"id": "bd8e4597",
"metadata": {},
"source": [
"To start, we'll need to install the requirements needed to train new openWakeWord models."
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "1ba07a8f",
"metadata": {},
"outputs": [],
"source": [
"# Install requirements (it's recommended that you do this in a new virtual environment)\n",
"\n",
"# !pip install openwakeword\n",
"# !pip install speechbrain\n",
"# !pip install datasets\n",
"# !pip install scipy matplotlib"
]
},
{
"cell_type": "code",
"execution_count": 1,
"id": "c914b0c9",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:26:26.308309Z",
"start_time": "2023-02-18T03:26:24.785801Z"
}
},
2023-10-13 09:18:39 +02:00
"outputs": [],
2023-02-17 22:56:15 -05:00
"source": [
"# Imports\n",
"\n",
"import os\n",
"import collections\n",
"import numpy as np\n",
"from numpy.lib.format import open_memmap\n",
"from pathlib import Path\n",
"from tqdm import tqdm\n",
"import openwakeword\n",
"import openwakeword.data\n",
"import openwakeword.utils\n",
"import openwakeword.metrics\n",
"\n",
"import scipy\n",
"import datasets\n",
"import matplotlib.pyplot as plt\n",
"import torch\n",
"from torch import nn\n",
"import IPython.display as ipd"
]
},
{
"cell_type": "markdown",
"id": "b40a7f25",
"metadata": {},
"source": [
"# Data Preparation"
]
},
{
"cell_type": "markdown",
"id": "aee94c6e",
"metadata": {},
"source": [
"## Download Data"
]
},
{
"cell_type": "markdown",
"id": "00ea736a",
"metadata": {},
"source": [
"Next we'll load the data used for training. For the purposes of this demonstration, we'll use a small set of positive and negative.\n",
"\n",
"For the positive data, there are ~3400 synthetic examples of the phrase \"turn on the office lights\" that were produced with the text-to-speech models documented in a [separate repo](https://github.com/dscripka/synthetic_speech_dataset_generation).\n",
"\n",
2023-02-17 23:07:41 -05:00
"These positive examples can be downloaded [here](https://f002.backblazeb2.com/file/openwakeword-resources/data/turn_on_the_office_lights.tar.gz).\n",
2023-02-17 22:56:15 -05:00
"\n",
"For negative data, we'll use small, already prepared samples of the [fma-large dataset](https://github.com/mdeff/fma) for music, the [FSD50k dataset](https://zenodo.org/record/4060432#.Y-hA2BzMJhE) for noise, and the [Common Voice 11](https://huggingface.co/datasets/mozilla-foundation/common_voice_11_0) dataset for speech.\n",
"\n",
"The fma-large sample can be downloaded [here](https://f002.backblazeb2.com/file/openwakeword-resources/data/fma_sample.zip), and then extracted into the working director.\n",
"\n",
"The FSD50k sample can be downloaded [here](https://f002.backblazeb2.com/file/openwakeword-resources/data/fsd50k_sample.zip), and then extracted into the working directory.\n",
"\n",
"And we'll use the HuggingFace Datasets library to get a portion of the test split of the Common Voice 11 (CV11) corpus.\n",
"\n",
"Note the data provided here is intended for non-commerical applications only; you will need to verify the license status of this (and other) data if you intend to use it for commerical purposes."
]
},
{
"cell_type": "code",
"execution_count": 381,
"id": "a31f760c",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-12T16:08:32.045657Z",
"start_time": "2023-02-12T16:07:24.104475Z"
}
},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
2023-10-13 09:18:39 +02:00
"Reading metadata...: 16354it [00:00, 26183.62it/s]\n",
"100%|██████████| 5000/5000 [00:44<00:00, 112.28it/s]\n"
2023-02-17 22:56:15 -05:00
]
}
],
"source": [
"# Download CV11 test split from HuggingFace, and convert the audio into 16 khz, 16-bit wav files\n",
"\n",
"cv_11 = datasets.load_dataset(\"mozilla-foundation/common_voice_11_0\", \"en\", split=\"test\", streaming=True)\n",
"cv_11 = cv_11.cast_column(\"audio\", datasets.Audio(sampling_rate=16000, mono=True)) # convert to 16-khz\n",
"cv_11 = iter(cv_11)\n",
"\n",
"# Convert and save clips (only first 5000)\n",
"limit = 5000\n",
"for i in tqdm(range(limit)):\n",
" example = next(cv_11)\n",
2023-10-13 09:18:39 +02:00
" output = os.path.join(\"cv11_test_clips\", example[\"path\"][0:-4] + \".wav\")\n",
" os.makedirs(os.path.dirname(output), exist_ok=True)\n",
"\n",
2023-02-17 22:56:15 -05:00
" wav_data = (example[\"audio\"][\"array\"]*32767).astype(np.int16) # convert to 16-bit PCM format\n",
2023-10-13 09:18:39 +02:00
" scipy.io.wavfile.write(output, 16000, wav_data)\n"
2023-02-17 22:56:15 -05:00
]
},
{
"cell_type": "markdown",
"id": "0a12ab3d",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-12T02:04:47.611837Z",
"start_time": "2023-02-12T02:04:47.606080Z"
}
},
"source": [
"## Compute Audio Embeddings"
]
},
{
"cell_type": "markdown",
"id": "bfea6a2b",
"metadata": {},
"source": [
"Once all the data is downloaded, we can now get the audio embeddings for the positive and negative clips. As this part of the openWakeWord model is frozen (i.e., not updated during training), it makes sense to pre-compute these features so that they only need to be prepared once."
]
},
{
"cell_type": "code",
"execution_count": 2,
"id": "473349ce",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:26:45.282504Z",
"start_time": "2023-02-18T03:26:45.093446Z"
}
},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"/home/dscripka/anaconda3/envs/torch_gpu/lib/python3.9/site-packages/onnxruntime/capi/onnxruntime_inference_collection.py:54: UserWarning: Specified provider 'CUDAExecutionProvider' is not in available provider names.Available providers: 'CPUExecutionProvider'\n",
" warnings.warn(\n"
]
}
],
"source": [
"# Create audio pre-processing object to get openWakeWord audio embeddings\n",
"\n",
"F = openwakeword.utils.AudioFeatures()"
]
},
{
"cell_type": "markdown",
"id": "9e757355",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-12T02:14:32.160470Z",
"start_time": "2023-02-12T02:14:32.154438Z"
}
},
"source": [
"### Negative Clips"
]
},
{
"cell_type": "code",
"execution_count": 3,
"id": "ab401215",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:27:23.911209Z",
"start_time": "2023-02-18T03:26:47.968057Z"
}
},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"200it [00:00, 2779.44it/s]\n",
"100%|██████████| 200/200 [00:01<00:00, 141.58it/s]\n",
"1000it [00:00, 2806.48it/s]\n",
"100%|██████████| 1000/1000 [00:05<00:00, 177.36it/s]\n",
"5000it [00:01, 2555.99it/s]\n",
"100%|██████████| 5000/5000 [00:26<00:00, 188.73it/s]"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"6096 negative clips after filtering, representing ~12.0 hours\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"\n"
]
}
],
"source": [
"# Get negative example paths, filtering out clips that are too long or too short\n",
"\n",
"negative_clips, negative_durations = openwakeword.data.filter_audio_paths(\n",
" [\n",
" \"fma_sample\",\n",
" \"fsd50k_sample\",\n",
" \"cv11_test_clips\"\n",
" ],\n",
" min_length_secs = 1.0, # minimum clip length in seconds\n",
" max_length_secs = 60*30, # maximum clip length in seconds\n",
" duration_method = \"header\" # use the file header to calculate duration\n",
")\n",
"\n",
"print(f\"{len(negative_clips)} negative clips after filtering, representing ~{sum(negative_durations)//3600} hours\")"
]
},
{
"cell_type": "code",
"execution_count": 4,
"id": "221d8662",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:28:06.568812Z",
"start_time": "2023-02-18T03:28:06.524651Z"
}
},
"outputs": [],
"source": [
"# Use HuggingFace datasets to load files from disk by batches\n",
"\n",
"audio_dataset = datasets.Dataset.from_dict({\"audio\": negative_clips})\n",
"audio_dataset = audio_dataset.cast_column(\"audio\", datasets.Audio(sampling_rate=16000))"
]
},
{
"cell_type": "code",
"execution_count": 6,
"id": "37ec1163",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:30:38.082245Z",
"start_time": "2023-02-18T03:29:08.355371Z"
}
},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
" 98%|█████████▊| 94/96 [01:26<00:01, 1.09it/s]\n",
"15it [00:03, 4.61it/s] \n"
]
}
],
"source": [
"# Get audio embeddings (features) for negative clips and save to .npy file\n",
"# Process files by batch and save to Numpy memory mapped file so that\n",
"# an array larger than the available system memory can be created\n",
"\n",
"batch_size = 64 # number of files to load, compute features, and write to mmap at a time\n",
"clip_size = 3 # the desired window size (in seconds) for the trained openWakeWord model\n",
"N_total = int(sum(negative_durations)//clip_size) # maximum number of rows in mmap file\n",
"n_feature_cols = F.get_embedding_shape(clip_size)\n",
"\n",
"output_file = \"negative_features.npy\"\n",
"output_array_shape = (N_total, n_feature_cols[0], n_feature_cols[1])\n",
"fp = open_memmap(output_file, mode='w+', dtype=np.float32, shape=output_array_shape)\n",
"\n",
"row_counter = 0\n",
"for i in tqdm(np.arange(0, audio_dataset.num_rows, batch_size)):\n",
" # Load data in batches and shape into rectangular array\n",
" wav_data = [(j[\"array\"]*32767).astype(np.int16) for j in audio_dataset[i:i+batch_size][\"audio\"]]\n",
" wav_data = openwakeword.data.stack_clips(wav_data, clip_size=16000*clip_size).astype(np.int16)\n",
" \n",
" # Compute features (increase ncpu argument for faster processing)\n",
" features = F.embed_clips(x=wav_data, batch_size=1024, ncpu=8)\n",
" \n",
" # Save computed features to mmap array file (stopping once the desired size is reached)\n",
" if row_counter + features.shape[0] > N_total:\n",
" fp[row_counter:min(row_counter+features.shape[0], N_total), :, :] = features[0:N_total - row_counter, :, :]\n",
" fp.flush()\n",
" break\n",
" else:\n",
" fp[row_counter:row_counter+features.shape[0], :, :] = features\n",
" row_counter += features.shape[0]\n",
" fp.flush()\n",
" \n",
"# Trip empty rows from the mmapped array\n",
"openwakeword.data.trim_mmap(output_file)"
]
},
{
"cell_type": "markdown",
"id": "c60aa86d",
"metadata": {},
"source": [
"Now we have all of the negative features prepared, and saved to fixed durations clips in a Numpy array. For this data, the array is small at ~160 MB, but in-practice the memory mapping allows the array to be very large (e.g., 100s of GBs)."
]
},
{
"cell_type": "markdown",
"id": "1b49b9c6",
"metadata": {},
"source": [
"### Positive Clips"
]
},
{
"cell_type": "markdown",
"id": "6f3f9ed0",
"metadata": {},
"source": [
2023-02-17 23:07:41 -05:00
"First, [download](https://f002.backblazeb2.com/file/openwakeword-resources/data/turn_on_the_office_lights.tar.gz) and extract the positive clips into the working directory.\n",
"\n",
"Then the positive clips will be prepared in two way:\n",
2023-02-17 22:56:15 -05:00
"\n",
"1) Mixing the synthetic positive clips with negative data at random SNRs to simulate noise data\n",
"\n",
"2) Aligning the positive clips with background data such that the end of the input window aligns with the end of the positive clip. This way the model will learn to predict the presence of the wakeword/phrase immediately after it is spoken.\n",
"\n",
"In practice, there are other possible ways to augment the positive data (e.g., creating reverberation with room impulse response files, mixing with synthetic noise, etc.) but in practice we have observed that mixing with realistic background data provides the best results. Again, see the [documentation](https://github.com/dscripka/openWakeWord/tree/main/docs/models) for the pre-trained openWakeWord models for more information about the types of data augmentation used.\n",
"\n",
"After this prepartion, the positive clips will be converted into the openWakeWord features in the same way as the negative files."
]
},
{
"cell_type": "code",
"execution_count": 7,
"id": "fe1964fb",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:31:01.912793Z",
"start_time": "2023-02-18T03:30:43.623741Z"
}
},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"3388it [00:01, 2771.26it/s]\n",
"100%|██████████| 3388/3388 [00:17<00:00, 198.61it/s]"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"3203 positive clips after filtering\n"
]
},
{
"name": "stderr",
"output_type": "stream",
"text": [
"\n"
]
}
],
"source": [
"# Get positive example paths, filtering out clips that are too long or too short\n",
"\n",
"positive_clips, durations = openwakeword.data.filter_audio_paths(\n",
" [\n",
" \"turn_on_the_office_lights\"\n",
" ],\n",
" min_length_secs = 1.0, # minimum clip length in seconds\n",
" max_length_secs = 2.0, # maximum clip length in seconds\n",
" duration_method = \"header\" # use the file header to calculate duration\n",
")\n",
"\n",
"print(f\"{len(positive_clips)} positive clips after filtering\")"
]
},
{
"cell_type": "code",
"execution_count": 8,
"id": "5d9dc47b",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:31:05.710564Z",
"start_time": "2023-02-18T03:31:05.699618Z"
}
},
"outputs": [],
"source": [
"# Define starting point for each positive clip based on its length, so that each one ends \n",
"# between 0-200 ms from the end of the total window size chosen for the model.\n",
"# This results in the model being most confident in the prediction right after the\n",
"# end of the wakeword in the audio stream, reducing latency in operation.\n",
"\n",
"# Get start and end positions for the positive audio in the full window\n",
"sr = 16000\n",
"total_length_seconds = 3 # must be the some window length as that used for the negative examples\n",
"total_length = int(sr*total_length_seconds)\n",
"\n",
"jitters = (np.random.uniform(0, 0.2, len(positive_clips))*sr).astype(np.int32)\n",
"starts = [total_length - (int(np.ceil(i*sr))+j) for i,j in zip(durations, jitters)]\n",
"ends = [int(i*sr) + j for i, j in zip(durations, starts)]\n",
"\n",
"# Create generator to mix the positive audio with background audio\n",
"batch_size = 8\n",
"mixing_generator = openwakeword.data.mix_clips_batch(\n",
" foreground_clips = positive_clips,\n",
" background_clips = negative_clips,\n",
" combined_size = total_length,\n",
" batch_size = batch_size,\n",
" snr_low = 5,\n",
" snr_high = 15,\n",
" start_index = starts,\n",
" volume_augmentation=True, # randomly scale the volume of the audio after mixing\n",
")\n"
]
},
{
"cell_type": "code",
"execution_count": null,
"id": "70898754",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-12T03:35:32.377177Z",
"start_time": "2023-02-12T03:35:32.349576Z"
}
},
"outputs": [],
"source": [
"# (Optionally) listen to mixed clips to confirm that the mixing appears correct\n",
"\n",
"mixed_clips, labels, background_clips = next(mixing_generator)\n",
"ipd.display(ipd.Audio(mixed_clips[0], rate=16000, normalize=True, autoplay=False))"
]
},
{
"cell_type": "code",
"execution_count": 10,
"id": "621c2ee6",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:33:35.853508Z",
"start_time": "2023-02-18T03:31:44.655774Z"
}
},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"100%|██████████| 400/400 [01:50<00:00, 3.62it/s]\n",
"4it [00:00, 5.66it/s] \n"
]
}
],
"source": [
"# Iterate through the mixing generator, computing audio features for positive examples and saving them\n",
"\n",
"N_total = len(positive_clips) # maximum number of rows in mmap file\n",
"n_feature_cols = F.get_embedding_shape(total_length_seconds)\n",
"\n",
"output_file = \"turn_on_the_office_lights_features.npy\"\n",
"output_array_shape = (N_total, n_feature_cols[0], n_feature_cols[1])\n",
"\n",
"fp = open_memmap(output_file, mode='w+', dtype=np.float32, shape=output_array_shape)\n",
"\n",
"row_counter = 0\n",
"for batch in tqdm(mixing_generator, total=N_total//batch_size):\n",
" batch, lbls, background = batch[0], batch[1], batch[2]\n",
" \n",
" # Compute audio features\n",
" features = F.embed_clips(batch, batch_size=256)\n",
"\n",
" # Save computed features\n",
" fp[row_counter:row_counter+features.shape[0], :, :] = features\n",
" row_counter += features.shape[0]\n",
" fp.flush()\n",
" \n",
" if row_counter >= N_total:\n",
" break\n",
"\n",
"# Trip empty rows from the mmapped array\n",
"openwakeword.data.trim_mmap(output_file)\n"
]
},
{
"cell_type": "markdown",
"id": "77a31736",
"metadata": {},
"source": [
"Alright! At this point the positive and negative features have been pre-computed and saved to disk, and now a model can be trained that takes these features and predicts whether the wakeword/phrase is present."
]
},
{
"cell_type": "markdown",
"id": "cab316e2",
"metadata": {},
"source": [
"# Training the Model"
]
},
{
"cell_type": "markdown",
"id": "a514773d",
"metadata": {},
"source": [
"At this point, you are free to use any type of model that you like, but in practice we've observed that a simple full-connected neural network can often perform quite well. For this example notebook, we'll create and train this network in Pytorch, but any framework that can export a model to the [ONNX](https://onnx.ai/) format will also work."
]
},
{
"cell_type": "markdown",
"id": "f0dab37e",
"metadata": {},
"source": [
"## Loading Data"
]
},
{
"cell_type": "code",
"execution_count": 11,
"id": "c7ed4de5",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:33:52.884948Z",
"start_time": "2023-02-18T03:33:51.191681Z"
}
},
"outputs": [],
"source": [
"# Load the data prepared in previous steps (it's small enough to load entirely in memory)\n",
"\n",
"negative_features = np.load(\"negative_features.npy\")\n",
"positive_features = np.load(\"turn_on_the_office_lights_features.npy\")\n",
"\n",
"X = np.vstack((negative_features, positive_features))\n",
"y = np.array([0]*len(negative_features) + [1]*len(positive_features)).astype(np.float32)[...,None]\n",
"\n",
"# Make Pytorch dataloader\n",
"batch_size = 512\n",
"training_data = torch.utils.data.DataLoader(\n",
" torch.utils.data.TensorDataset(torch.from_numpy(X), torch.from_numpy(y)),\n",
" batch_size = batch_size,\n",
" shuffle = True\n",
")\n"
]
},
{
"cell_type": "markdown",
"id": "1d1ba9e4",
"metadata": {},
"source": [
"## Define Model"
]
},
{
"cell_type": "code",
"execution_count": 12,
"id": "d7c71798",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:33:54.913238Z",
"start_time": "2023-02-18T03:33:54.896447Z"
}
},
"outputs": [],
"source": [
"# Define fully-connected network in PyTorch\n",
"\n",
"layer_dim = 32\n",
"fcn = nn.Sequential(\n",
" nn.Flatten(),\n",
" nn.Linear(X.shape[1]*X.shape[2], layer_dim), # since the input is flattened, it's timesteps*feature columns\n",
" nn.LayerNorm(layer_dim),\n",
" nn.ReLU(),\n",
" nn.Linear(layer_dim, layer_dim),\n",
" nn.LayerNorm(layer_dim),\n",
" nn.ReLU(),\n",
" nn.Linear(layer_dim, 1),\n",
" nn.Sigmoid(),\n",
" )\n",
"\n",
"loss_function = torch.nn.functional.binary_cross_entropy\n",
"optimizer = torch.optim.Adam(fcn.parameters(), lr=0.001)\n"
]
},
{
"cell_type": "markdown",
"id": "6bb834c1",
"metadata": {},
"source": [
"## Train Model"
]
},
{
"cell_type": "code",
"execution_count": 13,
"id": "5bc28f8b",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:33:59.286835Z",
"start_time": "2023-02-18T03:33:57.795926Z"
}
},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"100%|██████████| 10/10 [00:01<00:00, 6.74it/s]\n"
]
}
],
"source": [
"# Define training loop, metrics, and logging\n",
"\n",
"n_epochs = 10\n",
"history = collections.defaultdict(list)\n",
"for i in tqdm(range(n_epochs), total=n_epochs):\n",
" for batch in training_data:\n",
" # Get data for batch\n",
" x, y = batch[0], batch[1]\n",
" \n",
" # Get weights for classes, and assign 10x higher weight to negative class\n",
" # to help the model learn to not have too many false-positives\n",
" # As you have more data (both positive and negative), this is less important\n",
" weights = torch.ones(y.shape[0])\n",
" weights[y.flatten() == 1] = 0.1\n",
" \n",
" # Zero gradients\n",
" optimizer.zero_grad()\n",
" \n",
" # Run forward pass\n",
" predictions = fcn(x)\n",
" \n",
" # Update model parameters\n",
" loss = loss_function(predictions, y, weights[..., None])\n",
" loss.backward()\n",
" optimizer.step()\n",
" \n",
" # Log metrics\n",
" history['loss'].append(float(loss.detach().numpy()))\n",
" \n",
" tp = sum(predictions.flatten()[y.flatten() == 1] >= 0.5)\n",
" fn = sum(predictions.flatten()[y.flatten() == 1] < 0.5)\n",
" history['recall'].append(float(tp/(tp+fn).detach().numpy()))\n"
]
},
{
"cell_type": "code",
"execution_count": 14,
"id": "4231dd84",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:34:01.172348Z",
"start_time": "2023-02-18T03:34:01.043030Z"
}
},
"outputs": [
{
"data": {
"text/plain": [
"(0.0, 1.0)"
]
},
"execution_count": 14,
"metadata": {},
"output_type": "execute_result"
},
{
"data": {
2023-10-13 09:18:39 +02:00
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAiMAAAGiCAYAAAA1LsZRAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjUuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8qNh9FAAAACXBIWXMAAA9hAAAPYQGoP6dpAABjbUlEQVR4nO3dd3hUZdoG8HtaJr2QHgghQOhNQjEgRVAUFXVFRd1PYBVX7IAVu667qLuWtYAFsayusCq6rIssqPSiSIdQAoQkQEJI75nMzPn+eOdMSWbSSOZMcu7fdeXKzJkzM++Zgcw9z1uORpIkCUREREQK0SrdACIiIlI3hhEiIiJSFMMIERERKYphhIiIiBTFMEJERESKYhghIiIiRTGMEBERkaIYRoiIiEhRDCNERESkKIYRIiIiUlSLw8imTZswbdo0JCQkQKPR4LvvvmvyPhs3bkRqair8/f3Rs2dPvPfee61pKxEREXVCLQ4jlZWVGDp0KN55551m7Z+ZmYmrrroK48aNw549e/Dkk0/iwQcfxDfffNPixhIREVHno7mQE+VpNBp8++23uP766z3u8/jjj2PVqlU4fPiwfdvcuXOxb98+bN++vbVPTURERJ2Evr2fYPv27ZgyZYrLtiuuuAIfffQR6urqYDAYGtyntrYWtbW19utWqxVFRUWIjIyERqNp7yYTERFRG5AkCeXl5UhISIBW67kzpt3DSF5eHmJjY122xcbGwmw2o6CgAPHx8Q3us2jRIrzwwgvt3TQiIiLygpycHHTr1s3j7e0eRgA0qGbIPUOeqhwLFy7EggUL7NdLS0vRvXt35OTkIDQ0tP0aSkTeI0nAy4ni8uV/Akb8Qdn2NNcipz+oC08r1w5nFeeBty9yXB81Fxg5B3h3RMN9b/4M6DWp4faPrwLy9juuT/8I6HOFuLz6UWDfl677B8cCD+xy3ZaxDvja9j72uhy4+WNx2fm9ltvwr5n1GqAFYBUXL30KsFqAjS+77qLRA5JZXJ7xBdBzAvDRFCA/HbjiZSB5HPDeWAA6YOa3wGfXNjxO+9P5AVaT59tl8r/NkxuBFb8X27r0Bu7e0HDfH18Edn7gui12EHDuoOfHj+wD3PUT8Gov0R6/EMBU7rqPIQiYtQpYOhkwhgIL0oG/9QXqKptuf33DZwFj5wNvD3Ns6zcN+N2Slj9WM5SVlSExMREhISGN7tfuYSQuLg55eXku2/Lz86HX6xEZGen2PkajEUajscH20NBQhhGizqK6BDDavpB0iQba6v926WlgwyJg9FwgbnDj+5qqAL/Alj2+vx6QLOJyU22WJODEz0DsQKD0DHByPTB2HqDTAxYzoNECcum6skB80Oj9gMpCYOeHwLDbgOIsoKYE6D/Nc3ursh2vJQBoqwBNhes22cHPgIGXA35Brtvrzov9E0cDOb8AuVuAodcAp7YAJ79v+Fh15wF/nevj1J517Je7GTBqAWOw7dic7n/iPw0fT6MVrxcA1JwWr0+D9lsA2LZVZQGmfKD0sHhPRs4AgiKBrv2AgqPAoX+4P37HATgeqzE6k3ifz+90PF7lCQC2ILD5NRHi+lwB7P/Q9Tnjhoh/h/++1/Pjj74VCAsDohOAkiwAbt63hL5A17627eW2190KaFsxbGHoNCChJ5A4UIQ4ABg4ue3+/3nQ1BCLdg8jaWlp+M9//uOybe3atRgxYoTb8SJEpBKVBe3zuOueBQ5+A+z5HHi+1PN+p7YAn14LTHoaGLfA8371BUUBFefE5eoSICDc877Z24HPbwB6TQaqi4Gzu4GY/sDJDcDefwIh8cA9W8U352VTgb5XiqrBmseBA18Buz4BaisAUwVw3TvAqgeBK/4MXHyP6/NUnne9Xl0MlOe6b9OJn4DX+wP3/wYEx4htFjNQmS8uD58pwkjGOiA4zlGd0OiAiB5A0QnbA0nA+aNA1+GOxy487rhsqQWytooP6ZJs1zYc/Lphu+SABwDnj7ledyc/HbDYKhs9J4ggAgB9pogwUv85YgYC+Ycaf0x3am3/hk5usG3QAJCA1/sBWgNgrRObc34Rv7UGICJJvBb9rrFVl2z3CUsESnOA8CTx3p8/Agy9VdwvJM4WRmzCkxzXo/oA/qGAX7D4t3B2L2Cuafmx6P2BHpeIyz3GOcJIj/Etf6w21uKpvRUVFdi7dy/27t0LQEzd3bt3L7KzxT+2hQsXYuZMR/lt7ty5yMrKwoIFC3D48GEsW7YMH330ER555JG2OQIi6piqnMKIudbzfi1VkOG43NhkwcxN4gPv5PqWPb5kdVwuPOF5PwA4Z/vwyzvg+KDe/Rnw6wfiQ6UwA8jeAXwzR3x4p/9b7JNlm2lYnmsr2UvAL++L9h77X8PnaRBGSoCysw33S7AFh5pSEcac7y9ZReAYcD0ADVB2BjjyvWOf/tcAt/0LuHGZ+CADgHzHLEm3r0fpadff9SVdIrol6is4ChQcc38f2blDwLE14nK/qx3b+17tut+YB4BpfwdmfgcEdGn8MWX3bAMumS8u15QC+UeA3H3iuhweABFEEi8GrlgE9LxUBIgxDwDT3gKG/R8w+o8ivE5+Brjo/4Dr3hWvcepsYOa/gXkHgNAE8VghcY7H1RqA7hc7rkel2PaxjbHM2ua4TaMDdA17ElwEdAGCokXQlCtrPSeI36FdgchezXtd2lGLKyO//fYbLr30Uvt1eWzHrFmz8MknnyA3N9ceTAAgOTkZq1evxvz58/Huu+8iISEBb731FqZPn94GzSeiDsv5A9TSjL775grt6hj7UHgCiOrtfr/iLNffzSFJ4sNJVngc6JbqeX+5IiBXHQBRLXG2+1Og6KTrc4TEAWX1PsDlY3L3IS2/loGRQFWhrTLi2j2OwCjgj+tFdWX3p+LxBt0gbiu3BZeQONGt0iVZtEke6zDtLWDQdHFbVG8g51fg1GbxzXr/V0D6d8A1bzqCYNwQ8fhyu0pzxO+UKSIE1VWJ61NfER+0L8W4trW6uOEx1pe713G5z5WOy4mjgdBujtev2yhggG3syD1bRTXiH79z7P/HjaJi8bcUEfa0eiCqr+ODv6YM2PKGuNzvGlGZCusqutR6TgDih4rb0up1xfQY67g87mHH5afPiefQaACDv2N7sFMYCYkT1RCZfDk03hZgbWEkeTzwuw9E8FxqGwc04HrR9qjewH9tz9v7MmD6h67t6zMVmPQMkDhKtEVhLQ4jEydORGNLk3zyyScNtk2YMAG7d+9u6VO1iCRJMJvNsFiaKO1Rs+l0Ouj1ek6npvbh3E3TmpKzJ85hIXu75zAiB4WyM6KbQtfEn8NfPxRlcufg5NwtUVcDbHwF6DtV/IEHHB/C7toXGCWqQwfrLQBZW+76Lbm+sjPiA9LfqY9f/tCP6iOO2V03TZdk8Tt+iPidd8Bxmxxc5OeN7u8akAZcK4KILKa/+J1/GNhuWwCzKBOosD1O0hgRRipsIazE9jpE9xMDU0/8JK5H9gb0Rkf3Q306v6aDatxgIMxpULFWK0LWtrdsz9nXcVtoAmAIcLqzBogZIMbphHcHijPFb50e8A8Tu5w/4ghZ4xYAgV1E115r6TwMT3B+z4NjPYSRruK3XDmL6CECitxWQAyYnfCoeJ3XLBSvn/NjybRaYLzv9FB4ZTZNezOZTMjNzUVVVZXSTel0AgMDER8fDz8/P6WbQp1BXbX4Vp00pl4YacPKiEsVYgcw/Hb3+8lhxGoWH9zhiQ33qSoS3TjBccBqN3+4MzcBpvmi9P3Ti8COd4FtbwPPFrg+hzvDb3d843Zp/3nxIdyYggzXikyFHEZSPIeRCFsYibN9k891mjkj7ytXA2L6AUf/Ky4HxQABEa6PFTNA/HYONPJ4jKBoETIAx9gaOZSFdwd6ThRhJDLFURkI6OI+jHRPE5UUySLGO5hrAGhEV4gcNup3ywCuYUQ+bpl/uJidUlcpxszoba91l54ijHTpKa4bbWHv/BHxO3Yw0LWRKtiFkl97wLUyotE52iTvY64WvyN6iN9+gUBYd6A02xG+tTpR4Tl3QLyfPq7DhxGr1YrMzEzodDokJCTAz8+P3+TbgCR
2023-02-17 22:56:15 -05:00
"text/plain": [
"<Figure size 640x480 with 1 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"# Plot training metrics\n",
"\n",
"plt.figure()\n",
"plt.plot(history['loss'], label=\"loss\")\n",
"plt.plot(history['recall'], label=\"recall\")\n",
"plt.legend()\n",
"plt.ylim(0,1)\n"
]
},
{
"cell_type": "markdown",
"id": "03a3ae78",
"metadata": {},
"source": [
"## Try the Model on an Example Clip"
]
},
{
"cell_type": "markdown",
"id": "65c43bb6",
"metadata": {},
"source": [
"To confirm that the model is working as expected, let's test it on an example audio file (obtained from Youtube) of someone talking and then saying the phrase \"turn on the office lights\" at the end of the clip. We'll simulate how the model would be used in production, by predicting every 80 ms (1280 samples) and plotting the predictions over time.\n",
"\n",
"This clip is a good sanity test to confirm the model is performing in the right way, as it contains about ~30 seconds of speech that does no contain the target phrase, but does contain related words (e.g., \"lights\") that should not result in an activation. So ideally, the model scores are low up to the end of the recording, where there should then be a spike right after the target spoken phrase."
]
},
{
"cell_type": "code",
"execution_count": 16,
"id": "d6ad350e",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:38:17.596854Z",
"start_time": "2023-02-18T03:38:16.926888Z"
}
},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"100%|██████████| 394/394 [00:00<00:00, 19403.48it/s]\n"
]
},
{
"data": {
2023-10-13 09:18:39 +02:00
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAiMAAAGiCAYAAAA1LsZRAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjUuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8qNh9FAAAACXBIWXMAAA9hAAAPYQGoP6dpAAAyfklEQVR4nO3df3RU9Z3/8ddMflKEqICBKGC01WpRtGG1waKttrFo3fXUrWzdI9piT1NUxKhr0bNFXb8b1+93WftjQbuC1FNbWQu27rdUSVv5oehWIygCa/lWNNEmpsEagmggmc/3j+TeyWTy496YyWfmc5+Pc3JCZu4kn8u9TF58frw/MWOMEQAAgCVx2w0AAADRRhgBAABWEUYAAIBVhBEAAGAVYQQAAFhFGAEAAFYRRgAAgFWEEQAAYBVhBAAAWEUYAQAAVoUOI5s3b9Yll1yisrIyxWIx/eIXvxjyNZs2bVJFRYWKi4t1wgkn6P777x9OWwEAgINCh5H3339fM2fO1A9/+MNAx+/du1cXXXSR5syZo23btum2227TokWLtHbt2tCNBQAA7ol9lI3yYrGYHn/8cV166aUDHnPrrbfqiSee0O7du/3Hqqur9fLLL+u5554b7o8GAACOyM/0D3juuedUVVWV8tiFF16olStX6vDhwyooKEh7TUdHhzo6OvyvE4mE3n33XU2YMEGxWCzTTQYAACPAGKP29naVlZUpHh94MCbjYaS5uVmlpaUpj5WWlqqzs1Otra2aMmVK2mtqa2t15513ZrppAABgFDQ2Nuq4444b8PmMhxFJab0Z3sjQQL0cS5YsUU1Njf91W1ubpk2bpsbGRo0fPz5zDQUAACNm//79mjp1qsaNGzfocRkPI5MnT1Zzc3PKYy0tLcrPz9eECRP6fU1RUZGKiorSHh8/fjxhBACAHDPUFIuM1xmprKxUXV1dymMbNmzQrFmz+p0vAgAAoiV0GDlw4IC2b9+u7du3S+peurt9+3Y1NDRI6h5imT9/vn98dXW13nzzTdXU1Gj37t1atWqVVq5cqZtvvnlkzgAAgBz31l8O6j9fbNSHh7tsN8WK0GHkxRdf1JlnnqkzzzxTklRTU6MzzzxT3/3udyVJTU1NfjCRpPLycq1fv14bN27UGWecoX/6p3/S97//fV122WUjdAoAAOS22l//j/7h56/oypX/rbaDh9Oe/82ud3TtT19Sy/4PLbQu8z5SnZHRsn//fpWUlKitrY05IwAA51y58r+1ZU+rJOmk0iO0+utnqezIMZKkJ19tVvVP6iVJl886Tvf+7Uxr7Qwr6O9v9qYBAMAyr1sgFpP+8M4BXbZiqza+1qLv/3aPFj5S7x/38/q39NZfDlpqZeYQRgAAsMyoO43ccuHJOnHSWDW1fairH3pBy+r+oISR5s2aqrPLj1bCSGvr37bc2pFHGAEAwLJEovvzsUeO0dpvz9bls45TWUmxPlU2Xv/nqzN1z2Wn6YypR0qSDnSkzynJdaNS9AwAAAzM6xmJxWI68mOFOTUvZCTQMwIAgGX+nJHBDoqlHusSwggAAJZ5+SI+SKXSWE8acTCLEEYAALCu12qagcToGQEAAJnizxkZ5BjvOeNg3whhBAAAyxL0jAAAAJuSxdAH393WVYQRAAAs86PIYD0jDgcVwggAAJZ5HSODrqbxh2ncG6chjAAAYFmQQZrkBFb3EEYAALDM6+0YbJjGe9LBjhHCCAAAtpkAq2n8Yx3sGyGMAABgWbLOyGAVWN1FGAEAwLIgPSPUGQEAABmTDCPsTQMAACwItJqGnhEAAJApQVbTJJ9yL40QRgAAsMwfpnF6murACCMAAFjmr6ZhAisAALAh2Goaip4BAIAMSU5gHXqYhqJnAABgxCWCTGBlmAYAAGSMP4F1YNQZAQAAGeMP0wTZnMZBhBEAACwLVGeEYRoAAJApXr6IByh6xgRWAAAw4pK9HYPsTZNMI84hjAAAYFmg1TRMYAUAAJligqym8eeMuBdHCCMAAGQJVtMAAAArvN6OwSaw+sdmuC02EEYAALAsSDl49qYBAAAZE2wCazcHswhhBAAA24L0djCBFQAAZEyyHHzwY11CGAEAwLLk0t5B5oyMUltsIIwAAGBdz2qaQX4r+8t+HewaIYwAAGBZoJ4RP4u4l0YIIwAAWBZqNY17WYQwAgCAbck6I4OgzggAAMgUf5jG5VmqgyCMAABgmfGHaYZeTcOcEQAAMOKCDNMki55lujWjjzACAIBt/jDNYD0jsd6HOoUwAgCAZf5qmkGOoWcEAABkTJBy8Mmn3EsjhBEAACwLUvTMZYQRAAAs81bIDNozwjANAADIlCB1RpjACgAAMsYEWE0jv2fEvThCGAEAwDJ/mGaQY5JFz9xDGAEAwLJAwzTsTQMAADIlWYGV1TQAAMACbx5IPECdEQc7RggjAADY5geMQEt73YsjhBEAACwLUvRssPkkuY4wAgCARb17OgLVGXGvY4QwAgCATb3DRZDOD+PgrBHCCAAAFvWOFvFBukYYpgEAABkRdJgmeXwGG2PJsMLI8uXLVV5eruLiYlVUVGjLli2DHv/II49o5syZ+tjHPqYpU6bo61//uvbt2zesBgMA4JJEyjDNYD0jzBnxrVmzRosXL9btt9+ubdu2ac6cOZo7d64aGhr6Pf6ZZ57R/PnztWDBAu3cuVOPPfaYXnjhBV1zzTUfufEAAOS6lDkggeqMuJdGQoeRZcuWacGCBbrmmmt0yimn6L777tPUqVO1YsWKfo9//vnndfzxx2vRokUqLy/XZz/7WX3rW9/Siy+++JEbDwBArkuZwBqozkhm22NDqDBy6NAh1dfXq6qqKuXxqqoqbd26td/XzJ49W2+99ZbWr18vY4zeeecd/fznP9fFF1884M/p6OjQ/v37Uz4AAHCdw3NUBxUqjLS2tqqrq0ulpaUpj5eWlqq5ubnf18yePVuPPPKI5s2bp8LCQk2ePFlHHnmkfvCDHwz4c2pra1VSUuJ/TJ06NUwzAQDIGb17OgZdTePVGcl0gywY1gTWWJ+/LGNM2mOeXbt2adGiRfrud7+r+vp6Pfnkk9q7d6+qq6sH/P5LlixRW1ub/9HY2DicZgIAkPV6zwEJMkzjYhrJD3PwxIkTlZeXl9YL0tLSktZb4qmtrdU555yjW265RZJ0+umna+zYsZozZ47uvvtuTZkyJe01RUVFKioqCtM0AAByUuDVND2fIz+BtbCwUBUVFaqrq0t5vK6uTrNnz+73NQcPHlQ8nvpj8vLyJLm52Q8AAGEELgfPBNakmpoaPfjgg1q1apV2796tG2+8UQ0NDf6wy5IlSzR//nz/+EsuuUTr1q3TihUr9Prrr+vZZ5/VokWLdNZZZ6msrGzkzgQAgBwUPFu4O2ck1DCNJM2bN0/79u3TXXfdpaamJs2YMUPr16/X9OnTJUlNTU0pNUeuvvpqtbe364c//KFuuukmHXnkkTr//PP1L//yLyN3FgAA5KigE1hdFjqMSNLChQu1cOHCfp9bvXp12mPXX3+9rr/++uH8KAAA3Ba6zoh7fSPsTQMAgEWJ3nNGBjnO4cU0hBEAAGzqHS4GKpPR+zkHO0YIIwAA2GToGSGMAABgU2rPyMDHJYueuRdHCCMAAFiUulFeNFfTEEYAALDIq6g6VA7xV9NkuD02EEYAALDI6xkZqk/E3yjPwTRCGAEAwCI/jAzZNdJzvIN9I4QRAAAs8odphjjO4fmrhBEAAGxK9ozYbYdNhBEAACzyOjqGGqah6BkAAMiIRCLkME1GW2MHYQQAgCwQeGmvg10jhBEAACxKLu0dYphmyL6T3EUYAQDAotBFz9zrGCGMAABgkxcu4hFeTkMYAQDAIn81zRDHJSewutc1QhgBAMCiRPB68JIYpgEAACMs9N40mW2OFYQRAACs8iawDlX0rOdoB7tGCCMAAFgUtBw8Rc8AAEBGeOGC1TQAAMCKwHNG/HGajDbHCsIIAAAWeatpAhc9y3B7bCCMAABgUXI+6lDl4L3j3YsjhBEAACwKXQ4+w+2xgTACAIBFyXLwQx0ZSzneJYQRAACygMu78g6
2023-02-17 22:56:15 -05:00
"text/plain": [
"<Figure size 640x480 with 1 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"# Load data\n",
"sr, dat = scipy.io.wavfile.read(\"training_tutorial_data/turn_on_the_office_lights_test_clip.wav\")\n",
"\n",
"# Pre-compute audio features using helper function\n",
"features = F._get_embeddings(dat)\n",
"\n",
"# Get predictions for each window\n",
"scores = []\n",
"for i in tqdm(range(0, features.shape[0]-28)): # 28 is the number of timestep frames for this model\n",
" window = features[i:i+28][None,]\n",
" with torch.no_grad():\n",
" scores.append(float(fcn(torch.from_numpy(window)).detach().numpy()))\n",
" \n",
"plt.figure()\n",
"_ = plt.plot(scores)\n",
"_ = plt.ylim(0,1)\n"
]
},
{
"cell_type": "markdown",
"id": "6531eefd",
"metadata": {},
"source": [
"Overall, the model is working well on this test clip. There are a few spikes around the word \"lights\" spoken in other contexts, but the clear activation is around the entire phrase.\n",
"\n",
"To make this test a little more difficult, let's arbitrarily mix the test clip with some background music from the fma-large dataset at a low signal-to-noise ratio to simulate a more realistic (and challenging) scenario. Listen to the clip below to get a more intuitive feel of what the type of audio environment this represents."
]
},
{
"cell_type": "code",
"execution_count": 19,
"id": "2f72631b",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:40:11.733430Z",
"start_time": "2023-02-18T03:40:11.721719Z"
}
},
"outputs": [
{
"data": {
"text/html": [
"\n",
" <audio controls=\"controls\" >\n",
" <source src=\"data:audio/wav;base64,UklGRiTuAgBXQVZFZm10IBAAAAABAAEAgD4AAAB9AAACABAAZGF0YQDuAgBoOVQ9+DKuMvQ2MTBcKOwnjiL+G88a7BxvIMsnAC/qJysbnhdjI0gzojbPNwZBc0I5Puw/ukLqQdE92zu0OD40QDB/KCsgERhUEFkGBQFOCwcX8h4VKwQuESmgKHMo2SbRIZQdeReMDgD/bPOd9Cv85gcpGewhaxzkI4srXCOoH5AkgxzVEloSNhEtFA8XGRE8C70KswVm/OP1OfXA9KP19/VF+uH1dPJB+Qr/wQ7DE4oNqgvsC10Iige3CI4FAQPC/Vbybu/h8LTvl/gr/yD2FuYa4cLl8e0U7TnuQfVJ867zO/dq95HyBPDZ64fgh9X+yzvIe8cZw83A18P+xlTHhcOqvvbF2dGF2v7oqPS+6s3fQ+tS8Ofe6dKZzevFUMTbx5zOvtZg1nPOh8zEygTFRsaPzFbOc8wSz4XicvTC99L+xQeGAQD3dPT86s/kFOvp6oXmruRk28TTFNkE4Aroi/Qk/Ov/mgQnB04MDxB/DO4GhACs97boMeAE69f1GPZ28a7uTfJ5AUYRFRUCHKgg7hj4Fi8cfxuJHmMisBvDGqIiTCRcIokinBo0CpYFwQXt/b8Atw6eGKgnvzEVJtcerCAlGeoU0xgHFhMO8gUo+of0ZvlY/jACmgNtCN4O1xNWHskiIxy9FscOOgm7GN0noiLiH3UfUBqSDpH/AO9S4dXfvN382qTe7eFk4zXrT/r2Cm8T4BjiGscSTApZCGsEUfvC8dvlsueP7bbryuwI9Cj5uvPE8dD35fhW/DIJHxLgE/wSGReDIHkrhyqOHrMZhg+s96bxFPpu8lTtsPfp/r8AfP+H94P8/g06Fk4ZQh+gIX8fpCCyIowgsh4tIP4cqhX4CvgC5AS1A4v2N+z78Rz7QfyoAq8O/hNKF4kdaSLbIIwhNiPVHfwZyw8VAWT/fP288TPpxuQW30fhbu0Y+rcMSB8ZHegXbR8IIuQewSMjJRsczRJ9CSsEMgNJ/GL8VAy5ESUH9f7G/Rb9ZwEtDTwU+hGYDbkP0xEPEiUVohZYFM0OVAcLC28OQf7A9HT6//st/Pn9RgDBAHL9wPjA+Gj/GQaqCh8N4gxZDdMMxQQ1+03tdNsO0svR+NBgzXfR79ib4UvuyPC+6IHs6faZ+K78EwAo/CL8U/xk8W7iduem8JnmYt3824vRzceeydXEFr7wxnDSxtY75Uv2ifes9Jv3p/OB5wrW1cYlzenOLckn0FzYXt0I1/zISsV3yjHUjdzP2l7ae+NF7PX08AIjCIwBwPzx853iwtkO1nvOYspayO/K5dQI2SnT29JQ25/jl+t09ir9wwGeCfYMYw6xDjAIgQPmAaz6Q+mL3Tvde9mu2TnlP+7K8/P+3g3mGBkYfRKoFCsewSBKFkYLVAqICcMGPAj4CuQJRgDn9ML13/xV/Dn+WQR9Ag8EERakK9c0Ajm8Oag1TC9AJewcXRZzCg79OfYv+n7/oAHuBTwOzQ9hB9wCVwI2ClgYKyKDLYk9xz4VMFokgxWmBMUHCxeqFBcHmgjyEPIRExXgG5IexSMELeoqjCcdMQI3vDXnOMs3Dy4zJwAdjgaL8frqDu6n8yb4dP82C0QQxQ4lGGAo9i9gNYU5+DnqO0w6eTUpL9Mm6iPgIWsgRCRpJCsclg/sBXcDSAb+AeHyL+6MAcsUyyExMX8xAiqmK64pHyDXGykYog3Z/tfs7eH15t3wtPjkAxUO+hM0FJYQ7BbDHpgd4htKHn0fyRz4GgYZihUZED4LFQYe+tHmOdvK4MznK+wi8VP2m/gi+Df5of11BCL/aOxa47rqoe+b74n1dP05+2z0Fu5i7B7xzPAp66jkz+Ra7J3vQe7E7uvnDuLM6hzwl/P1+AT06er86dnp/uXv4nbhe+Jo477m0+PT2o/Trs2B0hTiuu6J977/igSoBjYG5AJB/7L6YO/p2gLHEsQlyoPR/N3158zpBumy4X3YwNzr5sDsNfJJ93z5/fwFA4wDNflS8QLxdPK28lrxtu5u6OPjheaX6FbiTt7f4xrrrvBu8KboUOCq6Y/86ANbDVwZvxc4DUYCau3n2zvhuufh393UHNXt3jvpHvJq/VYLyRgdJBMo3ySBI/wj6B8ZHJIaBBkrFOwOKwd0+nb1fvd/8Z3pBO4m8eXu4/YACl8YGRvwGlIj/C0XKSMgXxswDef96fXZ8KP7ew8PGX0WxxCYDnUO7BGmEAsFkgDXEisluSfbJ28qwyYxHrsZ9hV1D3MI9/uy8On0UgCZ/nrzIPPz+H75kfgk+bL9LAaEDBcYBCklLPonyymgJbAXewmn+KrpUOVg3t3UH9hH4yLtjf/mEGMSGxHLE10TqA65DJoHm/OP3gTguufj5+vukfZL83Tx/PBU5Sna09Z/0GzNYtoK7xECqgxZDZIN4g9pEOoMrQNU9Mblm9cGyuPBzcTtzUzUkd+X82D/p/3Q/On2zOmw6X/2M/78B5gSbQ5KCSUNJwlEAiwEU/+B7FzcXNuk2lTbHOo58+3xbvrx/QL8aQ3qH1IgXB6FH8EXVgv8AaP3O+344OvSvM3C2PPnFPLZ/BEF2gotEVQVcRYJGMEhxydUH4oaPhmSEBkKNA0jEfgTqhfPFzIRFQb1+t312/MO5+vXIdz+9OgIgRa0KsM4FzheNK4zHS81KbIiohV7Bp/6kemZ4T/yAAXwBzAHFwopCGkIShDDE4oQ0RZWJgozqjs5OzsyYyzsI6QUhAeh+Uvom+HX4knnR/mQB/QKlgnBBtgF6AyuGEwV2QsVFocl4iScIzglzSCFGqwQzP7C7cTmXOfp58jsIvfx95fxcPdhBrcQABQhEEwRqB3fIi8ciBiaFN4J2gERAnL+0Pi69wju2+VW6PHhddWs1/zmp/dKC+AcNh/TGHcWThdWFIoM1gDj7zPdBNRs0UjMQ9DZ4IPt5fRL+zv35/A19RjzFO8nAEgUShmiGk4gpB88FpoOmAVq+vvyP+0t6iDpuuO45aH4LQevCiEMjgmUABj3ZPFg8iT7awCtCAkY6hyKFxcYxxaQCGT8IO7f1z3NsM+40XXXSenDANwQQBmsIZAiEx5cIRMk3x2zF+oSNgj5+iTzavGD8HbwKfEv7EPnXOli63jkldnE2CXeh90W4ffwkAA4CoYQExU+F6wTfwx3CbEI1gCy+Sv1Xuk35M/r+fCL8brw+ert5c7rEvLX7tPyzv3mAhELphlAIYEZCxSQFKINggSu/RryReUM6XrvYuHC0fHPxc3yz8LaZODf7tgFGxffIn0nSiP8Gw8YOhDBAnrwK9130sLVQ+Ro9GT9SAB1AFP8gfg/+Nv1Xu5L517vNAZdENwMhRP4HYcZIRFYDUQDbPst+vv21fm6/Cf0su8/+7UGDQm9CYIH0wvNGRcgUh8tHwAc7BaIEVIMcwivBv3+/uwn5jHrPeaD3B7ald/+7X0E8hYRHjMiwyN9JegmUiMjHfwUJQe29dfsHOm24z3fh+Gj8CwGGQ4pBdn+JPo37jvnKem66S3uVf9vEEQTxQ9IDlkHbP5F+ZvxwOjf4indyt3A6l77gQZIA4D+lAA5/+f79/nT8+Px4AH+FIwb4htrGscRIQsPCoX+Q+wx3QjVXt2070YADQv8Bv36cvhw8crhR+Q79JP8ighxHX0j4BwvHp4etRZaEZwK7fqh6EPa19jz4lTp4e6b+6YLOBj6G1QUyQTA+oH6yvnv953+aQevDSkcTiZxH0IboB3sGTwbIR+SEmr+1fIO7U/q6fEY+4wAsQ9tJ6w3ljd9Me4tniPbFBcOzQ3aC0wGaPz+7I3nUfNN/UgDswfFB6YLVg5CB6YCvQwfHeQsoDqiPRkzMSklJiEbsw1dCvwJuQiIDFoUAhZ7FCcRogdV/Rj17+3b7KbvffXRCQIg4iGsHE4cdxX0EB8UEwzV+mjqL99o4ArxfwIHBFACMgu3EgsRoBBEFscZvRrgG5Ab6hhSD+oHggW8/XL1Qe1F2zH
" Your browser does not support the audio element.\n",
" </audio>\n",
" "
],
"text/plain": [
"<IPython.lib.display.Audio object>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"# Load two clips and mix them\n",
"_, dat = scipy.io.wavfile.read(\"turn_on_the_office_lights_test_clip.wav\")\n",
"_, dat_music = scipy.io.wavfile.read(\"fma_sample/000182.wav\")\n",
"dat[-20*16000:] = (dat[-20*16000:] + dat_music[0:20*16000]*.7)/2 #quick manual mixing\n",
"\n",
"ipd.display(ipd.Audio(dat[-16000*6:], rate=16000, normalize=True, autoplay=False))"
]
},
{
"cell_type": "code",
"execution_count": 20,
"id": "344383c7",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:40:23.392725Z",
"start_time": "2023-02-18T03:40:22.770119Z"
}
},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"100%|██████████| 394/394 [00:00<00:00, 19051.83it/s]\n"
]
},
{
"data": {
2023-10-13 09:18:39 +02:00
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAiMAAAGiCAYAAAA1LsZRAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjUuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8qNh9FAAAACXBIWXMAAA9hAAAPYQGoP6dpAAApJ0lEQVR4nO3df3RV1Z338c+5N8kNRRMVMD8UYrTSaqPMNKk2Uez4o3HQOuNqpzLjPAUtuMwoIkRdFllT1OWaOD5rKP0F2hF0XGMr04qVeaRK+oz8UOwzNSYWgcfyFDTRJmSCYxJAE3Lvfv5Izs39GbiBnHPuue/XWndBzj2X7M1B81l7f/feljHGCAAAwCUBtxsAAAByG2EEAAC4ijACAABcRRgBAACuIowAAABXEUYAAICrCCMAAMBVhBEAAOAqwggAAHAVYQQAALgq4zCybds23XDDDSovL5dlWfrlL395zM9s3bpV1dXVKiws1LnnnqvHH398PG0FAAA+lHEYOXz4sGbNmqUf/ehHx3X//v37dd1112n27NlqbW3VAw88oMWLF+v555/PuLEAAMB/rBM5KM+yLL3wwgu68cYb095z//33a+PGjdqzZ0/0WkNDg95++2298cYb4/3WAADAJ/Im+hu88cYbqq+vj7t27bXXau3atTp69Kjy8/OTPjMwMKCBgYHo15FIRB999JGmTJkiy7ImuskAAOAkMMaov79f5eXlCgTST8ZMeBjp6upSSUlJ3LWSkhINDQ2pp6dHZWVlSZ9pamrSQw89NNFNAwAADujo6NDZZ5+d9v0JDyOSkkYz7JmhdKMcy5YtU2NjY/Tr3t5ezZgxQx0dHSoqKpq4hgIAgJOmr69P06dP16mnnjrmfRMeRkpLS9XV1RV3rbu7W3l5eZoyZUrKz4RCIYVCoaTrRUVFhBEAALLMsUosJnyfkdraWjU3N8dd27x5s2pqalLWiwAAgNyScRg5dOiQ2tra1NbWJml46W5bW5va29slDU+xzJs3L3p/Q0OD3n//fTU2NmrPnj1at26d1q5dq3vvvffk9AAAAGS1jKdp3nzzTV155ZXRr+3ajvnz5+vpp59WZ2dnNJhIUmVlpTZt2qSlS5fqxz/+scrLy/WDH/xA3/jGN05C8wEAQLY7oX1GnNLX16fi4mL19vZSMwIAQJY43p/fnE0DAABcRRgBAACuIowAAABXEUYAAICrCCMAAMBVhBEAAOAqwggAAHAVYQQAALiKMAIAAFxFGAEAAK4ijAAAAFcRRgAAgKsIIwAAwFWEEQAA4CrCCAAAcBVhBAAAuIowAgAAXEUYAQAAriKMAAAAVxFGAACAqwgjAADAVYQRAADgKsIIAABwFWEEAAC4ijACAABcRRgBAACuIowAAABXEUYAAICrCCMAAMBVhBEAAOAqwggAAHAVYQQAALiKMAIAAFxFGAEAAK4ijAAAAFcRRgAAgKsIIwAAwFWEEQAA4CrCCAAAcBVhBAAAuIowAgAAXEUYAQAAriKMAAAAVxFGAACAqwgjAADAVYQRAADgKsIIAABwFWEEAAC4ijACAABcRRgBAACuIowAAABXEUYAAICrCCMAAMBVhBEAAOAqwggAAHAVYQQAALiKMAIAAFxFGAEAAK4ijAAAAFcRRgAAgKsIIwAAwFWEEQAA4CrCCAAAHhSJGBlj3G6GI8YVRlavXq3KykoVFhaqurpa27dvH/P+Z599VrNmzdJnPvMZlZWV6dZbb9XBgwfH1WAAAPxuYCisa763Vbc90+J2UxyRcRhZv369lixZouXLl6u1tVWzZ8/WnDlz1N7envL+1157TfPmzdOCBQu0a9cu/fznP9dvf/tbLVy48IQbDwCAH+34fwe1778O69d7DrjdFEdkHEZWrlypBQsWaOHChbrgggu0atUqTZ8+XWvWrEl5/29+8xudc845Wrx4sSorK3X55Zfr9ttv15tvvnnCjQcAwI8ODQy53QRHZRRGBgcH1dLSovr6+rjr9fX12rFjR8rP1NXV6YMPPtCmTZtkjNGBAwf0i1/8Qtdff33a7zMwMKC+vr64FwAAueIwYSS9np4ehcNhlZSUxF0vKSlRV1dXys/U1dXp2Wef1dy5c1VQUKDS0lKddtpp+uEPf5j2+zQ1Nam4uDj6mj59eibNBAAgqx0eDLvdBEeNq4DVsqy4r40xSddsu3fv1uLFi/Xd735XLS0tevnll7V//341NDSk/fOXLVum3t7e6Kujo2M8zQQAICsdybGRkbxMbp46daqCwWDSKEh3d3fSaImtqalJl112me677z5J0sUXX6zJkydr9uzZeuSRR1RWVpb0mVAopFAolEnTAADwjUODuRVGMhoZKSgoUHV1tZqbm+OuNzc3q66uLuVnjhw5okAg/tsEg0FJypn10wAAZOLIANM0Y2psbNSTTz6pdevWac+ePVq6dKna29uj0y7Lli3TvHnzovffcMMN2rBhg9asWaN9+/bp9ddf1+LFi3XJJZeovLz85PUEAACfOJxjIyMZTdNI0ty5c3Xw4EE9/PDD6uzsVFVVlTZt2qSKigpJUmdnZ9yeI7fccov6+/v1ox/9SPfcc49OO+00XXXVVfrHf/zHk9cLAAB8JNdW01gmC+ZK+vr6VFxcrN7eXhUVFbndHAAAJtT/ePL/6LX/1yNJeu/R9FtheN3x/vzmbBoAADwmdpomC8YMThhhBAAAj4ktYI34P4sQRgAA8JrY7eAZGQEAAI47EjtN42I7nEIYAQDAYw7HTdP4P44QRgAA8JjBcCT6+xzIIoQRAAC8ZHAocuybfIYwAgCAhyRueMY0DQAAcNShhDCSA1mEMAIAgJf0f5oQRlxqh5MIIwAAeEjyyIj/4whhBAAADzkymFgz4lJDHEQYAQDAQ5IKVgkjAADASeGElb0mB9IIYQQAAA8JJ8zLME0DAAAclThNQwErAABwVOLIiP+jCGEEAABPSR4ZcakhDiKMAADgIUkjIzmQRggjAAB4CNM0AADAVUzTAAAAVyXuM8KpvQAAwFHhxJERl9rhJMIIAAAeEqGAFQAAuCl5NY1LDXEQYQQAAA+hgBUAALgqeWmv/9MIYQQAAA9JPBiPkREAAOCoxGkalvYCAABHsQMrAABwFatpAACAq5JX0/g/jRBGAADwEKZpAACAq5K2g8+BNEIYAQDAQ5K2g8+BsRHCCAAAHpJ0am8k9X1+QhgBAMBDkgpYGRkBAABOYmkvAABwFQWsAADAVRSwAgAAVzFNAwAAXJU4TcNBeQAAwFGJ2cP/UYQwAgCApzBNAwAAXJU4TZMLYyOEEQAAPCRxNU3E/1mEMAIAgJcwTQMAAFyVtB18DqQRwggAAB6SODLCNA0AAHBUOGlpr//TCGEEAAAPSSxgzYEsQhgBAMBLkgpYXWqHkwgjAAB4CNvBAwAAVyWd2uv/LEIYAQDASxJHRnIgixBGAADwkuQdWP0fRwgjAAB4SNLZNP7PIoQRAAC8JBKJ/5p9RgAAgKMSp2USw4kfEUYAAPAQ9hkBAACuSlpNQwErAABwUtI+Iy61w0mEEQAAPMQeGQkGLEmMjKS1evVqVVZWqrCwUNXV1dq+ffuY9w8MDGj58uWqqKhQKBTSeeedp3Xr1o2rwQAA+JldsDoaRlxsjEPyMv3A+vXrtWTJEq1evVqXXXaZnnjiCc2ZM0e7d+/WjBkzUn7mpptu0oEDB7R27Vp99rOfVXd3t4aGhk648QAA+I1dwJoXsDSo3JimyTiMrFy5UgsWLNDChQslSatWrdIrr7yiNWvWqKmpKen+l19+WVu3btW+fft0xhlnSJLOOeecE2s1AAA+ZU/T5I2MjLADa4LBwUG1tLSovr4+7np9fb127NiR8jMbN25UTU2NHnvsMZ111lmaOXOm7r33Xn3yySdpv8/AwID6+vriXgAA5AK7gDUvOPwjOgeySGYjIz09PQqHwyopKYm7XlJSoq6urpSf2bdvn1577TUVFhbqhRdeUE9Pj+644w599NFHaetGmpqa9NBDD2XSNAAAfCGpgNXNxjhkXAWslmXFfW2MSbpmi0QisixLzz77rC655BJdd91
2023-02-17 22:56:15 -05:00
"text/plain": [
"<Figure size 640x480 with 1 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"# Pre-compute audio features using helper function\n",
"features = F._get_embeddings(dat)\n",
"\n",
"# Get predictions for each window\n",
"scores = []\n",
"for i in tqdm(range(0, features.shape[0]-28)): # 28 is the number of timestep frames for this model\n",
" window = features[i:i+28][None,]\n",
" with torch.no_grad():\n",
" scores.append(float(fcn(torch.from_numpy(window)).detach().numpy()))\n",
" \n",
"plt.figure()\n",
"_ = plt.plot(scores)\n",
"_ = plt.ylim(0,1)"
]
},
{
"cell_type": "markdown",
"id": "078a1bc2",
"metadata": {},
"source": [
"The model is now less confidant in it's prediction than before, but the score is still above a default score of 0.5 which confirms that the model at least represents a good starting point."
]
},
{
"cell_type": "markdown",
"id": "0367ee93",
"metadata": {},
"source": [
"# Export the Model"
]
},
{
"cell_type": "markdown",
"id": "f4eb66ce",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-12T20:00:58.877229Z",
"start_time": "2023-02-12T20:00:58.868037Z"
}
},
"source": [
"Now that the model is trained and passes basic performance validation tests, it can be exported to ONNX so that it can be used by the openWakeWord inference engine. With Torch, this process is quite simple."
]
},
{
"cell_type": "code",
"execution_count": 21,
"id": "75acc3bb",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:40:40.474020Z",
"start_time": "2023-02-18T03:40:40.401258Z"
}
},
"outputs": [],
"source": [
"# Export model to ONNX format\n",
"\n",
"output_path = \"turn_on_the_office_lights.onnx\"\n",
"torch.onnx.export(fcn, args=torch.zeros((1, 28, 96)), f=output_path) # the 'args' is the shape of a single example"
]
},
{
"cell_type": "markdown",
"id": "03bd9f1e",
"metadata": {},
"source": [
"# Evaluate the Model"
]
},
{
"cell_type": "markdown",
"id": "8c671f55",
"metadata": {},
"source": [
"Let's now load in the ONNX model with openWakeWord, and use that to run some more rigorous testing. First, let's just confirm that the ONNX model works as expected."
]
},
{
"cell_type": "code",
"execution_count": 22,
"id": "295fd55c",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:41:08.824352Z",
"start_time": "2023-02-18T03:41:08.615724Z"
}
},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"/home/dscripka/anaconda3/envs/torch_gpu/lib/python3.9/site-packages/onnxruntime/capi/onnxruntime_inference_collection.py:54: UserWarning: Specified provider 'CUDAExecutionProvider' is not in available provider names.Available providers: 'CPUExecutionProvider'\n",
" warnings.warn(\n"
]
}
],
"source": [
"# Create openWakeWord instance\n",
"\n",
"oww = openwakeword.Model(\n",
" wakeword_model_paths=[\"turn_on_the_office_lights.onnx\"],\n",
" enable_speex_noise_suppression=True,\n",
" vad_threshold=0.5\n",
")\n"
]
},
{
"cell_type": "code",
"execution_count": 23,
"id": "6f009d79",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:41:13.924834Z",
"start_time": "2023-02-18T03:41:12.995975Z"
}
},
"outputs": [
{
"data": {
2023-10-13 09:18:39 +02:00
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAiMAAAGdCAYAAADAAnMpAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjUuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8qNh9FAAAACXBIWXMAAA9hAAAPYQGoP6dpAAAw3UlEQVR4nO3df3BU9b3/8dduNj/4lSAggUDAWH/RcqU1aAsWW380Fq39tuMMfGun2BZmmqIg5Np+RWaqZTqN7bQM2gq0V5BxxquMF+y1t6mS3lZAsVMJSUWhaiuSAAkxCEn4lR+7n+8fm3M2mx+754Sc7IbzfMzsJJz97OazORvyyufzPp9PwBhjBAAAkCLBVHcAAAD4G2EEAACkFGEEAACkFGEEAACkFGEEAACkFGEEAACkFGEEAACkFGEEAACkVCjVHXAiEono2LFjGjNmjAKBQKq7AwAAHDDGqLW1VQUFBQoG+x//GBZh5NixYyosLEx1NwAAwADU1dVp6tSp/d4/LMLImDFjJEVfTG5ubop7AwAAnGhpaVFhYaH9e7w/wyKMWFMzubm5hBEAAIaZZCUWFLACAICUIowAAICUIowAAICUIowAAICUch1Gdu3apbvuuksFBQUKBAL63e9+l/QxO3fuVHFxsXJycnT55Zdr48aNA+krAAC4CLkOI2fOnNGsWbP061//2lH7Q4cO6Y477tC8efNUXV2thx9+WMuXL9e2bdtcdxYAAFx8XF/aO3/+fM2fP99x+40bN2ratGlat26dJGnGjBnau3evfvGLX+juu+92++UBAMBFxvOakTfeeEMlJSVxx26//Xbt3btXHR0dfT6mra1NLS0tcTcAAHBx8jyMNDQ0KD8/P+5Yfn6+Ojs71dTU1OdjysvLlZeXZ99YCh4AgIvXkFxN03PlNWNMn8ctq1atUnNzs32rq6vzvI8AACA1PF8OftKkSWpoaIg71tjYqFAopPHjx/f5mOzsbGVnZ3vdNQAAkAY8HxmZM2eOKisr447t2LFDs2fPVmZmptdfHgAApDnXYeT06dOqqalRTU2NpOiluzU1NaqtrZUUnWJZtGiR3b60tFSHDx9WWVmZDh48qM2bN2vTpk168MEHB+cVAACQxowx+mfjaXWEI73uaz7bod/u+pfqPj6bgp6lD9fTNHv37tXNN99s/7usrEySdO+992rLli2qr6+3g4kkFRUVqaKiQitXrtSTTz6pgoICPfHEE1zWCwC46IUjRg9te0svVB3RrKl5embxZ5U3Ijor0BGO6PM/+7Na2zr14Ymz+unX/y3FvU2dgLGqSdNYS0uL8vLy1NzcrNzc3FR3BwAAR/78j+P67pa99r+tQJIdCmr5c9XaceC4JOm6aWO1femNqeqmZ5z+/va8gBUAAL9qOt0uSZo4Jlsd4Yj+fqRZt63dqayMoI6eOme3m3LJyFR1MS2wUR4AAB77VEGunl3yOU0ZO0Iftbbp6KlzGj8qS1/+1CRJUiT9Jyk8xcgIAABe6coYgUBAnyzI1Y6VN+kP++tljNGtM/JVsb9eL7/TYLfzK8IIAAAeMV0pw1ric1R2SAtmx1YVtxb/9PvICNM0AAB4xNgjI33fbx0mjAAAgJQIdqUUn2cRwggAAF6JZYy+h0aCXYcjhBEAAOCFpNM0Aaudv9MIYQQAAI/0LGDtiQLWKMIIAAAeSZYx7JqRIehLOiOMAADgseRX0wxZV9ISYQQAAI9YGSPQXwFr129hakYAAIA3ukJGfyMjXNobRRgBAMAjyTIGBaxRhBEAADzGCqyJEUYAAPCIvc5Iv4ueMU0jEUYAAPCMiaWRPgXtRc+Gpj/pijACAIBHYlfT9C1gLwfv7zRCGAEAwCPJMkaARc8kEUYAAPBcoJ8K1iBX00gijAAA4Jmk0zRdH1mBFQAAeMIkW/QsaDccmg6lKcIIAAApElv0LMUdSTHCCAAAHks+TePvNEIYAQDAI/YyI0kKWH2eRQgjAAB4xXSVsPY3MsLVNFGEEQAAPGKSXE4TYAVWSYQRAAA8k3zXXqudv9MIYQQAAI8l2yiPq2kAAIAnYgWsfd/P1TRRhBEAADyStIA1yNU0EmEEAADPJAsZQbuA1d9phDACAIDH+pumscZMqBkBAACe6r+ANfqRmhEAAOCJpBvlsQKrJMIIAACeSXo1DTUjkggjAAB4JlnEYJ2RKMIIAACe63tohBVYowgjAAB4JPmiZ4yMSIQRAAA8k3zRs6521IwAAAAvJF/0jKtpJMIIAACeY2+axAgjAAB4xIoY/S16FuBqGkmEEQAAvJN00TOrmb/TCGEEAACPJIsYAWpGJBFGAADwjH1pbz/3szdNFGEEAACPBfqZp7GvphnKzqQhwggAAB5xurIqIyMAAMATyVZgDQa5mkYijAAA4JnkG+V1tWNkBAAAeKm/dUZYgTWKMAIAgEeSb5QXRc0IAADwRLKN8liBNYowAgCAV5JulNetqY9HRwgjAAB4rN9pmm53+DiLEEYAAPCKvVFev4uexT73c90IYQQAAI9YUy/JakYkf6/COqAwsn79ehUVFSknJ0fFxcXavXt3wvbPPvusZs2apZEjR2ry5Mn6zne+oxMnTgyowwAADBf2YEe/0zSxzxkZcWHr1q1asWKFVq9ererqas2bN0/z589XbW1tn+1fe+01LVq0SIsXL9Y777yjF154QW+++aaWLFlywZ0HACCdJV/0jJoRaQBhZO3atVq8eLGWLFmiGTNmaN26dSosLNSGDRv6bP/Xv/5Vl112mZYvX66ioiJ9/vOf1/e+9z3t3bv3gjsPAMBw0P+iZ7HPCSMOtbe3q6qqSiUlJXHHS0pKtGfPnj4fM3fuXB05ckQVFRUyxuj48eP6r//6L9155539fp22tja1tLTE3QAAGG6SL3oWu4NpGoeampoUDoeVn58fdzw/P18NDQ19Pmbu3Ll69tlntXDhQmVlZWnSpEkaO3asfvWrX/X7dcrLy5WXl2ffCgsL3XQTAIC0kHzRs9jnhBGXel6iZIzp97KlAwcOaPny5frRj36kqqoqvfzyyzp06JBKS0v7ff5Vq1apubnZvtXV1Q2kmwAApFSyfBHkahpJUshN4wkTJigjI6PXKEhjY2Ov0RJLeXm5brzxRv3gBz+QJF177bUaNWqU5s2bp5/85CeaPHlyr8dkZ2crOzvbTdcAAEhb/S96FvvcRIamL+nI1chIVlaWiouLVVlZGXe8srJSc+fO7fMxZ8+eVTAY/2UyMjIk+XvpWwCAfyTbtVdimsaVsrIyPfXUU9q8ebMOHjyolStXqra21p52WbVqlRYtWmS3v+uuu7R9+3Zt2LBBH3zwgV5//XUtX75cN9xwgwoKCgbvlQAAkGbsRc/6GRmJu5pmCPqTrlxN00jSwoULdeLECa1Zs0b19fWaOXOmKioqNH36dElSfX193Joj3/72t9Xa2qpf//rX+vd//3eNHTtWt9xyi372s58N3qsAACANJVnzLK7e0s8jI67DiCQtXbpUS5cu7fO+LVu29Dq2bNkyLVu2bCBfCgCAYctJvggEou38HEbYmwYAAK/1N0+jbnUj/s0ihBEAALySbJ2R7vdFCCMAAGCwJVuBVYqNjDBNAwAABp2TeGEFFcIIAADwTH/rjEixMOLjLEIYAQDAK26maQgjAADAA8kLWKkZIYwAAOAZR+uMWG097Ul6I4wAAOARJ9M0FLASRgAA8Fwg0aJnQatmhDACAAAGmXEw+WJP0/g3ixBGAADwirtFz4agQ2mKMAIAgEecLXrG1TSEEQAAPMaiZ4kRRgAA8IizaZroR0ZGAADAoHOyay8rsBJGAADwjqtFz/ybRggjAAB4xIoXiRc942oawggAAB5LVMAa7PpNTM0IAAAYdNaqquzamxhhBAAAjzhaZ8Rq6+M0QhgBAMAjTvIFK7ASRgAA8FyijfJii575N40QRgAA8Ih9NU2CNlxNQxgBAMAzzgpY49v6EWEEAACPOIkX1IwQRgAA8I61N42jpv5NI4QRAAA8lqiAlZERwggAAJ6xN8pLVDPCCqyEEQAAvOIkX9hLxfs3ixBGAADwinFQM2J
2023-02-17 22:56:15 -05:00
"text/plain": [
"<Figure size 640x480 with 1 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"# Do a quick test prediction on the test clip to confirm that the behavior is still as expected\n",
"\n",
"scores = oww.predict_clip(\"turn_on_the_office_lights_test_clip.wav\")\n",
"\n",
"plt.figure()\n",
"_ = plt.plot([i[\"turn_on_the_office_lights\"] for i in scores])"
]
},
{
"cell_type": "markdown",
"id": "d26f14af",
"metadata": {},
"source": [
"Since that looks fine, we can now conduct a more rigorous test to evaluate the false-accept rate in something closer to a production scenario. Specifically, we want the openWakeWord models to respond consistently when a user speaks the target wake word/phrase, but also does not activate even in the presence of many hours of continuous background noise and un-related speech.\n",
"\n",
"To test that, we'll use a few clips (for a total of ~ 1 hour) from the [Santa Barbara Corpus of Spoken American English](https://www.linguistics.ucsb.edu/research/santa-barbara-corpus) to produce a more realistic metric for the false-activation rate per hour.\n",
"\n",
"The combined clip (already converted to a single-channel, 16khz, 16-bit WAV file) can be downloaded [here](https://f002.backblazeb2.com/file/openwakeword-resources/data/santa_barbara_corpus_test_clip.wav).\n"
]
},
{
"cell_type": "code",
"execution_count": 24,
"id": "756d7100",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:43:13.392821Z",
"start_time": "2023-02-18T03:41:57.091857Z"
}
},
"outputs": [
{
"data": {
2023-10-13 09:18:39 +02:00
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAiMAAAGiCAYAAAA1LsZRAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjUuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8qNh9FAAAACXBIWXMAAA9hAAAPYQGoP6dpAAA9bElEQVR4nO3de3hU1aH//0+4BY/FqCABFCltj61t1NpQNVhs1ZaWIj3+6u/IqX0AW+hXSpGDaL8WbQtaj7H1SNEqeEGktqjYgooVhaDcAyoh3AOCXBIgISRAEi65r+8fmJDLTDKXPfv6fj1PfGRmz8xae+9Z67PXXntPkjHGCAAAwCEdnC4AAAAINsIIAABwFGEEAAA4ijACAAAcRRgBAACOIowAAABHEUYAAICjCCMAAMBRhBEAAOAowggAAHBU1GFk5cqVGjZsmPr06aOkpCS9+eab7b5mxYoVSk9PV9euXfWFL3xBzz77bCxlBQAAPhR1GDl58qSuuuoqPf300xEtv3fvXv3whz/UoEGDlJubqwceeEATJkzQ/Pnzoy4sAADwn6R4figvKSlJb7zxhm699dawy9x///1auHCh8vLyGh8bO3asNm3apLVr18b60QAAwCc6JfoD1q5dq8GDBzd77Pvf/75efPFF1dTUqHPnzq1eU1VVpaqqqsZ/19fX6+jRo+revbuSkpISXWQAAGABY4wqKirUp08fdegQ/mRMwsNIUVGRUlNTmz2Wmpqq2tpalZSUqHfv3q1ek5mZqYceeijRRQMAADYoKCjQJZdcEvb5hIcRSa1GMxrODIUb5Zg8ebImTZrU+O+ysjJdeumlKigo0HnnnZe4ggIAAMuUl5erb9++6tatW5vLJTyM9OrVS0VFRc0eKy4uVqdOndS9e/eQr0lOTlZycnKrx8877zzCCAAAHtPeFIuE32ckIyNDWVlZzR5bsmSJBgwYEHK+CAAACJaow8iJEye0ceNGbdy4UdKZS3c3btyo/Px8SWdOsYwcObJx+bFjx2r//v2aNGmS8vLyNHv2bL344ou67777rKkBgJgt2lKonP3HnC6GrxWWnVZ1bb3TxQBcLeowsn79el199dW6+uqrJUmTJk3S1Vdfrd///veSpMLCwsZgIkn9+/fXokWLtHz5cn3961/XH/7wBz311FO67bbbLKoCgFjsKCrXuLkbdNvMbKeLEpWjJ6sVxx0JbLXtUJkyMj/Q7c/ZexuD4vJK1dd7Yx0BUgxzRr7zne+02RDMmTOn1WPf/va3tWHDhmg/CkAC5ZeecroIUVu0pVDj5m7QnQM/r6k/+prTxWnX/JyDkqSNBcdt+8zlO4t150sf6wdf66VnR6Tb9rlAPPhtGgCe8eiiMzdPnJO9z9mCRKiDA7dFen7lHknSe9uK2lkScA/CCICw9peeVHF5pdPFAOBzhBH4wp4jJ1RT579JggVHT2nG8t0qO11j+2cfO1mtbz++XNc8+r7tn+0X5ZX2bzfAi2y56RmQSP/afEjjX8nV9V/qrrljrnO6OJb60dOrdexUjfIKK/SXn1xt62fvKTlp6+f5zUd7j+r19QecLgbgCYyMwPP++tn8gTW7S50tSAIcO3XmyHrtpyUOlwTR+ssHu5wuAuAZhBEAjjhdXaefvfSR5n643+miAHAYYQSAI/66dp+W7TyiB9/Y6nRRADiMMALAERVM7gTwGcII4IDD5ZXKzec27IAb1dbV6z+eXq2Jr+U6Vob9pSf1zLLdgbkiizACOODaR9/X/zcjW1sOlDldFAAtrN9/TJsOlOnNjYccK8OQJ1fp8cU7NXXhNsfKYCfCCOCgj/cddboIAFqod8FvH52qrpN05hLxICCMIGFq6+r13IpPtfnAcaeLAgBwMcIIEuaVj/KV+e4O/ejpNU4XBQ7bXXxCs1btUWVNnWXvmbkoTzc/sVwnqmote08AzuAOrEiYvMIKp4sAl/jutBWSpBNVtZr43cssec/nPvtBuNc+yteYQV+w5D0BOIOREQC2yc0/bvl7uuD0PoA4EUYAAICjCCMA4COMFMGLCCOAj0xf+onu+tt61dXTIwHwDsKIR3x65ISqa+udLgZcbvrSXVq87bBW7+ZXfgF4B2HEAxZtKdTNT6zQiBc/dLoo8IgqCy+hBYBEI4x4wMtr90mSPgzInfjQmhPzAJKS7P9MAMFEGAEAAI4ijAAAAEcRRgAAgKMIIwAAwFGEEQAA4CjCCCBpb8lJ7Szih/0AwAn8ai88L97LXo0xuvF/l0uSNk0ZrJRzOsdfKABAxBgZQeA1vXP6kYpK5wqCdnHvE8CfCCOAg4L2CzJN65skkgXciX3TfoQRAADgKMIIEoYhdQBAJAgj8DxCD3CWCdzJP/gBYQQAADiKMAIAABxFGAECKonzWwBcgjDiAVxmhkQw7dwtzit7HXMk4Gfx3tTRKwgjgAcEpD0CEFCEEQAA4CjCCADPCMqQNRA0hBEAAOAowggAAHAUYQQAADiKMAIAABxFGAF8iHmeQOy4d439CCOAg7xyYzGrtHejNQDBRBhB4DUNBPSVAGA/wggAAC4Vy09I1dcb/TPngHYXn7C+QAnSyekCAEHGQAysxuge3t58SPf9Y5Mkad9jQx0uTWQYGQEAoAmv/zhpbv5xp4sQNcIIEsaurzNHggDgbYQRD4jlnCHgduzXABoQRgCElERaAGATwgiAkLgnCAC7EEaAAKmsqSNkAHAdwggQELuLK/SV372ne1/f5HRRAKAZwggQo7kf7tdLa/Y6XYyIzVp1pqwLcg86XJLYMY0F8CduegbEoKq2Tg++sVWSNOyqPurxuWQZY5j0CQAxYGQEiEFd/dl5F5U1dXpm2W5983+WquDoKQdLBQDeRBiB57lhMOLxxTtVcqJaj723w+miBA7zcQHvI4wAgI+4IZwD0SKMAAAARxFGAACAo2IKIzNmzFD//v3VtWtXpaena9WqVW0uP3fuXF111VX6t3/7N/Xu3Vs/+9nPVFpaGlOBAfgDcz2A9gXlexJ1GJk3b54mTpyoBx98ULm5uRo0aJCGDBmi/Pz8kMuvXr1aI0eO1OjRo7Vt2zb94x//0Mcff6wxY8bEXXjAagH53gOAq0QdRqZNm6bRo0drzJgxuvzyyzV9+nT17dtXM2fODLn8unXr9PnPf14TJkxQ//799a1vfUt33XWX1q9fH3fhAQDNBeVIGv4SVRiprq5WTk6OBg8e3OzxwYMHKzs7O+RrBg4cqAMHDmjRokUyxujw4cP65z//qaFDh4b9nKqqKpWXlzf7AwAA/hRVGCkpKVFdXZ1SU1ObPZ6amqqioqKQrxk4cKDmzp2r4cOHq0uXLurVq5fOP/98/eUvfwn7OZmZmUpJSWn869u3bzTF9B0u1QMA+FlME1hb3vK6rdtgb9++XRMmTNDvf/975eTk6L333tPevXs1duzYsO8/efJklZWVNf4VFBTEUkw4jBDVPjf/gi63tkdQsevbL6rfpunRo4c6duzYahSkuLi41WhJg8zMTF1//fX69a9/LUm68sorde6552rQoEF65JFH1Lt371avSU5OVnJycjRFAxBQdByA90U1MtKlSxelp6crKyur2eNZWVkaOHBgyNecOnVKHTo0/5iOHTtKcvdRIQAAsEfUp2kmTZqkWbNmafbs2crLy9M999yj/Pz8xtMukydP1siRIxuXHzZsmBYsWKCZM2dqz549WrNmjSZMmKBrrrlGffr0sa4mgEPsyNTRBndyPgAvieo0jSQNHz5cpaWlevjhh1VYWKi0tDQtWrRI/fr1kyQVFhY2u+fInXfeqYqKCj399NO69957df755+umm27SH//4R+tqAdgsSWHODbg4BLQ8ncGcEABuEXUYkaRx48Zp3LhxIZ+bM2dOq8fuvvtu3X333bF8FACfIgsBaMBv08DzOCUBAN5GGAF8KJJRByaQA3ALwogH0GfYhzMHALzOiwcahBFExRij3cUVqqmrd7ooAACfIIygXcYYLdpSqPzSU1qw4aC+O22lRv+VHzoEAFgjpqtpECzvbS3SuLkbJElpF58nSVr5yREniwQACePBsxyex8gI2vX
2023-02-17 22:56:15 -05:00
"text/plain": [
"<Figure size 640x480 with 1 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"# Estimate the false-accept rate on realistic test data (will take up to several minutes an a desktop-grade CPU)\n",
"\n",
"scores = oww.predict_clip(\"santa_barbara_corpus_test_clip.wav\")\n",
"\n",
"plt.figure()\n",
"_ = plt.plot([i[\"turn_on_the_office_lights\"] for i in scores])\n",
"_ = plt.ylim(0,1)"
]
},
{
"cell_type": "code",
"execution_count": 25,
"id": "b9e2f282",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:43:20.326550Z",
"start_time": "2023-02-18T03:43:20.309552Z"
}
},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"False-accept rate per hour: 94.0\n"
]
}
],
"source": [
"# Calculate the false-accept rate per hour from this result\n",
"\n",
"false_accepts = openwakeword.metrics.get_false_positives(\n",
" [i[\"turn_on_the_office_lights\"] for i in scores], threshold=0.5\n",
")\n",
"\n",
"print(f\"False-accept rate per hour: {false_accepts/1}\")"
]
},
{
"cell_type": "markdown",
"id": "04d294cf",
"metadata": {},
"source": [
"It looks like the false-accept rate for this model is very high, and would need to be reduced quite significantly to be viable for a production deployment. Of course, this was expected as the model was trained on a very small amount of positive and negative data."
]
},
{
"cell_type": "markdown",
"id": "c758890a",
"metadata": {},
"source": [
"# Create a User-specific Verifier Model"
]
},
{
"cell_type": "markdown",
"id": "362c6812",
"metadata": {},
"source": [
"As we saw, the simple model trained on this dataset performs quite well at detecting the presence of the wakeword/phrase, but often activates when it shouldn't, leading to an unnacceptably high false-accept rate.\n",
"\n",
"In practice, there are two ways to improve the performance of the model:\n",
"\n",
"1) Train on much larger amounts of positive and negative examples. The models released with openWakeWord are often trained on >100,000 positive examples, and over 30,000 hours of negative data.\n",
"\n",
"2) Create a user-specific \"verifier\" model based on examples of a specific person speaking the both the wake word/phrase and unrelated speech. The openWakeWord inference engine uses this verifier model to filter out likely false activations by focusing on only known speakers.\n",
"\n",
"We'll demonstrate the 2nd option here, as it's a very quick way to significantly improve performance at the cost of making the model far less likely to work well with other voices. The approach behind this verifier model are relatively simple, and are discussed in more detail in the openWakeWord documentation [here](https://github.com/dscripka/openWakeWord/docs/custom_verifier_models.md).\n",
"\n",
"For test data, we'll use 20 examples of the wake phrase generated with the [Tortoise](https://github.com/neonbjb/tortoise-tts) TTS model. While this is also synthetic data, it's very high quality and trained on different data than the TTS models used to generate training data. For unrelated speech, an English [phonetic pangram](https://www.liquisearch.com/list_of_pangrams/english_phonetic_pangrams) sentence was generated with the same TTS voice.\n",
"\n",
"This example wake phrase clips and reference negative speech (phonetic pangram) are included in the openWakeWord repo (in the `notebooks/training_tutorial_data` directory)."
]
},
{
"cell_type": "code",
"execution_count": 26,
"id": "f6b88046",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:45:18.538259Z",
"start_time": "2023-02-18T03:45:18.525190Z"
}
},
"outputs": [
{
"data": {
"text/html": [
"\n",
" <audio controls=\"controls\" >\n",
" <source src=\"data:audio/x-wav;base64,UklGRs7yAABXQVZFZm10IBAAAAABAAEAgD4AAAB9AAACABAAZGF0YaryAAD9//z/+f/3//j/9v/1//X/8v/0//P/8//z//T/9f/2//X/9P/0//L/8f/w//D/7//t/+3/6//s/+z/6//r/+r/6//r/+r/6v/p/+n/6f/p/+j/5//n/+n/5//m/+f/5//n/+j/5//o/+j/6P/o/+n/6f/q/+r/6//s/+z/7P/s/+z/7v/v/+//8P/y//P/9P/1//b/9f/2//j/+f/6//v//P/8//3//v////7///8AAAEAAgADAAQAAwACAAMAAgACAAMAAwACAAAA/v/+////AAACAAAA/v/+//z//v/+//7//v/+//7////+//3//f/8//z/+v/7//r/+//7//v/+//6//n/+v/8//z/+//6//r/+//7//z//P/9////AAABAAEAAgADAAMABAADAAMABAAEAAMAAwADAAMAAwABAAAAAgABAAAAAQAAAP//AQAAAP7//f////7//f/8//3//f/9//v/+P/2//b/8//z//X/9v/1//T/8//z//D/7f/u/+z/6//r/+v/7P/r/+n/6f/p/+r/6v/p/+j/6P/q/+r/6v/m/+X/5P/l/+b/5v/m/+P/4//i/+H/4P/e/97/3v/e/97/3//f/97/3P/c/9z/2//a/9n/2f/Y/9j/1P/T/9T/1P/T/9H/zf/O/87/y//K/8f/xP/F/8P/wP++/7z/uv+6/7j/uP+3/7P/sv+w/7D/sf+u/6v/qf+n/6f/pf+j/6L/of+g/57/nf+c/53/mf+W/5X/k/+R/5D/j/+N/4r/h/+F/4L/gf9+/3v/fP9+/37/ff98/3r/ev96/3n/dv9z/3H/cf9w/27/bv9r/2j/Z/9k/2X/Zv9m/2b/Zf9l/2T/Yv9i/2D/Xv9b/1z/W/9Z/1j/Vv9V/1T/Uv9R/1H/UP9Q/03/TP9G/0X/Rv9E/0X/Rf9E/0L/Q/9D/0P/Q/9H/0f/R/9H/0b/Rv9F/0D/Pv8//0D/Qf9A/z//QP9B/0T/Rv9C/0P/Rf9G/0n/S/9N/07/Tv9P/0//UP9S/1H/Uv9S/1b/WP9Y/1r/Wf9c/17/YP9f/2P/Zf9p/2z/bv9v/3L/dP91/3j/ff+D/4f/iP+L/4//lP+Y/5v/nv+i/6j/rf+u/6//s/+7/8H/xv/K/87/0P/T/9n/3P/i/+j/7f/v/+//8f/z//f/+v/9////AwAHAA0AFQAbAB4AIAAjACgALQAxADUANwA8AEMASABJAEwATwBTAFYAWwBfAGIAZABlAGoAbgBxAHUAeAB7AH0AfQB+AH8AgACFAIkAiwCPAJEAkgCTAJYAmACbAJ0AngCgAKEApAClAKkAqwCuALAArgCwALIAswC1ALQAswC1ALcAuAC6ALsAuwC+AL8AwADCAMQAxwDHAMcAxgDGAMYAxQDJAMsAywDKAMkAygDLAMoAywDJAMcAygDKAMwAzQDOAMwAzADNAMwAzgDLAM0AzgDPAM4AywDMAM0AzADLAMsAzQDLAMoAyQDGAMQAxADDAMMAxADCAMIAwQC/ALwAuQC5ALoAuAC1ALIArwCvAKwAqQCmAKYAogChAJ8AnACYAJYAkgCNAIoAhgCDAIUAggB7AHgAcQBsAGoAZwBhAFoAVABRAFMATwBLAEcAQwA8ADgANQAwACwAKQAkACAAHgAXABMAEQALAAUA/v/3//P/7//u/+z/6P/i/97/2v/a/9j/1v/V/9H/z//N/8r/xv/G/8X/xf/C/7z/uf+5/7j/tf+y/7L/sv+v/63/rv+q/6n/qv+t/6r/q/+r/6n/qP+k/6P/pP+n/6f/qP+r/6v/qv+q/6r/qv+q/6f/qf+t/7D/s/+y/7L/s/+w/7D/sf+y/7L/s/+z/7L/sv+z/7P/uf+5/73/v/+8/7n/v//D/8P/xP/E/8T/wv/D/8H/xf/G/8j/xv/D/8X/yv/O/9L/0f/Q/9H/0//V/9j/1//V/9f/2v/a/9r/2//e/97/4v/h/+L/5P/j/+D/4//q/+3/6//s/+//7//w/+7/8f/w/+//7//1//X/+f/8//7//v/+//v//f8AAP7/AgAFAAcABwAKAA0AEAAQAA8AEAANAAwAEQAQABAAEAANAA4ADQAPAA8AEgARABMAFQAVABMAEgARAA4ADQAPABEAFQAXABIAEQAQABAADwAMAAsADgAPABIAEAASABIAEgASABAADQANAA8ADQALAAgACAAIAAcABAAFAAUABQACAAAAAwAFAAMAAgD///7////9//z//f/7//3//v/6//X/9f/2//b/8v/t/+v/6v/n/+r/6//k/+P/4//f/+P/4v/f/+P/5f/k/97/3P/d/9r/2//b/9//3P/Z/9v/2P/X/9f/2f/b/9v/2//c/9r/2f/a/9r/2v/b/93/3v/e/97/4v/g/+H/4P/f/9z/3//g/+D/4v/i/+H/3//c/9r/3v/j/+f/4//h/9//4f/o/+j/6P/p/+b/5f/r/+r/6P/q/+j/5f/p/+f/5//o/+v/7P/r/+7/7P/u//T/8//r/+b/5v/s/+7/7//v//H/8P/u/+7/7f/q/+v/7f/t/+3/8P/w/+//7v/t/+7/8P/s/+f/7//v/+n/5v/m/+H/4P/k/+L/5P/l/+L/5f/j/+P/5f/p/+b/5f/o/+f/5f/l/+T/5f/l/+n/5//k/+b/4f/j/+b/4//i/+L/5f/j/9//2//c/9v/3P/f/93/3//i/+L/5P/j/93/3v/d/97/2//e/9//3v/e/9//3f/c/9z/3f/f/9//3v/d/9//2f/X/9r/2v/a/9j/1//X/9X/1P/X/9f/1f/W/9T/0v/U/9H/z//R/9D/zf/H/8r/zf/M/8z/yv/M/83/zP/P/87/yf/K/8z/zf/L/8n/yP/J/8f/x//J/8n/yf/D/8X/xf/I/8j/w//H/8L/v//A/7z/vv+9/7n/u//A/8P/w//A/7v/vv+7/7z/wf++/73/vf+4/7b/u/+4/7v/u/+3/7f/wP+3/7X/uv/B/8X/wP/F/7j/uf++/8L/wP+5/7z/v/+5/7L/tv+3/7T/tf+6/77/wf/E/8f/yP/D/8H/v/+7/73/v//E/8X/v//A/8P/wv/C/8P/xP/D/8X/w//E/8L/vP/B/8H/xv/A/8P/vf+9/7v/uf/H/7z/u//A/7//wf/E/7//vf/B/8f/xf/A/8X/xf/C/8v/x//C/8f/w//E/8P/wv/E/8f/wP+9/73/v//F/7//xv/C/7z/wf++/8H/wP++/8X/x//B/8P/w//C/8L/y//K/8b/y//M/8v/y//G/8j/yf/K/8j/y//M/8n/yf/K/9H/zP/Q/9D/z//P/9L/1P/O/9j/1v/U/87/0v/b/97/3f/b/+P/3//o/+X/3//m/+D/4v/o/+b/6v/p/+7/8f/t//X/9f/y/+r/8f/3//n/9v/z//b/9v/3//3//////wAAAQAGAAAA/f8LAAwACgAVABMACQAOABgAEwASAA4AGQAaAB0AHAAjACUAHgAkACgALAAfACIAHwAcAB4AHwAlACoAKQAoACoAJQAmACsALgAmACUAMQA1AC4ALAAyADoAOABAADcAOgA/ADUAPgBCAEoAQgBAAEIAOwBGAEkASwBCAEwASwBKAEsAOwA1ADcAVwBpAGEAWgBYAEwASwBOAEMAPgA2AC0ALQArACQAJgAkABkAEwAtAB8AGAAnABwADwAUAAUA9/8FAPL/8//r//D/6v/2//j/AgDk//T/5P/l/9P/zv/E/8D/1f/h/+T/3v/xAGgABv3P/ScE6AIj/SH+XAJZAVf/i/9XAcYDaAcsCvsDpABdA60FywEu/YH/7QFMANT+R/96APUBXgE8/x0AeQFJAMX+a//uANH+6f+eAs4B2gE5Au7+dAJcBDz/IP/cBKMBEgD1/hAGMQmi+Qb/lAgLBS36rP8qBZ4CzP4CAAsHDgK3/rMAqgPFBXb+1/54BZcFdf5a/
" Your browser does not support the audio element.\n",
" </audio>\n",
" "
],
"text/plain": [
"<IPython.lib.display.Audio object>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"# Provide paths to positive and negative speech from the target speaker for training a custom\n",
"# verifier model\n",
"\n",
"reference_clips = [str(i) for i in Path(\"training_tutorial_data/positive/\").glob(\"*.wav\")]\n",
"negative_clips = [str(i) for i in Path(\"training_tutorial_data/negative/\").glob(\"*.wav\")]\n",
"\n",
"# Listen to one of the clips\n",
"ipd.display(ipd.Audio(reference_clips[0], rate=16000, normalize=True, autoplay=False))"
]
},
{
"cell_type": "markdown",
2023-02-17 23:07:41 -05:00
"id": "937e0df5",
2023-02-17 22:56:15 -05:00
"metadata": {},
"source": [
"Now that we have the data (note that all of the clips *must* be 16 khz, 16-bit PCM WAV files), we can train a custom verifier model. This is simply a scikit-learn logistic regression model, using the same audio features from the normal openWakeWord pre-processor, so it is very fast to train and adds negligble time to the openWakeWord inference engine."
]
},
{
"cell_type": "code",
"execution_count": 27,
"id": "7f965293",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:45:57.606472Z",
"start_time": "2023-02-18T03:45:56.587146Z"
}
},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"Processing positive reference clips: 100%|██████████| 3/3 [00:00<00:00, 4.77it/s]\n",
"Processing negative reference clips: 100%|██████████| 1/1 [00:00<00:00, 5.53it/s]\n"
]
},
{
"name": "stdout",
"output_type": "stream",
"text": [
"Training and saving verifier model...\n",
"Done!\n"
]
}
],
"source": [
"# Train verifier model on the reference clips\n",
"\n",
"output_model_path = \"turn_on_the_office_lights_verifier.pkl\"\n",
"openwakeword.train_custom_verifier(\n",
" positive_reference_clips = reference_clips[0:3], # use 3 reference examples for the wake phrase\n",
" negative_reference_clips = negative_clips,\n",
" output_path = output_model_path,\n",
" model_name = \"turn_on_the_office_lights.onnx\"\n",
")\n"
]
},
{
"cell_type": "markdown",
2023-02-17 23:07:41 -05:00
"id": "4ad4d2ec",
2023-02-17 22:56:15 -05:00
"metadata": {},
"source": [
"After the model is trained, we can instantiate a new openWakeWord instance and include the path to the trained verifier model, as well as set the threshold score from the base model required to invoke the verifier. In practice, you can set this threshold score a bit lower than normal, though as usual actual testing in the deployment environment is recommended."
]
},
{
"cell_type": "code",
"execution_count": 28,
"id": "50b99677",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:46:01.357976Z",
"start_time": "2023-02-18T03:46:01.146646Z"
}
},
"outputs": [
{
"name": "stderr",
"output_type": "stream",
"text": [
"/home/dscripka/anaconda3/envs/torch_gpu/lib/python3.9/site-packages/onnxruntime/capi/onnxruntime_inference_collection.py:54: UserWarning: Specified provider 'CUDAExecutionProvider' is not in available provider names.Available providers: 'CPUExecutionProvider'\n",
" warnings.warn(\n"
]
}
],
"source": [
"# Create openWakeWord instance with verifier model\n",
"\n",
"oww = openwakeword.Model(\n",
" wakeword_model_paths=[\"turn_on_the_office_lights.onnx\"],\n",
" enable_speex_noise_suppression=True,\n",
" vad_threshold=0.5,\n",
" custom_verifier_models={\"turn_on_the_office_lights\": \"turn_on_the_office_lights_verifier.pkl\"},\n",
" custom_verifier_threshold=0.3,\n",
")\n"
]
},
{
"cell_type": "markdown",
2023-02-17 23:07:41 -05:00
"id": "6666dadb",
2023-02-17 22:56:15 -05:00
"metadata": {},
"source": [
"Finally, we can run the model on our test clip from the Santa Barbara corpus and see if the false activation rate has decreased to an acceptable level."
]
},
{
"cell_type": "code",
"execution_count": 29,
"id": "19fd5f51",
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:47:25.076640Z",
"start_time": "2023-02-18T03:46:07.021368Z"
}
},
"outputs": [
{
"data": {
2023-10-13 09:18:39 +02:00
"image/png": "iVBORw0KGgoAAAANSUhEUgAAAiMAAAGiCAYAAAA1LsZRAAAAOXRFWHRTb2Z0d2FyZQBNYXRwbG90bGliIHZlcnNpb24zLjUuMiwgaHR0cHM6Ly9tYXRwbG90bGliLm9yZy8qNh9FAAAACXBIWXMAAA9hAAAPYQGoP6dpAAA7XUlEQVR4nO3deXxU5aH/8S9r8FqJVTSAIpe219beqK3hqsGtasVS9dbb/iqt/QlW7c/UIlW016JtUes1aiu1rYIoIPWWIlpAbUUkVNnDFgKyhEUIScCEkEAmgZD9+f1BM2SZySw5M885M5/365WXMnOW5znLc77nOcv0MMYYAQAAWNLTdgEAAEByI4wAAACrCCMAAMAqwggAALCKMAIAAKwijAAAAKsIIwAAwCrCCAAAsIowAgAArCKMAAAAqyIOI8uXL9ctt9yiwYMHq0ePHnr77bdDjrNs2TJlZGSoX79++tznPqeXX345mrICAIAEFHEYOXbsmC6++GK9+OKLYQ1fWFiob37zm7rqqquUn5+vRx99VOPHj9e8efMiLiwAAEg8PbrzQ3k9evTQggULdOuttwYd5pFHHtG7776rgoIC/2dZWVnavHmzcnNzo501AABIEL1jPYPc3FyNHDmy3Wc33nijZsyYocbGRvXp06fTOPX19aqvr/f/u6WlRYcPH9aZZ56pHj16xLrIAADAAcYY1dTUaPDgwerZM/jFmJiHkbKyMqWlpbX7LC0tTU1NTaqoqNCgQYM6jZOdna0nnngi1kUDAABxUFJSonPPPTfo9zEPI5I69Wa0XhkK1ssxceJETZgwwf9vn8+n8847TyUlJerfv3/sCgoAABxTXV2tIUOG6LTTTutyuJiHkYEDB6qsrKzdZ+Xl5erdu7fOPPPMgOOkpKQoJSWl0+f9+/cnjAAA4DGhbrGI+XtGMjMzlZOT0+6zxYsXa/jw4QHvFwEAAMkl4jBy9OhRbdq0SZs2bZJ04tHdTZs2qbi4WNKJSyxjxozxD5+VlaWioiJNmDBBBQUFmjlzpmbMmKGHH37YmRoAAABPi/gyzYYNG3Tttdf6/916b8fYsWM1a9YslZaW+oOJJA0bNkwLFy7Ugw8+qJdeekmDBw/WH/7wB33nO99xoPgAAMDruvWekXiprq5WamqqfD4f94wAAOAR4R6/+W0aAABgFWEEAABYRRgBAABWEUYAAIBVhBEAAGAVYQQAAFhFGAEAAFYRRgAAgFWEEQAAYBVhBAAAWEUYAQAAVhFGAACAVYQRAABgFWEEAABYRRgBAABWEUYAAIBVhBEAAGAVYQQAAFhFGAEAAFYRRgAAgFWEEQAAYBVhBAAAWEUYAQAAVhFGAACAVYQRAABgFWEEAABYRRgBAABWEUYAAIBVhBEAAGAVYQQAAFhFGAEAAFYRRgAAgFWEEQAAYBVhBAAAWEUYAQAAVhFGAACAVYQRAABgFWEEAABYRRgBAABWEUYAAIBVhBEAAGAVYQQAAFhFGAEAAFYRRgAAgFWEEQAAYBVhBAAAWEUYAQAAVhFGAACAVYQRAABgFWEEAABYRRgBAABWEUYAAIBVhBEAAGAVYQQAAFhFGAEAAFYRRgAAgFWEEQAAYBVhBAAAWEUYAQAAVhFGAACAVYQRAABgFWEEAABYRRgBAABWRRVGpkyZomHDhqlfv37KyMjQihUruhx+9uzZuvjii/Uv//IvGjRokH74wx+qsrIyqgIDAIDEEnEYmTt3rh544AE99thjys/P11VXXaVRo0apuLg44PArV67UmDFjdPfdd2vbtm166623tH79et1zzz3dLjwAAPC+iMPI5MmTdffdd+uee+7RBRdcoBdeeEFDhgzR1KlTAw6/Zs0a/eu//qvGjx+vYcOG6corr9S9996rDRs2dLvwAADA+yIKIw0NDcrLy9PIkSPbfT5y5EitXr064DgjRozQ/v37tXDhQhljdPDgQf31r3/VTTfdFHQ+9fX1qq6ubvcHAAASU0RhpKKiQs3NzUpLS2v3eVpamsrKygKOM2LECM2ePVujR49W3759NXDgQJ1++un64x//GHQ+2dnZSk1N9f8NGTIkkmICAAAPieoG1h49erT7tzGm02ettm/frvHjx+tXv/qV8vLytGjRIhUWFiorKyvo9CdOnCifz+f/KykpiaaYAADAA3pHMvCAAQPUq1evTr0g5eXlnXpLWmVnZ+uKK67Qz372M0nSRRddpFNPPVVXXXWVnnrqKQ0aNKjTOCkpKUpJSYmkaAAAwKMi6hnp27evMjIylJOT0+7znJwcjRgxIuA4tbW16tmz/Wx69eol6USPCgAASG4RX6aZMGGCpk+frpkzZ6qgoEAPPvigiouL/ZddJk6cqDFjxviHv+WWWzR//nxNnTpVe/fu1apVqzR+/HhdeumlGjx4sHM1AQAAnhTRZRpJGj16tCorK/Xkk0+qtLRU6enpWrhwoYYOHSpJKi0tbffOkTvvvFM1NTV68cUX9dBDD+n000/Xddddp2effda5WgAAAM/qYTxwraS6ulqpqany+Xzq37+/7eIAAIAwhHv85rdpAACAVYQRAABgFWEEAABYRRgBAABWEUYAAIBVhBEAAGAVYQQAAFhFGAEAAFYRRgAAgFWEEQAAYBVhBAAAWEUYAQAAVhFGAACAVYQRAABgFWEEAABYRRgBAABWEUYAAIBVhBEAAGAVYQQAAFhFGAEAAFYRRgAAgFWEEQAAYBVhBAAAWEUYAQAAVhFGAACAVYQRAABgFWEEAABYRRgBAABWEUYAAIBVhBEAAGAVYQQAAFhFGAEAAFYRRgAAgFWEEQAAYBVhBAAAWEUYAQAAVhFGAACAVYQRAABgFWEEAABYRRgBAABWEUYAAIBVhBEAAGAVYQQAAFhFGAEAAFYRRgAAgFWEEQAAYBVhBAAAWEUYAQAAVhFGAACAVYQRAABgFWEEAABYRRgBAABWEUYAAIBVhBEAAGAVYQQAAFhFGAEAAFYRRgAAgFWEEQAAYBVhBAAAWEUYAQAAVhFGAACAVYQRAABgFWEEAABYFVUYmTJlioYNG6Z+/fopIyNDK1as6HL4+vp6PfbYYxo6dKhSUlL0+c9/XjNnzoyqwAAAILH0jnSEuXPn6oEHHtCUKVN0xRVXaNq0aRo1apS2b9+u8847L+A4t912mw4ePKgZM2boC1/4gsrLy9XU1NTtwgMAAO/rYYwxkYxw2WWX6ZJLLtHUqVP9n11wwQW69dZblZ2d3Wn4RYsW6Xvf+5727t2rM844I6pCVldXKzU1VT6fT/37949qGgAAIL7CPX5HdJmmoaFBeXl5GjlyZLvPR44cqdWrVwcc591339Xw4cP13HPP6ZxzztH555+vhx9+WMePHw86n/r6elVXV7f7AwAAiSmiyzQVFRVqbm5WWlpau8/T0tJUVlYWcJy9e/dq5cqV6tevnxYsWKCKigrdd999Onz4cND7RrKzs/XEE09EUjQAAOBRUd3A2qNHj3b/NsZ0+qxVS0uLevToodmzZ+vSSy/VN7/5TU2ePFmzZs0K2jsyceJE+Xw+/19JSUk0xQQAAB4QUc/IgAED1KtXr069IOXl5Z16S1oNGjRI55xzjlJTU/2fXXDBBTLGaP/+/fq3f/u3TuOkpKQoJSUlkqIBAACPiqhnpG/fvsrIyFBOTk67z3NycjRixIiA41xxxRX69NNPdfToUf9nu3btUs+ePXXuuedGUWQAAJBIIr5MM2HCBE2fPl0zZ85UQUGBHnzwQRUXFysrK0vSiUssY8aM8Q9/++2368wzz9QPf/hDbd++XcuXL9fPfvYz3XXXXTrllFOcqwkAAPCkiN8zMnr0aFVWVurJJ59UaWmp0tPTtXDhQg0dOlSSVFpaquLiYv/wn/nMZ5STk6P7779fw4cP15lnnqnbbrtNTz31lHO1AAAAnhXxe0Zs4D0jAAB4T0zeMwIAAOA0wggAALCKMAIAAKwijAAAAKsIIwAAwCrCCAAAsIowAgAArCKMAAAAqwgjAADAKsIIAACwijACAACsIowAAACrCCMAAMAqwggAALCKMAIAAKwijAAAAKsIIwAAwCr
2023-02-17 22:56:15 -05:00
"text/plain": [
"<Figure size 640x480 with 1 Axes>"
]
},
"metadata": {},
"output_type": "display_data"
}
],
"source": [
"# Run false-accept rate test again, now with the verifier model\n",
"\n",
"scores = oww.predict_clip(\"santa_barbara_corpus_test_clip.wav\")\n",
"\n",
"plt.figure()\n",
"plt.plot([i[\"turn_on_the_office_lights\"] for i in scores])\n",
"_ = plt.ylim(0,1)\n"
]
},
{
"cell_type": "code",
"execution_count": 30,
2023-02-17 23:07:41 -05:00
"id": "c59007d2",
2023-02-17 22:56:15 -05:00
"metadata": {
"ExecuteTime": {
"end_time": "2023-02-18T03:47:29.856142Z",
"start_time": "2023-02-18T03:47:29.836354Z"
}
},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"False-accept rate per hour: 0.0\n"
]
}
],
"source": [
"# Calculate the false-accept rate per hour from this new result\n",
"\n",
"false_accepts = openwakeword.metrics.get_false_positives(\n",
" [i[\"turn_on_the_office_lights\"] for i in scores], threshold=0.5\n",
")\n",
"\n",
"print(f\"False-accept rate per hour: {false_accepts/1}\")"
]
},
{
"cell_type": "markdown",
2023-02-17 23:07:41 -05:00
"id": "7c7375c6",
2023-02-17 22:56:15 -05:00
"metadata": {},
"source": [
"Sucess! Now the false-activation rate is at most <1 per hour given that there weren't any false-positives in our ~1 hour test clip, which is an orders of magnitude decrease! This model is now much closer to being ready for a production deployment, assuming that each user is known and can provide the neccessary data to train the verifier model."
]
}
],
"metadata": {
"kernelspec": {
"display_name": "torch_gpu",
"language": "python",
"name": "torch_gpu"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.9.13"
},
"toc": {
"base_numbering": 1,
"nav_menu": {},
"number_sections": true,
"sideBar": true,
"skip_h1_title": false,
"title_cell": "Table of Contents",
"title_sidebar": "Contents",
"toc_cell": false,
"toc_position": {
"height": "calc(100% - 180px)",
"left": "10px",
"top": "150px",
"width": "384px"
},
"toc_section_display": true,
"toc_window_display": true
}
},
"nbformat": 4,
"nbformat_minor": 5
}