From 68e88c1350113a1e70afc8cb64e8d9d120db619f Mon Sep 17 00:00:00 2001 From: dscripka Date: Sun, 11 Feb 2024 12:08:03 -0500 Subject: [PATCH] Added/fixed reset methods --- openwakeword/model.py | 4 +++- openwakeword/utils.py | 10 +++++++++- 2 files changed, 12 insertions(+), 2 deletions(-) diff --git a/openwakeword/model.py b/openwakeword/model.py index 97d1303..8f2ef42 100755 --- a/openwakeword/model.py +++ b/openwakeword/model.py @@ -224,8 +224,10 @@ class Model(): return parent_model def reset(self): - """Reset the prediction buffer""" + """Reset the prediction and audio feature buffers. Useful for re-initializing the model, though may not be efficient + when called too frequently.""" self.prediction_buffer = defaultdict(partial(deque, maxlen=30)) + self.preprocessor.reset() def predict(self, x: np.ndarray, patience: dict = {}, threshold: dict = {}, debounce_time: float = 0.0, timing: bool = False): diff --git a/openwakeword/utils.py b/openwakeword/utils.py index 8da8048..4964706 100644 --- a/openwakeword/utils.py +++ b/openwakeword/utils.py @@ -160,7 +160,7 @@ class AudioFeatures(): self.embedding_model_predict = tflite_embedding_predict - # Create databuffers + # Create databuffers with empty/random data self.raw_data_buffer: Deque = deque(maxlen=sr*10) self.melspectrogram_buffer = np.ones((76, 32)) # n_frames x num_features self.melspectrogram_max_len = 10*97 # 97 is the number of frames in 1 second of 16hz audio @@ -169,6 +169,14 @@ class AudioFeatures(): self.feature_buffer = self._get_embeddings(np.random.randint(-1000, 1000, 16000*4).astype(np.int16)) self.feature_buffer_max_len = 120 # ~10 seconds of feature buffer history + def reset(self): + """Reset the internal buffers""" + self.raw_data_buffer.clear() + self.melspectrogram_buffer = np.ones((76, 32)) + self.accumulated_samples = 0 + self.raw_data_remainder = np.empty(0) + self.feature_buffer = self._get_embeddings(np.random.randint(-1000, 1000, 16000*4).astype(np.int16)) + def _get_melspectrogram(self, x: Union[np.ndarray, List], melspec_transform: Callable = lambda x: x/10 + 2): """ Function to compute the mel-spectrogram of the provided audio samples.