Added VAD model, adjusted download function behavior to skip files that already exist [skip ci]

This commit is contained in:
dscripka 2023-10-08 21:29:40 -04:00
parent 74839d5ca2
commit 9f394d7abc
2 changed files with 21 additions and 6 deletions

View file

@ -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 = {

View file

@ -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)