feat: integrate JV Voice Profile cloning, Voicebox research, and ensure session voice/markdown prompt nudge
This commit is contained in:
@@ -0,0 +1,97 @@
|
||||
"""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()
|
||||
Reference in New Issue
Block a user