Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -216,6 +216,11 @@ def __init__(
ws_base_url: str,
session: aiohttp.ClientSession,
language: LanguageCode,
turn_start_threshold: NotGivenOr[float] = NOT_GIVEN,
turn_eager_end_threshold: NotGivenOr[float] = NOT_GIVEN,
turn_end_threshold: NotGivenOr[float] = NOT_GIVEN,
turn_end_timeout_ms: NotGivenOr[int] = NOT_GIVEN,
keyterm: NotGivenOr[list[str]] = NOT_GIVEN,
) -> None:
super().__init__(stt=stt, conn_options=conn_options, sample_rate=sample_rate)
self._encoding = encoding
Expand All @@ -226,6 +231,11 @@ def __init__(
self._ws_base_url = ws_base_url
self._session = session
self._language = language
self._turn_start_threshold = turn_start_threshold
self._turn_eager_end_threshold = turn_eager_end_threshold
self._turn_end_threshold = turn_end_threshold
self._turn_end_timeout_ms = turn_end_timeout_ms
self._keyterm = keyterm
self._request_id = ""
self._speaking = False
self._speech_duration: float = 0.0
Expand Down Expand Up @@ -359,11 +369,22 @@ async def recv_task(ws: aiohttp.ClientWebSocketResponse) -> None:
await ws.close()

async def _connect_ws(self) -> aiohttp.ClientWebSocketResponse:
params = {
"model": self._model,
"sample_rate": str(self._sample_rate),
"encoding": self._encoding,
}
# keyterm may repeat, so params is a list of pairs rather than a dict
params: list[tuple[str, str]] = [
("model", str(self._model)),
("sample_rate", str(self._sample_rate)),
("encoding", str(self._encoding)),
]
if utils.is_given(self._turn_start_threshold):
params.append(("turn_start_threshold", str(self._turn_start_threshold)))
if utils.is_given(self._turn_eager_end_threshold):
params.append(("turn_eager_end_threshold", str(self._turn_eager_end_threshold)))
if utils.is_given(self._turn_end_threshold):
params.append(("turn_end_threshold", str(self._turn_end_threshold)))
if utils.is_given(self._turn_end_timeout_ms):
params.append(("turn_end_timeout_ms", str(self._turn_end_timeout_ms)))
if utils.is_given(self._keyterm):
params.extend(("keyterm", term) for term in self._keyterm)

ws_url = f"{self._ws_base_url}/stt/turns/websocket?{urlencode(params)}"

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -104,6 +104,11 @@ def __init__(
base_url: str = "https://api.cartesia.ai",
language: STTLanguages | str | None = None,
encoding: STTEncoding = AUDIO_ENCODING,
turn_start_threshold: NotGivenOr[float] = NOT_GIVEN,
turn_eager_end_threshold: NotGivenOr[float] = NOT_GIVEN,
turn_end_threshold: NotGivenOr[float] = NOT_GIVEN,
turn_end_timeout_ms: NotGivenOr[int] = NOT_GIVEN,
keyterm: NotGivenOr[list[str]] = NOT_GIVEN,
) -> None:
"""
Create a new instance of Cartesia STT.
Expand Down Expand Up @@ -135,9 +140,27 @@ def __init__(
language: The language code for recognition.
This plugin only supports ``en`` for ``ink-2``.
encoding: The audio encoding format. Must be ``pcm_s16le``.
turn_start_threshold: Likelihood above which the model starts a turn.
Range 0.5-0.9, Cartesia default 0.8. Must stay above
``turn_eager_end_threshold``. Turn-detecting models (e.g. ``ink-2``) only.
turn_eager_end_threshold: Likelihood below which the model emits
``turn.eager_end``, the early might-be-done signal. Range 0.3-0.6,
Cartesia default 0.4. Must stay between the end and start thresholds.
Turn-detecting models only.
turn_end_threshold: Likelihood below which the model ends the turn.
Range 0.05-0.5, Cartesia default 0.2. Must stay below
``turn_eager_end_threshold``. Turn-detecting models only.
turn_end_timeout_ms: Maximum time in milliseconds to wait after the user
stops speaking before ending the turn even if the likelihood never
falls below ``turn_end_threshold``. Range 640-11200, Cartesia default
5600. Turn-detecting models only.
keyterm: Key terms to improve recall of specific words and phrases
(up to 100 terms totaling 1200 characters). Turn-detecting models only.

Raises:
ValueError: If no API key is provided or found in environment variables.
ValueError: If no API key is provided or found in environment variables,
or if a turn-detection parameter is passed with a model that does not
support turn detection (e.g. ``ink-whisper``).

Examples:

Expand Down Expand Up @@ -180,6 +203,21 @@ def __init__(
else:
resolved_final_transcript_mode = "auto"

if resolved_final_transcript_mode == "legacy":
_turns_only_params = {
"turn_start_threshold": turn_start_threshold,
"turn_eager_end_threshold": turn_eager_end_threshold,
"turn_end_threshold": turn_end_threshold,
"turn_end_timeout_ms": turn_end_timeout_ms,
"keyterm": keyterm,
}
for _param_name, _param_value in _turns_only_params.items():
if utils.is_given(_param_value):
raise ValueError(
f"The {_param_name!r} parameter is only supported by turn-detecting"
f" models (e.g. ink-2); model {resolved_model!r} does not support it."
)

super().__init__(
capabilities=stt.STTCapabilities(
streaming=True,
Expand All @@ -199,6 +237,11 @@ def __init__(
self._sample_rate = sample_rate
self._session = http_session
self._ws_base_url = _base_url_to_ws_base_url(base_url=base_url)
self._turn_start_threshold = turn_start_threshold
self._turn_eager_end_threshold = turn_eager_end_threshold
self._turn_end_threshold = turn_end_threshold
self._turn_end_timeout_ms = turn_end_timeout_ms
self._keyterm = keyterm

self._streams = weakref.WeakSet[CartesiaRecognizeStream]()

Expand Down Expand Up @@ -258,6 +301,11 @@ def stream(
ws_base_url=self._ws_base_url,
session=session,
language=resolved_language or LanguageCode("en"),
turn_start_threshold=self._turn_start_threshold,
turn_eager_end_threshold=self._turn_eager_end_threshold,
turn_end_threshold=self._turn_end_threshold,
turn_end_timeout_ms=self._turn_end_timeout_ms,
keyterm=self._keyterm,
)
case "legacy":
stream = LegacyRecognizeStream(
Expand Down