mirror of
https://github.com/dscripka/openWakeWord.git
synced 2026-08-27 18:17:20 -04:00
Test coverage for model.py back to 100%
This commit is contained in:
parent
4d9d549930
commit
054f7578d0
1 changed files with 48 additions and 0 deletions
|
|
@ -37,6 +37,26 @@ import pytest
|
|||
|
||||
# Tests
|
||||
class TestModels:
|
||||
def test_load_models_by_path(self):
|
||||
# Load model with defaults
|
||||
owwModel = openwakeword.Model(wakeword_model_paths=[
|
||||
os.path.join("openwakeword", "resources", "models", "alexa_v0.1.onnx")
|
||||
])
|
||||
|
||||
# Prediction on random data
|
||||
owwModel.predict(np.random.randint(-1000, 1000, 1280).astype(np.int16))
|
||||
|
||||
def test_custom_model_label_mapping_dict(self):
|
||||
# Load model with model path
|
||||
owwModel = openwakeword.Model(wakeword_model_paths=[
|
||||
os.path.join("openwakeword", "resources", "models", "alexa_v0.1.onnx")
|
||||
],
|
||||
class_mapping_dicts=[{"alexa_v0.1": {"0": "positive"}}]
|
||||
)
|
||||
|
||||
# Prediction on random data
|
||||
owwModel.predict(np.random.randint(-1000, 1000, 1280).astype(np.int16))
|
||||
|
||||
def test_models(self):
|
||||
# Load model with defaults
|
||||
owwModel = openwakeword.Model()
|
||||
|
|
@ -65,6 +85,34 @@ class TestModels:
|
|||
else:
|
||||
assert max(predictions_flat[key]) < 0.5
|
||||
|
||||
def test_models_with_speex_noise_cancellation(self):
|
||||
# Load model with defaults
|
||||
owwModel = openwakeword.Model(enable_speex_noise_suppression=True)
|
||||
|
||||
# Get clips for each model (assumes that test clips will have the model name in the filename)
|
||||
test_dict = {}
|
||||
for mdl_name in owwModel.models.keys():
|
||||
all_clips = [str(i) for i in Path(os.path.join("tests", "data")).glob("*.wav")]
|
||||
test_dict[mdl_name] = [i for i in all_clips if mdl_name in i]
|
||||
|
||||
# Predict
|
||||
for model, clips in test_dict.items():
|
||||
for clip in clips:
|
||||
# Get predictions for reach frame in the clip
|
||||
predictions = owwModel.predict_clip(clip)
|
||||
owwModel.reset() # reset after each clip to ensure independent results
|
||||
|
||||
# Make predictions dictionary flatter
|
||||
predictions_flat = collections.defaultdict(list)
|
||||
[predictions_flat[key].append(i[key]) for i in predictions for key in i.keys()]
|
||||
|
||||
# Check scores against default threshold (0.5)
|
||||
for key in predictions_flat.keys():
|
||||
if key in clip:
|
||||
assert max(predictions_flat[key]) >= 0.5
|
||||
else:
|
||||
assert max(predictions_flat[key]) < 0.5
|
||||
|
||||
def test_predict_clip_with_array(self):
|
||||
# Load model with defaults
|
||||
owwModel = openwakeword.Model()
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue