diff --git a/openwakeword/model.py b/openwakeword/model.py index f6036a9..59568ce 100755 --- a/openwakeword/model.py +++ b/openwakeword/model.py @@ -416,7 +416,7 @@ class Model(): else: return predictions - def self_confirm(self, last_n_seconds: float = 1.5, background=False): + def self_confirm(self, last_n_seconds: float = 1.5, delay_time: float = 0.250): """ Use the confirmation model to confirm the predictions from the main model. This is a form of test-time augmentation that can significantly reduce false detections, but significantly increases @@ -431,18 +431,22 @@ class Model(): You are encouraged to experiment with the `last_n_seconds` argument to find the best balance between true-positive and false-positive detections for your use case. + This is a background task to not block the main model from processing audio, so the results + of the confirmation model are stored in the `confirmation_results` class attribute once available. + This is a dictionary with the same format as the output of the `predict` method, containing the + maximum score from the confirmation model over the last `last_n_seconds` seconds of audio, giving a + "confirmation" score for each model, indicating if a detection in the `last_n_seconds` seconds of audio + was likely valid or not. + Args: last_n_seconds (float): The number of seconds of audio to use for confirmation. The default (1.5) should be sufficient for most use cases, but increase if your target wake-word/phrase is long, or decrease if short. - background (bool): Whether to run the confirmation model in a background thread. If True, the results of - the function will be returned asynchronously and stored in the - `self.confirmation_results` attribute. Until the results are available, this attribute - will be None. + delay_time (float): The time (in seconds) to wait before running the confirmation model. This allows the + main model to process enough audio after a detection to ensure that the confirmation + model has enough audio context. Returns: - dict: A dictionary of scores between 0 and 1 for each model, representing the maximum - score from the confirmation model over the last `last_n_seconds` seconds of audio. - If background=True, returns None and stores results in self.confirmation_results when ready. + concurrent.futures.Future: A futures object representing the threading task running the confirmation model. """ # Check for self-confirm functionality if self.self_confirm_enabled is False: @@ -450,11 +454,15 @@ class Model(): # Check for at least two cores cpu_count = os.cpu_count() - if (cpu_count is None or cpu_count < 2) and background is True: + if (cpu_count is None or cpu_count < 2): raise ValueError("The self-confirm functionality requires at least two CPU cores, as it uses threading.") # Define the function to run predictions def _run_confirmation_predictions(): + # Wait to allow main model to process audio + if delay_time > 0: + time.sleep(delay_time) + # Get the last n seconds of audio from the audio buffer of the main model, and get the features # with the self-confirmation model preprocessor n_samples = int(last_n_seconds*16000) @@ -480,15 +488,11 @@ class Model(): # Store results asynchronously self.confirmation_results = predictions_dict - # Run in background thread if requested - if background: - self.confirmation_results = None - self.confirmation_executor.submit(_run_confirmation_predictions) - return None - else: - # Run synchronously - _run_confirmation_predictions() - return self.confirmation_results + # Submit confirmation prediction task to thread pool + self.confirmation_results = None # reset previous results + future = self.confirmation_executor.submit(_run_confirmation_predictions) + + return future def predict_clip(self, clip: Union[str, np.ndarray], padding: int = 1, chunk_size=1280, **kwargs): """Predict on an full audio clip, simulating streaming prediction. diff --git a/tests/test_self_confirm.py b/tests/test_self_confirm.py index bc2fd5d..38a4b7e 100644 --- a/tests/test_self_confirm.py +++ b/tests/test_self_confirm.py @@ -42,7 +42,14 @@ class TestSelfConfirm: owwModel.predict(random_audio) # Run the self-confirm function - predictions_dict = owwModel.self_confirm(last_n_seconds=1.5) + owwModel.self_confirm(last_n_seconds=1.5) + + # Poll for results with a timeout (max 10 seconds) + max_wait_time = 10 + start_time = time.time() + while owwModel.confirmation_results is None and (time.time() - start_time) < max_wait_time: + time.sleep(0.1) + predictions_dict = owwModel.confirmation_results # Verify predictions_dict is properly formed assert isinstance(predictions_dict, dict), "predictions_dict should be a dictionary" @@ -76,7 +83,14 @@ class TestSelfConfirm: owwModel.predict(random_audio) # Run self-confirm - predictions_dict = owwModel.self_confirm(last_n_seconds=1.5) + owwModel.self_confirm(last_n_seconds=1.5) + + # Poll for results with a timeout (max 10 seconds) + max_wait_time = 10 + start_time = time.time() + while owwModel.confirmation_results is None and (time.time() - start_time) < max_wait_time: + time.sleep(0.1) + predictions_dict = owwModel.confirmation_results # Verify all models have predictions assert len(predictions_dict) >= 2, "predictions_dict should have at least 2 models" @@ -120,7 +134,8 @@ class TestSelfConfirm: # Attempting to call self_confirm should raise ValueError with pytest.raises(ValueError, match="Not enough audio data"): - owwModel.self_confirm(last_n_seconds=1.5) + future = owwModel.self_confirm(last_n_seconds=1.5) + future.result() def test_self_confirm_with_tflite_models(self): """Test self_confirm with tflite inference framework""" @@ -139,7 +154,14 @@ class TestSelfConfirm: owwModel.predict(random_audio) # Run self-confirm - predictions_dict = owwModel.self_confirm(last_n_seconds=1.5) + owwModel.self_confirm(last_n_seconds=1.5) + + # Poll for results with a timeout (max 10 seconds) + max_wait_time = 10 + start_time = time.time() + while owwModel.confirmation_results is None and (time.time() - start_time) < max_wait_time: + time.sleep(0.1) + predictions_dict = owwModel.confirmation_results # Verify predictions_dict is properly formed assert isinstance(predictions_dict, dict) @@ -163,7 +185,14 @@ class TestSelfConfirm: owwModel.predict(random_audio) # Run self-confirm - predictions_dict = owwModel.self_confirm(last_n_seconds=1.5) + owwModel.self_confirm(last_n_seconds=1.5) + + # Poll for results with a timeout (max 10 seconds) + max_wait_time = 10 + start_time = time.time() + while owwModel.confirmation_results is None and (time.time() - start_time) < max_wait_time: + time.sleep(0.1) + predictions_dict = owwModel.confirmation_results # Verify predictions_dict is properly formed assert isinstance(predictions_dict, dict) @@ -172,48 +201,3 @@ class TestSelfConfirm: for model_name, score in predictions_dict.items(): assert isinstance(score, (float, np.floating)), f"Score for {model_name} should be a float" assert 0 <= score <= 1, f"Score for {model_name} should be between 0 and 1, got {score}" - - def test_self_confirm_background_true(self): - """Test self_confirm with background=True returns None and populates confirmation_results""" - owwModel = openwakeword.Model( - wakeword_models=[os.path.join("openwakeword", "resources", "models", "alexa_v0.1.onnx")], - inference_framework="onnx", - self_confirm=True - ) - - # Feed in ~10 seconds of random data to fill the audio buffer - chunk_size = 1280 - n_samples = 160000 - - for i in range(0, n_samples, chunk_size): - random_audio = np.random.randint(-1000, 1000, chunk_size).astype(np.int16) - owwModel.predict(random_audio) - - # Run self-confirm in background mode - result = owwModel.self_confirm(last_n_seconds=1.5, background=True) - - # When background=True, should return None immediately - assert result is None, "self_confirm with background=True should return None" - - # confirmation_results should eventually be populated - # Poll for results with a timeout (max 10 seconds) - max_wait_time = 10 - start_time = time.time() - while owwModel.confirmation_results is None and (time.time() - start_time) < max_wait_time: - time.sleep(0.1) - - # Verify that confirmation_results has been populated - assert owwModel.confirmation_results is not None, "confirmation_results should be populated after background execution" - - # Verify confirmation_results is properly formed - predictions_dict = owwModel.confirmation_results - assert isinstance(predictions_dict, dict), "confirmation_results should be a dictionary" - - expected_models = list(owwModel.models.keys()) - assert len(predictions_dict) == len(expected_models), f"confirmation_results should have {len(expected_models)} key(s)" - - for model_name in expected_models: - assert model_name in predictions_dict, f"confirmation_results should contain key '{model_name}'" - score = predictions_dict[model_name] - assert isinstance(score, (float, np.floating)), f"Score for {model_name} should be a float" - assert 0 <= score <= 1, f"Score for {model_name} should be between 0 and 1, got {score}"