Files
vibe-bot/vibe_bot/tts.py
T
2026-08-19 13:16:43 -04:00

129 lines
3.6 KiB
Python

"""Text-to-speech engine using Kokoro TTS."""
from __future__ import annotations
import logging
from dataclasses import dataclass
from io import BytesIO
import numpy as np
# kokoro-tts and soundfile ship no type stubs upstream.
import soundfile as sf # type: ignore[import-untyped]
from kokoro_tts import ( # type: ignore[import-untyped]
Kokoro,
chunk_text,
process_chunk_sequential,
)
from vibe_bot.config import TTS_SPEED, TTS_VOICE
logger = logging.getLogger(__name__)
# Default voice settings (single source of truth: vibe_bot.config).
DEFAULT_VOICE = TTS_VOICE
DEFAULT_SPEED = TTS_SPEED
DEFAULT_LANG = "en-us"
@dataclass
class AudioResult:
"""Audio output from a TTS generation.
Attributes:
audio: The encoded audio (MP3) as a seekable BytesIO.
partial: True if one or more text chunks failed to produce audio.
failed_chunks: Number of text chunks that failed.
"""
audio: BytesIO
partial: bool
failed_chunks: int
class TTSEngine:
"""Text-to-speech engine wrapper around Kokoro TTS."""
def __init__(self, model_path: str, voices_path: str) -> None:
"""Initialize the TTS engine with model and voices paths.
Args:
model_path: Path to the Kokoro model file.
voices_path: Path to the voices file.
"""
self.model_path = model_path
self.voices_path = voices_path
self.kokoro = Kokoro(model_path, voices_path)
logger.info("Kokoro TTS engine initialized")
def generate_audio(
self,
text: str,
voice: str = DEFAULT_VOICE,
speed: float = DEFAULT_SPEED,
lang: str = DEFAULT_LANG,
) -> AudioResult:
"""Convert text to audio and return an AudioResult (MP3 in .audio)."""
all_samples: list[np.ndarray] = []
sample_rate: int | None = None
failed_chunks = 0
chunks: list[str] = list(chunk_text(text))
logger.debug("Split text into %d chunks", len(chunks))
for i, chunk in enumerate(chunks):
try:
samples, sr = process_chunk_sequential(
chunk,
self.kokoro,
voice,
speed,
lang,
)
except Exception:
logger.exception("Error processing chunk %d", i + 1)
failed_chunks += 1
continue
if samples is None:
logger.warning("Chunk %d/%d produced no audio", i + 1, len(chunks))
failed_chunks += 1
continue
if sample_rate is None:
sample_rate = sr
all_samples.append(np.asarray(samples))
logger.debug("Processed chunk %d/%d", i + 1, len(chunks))
if not all_samples:
msg = "No audio samples generated - text may be invalid or too long"
raise ValueError(msg)
combined = np.concatenate(all_samples)
buffer = BytesIO()
sf.write( # pyright: ignore[reportUnknownMemberType]
buffer,
combined,
sample_rate,
format="MP3",
subtype="MPEG_LAYER_III",
)
buffer.seek(0)
partial = failed_chunks > 0
if partial:
logger.warning(
"TTS produced partial audio: %d of %d chunks failed",
failed_chunks,
len(chunks),
)
logger.debug(
"Generated MP3 audio: %d samples at %dHz",
len(combined),
sample_rate or 0,
)
return AudioResult(audio=buffer, partial=partial, failed_chunks=failed_chunks)