Files
VoiceAgent/pocket_tts_service.py

98 lines
3.3 KiB
Python

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