diff --git a/openwakeword/__init__.py b/openwakeword/__init__.py index b74f8bc..6ad8f3f 100755 --- a/openwakeword/__init__.py +++ b/openwakeword/__init__.py @@ -11,9 +11,16 @@ FEATURE_MODELS = { "download_url": "https://github.com/dscripka/openWakeWord/releases/download/v0.5.1/embedding_model.tflite" }, "melspectrogram": { - "model_path": os.path.join(os.path.dirname(os.path.abspath(__file__)), "resources/models/melspectrogram_model.tflite"), - "download_url": "https://github.com/dscripka/openWakeWord/releases/download/v0.5.1/melspectrogram_model.tflite" - }, + "model_path": os.path.join(os.path.dirname(os.path.abspath(__file__)), "resources/models/melspectrogram.tflite"), + "download_url": "https://github.com/dscripka/openWakeWord/releases/download/v0.5.1/melspectrogram.tflite" + } +} + +VAD_MODELS = { + "silero_vad": { + "model_path": os.path.join(os.path.dirname(os.path.abspath(__file__)), "resources/models/silero_vad.onnx"), + "download_url": "https://github.com/dscripka/openWakeWord/releases/download/v0.5.1/silero_vad.onnx" + } } MODELS = { diff --git a/openwakeword/utils.py b/openwakeword/utils.py index 4aa932d..ff89c42 100644 --- a/openwakeword/utils.py +++ b/openwakeword/utils.py @@ -573,6 +573,11 @@ def download_models( download_file(feature_model["download_url"], target_directory) download_file(feature_model["download_url"].replace(".tflite", ".onnx"), target_directory) + # Always download VAD models, if they don't already exist + for vad_model in openwakeword.VAD_MODELS.values(): + if not os.path.exists(os.path.join(target_directory, vad_model["download_url"].split("/")[-1])): + download_file(vad_model["download_url"], target_directory) + # Get all model urls official_model_urls = [i["download_url"] for i in openwakeword.MODELS.values()] official_model_names = [i["download_url"].split("/")[-1] for i in openwakeword.MODELS.values()] @@ -581,11 +586,14 @@ def download_models( for model_name in model_names: url = [i for i, j in zip(official_model_urls, official_model_names) if model_name in j] if url != []: - download_file(url[0], target_directory) + if not os.path.exists(os.path.join(target_directory, url[0].split("/")[-1])): + download_file(url[0], target_directory) else: + print(official_model_urls) for official_model_url in official_model_urls: - download_file(official_model_url, target_directory) - download_file(official_model_url.replace(".tflite", ".onnx"), target_directory) + if not os.path.exists(os.path.join(target_directory, official_model_url.split("/")[-1])): + download_file(official_model_url, target_directory) + download_file(official_model_url.replace(".tflite", ".onnx"), target_directory) # Handle deprecated arguments and naming (thanks to https://stackoverflow.com/a/74564394)