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:
Mark Backman
2026-02-12 14:01:41 -05:00
parent 794811fbdb
commit 2b9777b812
2 changed files with 141 additions and 30 deletions

View File

@@ -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"))

View File

@@ -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)