Merge pull request #3642 from pipecat-ai/cb/rime-arcana-v3
Update RimeTTSService for arcana and mistv2 model support
This commit is contained in:
1
changelog/3642.added.md
Normal file
1
changelog/3642.added.md
Normal file
@@ -0,0 +1 @@
|
|||||||
|
- Added model-specific `InputParams` to `RimeTTSService`: arcana params (`repetition_penalty`, `temperature`, `top_p`) and mistv2 params (`no_text_normalization`, `save_oovs`, `segment`). Model, voice, and param changes now trigger WebSocket reconnection.
|
||||||
1
changelog/3642.changed.md
Normal file
1
changelog/3642.changed.md
Normal file
@@ -0,0 +1 @@
|
|||||||
|
- ⚠️ `RimeTTSService` now defaults to `model="arcana"` and the `wss://users-ws.rime.ai/ws3` endpoint. `InputParams` defaults changed from mistv2-specific values to `None` — only explicitly-set params are sent as query params.
|
||||||
@@ -56,7 +56,7 @@ async def run_bot(transport: BaseTransport, runner_args: RunnerArguments):
|
|||||||
|
|
||||||
tts = RimeTTSService(
|
tts = RimeTTSService(
|
||||||
api_key=os.getenv("RIME_API_KEY", ""),
|
api_key=os.getenv("RIME_API_KEY", ""),
|
||||||
voice_id="rex",
|
voice_id="luna",
|
||||||
)
|
)
|
||||||
|
|
||||||
llm = OpenAILLMService(api_key=os.getenv("OPENAI_API_KEY"))
|
llm = OpenAILLMService(api_key=os.getenv("OPENAI_API_KEY"))
|
||||||
|
|||||||
@@ -82,25 +82,39 @@ class RimeTTSService(AudioContextWordTTSService):
|
|||||||
|
|
||||||
Parameters:
|
Parameters:
|
||||||
language: Language for synthesis. Defaults to English.
|
language: Language for synthesis. Defaults to English.
|
||||||
speed_alpha: Speech speed multiplier. Defaults to 1.0.
|
segment: Text segmentation mode ("immediate", "bySentence", "never").
|
||||||
reduce_latency: Whether to reduce latency at potential quality cost.
|
repetition_penalty: Token repetition penalty (arcana only).
|
||||||
pause_between_brackets: Whether to add pauses between bracketed content.
|
temperature: Sampling temperature (arcana only).
|
||||||
phonemize_between_brackets: Whether to phonemize bracketed content.
|
top_p: Cumulative probability threshold (arcana only).
|
||||||
|
speed_alpha: Speech speed multiplier (mistv2 only).
|
||||||
|
reduce_latency: Whether to reduce latency at potential quality cost (mistv2 only).
|
||||||
|
pause_between_brackets: Whether to add pauses between bracketed content (mistv2 only).
|
||||||
|
phonemize_between_brackets: Whether to phonemize bracketed content (mistv2 only).
|
||||||
|
no_text_normalization: Whether to disable text normalization (mistv2 only).
|
||||||
|
save_oovs: Whether to save out-of-vocabulary words (mistv2 only).
|
||||||
"""
|
"""
|
||||||
|
|
||||||
language: Optional[Language] = Language.EN
|
language: Optional[Language] = Language.EN
|
||||||
speed_alpha: Optional[float] = 1.0
|
segment: Optional[str] = None
|
||||||
reduce_latency: Optional[bool] = False
|
# Arcana params
|
||||||
pause_between_brackets: Optional[bool] = False
|
repetition_penalty: Optional[float] = None
|
||||||
phonemize_between_brackets: Optional[bool] = False
|
temperature: Optional[float] = None
|
||||||
|
top_p: Optional[float] = None
|
||||||
|
# Mistv2 params
|
||||||
|
speed_alpha: Optional[float] = None
|
||||||
|
reduce_latency: Optional[bool] = None
|
||||||
|
pause_between_brackets: Optional[bool] = None
|
||||||
|
phonemize_between_brackets: Optional[bool] = None
|
||||||
|
no_text_normalization: Optional[bool] = None
|
||||||
|
save_oovs: Optional[bool] = None
|
||||||
|
|
||||||
def __init__(
|
def __init__(
|
||||||
self,
|
self,
|
||||||
*,
|
*,
|
||||||
api_key: str,
|
api_key: str,
|
||||||
voice_id: str,
|
voice_id: str,
|
||||||
url: str = "wss://users.rime.ai/ws2",
|
url: str = "wss://users-ws.rime.ai/ws3",
|
||||||
model: str = "mistv2",
|
model: str = "arcana",
|
||||||
sample_rate: Optional[int] = None,
|
sample_rate: Optional[int] = None,
|
||||||
params: Optional[InputParams] = None,
|
params: Optional[InputParams] = None,
|
||||||
text_aggregator: Optional[BaseTextAggregator] = None,
|
text_aggregator: Optional[BaseTextAggregator] = None,
|
||||||
@@ -143,26 +157,14 @@ class RimeTTSService(AudioContextWordTTSService):
|
|||||||
# and insert these tags for the purpose of the TTS service alone.
|
# and insert these tags for the purpose of the TTS service alone.
|
||||||
self._text_aggregator = SkipTagsAggregator([("spell(", ")")])
|
self._text_aggregator = SkipTagsAggregator([("spell(", ")")])
|
||||||
|
|
||||||
params = params or RimeTTSService.InputParams()
|
self._params = params or RimeTTSService.InputParams()
|
||||||
|
|
||||||
# Store service configuration
|
# Store service configuration
|
||||||
self._api_key = api_key
|
self._api_key = api_key
|
||||||
self._url = url
|
self._url = url
|
||||||
self._voice_id = voice_id
|
self._voice_id = voice_id
|
||||||
self._model = model
|
self._model = model
|
||||||
self._settings = {
|
self._settings = self._build_settings()
|
||||||
"speaker": voice_id,
|
|
||||||
"modelId": model,
|
|
||||||
"audioFormat": "pcm",
|
|
||||||
"samplingRate": 0,
|
|
||||||
"lang": self.language_to_service_language(params.language)
|
|
||||||
if params.language
|
|
||||||
else "eng",
|
|
||||||
"speedAlpha": params.speed_alpha,
|
|
||||||
"reduceLatency": params.reduce_latency,
|
|
||||||
"pauseBetweenBrackets": json.dumps(params.pause_between_brackets),
|
|
||||||
"phonemizeBetweenBrackets": json.dumps(params.phonemize_between_brackets),
|
|
||||||
}
|
|
||||||
|
|
||||||
# State tracking
|
# State tracking
|
||||||
self._context_id = None # Tracks current turn
|
self._context_id = None # Tracks current turn
|
||||||
@@ -189,14 +191,60 @@ class RimeTTSService(AudioContextWordTTSService):
|
|||||||
"""
|
"""
|
||||||
return language_to_rime_language(language)
|
return language_to_rime_language(language)
|
||||||
|
|
||||||
|
def _build_settings(self) -> dict:
|
||||||
|
"""Build query params for the WebSocket URL based on the current model and params.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dictionary of query parameters. Only explicitly-set values are included.
|
||||||
|
"""
|
||||||
|
settings = {
|
||||||
|
"speaker": self._voice_id,
|
||||||
|
"modelId": self._model,
|
||||||
|
"audioFormat": "pcm",
|
||||||
|
"samplingRate": self.sample_rate or 0,
|
||||||
|
}
|
||||||
|
if self._params.language:
|
||||||
|
settings["lang"] = self.language_to_service_language(self._params.language) or "eng"
|
||||||
|
if self._params.segment is not None:
|
||||||
|
settings["segment"] = self._params.segment
|
||||||
|
|
||||||
|
if self._model == "arcana":
|
||||||
|
if self._params.repetition_penalty is not None:
|
||||||
|
settings["repetition_penalty"] = self._params.repetition_penalty
|
||||||
|
if self._params.temperature is not None:
|
||||||
|
settings["temperature"] = self._params.temperature
|
||||||
|
if self._params.top_p is not None:
|
||||||
|
settings["top_p"] = self._params.top_p
|
||||||
|
else: # mistv2/mist
|
||||||
|
if self._params.speed_alpha is not None:
|
||||||
|
settings["speedAlpha"] = self._params.speed_alpha
|
||||||
|
if self._params.reduce_latency is not None:
|
||||||
|
settings["reduceLatency"] = self._params.reduce_latency
|
||||||
|
if self._params.pause_between_brackets is not None:
|
||||||
|
settings["pauseBetweenBrackets"] = json.dumps(self._params.pause_between_brackets)
|
||||||
|
if self._params.phonemize_between_brackets is not None:
|
||||||
|
settings["phonemizeBetweenBrackets"] = json.dumps(
|
||||||
|
self._params.phonemize_between_brackets
|
||||||
|
)
|
||||||
|
if self._params.no_text_normalization is not None:
|
||||||
|
settings["noTextNormalization"] = json.dumps(self._params.no_text_normalization)
|
||||||
|
if self._params.save_oovs is not None:
|
||||||
|
settings["saveOovs"] = json.dumps(self._params.save_oovs)
|
||||||
|
|
||||||
|
return settings
|
||||||
|
|
||||||
async def set_model(self, model: str):
|
async def set_model(self, model: str):
|
||||||
"""Update the TTS model.
|
"""Update the TTS model and reconnect.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
model: The model name to use for synthesis.
|
model: The model name to use for synthesis.
|
||||||
"""
|
"""
|
||||||
self._model = model
|
self._model = model
|
||||||
|
self._settings = self._build_settings()
|
||||||
await super().set_model(model)
|
await super().set_model(model)
|
||||||
|
if self._websocket:
|
||||||
|
await self._disconnect()
|
||||||
|
await self._connect()
|
||||||
|
|
||||||
# A set of Rime-specific helpers for text transformations
|
# A set of Rime-specific helpers for text transformations
|
||||||
def SPELL(text: str) -> str:
|
def SPELL(text: str) -> str:
|
||||||
@@ -224,12 +272,68 @@ class RimeTTSService(AudioContextWordTTSService):
|
|||||||
return f"[{text}]"
|
return f"[{text}]"
|
||||||
|
|
||||||
async def _update_settings(self, settings: Mapping[str, Any]):
|
async def _update_settings(self, settings: Mapping[str, Any]):
|
||||||
"""Update service settings and reconnect if voice changed."""
|
"""Update service settings and reconnect if necessary.
|
||||||
prev_voice = self._voice_id
|
|
||||||
|
Since all settings are WebSocket URL query parameters,
|
||||||
|
any setting change requires reconnecting to apply the new values.
|
||||||
|
"""
|
||||||
|
prev_settings = self._settings.copy()
|
||||||
await super()._update_settings(settings)
|
await super()._update_settings(settings)
|
||||||
if not prev_voice == self._voice_id:
|
|
||||||
|
needs_reconnect = False
|
||||||
|
|
||||||
|
if "voice" in settings or "voice_id" in settings:
|
||||||
self._settings["speaker"] = self._voice_id
|
self._settings["speaker"] = self._voice_id
|
||||||
logger.info(f"Switching TTS voice to: [{self._voice_id}]")
|
if prev_settings.get("speaker") != self._voice_id:
|
||||||
|
logger.info(f"Switching TTS voice to: [{self._voice_id}]")
|
||||||
|
needs_reconnect = True
|
||||||
|
|
||||||
|
if "model" in settings:
|
||||||
|
self._settings = self._build_settings()
|
||||||
|
needs_reconnect = True
|
||||||
|
|
||||||
|
if "language" in settings:
|
||||||
|
new_lang = self.language_to_service_language(settings["language"])
|
||||||
|
if new_lang and new_lang != prev_settings.get("lang"):
|
||||||
|
logger.info(f"Updating language to: [{new_lang}]")
|
||||||
|
self._settings["lang"] = new_lang
|
||||||
|
needs_reconnect = True
|
||||||
|
|
||||||
|
# Arcana params
|
||||||
|
for key, settings_key in [
|
||||||
|
("repetition_penalty", "repetition_penalty"),
|
||||||
|
("temperature", "temperature"),
|
||||||
|
("top_p", "top_p"),
|
||||||
|
]:
|
||||||
|
if key in settings and settings[key] != prev_settings.get(settings_key):
|
||||||
|
self._settings[settings_key] = settings[key]
|
||||||
|
needs_reconnect = True
|
||||||
|
|
||||||
|
# Mistv2 params
|
||||||
|
for key, settings_key in [
|
||||||
|
("speed_alpha", "speedAlpha"),
|
||||||
|
("reduce_latency", "reduceLatency"),
|
||||||
|
]:
|
||||||
|
if key in settings and settings[key] != prev_settings.get(settings_key):
|
||||||
|
self._settings[settings_key] = settings[key]
|
||||||
|
needs_reconnect = True
|
||||||
|
|
||||||
|
# Mistv2 boolean params (need json.dumps)
|
||||||
|
for key, settings_key in [
|
||||||
|
("pause_between_brackets", "pauseBetweenBrackets"),
|
||||||
|
("phonemize_between_brackets", "phonemizeBetweenBrackets"),
|
||||||
|
("no_text_normalization", "noTextNormalization"),
|
||||||
|
("save_oovs", "saveOovs"),
|
||||||
|
]:
|
||||||
|
if key in settings and json.dumps(settings[key]) != prev_settings.get(settings_key):
|
||||||
|
self._settings[settings_key] = json.dumps(settings[key])
|
||||||
|
needs_reconnect = True
|
||||||
|
|
||||||
|
if "segment" in settings and settings["segment"] != prev_settings.get("segment"):
|
||||||
|
self._settings["segment"] = settings["segment"]
|
||||||
|
needs_reconnect = True
|
||||||
|
|
||||||
|
if needs_reconnect and self._websocket:
|
||||||
await self._disconnect()
|
await self._disconnect()
|
||||||
await self._connect()
|
await self._connect()
|
||||||
|
|
||||||
@@ -256,7 +360,7 @@ class RimeTTSService(AudioContextWordTTSService):
|
|||||||
frame: The start frame containing initialization parameters.
|
frame: The start frame containing initialization parameters.
|
||||||
"""
|
"""
|
||||||
await super().start(frame)
|
await super().start(frame)
|
||||||
self._settings["samplingRate"] = self.sample_rate
|
self._settings = self._build_settings()
|
||||||
await self._connect()
|
await self._connect()
|
||||||
|
|
||||||
async def stop(self, frame: EndFrame):
|
async def stop(self, frame: EndFrame):
|
||||||
@@ -302,7 +406,7 @@ class RimeTTSService(AudioContextWordTTSService):
|
|||||||
if self._websocket and self._websocket.state is State.OPEN:
|
if self._websocket and self._websocket.state is State.OPEN:
|
||||||
return
|
return
|
||||||
|
|
||||||
params = "&".join(f"{k}={v}" for k, v in self._settings.items())
|
params = "&".join(f"{k}={v}" for k, v in self._settings.items() if v is not None)
|
||||||
url = f"{self._url}?{params}"
|
url = f"{self._url}?{params}"
|
||||||
headers = {"Authorization": f"Bearer {self._api_key}"}
|
headers = {"Authorization": f"Bearer {self._api_key}"}
|
||||||
self._websocket = await websocket_connect(url, additional_headers=headers)
|
self._websocket = await websocket_connect(url, additional_headers=headers)
|
||||||
@@ -398,7 +502,7 @@ class RimeTTSService(AudioContextWordTTSService):
|
|||||||
async for message in self._get_websocket():
|
async for message in self._get_websocket():
|
||||||
msg = json.loads(message)
|
msg = json.loads(message)
|
||||||
|
|
||||||
if not msg or not self.audio_context_available(msg["contextId"]):
|
if not msg or not self.audio_context_available(msg.get("contextId")):
|
||||||
continue
|
continue
|
||||||
|
|
||||||
context_id = msg["contextId"]
|
context_id = msg["contextId"]
|
||||||
@@ -641,20 +745,18 @@ class RimeHttpTTSService(TTSService):
|
|||||||
class RimeNonJsonTTSService(InterruptibleTTSService):
|
class RimeNonJsonTTSService(InterruptibleTTSService):
|
||||||
"""Pipecat TTS service for Rime's non-JSON WebSocket API.
|
"""Pipecat TTS service for Rime's non-JSON WebSocket API.
|
||||||
|
|
||||||
|
.. deprecated:: 0.0.102
|
||||||
|
Arcana now supports JSON WebSocket with word-level timestamps via the
|
||||||
|
``wss://users-ws.rime.ai/ws3`` endpoint. Use :class:`RimeTTSService`
|
||||||
|
with ``model="arcana"`` instead.
|
||||||
|
|
||||||
This service enables Text-to-Speech synthesis over WebSocket endpoints
|
This service enables Text-to-Speech synthesis over WebSocket endpoints
|
||||||
that require plain text (not JSON) messages and return raw audio bytes.
|
that require plain text (not JSON) messages and return raw audio bytes.
|
||||||
It is designed for use with TTS models like Arcana, which currently do
|
|
||||||
not support JSON-based WebSocket protocols (though this may change in
|
|
||||||
the future).
|
|
||||||
|
|
||||||
Limitations:
|
Limitations:
|
||||||
- Does not support word-level timestamps or context IDs.
|
- Does not support word-level timestamps or context IDs.
|
||||||
- Intended specifically for integrations where the TTS provider only
|
- Intended specifically for integrations where the TTS provider only
|
||||||
accepts and returns non-JSON messages.
|
accepts and returns non-JSON messages.
|
||||||
|
|
||||||
Note:
|
|
||||||
- Arcana and similar models may add JSON WebSocket support in the
|
|
||||||
future. This service focuses on the current plain text protocol.
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
class InputParams(BaseModel):
|
class InputParams(BaseModel):
|
||||||
|
|||||||
Reference in New Issue
Block a user