Files
VoiceAgent/qwen_tts_service.py
T

119 lines
4.3 KiB
Python

"""qwen_tts_service.py
Qwen3-TTS Service for Pipecat using MLX on Apple Silicon.
Provides local speech synthesis using Alibaba Qwen3-TTS 0.6B
with zero-shot voice cloning capabilities via mlx-audio.
"""
import sys
import os
import asyncio
import numpy as np
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"
JV_SAMPLE_WAV = Path.home() / "Library/Application Support/sh.voicebox.app/profiles/2f4e8f2e-dbc4-43fc-b940-c6ca7ac694c7/c496be86-267d-4793-b26a-7beead57fbd4.wav"
JV_SAMPLE_TEXT = "I have completed a diagnostic scan of your current schedule, and it appears several conflicts have arisen. While I have taken the liberty of reorganizing your morning appointments to ensure maximum efficiency, I cannot account for human fatigue. Perhaps a second cup of coffee would be a logical next step."
MODEL_ID = "mlx-community/Qwen3-TTS-12Hz-0.6B-Base-bf16"
class QwenTTSService(TTSService):
def __init__(
self,
*,
voice: str = "qwen_jv",
model_id: str = MODEL_ID,
sample_rate: int = 24000,
temperature: float = 0.3,
**kwargs,
):
super().__init__(sample_rate=sample_rate, **kwargs)
self._voice_name = voice
self._model_id = model_id
self._temperature = temperature
self._model = None
def _ensure_model_loaded(self):
if self._model is not None:
return
from mlx_audio.tts import load_model
logger.info(f"Loading local MLX Qwen3-TTS model ({self._model_id})...")
self._model = load_model(self._model_id)
logger.info("Qwen3-TTS model loaded successfully on Metal GPU.")
def set_voice(self, voice: str):
self._voice_name = voice
logger.info(f"QwenTTSService 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)
self._ensure_model_loaded()
from mlx_audio.tts.generate import generate_audio
ref_audio = str(JV_SAMPLE_WAV) if JV_SAMPLE_WAV.exists() else None
ref_text = JV_SAMPLE_TEXT if ref_audio else None
# Generate audio using MLX Qwen3-TTS in background thread
loop = asyncio.get_running_loop()
output_file = "/tmp/qwen_gen_temp.wav"
await loop.run_in_executor(
None,
lambda: generate_audio(
text=text,
model=self._model,
ref_audio=ref_audio,
ref_text=ref_text,
temperature=self._temperature,
output_path=output_file,
verbose=False,
),
)
await self.stop_ttfb_metrics()
actual_wav = None
if os.path.isdir(output_file):
sub_wav = os.path.join(output_file, "audio_000.wav")
if os.path.exists(sub_wav):
actual_wav = sub_wav
elif os.path.exists(output_file):
actual_wav = output_file
if actual_wav:
import scipy.io.wavfile as wavfile
sr, audio_data = wavfile.read(actual_wav)
if audio_data.dtype != np.int16:
audio_data = (np.clip(audio_data, -1.0, 1.0) * 32767).astype(np.int16)
audio_bytes = audio_data.tobytes()
if os.path.isdir(output_file):
import shutil
shutil.rmtree(output_file, ignore_errors=True)
else:
try:
os.remove(output_file)
except Exception:
pass
yield TTSAudioRawFrame(
audio=audio_bytes,
sample_rate=sr,
num_channels=1,
context_id=context_id,
)
except Exception as e:
logger.error(f"Error in QwenTTSService: {e}")
yield ErrorFrame(error=f"Qwen TTS error: {e}")
finally:
await self.stop_ttfb_metrics()