Fixed bug for onnx models, removed tflite install for windows

This commit is contained in:
dscripka 2023-06-10 11:41:01 -04:00
parent 3f79f3e3c0
commit 8bf101aa22
2 changed files with 26 additions and 13 deletions

View file

@ -19,6 +19,7 @@ from openwakeword.utils import AudioFeatures, re_arg
import wave
import os
import logging
import functools
import pickle
from collections import deque, defaultdict
@ -108,11 +109,6 @@ class Model():
self.custom_verifier_threshold = custom_verifier_threshold
# Do imports for inference framework
if inference_framework == "onnx":
try:
import onnxruntime as ort
except ImportError:
raise ValueError("Tried to import onnxruntime, but it was not found. Please install it using `pip install onnxruntime`")
if inference_framework == "tflite":
try:
import tflite_runtime.interpreter as tflite
@ -123,8 +119,26 @@ class Model():
return tflite_interpreter.get_tensor(output_index)[None, ]
except ImportError:
raise ValueError("Tried to import the TFLite runtime, but it was not found."
"Please install it using `pip install tflite-runtime`")
logging.warning("Tried to import the tflite runtime, but it was not found. "
"Trying to switching to onnxruntime instead, if appropriate models are available.")
if wakeword_models != [] and all(['.onnx' in i for i in wakeword_models]):
inference_framework = "onnx"
elif wakeword_models != [] and all([os.path.exists(i.replace('.tflite', '.onnx')) for i in wakeword_models]):
inference_framework = "onnx"
wakeword_models = [i.replace('.tflite', '.onnx') for i in wakeword_models]
else:
raise ValueError("Tried to import the tflite runtime for provided tflite models, but it was not found. "
"Please install it using `pip install tflite-runtime`")
if inference_framework == "onnx":
try:
import onnxruntime as ort
def onnx_predict(onnx_model, x):
return onnx_model.run(None, {onnx_model.get_inputs()[0].name: x})
except ImportError:
raise ValueError("Tried to import onnxruntime, but it was not found. Please install it using `pip install onnxruntime`")
for mdl_path, mdl_name in zip(wakeword_models, wakeword_model_names):
# Load openwakeword models
@ -141,9 +155,8 @@ class Model():
self.model_inputs[mdl_name] = self.models[mdl_name].get_inputs()[0].shape[1]
self.model_outputs[mdl_name] = self.models[mdl_name].get_outputs()[0].shape[1]
self.model_prediction_function[mdl_name] = lambda x: self.models[mdl_name].run(
None, {self.models[mdl_name].get_inputs()[0].name: x}
)
pred_function = functools.partial(onnx_predict, self.models[mdl_name])
self.model_prediction_function[mdl_name] = pred_function
if inference_framework == "tflite":
if ".onnx" in mdl_path:
@ -158,8 +171,8 @@ class Model():
tflite_input_index = self.models[mdl_name].get_input_details()[0]['index']
tflite_output_index = self.models[mdl_name].get_output_details()[0]['index']
foo = functools.partial(tflite_predict, self.models[mdl_name], tflite_input_index, tflite_output_index)
self.model_prediction_function[mdl_name] = foo
pred_function = functools.partial(tflite_predict, self.models[mdl_name], tflite_input_index, tflite_output_index)
self.model_prediction_function[mdl_name] = pred_function
if class_mapping_dicts and class_mapping_dicts[wakeword_models.index(mdl_path)].get(mdl_name, None):
self.class_mapping[mdl_name] = class_mapping_dicts[wakeword_models.index(mdl_path)]

View file

@ -29,7 +29,7 @@ setuptools.setup(
version="0.4.0",
install_requires=[
'onnxruntime>=1.10.0,<2',
'tflite-runtime>=2.8.0,<3',
'tflite-runtime>=2.8.0,<3; platform_system == "Linux"',
'tqdm>=4.0,<5.0',
'scipy>=1.3,<2',
'scikit-learn>=1,<2'