Tests passing for online custom verifier model training

This commit is contained in:
dscripka 2023-04-02 11:15:54 -04:00
parent b66d1d4469
commit 2cb63d8296
4 changed files with 103 additions and 29 deletions

View file

@ -22,7 +22,6 @@ import scipy
import pickle
from sklearn.linear_model import LogisticRegression
from sklearn.svm import SVC
from sklearn.pipeline import make_pipeline
from sklearn.preprocessing import FunctionTransformer, StandardScaler
@ -90,12 +89,14 @@ def get_reference_clip_features(
def flatten_features(x):
return [i.flatten() for i in x]
def make_sklearn_pipeline():
# clf = SVC(gamma='auto', probability=True)
clf = LogisticRegression(random_state=0, max_iter=2000, C=0.01)
pipeline = make_pipeline(FunctionTransformer(flatten_features), StandardScaler(), clf)
return pipeline
def train_verifier_model(features: np.ndarray, labels: np.ndarray, **kwargs):
"""
Train a logistic regression binary classifier model on the provided features and labels

View file

@ -20,11 +20,12 @@ from openwakeword.utils import AudioFeatures
import wave
import os
from pathlib import Path
import pickle
from collections import deque, defaultdict
from functools import partial
import time
from typing import List, Union, DefaultDict, Dict
from typing import List, Union, DefaultDict, Dict, Any
# Define main model class
@ -41,7 +42,8 @@ class Model():
vad_threshold: float = 0,
custom_verifier_models: dict = {},
custom_verifier_threshold: float = 0.1,
custom_verifier_model_online_learning: bool=False,
custom_verifier_model_online_learning: bool = False,
cache_directory: str = "",
**kwargs
):
"""Initialize the openWakeWord model object.
@ -71,7 +73,17 @@ class Model():
from a model for a given frame is greater than this value, the
associated custom verifier model will also predict on that frame, and
the verifier score will be returned.
kwargs (dict): Any other keyword arguments to pass the the preprocessor instance
custom_verifier_model_online_learning (bool): Whether to enable online learning for custom verifier models.
When enabled, will dynamically train custom verifier models
based on usage data, re-training after every positive activation
or fixed period of time. This can be a good way to leverage most of
the benefits of verifier models without having to manually manage
the data, training, and usage of the verifier models. Though in
some cases performance of the verifier models will be worse
as the data is less curated.
cache_directory (str): The cache directory to use for openWakeWord. Defaults to the home directory for each
supported platform (/home/<user>/.openwakeword for linux, C:\\Users\\<user> on Windows, etc.)
kwargs (dict): Any other keyword arguments to pass to the preprocessor instance
"""
# Initialize the ONNX models and store them
@ -95,9 +107,9 @@ class Model():
self.model_input_names = {}
self.custom_verifier_models = {}
self.custom_verifier_threshold = custom_verifier_threshold
self.custom_verifier_data = {}
self.custom_verifier_data: Dict[str, Any] = {}
self.custom_verifier_model_online_learning = custom_verifier_model_online_learning
self.custom_verifier_model_last_train_time = time.time()
self.samples_processed = 0
for mdl_path, mdl_name in zip(wakeword_model_paths, wakeword_model_names):
# Load openwakeword models
self.models[mdl_name] = ort.InferenceSession(mdl_path, sess_options=sessionOptions,
@ -112,19 +124,35 @@ class Model():
self.class_mapping[mdl_name] = {str(i): str(i) for i in range(0, self.model_outputs[mdl_name])}
self.model_input_names[mdl_name] = self.models[mdl_name].get_inputs()[0].name
# Load custom verifier models
# Create attributes to store realtime usage data for verifier models
self.custom_verifier_data[mdl_name] = {
"features": {
"positive": deque(maxlen=1000),
"negative": deque(maxlen=1000)
}
}
# Create filesystem cache locations, or load existing data in the cache
if cache_directory == "":
home_dir = str(Path.home())
self.cache_dir = os.path.join(home_dir, ".openwakeword")
else:
self.cache_dir = cache_directory
if not os.path.exists(self.cache_dir):
os.mkdir(self.cache_dir)
else:
if os.path.exists(os.path.join(self.cache_dir, "cached_features.pkl")):
# Load existing data in the cache
self.custom_verifier_data = pickle.load(open(os.path.join(self.cache_dir, "cached_features.pkl"), 'rb'))
# Load/create verifier models depending on initialization arguments
if isinstance(custom_verifier_models, dict):
if custom_verifier_models.get(mdl_name, False):
self.custom_verifier_models[mdl_name] = pickle.load(open(custom_verifier_models[mdl_name], 'rb'))
else:
if self.custom_verifier_model_online_learning:
if not self.custom_verifier_models.get(mdl_name, False):
self.custom_verifier_models[mdl_name] = openwakeword.custom_verifier_model.make_sklearn_pipeline()
self.custom_verifier_data[mdl_name] = {
"features": {
"positive": deque(maxlen=1000),
"negative": deque(maxlen=1000)
}
}
if len(self.custom_verifier_models.keys()) < len(custom_verifier_models.keys()):
raise ValueError(
@ -207,6 +235,9 @@ class Model():
if timing:
timing_dict["models"]["preprocessor"] = time.time() - feature_start
# Increment counter
self.samples_processed += len(x)
# Get predictions from model(s)
predictions = {}
for mdl in self.models.keys():
@ -230,7 +261,7 @@ class Model():
group_features.extend(frame_features)
prediction = np.array(group_predictions).max(axis=0)[None, ]
frame_features = group_features[np.array(group_predictions).argmax(axis=0)]
frame_features = group_features[np.array(group_predictions).argmax(axis=0)[0][0]]
else:
frame_features = self.preprocessor.get_features(self.model_inputs[mdl])
prediction = self.models[mdl].run(
@ -249,10 +280,8 @@ class Model():
for cls in predictions.keys():
# Update cached negative data for verifier model (if using online learning) if score is above arbitrary level
if self.custom_verifier_model_online_learning and predictions[cls] >= .01 and \
predictions[cls] < 0.05:
predictions[cls] < 0.05:
self.custom_verifier_data[cls]["features"]["negative"].append(frame_features.flatten())
# elif self.custom_verifier_model_online_learning and predictions[cls] > 0.01 and predictions[cls] < 0.1:
# self.custom_verifier_data[cls]["features"]["negative"].append(np.random.random((16,96)).flatten())
# Update scores of positive predictions
positive_examples_added = False
@ -271,20 +300,22 @@ class Model():
self.custom_verifier_data[cls]["features"]["positive"].append(frame_features.flatten())
positive_examples_added = True
# Save feature data to cache
pickle.dump(self.custom_verifier_data, open(os.path.join(self.cache_dir, "cached_features.pkl"), 'wb'))
# Train verifier model on latest data at most every 10 seconds or after every positive detection
if self.custom_verifier_model_online_learning and \
((time.time() - self.custom_verifier_model_last_train_time) >= 5 or positive_examples_added):
self.custom_verifier_model_last_train_time = time.time()
(self.samples_processed >= 16000*10 or positive_examples_added):
self.samples_processed = 0
parent_model = self.get_parent_model_from_label(cls)
y = np.array(
[1]*np.array(self.custom_verifier_data[cls]["features"]["positive"]).shape[0] + \
[1]*np.array(self.custom_verifier_data[cls]["features"]["positive"]).shape[0] +
[0]*np.array(self.custom_verifier_data[cls]["features"]["negative"]).shape[0]
)
# need a minimum number of examples to train
if len(self.custom_verifier_data[cls]["features"]["positive"]) > 5 and \
len(self.custom_verifier_data[cls]["features"]["negative"]) > 5 :
# need a minimum number of examples to train (5)
if len(self.custom_verifier_data[cls]["features"]["positive"]) > 3 and \
len(self.custom_verifier_data[cls]["features"]["negative"]) > 5:
print("Updating custom verifier model")
x = np.vstack((
np.array(self.custom_verifier_data[cls]["features"]["positive"]),

View file

@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta"
[tool.pytest.ini_options]
addopts = "--cov=openwakeword --cov-report term-missing --flake8 --mypy --mypy-ignore-missing-imports"
flake8-max-line-length = "120"
flake8-max-line-length = "140"
testpaths = [
"tests",
"openwakeword"

View file

@ -43,7 +43,7 @@ class TestModels:
scipy.io.wavfile.write(os.path.join(tmp_dir, "negative_reference.wav"),
16000, np.random.randint(-1000, 1000, 16000*4).astype(np.int16))
# Load random clips
# Load reference clips
reference_clips = [os.path.join("tests", "data", "hey_mycroft_test.wav")]
negative_clips = [os.path.join(tmp_dir, "negative_reference.wav")]
@ -80,7 +80,7 @@ class TestModels:
custom_verifier_threshold=0.3,
)
# Load model with verifier model incorrectly to catch ValueError
# Load model with verifier model correctly
owwModel = openwakeword.Model(
wakeword_model_paths=[os.path.join("openwakeword", "resources", "models", "hey_mycroft_v0.1.onnx")],
custom_verifier_models={"hey_mycroft_v0.1": os.path.join(tmp_dir, "verifier_model.pkl")},
@ -89,3 +89,45 @@ class TestModels:
# Prediction on random data
owwModel.predict_clip(reference_clips[0])
def test_train_verifier_model_online(self):
with tempfile.TemporaryDirectory() as tmp_dir:
# Load model with online custom verifier model training
owwModel = openwakeword.Model(
wakeword_model_paths=[os.path.join("openwakeword", "resources", "models", "hey_mycroft_v0.1.onnx")],
custom_verifier_model_online_learning=True,
cache_directory=tmp_dir
)
# Load positive clip for online training
sr, reference_clip = scipy.io.wavfile.read(os.path.join("tests", "data", "hey_mycroft_test.wav"))
# Make combined clip
combined_clip = np.concatenate((
np.random.randint(-1000, 1000, 16000*10).astype(np.int16),
reference_clip,
np.random.randint(-1000, 1000, 16000*10).astype(np.int16),
reference_clip,
np.random.randint(-1000, 1000, 16000*10).astype(np.int16),
reference_clip,
np.random.randint(-1000, 1000, 16000*10).astype(np.int16),
reference_clip,
np.random.randint(-1000, 1000, 16000*10).astype(np.int16),
reference_clip,
np.random.randint(-1000, 1000, 16000*10).astype(np.int16),
reference_clip,
np.random.randint(-1000, 1000, 16000*10).astype(np.int16),
reference_clip,
))
# Run prediction for online training
owwModel.predict_clip(combined_clip)
# Instantiate new model to load cache directory
owwModel = openwakeword.Model(
wakeword_model_paths=[os.path.join("openwakeword", "resources", "models", "hey_mycroft_v0.1.onnx")],
custom_verifier_model_online_learning=True,
cache_directory=tmp_dir
)
assert owwModel.custom_verifier_data != {}