Update RimeTTSService InputParams for arcana and mistv2 model support
Add model-specific params (arcana: repetition_penalty, temperature, top_p; mistv2: no_text_normalization, save_oovs, segment) with dynamic query param building via _build_settings(). Model/voice/param changes now trigger WebSocket reconnection since all settings are URL query params. Co-Authored-By: Claude Opus 4.6 <noreply@anthropic.com>
This commit is contained in:
@@ -25,6 +25,7 @@ from pipecat.runner.utils import create_transport
|
|||||||
from pipecat.services.deepgram.stt import DeepgramSTTService
|
from pipecat.services.deepgram.stt import DeepgramSTTService
|
||||||
from pipecat.services.openai.llm import OpenAILLMService
|
from pipecat.services.openai.llm import OpenAILLMService
|
||||||
from pipecat.services.rime.tts import RimeTTSService
|
from pipecat.services.rime.tts import RimeTTSService
|
||||||
|
from pipecat.transcriptions.language import Language
|
||||||
from pipecat.transports.base_transport import BaseTransport, TransportParams
|
from pipecat.transports.base_transport import BaseTransport, TransportParams
|
||||||
from pipecat.transports.daily.transport import DailyParams
|
from pipecat.transports.daily.transport import DailyParams
|
||||||
from pipecat.transports.websocket.fastapi import FastAPIWebsocketParams
|
from pipecat.transports.websocket.fastapi import FastAPIWebsocketParams
|
||||||
@@ -56,7 +57,13 @@ 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",
|
||||||
|
params=RimeTTSService.InputParams(
|
||||||
|
language=Language.EN,
|
||||||
|
repetition_penalty=1.0,
|
||||||
|
temperature=0.5,
|
||||||
|
top_p=0.9,
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
llm = OpenAILLMService(api_key=os.getenv("OPENAI_API_KEY"))
|
llm = OpenAILLMService(api_key=os.getenv("OPENAI_API_KEY"))
|
||||||
|
|||||||
@@ -81,17 +81,31 @@ 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,
|
||||||
@@ -142,26 +156,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
|
||||||
@@ -188,14 +190,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:
|
||||||
@@ -223,12 +271,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()
|
||||||
|
|
||||||
@@ -255,7 +359,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):
|
||||||
@@ -301,7 +405,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)
|
||||||
|
|||||||
Reference in New Issue
Block a user