Bundle Smart Turn v3 with Pipecat

This commit is contained in:
marcus-daily
2025-09-11 20:03:55 +01:00
committed by Marcus
parent dac58deffc
commit d2f210e960
4 changed files with 25 additions and 6 deletions

View File

@@ -155,6 +155,7 @@ where = ["src"]
"src/pipecat/audio/dtmf/dtmf-star.wav", "src/pipecat/audio/dtmf/dtmf-star.wav",
] ]
"pipecat.services.aws_nova_sonic" = ["src/pipecat/services/aws_nova_sonic/ready.wav"] "pipecat.services.aws_nova_sonic" = ["src/pipecat/services/aws_nova_sonic/ready.wav"]
"pipecat.audio.turn.smart_turn.data" = ["src/pipecat/audio/turn/smart_turn/data/smart-turn-v3.0.onnx"]
[tool.pytest.ini_options] [tool.pytest.ini_options]
addopts = "--verbose" addopts = "--verbose"

View File

@@ -10,7 +10,7 @@ This module provides a smart turn analyzer that uses an ONNX model for
local end-of-turn detection without requiring network connectivity. local end-of-turn detection without requiring network connectivity.
""" """
from typing import Any, Dict from typing import Any, Dict, Optional
import numpy as np import numpy as np
from loguru import logger from loguru import logger
@@ -35,20 +35,38 @@ class LocalSmartTurnAnalyzerV3(BaseSmartTurn):
enabling offline operation without network dependencies. enabling offline operation without network dependencies.
""" """
def __init__(self, *, smart_turn_model_path: str, **kwargs): def __init__(self, *, smart_turn_model_path: Optional[str] = None, **kwargs):
"""Initialize the local ONNX smart-turn-v3 analyzer. """Initialize the local ONNX smart-turn-v3 analyzer.
Args: Args:
smart_turn_model_path: Path to the ONNX model file. smart_turn_model_path: Path to the ONNX model file. If this is not
set, the bundled smart-turn-v3.0 model will be used.
**kwargs: Additional arguments passed to BaseSmartTurn. **kwargs: Additional arguments passed to BaseSmartTurn.
""" """
super().__init__(**kwargs) super().__init__(**kwargs)
if not smart_turn_model_path:
raise ValueError("smart_turn_model_path must be provided")
logger.debug("Loading Local Smart Turn v3 model...") logger.debug("Loading Local Smart Turn v3 model...")
if not smart_turn_model_path:
# Load bundled model
model_name = "smart-turn-v3.0.onnx"
package_path = "pipecat.audio.turn.smart_turn.data"
try:
import importlib_resources as impresources
smart_turn_model_path = str(impresources.files(package_path).joinpath(model_name))
except BaseException:
from importlib import resources as impresources
try:
with impresources.path(package_path, model_name) as f:
smart_turn_model_path = f
except BaseException:
smart_turn_model_path = str(
impresources.files(package_path).joinpath(model_name)
)
so = ort.SessionOptions() so = ort.SessionOptions()
so.execution_mode = ort.ExecutionMode.ORT_SEQUENTIAL so.execution_mode = ort.ExecutionMode.ORT_SEQUENTIAL
so.inter_op_num_threads = 1 so.inter_op_num_threads = 1