mirror of
https://github.com/dscripka/openWakeWord.git
synced 2026-08-27 18:17:20 -04:00
Added basic debounce logic for model.predict
This commit is contained in:
parent
dc5a234218
commit
c63384489e
1 changed files with 12 additions and 3 deletions
|
|
@ -227,7 +227,8 @@ class Model():
|
|||
"""Reset the prediction buffer"""
|
||||
self.prediction_buffer = defaultdict(partial(deque, maxlen=30))
|
||||
|
||||
def predict(self, x: np.ndarray, patience: dict = {}, threshold: dict = {}, timing: bool = False):
|
||||
def predict(self, x: np.ndarray, patience: dict = {},
|
||||
threshold: dict = {}, debounce_time: float = 0.0, timing: bool = False):
|
||||
"""Predict with all of the wakeword models on the input audio frames
|
||||
|
||||
Args:
|
||||
|
|
@ -242,9 +243,11 @@ class Model():
|
|||
model names and the values are the number of frames. Can reduce false-positive
|
||||
detections at the cost of a lower true-positive rate.
|
||||
By default, this behavior is disabled.
|
||||
threshold (dict): The threshold values to use when the `patience` behavior is enabled.
|
||||
threshold (dict): The threshold values to use when the `patience` or `debounce_time` behavior is enabled.
|
||||
Must be provided as an a dictionary where the keys are the
|
||||
model names and the values are the thresholds.
|
||||
debounce_time (float): The time (in seconds) to wait before returning another non-zero prediction
|
||||
after a non-zero prediction. Can preven multiple detections of the same wake-word.
|
||||
timing (bool): Whether to return timing information of the models. Can be useful to debug and
|
||||
assess how efficiently models are running on the current hardware.
|
||||
|
||||
|
|
@ -333,16 +336,22 @@ class Model():
|
|||
timing_dict["models"][mdl] = time.time() - model_start
|
||||
|
||||
# Update scores based on thresholds or patience arguments
|
||||
if patience != {}:
|
||||
if patience != {} or debounce_time > 0:
|
||||
if threshold == {}:
|
||||
raise ValueError("Error! When using the `patience` argument, threshold "
|
||||
"values must be provided via the `threshold` argument!")
|
||||
if patience != {} and debounce_time > 0:
|
||||
raise ValueError("Error! The `patience` and `debounce_time` arguments cannot be used together!")
|
||||
for mdl in predictions.keys():
|
||||
parent_model = self.get_parent_model_from_label(mdl)
|
||||
if parent_model in patience.keys():
|
||||
scores = np.array(self.prediction_buffer[mdl])[-patience[parent_model]:]
|
||||
if (scores >= threshold[parent_model]).sum() < patience[parent_model]:
|
||||
predictions[mdl] = 0.0
|
||||
if debounce_time > 0:
|
||||
n_frames = int(debounce_time*1000/80)
|
||||
if (np.array(self.prediction_buffer[mdl])[-n_frames:] >= threshold[parent_model]).sum() > 0:
|
||||
predictions[mdl] = 0.0
|
||||
|
||||
# (optionally) get voice activity detection scores and update model scores
|
||||
if self.vad_threshold > 0:
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue