Ruff fixes
This commit is contained in:
@@ -18,8 +18,8 @@ from loguru import logger
|
|||||||
from pipecat.audio.turn.smart_turn.base_smart_turn import BaseSmartTurn
|
from pipecat.audio.turn.smart_turn.base_smart_turn import BaseSmartTurn
|
||||||
|
|
||||||
try:
|
try:
|
||||||
from transformers import WhisperFeatureExtractor
|
|
||||||
import onnxruntime as ort
|
import onnxruntime as ort
|
||||||
|
from transformers import WhisperFeatureExtractor
|
||||||
except ModuleNotFoundError as e:
|
except ModuleNotFoundError as e:
|
||||||
logger.error(f"Exception: {e}")
|
logger.error(f"Exception: {e}")
|
||||||
logger.error(
|
logger.error(
|
||||||
@@ -28,14 +28,6 @@ except ModuleNotFoundError as e:
|
|||||||
raise Exception(f"Missing module: {e}")
|
raise Exception(f"Missing module: {e}")
|
||||||
|
|
||||||
|
|
||||||
def build_session(onnx_path):
|
|
||||||
so = ort.SessionOptions()
|
|
||||||
so.execution_mode = ort.ExecutionMode.ORT_SEQUENTIAL
|
|
||||||
so.inter_op_num_threads = 1
|
|
||||||
so.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
|
|
||||||
return ort.InferenceSession(onnx_path, sess_options=so)
|
|
||||||
|
|
||||||
|
|
||||||
class LocalSmartTurnAnalyzerV3(BaseSmartTurn):
|
class LocalSmartTurnAnalyzerV3(BaseSmartTurn):
|
||||||
"""Local turn analyzer using the smart-turn-v3 ONNX model.
|
"""Local turn analyzer using the smart-turn-v3 ONNX model.
|
||||||
|
|
||||||
@@ -57,8 +49,13 @@ class LocalSmartTurnAnalyzerV3(BaseSmartTurn):
|
|||||||
|
|
||||||
logger.debug("Loading Local Smart Turn v3 model...")
|
logger.debug("Loading Local Smart Turn v3 model...")
|
||||||
|
|
||||||
|
so = ort.SessionOptions()
|
||||||
|
so.execution_mode = ort.ExecutionMode.ORT_SEQUENTIAL
|
||||||
|
so.inter_op_num_threads = 1
|
||||||
|
so.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
|
||||||
|
|
||||||
self._feature_extractor = WhisperFeatureExtractor(chunk_length=8)
|
self._feature_extractor = WhisperFeatureExtractor(chunk_length=8)
|
||||||
self._session = build_session(smart_turn_model_path)
|
self._session = ort.InferenceSession(smart_turn_model_path, sess_options=so)
|
||||||
|
|
||||||
logger.debug("Loaded Local Smart Turn v3")
|
logger.debug("Loaded Local Smart Turn v3")
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user