"""pocket_tts_service.py Kyutai Pocket TTS Service for Pipecat. Provides real-time local speech synthesis using Kyutai Pocket TTS with zero-shot voice cloning capabilities. """ import sys import os import asyncio import numpy as np import torch from pathlib import Path from collections.abc import AsyncGenerator from loguru import logger from pipecat.frames.frames import ErrorFrame, Frame, TTSAudioRawFrame from pipecat.services.tts_service import TTSService CUSTOM_VOICES_DIR = Path(__file__).resolve().parent / "custom_voices" class PocketTTSService(TTSService): def __init__( self, *, voice: str = "custom_pocket", sample_rate: int = 24000, **kwargs, ): super().__init__(sample_rate=sample_rate, **kwargs) self._voice_name = voice self._model = None self._voice_states = {} def _ensure_model_loaded(self): if self._model is not None: return from pocket_tts import TTSModel logger.info("Initializing Kyutai Pocket TTS model (temp=0.5, lsd_decode_steps=2)...") self._model = TTSModel.load_model(temp=0.5, lsd_decode_steps=2) logger.info("Kyutai Pocket TTS model loaded successfully.") def _get_voice_state(self, voice_name: str): self._ensure_model_loaded() if voice_name in self._voice_states: return self._voice_states[voice_name] # Check for saved custom voice clone state file (.pt) custom_file = CUSTOM_VOICES_DIR / f"{voice_name}.pt" if custom_file.exists(): logger.info(f"Loading custom Pocket TTS voice state from {custom_file.name}...") state = torch.load(custom_file) self._voice_states[voice_name] = state return state # Fallback to Pocket TTS built-in catalog voice logger.info(f"Loading Pocket TTS catalog voice '{voice_name}'...") state = self._model.get_state_for_audio_prompt(voice_name) self._voice_states[voice_name] = state return state def set_voice(self, voice: str): self._voice_name = voice logger.info(f"PocketTTSService active voice set to '{voice}'") async def run_tts(self, text: str, context_id: str) -> AsyncGenerator[Frame, None]: try: await self.start_tts_usage_metrics(text) voice_name = self._voice_name or "custom_pocket" state = self._get_voice_state(voice_name) # Generate audio tensor using Pocket TTS loop = asyncio.get_running_loop() audio_tensor = await loop.run_in_executor( None, lambda: self._model.generate_audio(state, text) ) await self.stop_ttfb_metrics() # Convert float tensor to 16-bit PCM bytes audio_np = audio_tensor.cpu().numpy() audio_int16 = (np.clip(audio_np, -1.0, 1.0) * 32767).astype(np.int16) audio_bytes = audio_int16.tobytes() yield TTSAudioRawFrame( audio=audio_bytes, sample_rate=self.sample_rate, num_channels=1, context_id=context_id, ) except Exception as e: logger.error(f"Error in PocketTTSService: {e}") yield ErrorFrame(error=f"Pocket TTS error: {e}") finally: await self.stop_ttfb_metrics()