263 lines
11 KiB
Python
263 lines
11 KiB
Python
"""Local on-device TTS engines for ``tools.tts_tool``: NeuTTS, Piper, KittenTTS.
|
|
|
|
All three synthesize WAV natively; :func:`_finalize_wav_output` (shared) then
|
|
converts/renames to the caller's requested container. Piper and KittenTTS keep
|
|
their loaded models in small LRU caches registered in
|
|
``_LOCAL_TTS_MODEL_CACHES`` so the origin module's warm/release lifecycle can
|
|
pre-load or drop them. ``_import_piper`` / ``_import_kittentts`` are resolved
|
|
through the origin module at call time so test monkeypatches there apply.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import logging
|
|
import subprocess
|
|
import sys
|
|
from pathlib import Path
|
|
from typing import Any, Callable, Dict, Tuple
|
|
|
|
from tools.tts_tool_delivery import _finalize_wav_output, _wav_sidecar_path
|
|
|
|
logger = logging.getLogger("tools.tts_tool")
|
|
|
|
DEFAULT_KITTENTTS_MODEL = "KittenML/kitten-tts-nano-0.8-int8" # 25MB
|
|
DEFAULT_KITTENTTS_VOICE = "Jasper"
|
|
DEFAULT_PIPER_VOICE = "en_US-lessac-medium" # balanced size/quality
|
|
|
|
|
|
def _origin():
|
|
from tools import tts_tool
|
|
|
|
return tts_tool
|
|
|
|
|
|
# ===========================================================================
|
|
# Bounded model caches
|
|
# ===========================================================================
|
|
# Each cached entry is a whole loaded model (tens of MB); an unbounded dict
|
|
# would pin one per distinct voice for the process lifetime. Small LRU: most
|
|
# sessions use one or two voices and a cold reload is cheap.
|
|
_TTS_MODEL_CACHE_MAX = 3
|
|
|
|
# Provider name → the model cache it populates. Consulted by
|
|
# warm_tts_provider() / release_tts_provider() in the origin module; a new
|
|
# local engine adds one row here plus a loader in _local_tts_warmers().
|
|
_LOCAL_TTS_MODEL_CACHES: Dict[str, Dict[str, Any]] = {}
|
|
|
|
# Piper voices keyed on absolute .onnx path (+cuda flag); KittenTTS on model name.
|
|
_piper_voice_cache: Dict[str, Any] = {}
|
|
_kittentts_model_cache: Dict[str, Any] = {}
|
|
_LOCAL_TTS_MODEL_CACHES["piper"] = _piper_voice_cache
|
|
_LOCAL_TTS_MODEL_CACHES["kittentts"] = _kittentts_model_cache
|
|
|
|
|
|
def _tts_cache_get_or_load(cache: Dict[str, Any], key: str, load: Callable[[], Any]) -> Any:
|
|
"""Get ``key`` from ``cache`` or load it, keeping the cache LRU-bounded.
|
|
|
|
A hit refreshes recency (pop + reinsert on the insertion-ordered dict); a
|
|
miss loads then evicts LRU entries beyond ``_TTS_MODEL_CACHE_MAX``. Callers
|
|
holding an evicted reference keep it alive; only the slot is released.
|
|
"""
|
|
if key in cache:
|
|
cache[key] = cache.pop(key)
|
|
return cache[key]
|
|
value = load()
|
|
cache[key] = value
|
|
while len(cache) > _TTS_MODEL_CACHE_MAX:
|
|
cache.pop(next(iter(cache)), None)
|
|
return value
|
|
|
|
|
|
# ===========================================================================
|
|
# NeuTTS (subprocess via tools/neutts_synth.py so the ~500MB model exits after use)
|
|
# ===========================================================================
|
|
|
|
def _default_neutts_ref_audio() -> str:
|
|
return str(Path(__file__).parent / "neutts_samples" / "jo.wav")
|
|
|
|
|
|
def _default_neutts_ref_text() -> str:
|
|
return str(Path(__file__).parent / "neutts_samples" / "jo.txt")
|
|
|
|
|
|
def _generate_neutts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str:
|
|
neutts_config = tts_config.get("neutts") or {}
|
|
wav_path = _wav_sidecar_path(output_path)
|
|
cmd = [
|
|
sys.executable, str(Path(__file__).parent / "neutts_synth.py"),
|
|
"--text", text,
|
|
"--out", wav_path,
|
|
"--ref-audio", neutts_config.get("ref_audio", "") or _default_neutts_ref_audio(),
|
|
"--ref-text", neutts_config.get("ref_text", "") or _default_neutts_ref_text(),
|
|
"--model", neutts_config.get("model", "neuphonic/neutts-air-q4-gguf"),
|
|
"--device", neutts_config.get("device", "cpu"),
|
|
]
|
|
result = subprocess.run(cmd, capture_output=True, text=True, encoding='utf-8', errors='replace', timeout=120, stdin=subprocess.DEVNULL)
|
|
if result.returncode != 0:
|
|
# The synth script reports success lines as "OK:" on stderr too.
|
|
error_lines = [l for l in result.stderr.strip().splitlines() if not l.startswith("OK:")]
|
|
raise RuntimeError(f"NeuTTS synthesis failed: {chr(10).join(error_lines) or 'unknown error'}")
|
|
return _finalize_wav_output(wav_path, output_path)
|
|
|
|
|
|
# ===========================================================================
|
|
# Piper (local neural VITS, 44 languages)
|
|
# ===========================================================================
|
|
|
|
def _get_piper_voices_dir() -> Path:
|
|
"""``<HERMES_HOME>/cache/piper-voices/`` so voice downloads follow profile boundaries."""
|
|
from hermes_constants import get_hermes_dir
|
|
root = Path(get_hermes_dir("cache/piper-voices", "piper_voices_cache"))
|
|
root.mkdir(parents=True, exist_ok=True)
|
|
return root
|
|
|
|
|
|
def _resolve_piper_voice_path(voice: str, download_dir: Path) -> str:
|
|
"""Resolve *voice* (an .onnx path or a voice name) to a concrete .onnx file.
|
|
|
|
Names like ``en_US-lessac-medium`` are downloaded into *download_dir* on
|
|
first use via ``python -m piper.download_voices``. Raises RuntimeError
|
|
when the model can't be located or downloaded.
|
|
"""
|
|
if not voice:
|
|
voice = DEFAULT_PIPER_VOICE
|
|
|
|
candidate = Path(voice).expanduser()
|
|
if candidate.suffix.lower() == ".onnx" and candidate.exists():
|
|
return str(candidate)
|
|
|
|
cached = download_dir / f"{voice}.onnx"
|
|
if cached.exists() and (download_dir / f"{voice}.onnx.json").exists():
|
|
return str(cached)
|
|
|
|
logger.info("[Piper] Downloading voice '%s' to %s (first use)", voice, download_dir)
|
|
try:
|
|
result = subprocess.run(
|
|
[sys.executable, "-m", "piper.download_voices", voice,
|
|
"--download-dir", str(download_dir)],
|
|
capture_output=True, text=True, encoding='utf-8', errors='replace', timeout=300,
|
|
stdin=subprocess.DEVNULL,
|
|
)
|
|
except subprocess.TimeoutExpired as exc:
|
|
raise RuntimeError(f"Piper voice download timed out after 300s for '{voice}'") from exc
|
|
|
|
if result.returncode != 0:
|
|
stderr = (result.stderr or "").strip() or "no stderr output"
|
|
raise RuntimeError(f"Piper voice download failed for '{voice}': {stderr[:400]}")
|
|
|
|
if not cached.exists():
|
|
raise RuntimeError(
|
|
f"Piper voice download completed but {cached} is missing — "
|
|
f"check voice name (see: https://github.com/OHF-Voice/piper1-gpl/"
|
|
f"blob/main/docs/VOICES.md)"
|
|
)
|
|
return str(cached)
|
|
|
|
|
|
def _load_piper_voice_for_config(tts_config: Dict[str, Any]) -> Tuple[Any, Dict[str, Any]]:
|
|
"""Resolve + load (or fetch from cache) the Piper voice ``tts_config`` selects.
|
|
|
|
Shared by synthesis and ``warm_tts_provider`` so a warm-up fills exactly
|
|
the cache slot the next synthesis hits. Returns ``(voice, piper_config)``.
|
|
"""
|
|
PiperVoice = _origin()._import_piper()
|
|
|
|
piper_config = tts_config.get("piper") or {} if isinstance(tts_config, dict) else {}
|
|
voice_name = piper_config.get("voice") or DEFAULT_PIPER_VOICE
|
|
download_dir = Path(piper_config.get("voices_dir") or _get_piper_voices_dir()).expanduser()
|
|
download_dir.mkdir(parents=True, exist_ok=True)
|
|
use_cuda = bool(piper_config.get("use_cuda", False))
|
|
|
|
model_path = _resolve_piper_voice_path(voice_name, download_dir)
|
|
# speaker_id is applied per call via syn_config, so one PiperVoice
|
|
# instance serves every speaker and stays out of the cache key.
|
|
cache_key = f"{model_path}::cuda={use_cuda}"
|
|
|
|
def _load_piper_voice():
|
|
logger.info("[Piper] Loading voice: %s", model_path)
|
|
v = PiperVoice.load(model_path, use_cuda=use_cuda)
|
|
logger.info("[Piper] Voice loaded")
|
|
return v
|
|
|
|
voice = _tts_cache_get_or_load(_piper_voice_cache, cache_key, _load_piper_voice)
|
|
return voice, piper_config
|
|
|
|
|
|
def _generate_piper_tts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str:
|
|
import wave
|
|
|
|
voice, piper_config = _load_piper_voice_for_config(tts_config)
|
|
|
|
# Bad speaker_id input drops to 0 (Piper's default); booleans are rejected
|
|
# outright since True/False would silently coerce to 1/0.
|
|
_raw_speaker = piper_config.get("speaker_id", 0)
|
|
speaker_id = 0 if isinstance(_raw_speaker, bool) or not isinstance(_raw_speaker, int) else _raw_speaker
|
|
|
|
# Only build a SynthesisConfig when an advanced knob is configured, so we
|
|
# don't depend on a newer piper-tts than the user's unless we must.
|
|
syn_config = None
|
|
has_advanced = any(
|
|
k in piper_config
|
|
for k in ("length_scale", "noise_scale", "noise_w_scale", "volume", "normalize_audio", "speaker_id")
|
|
)
|
|
if has_advanced:
|
|
try:
|
|
from piper import SynthesisConfig # type: ignore
|
|
syn_config = SynthesisConfig(
|
|
length_scale=float(piper_config.get("length_scale", 1.0)),
|
|
noise_scale=float(piper_config.get("noise_scale", 0.667)),
|
|
noise_w_scale=float(piper_config.get("noise_w_scale", 0.8)),
|
|
volume=float(piper_config.get("volume", 1.0)),
|
|
normalize_audio=bool(piper_config.get("normalize_audio", True)),
|
|
speaker_id=speaker_id,
|
|
)
|
|
except ImportError:
|
|
logger.warning(
|
|
"[Piper] SynthesisConfig not available in this piper-tts "
|
|
"version — advanced knobs ignored"
|
|
)
|
|
|
|
wav_path = _wav_sidecar_path(output_path)
|
|
with wave.open(wav_path, "wb") as wav_file:
|
|
if syn_config is not None:
|
|
voice.synthesize_wav(text, wav_file, syn_config=syn_config)
|
|
else:
|
|
voice.synthesize_wav(text, wav_file)
|
|
return _finalize_wav_output(wav_path, output_path)
|
|
|
|
|
|
# ===========================================================================
|
|
# KittenTTS (local ONNX, 25-80MB models, CPU only)
|
|
# ===========================================================================
|
|
|
|
def _load_kittentts_model_for_config(tts_config: Dict[str, Any]) -> Tuple[Any, Dict[str, Any]]:
|
|
"""Load (or fetch from cache) the KittenTTS model; returns ``(model, kittentts_config)``."""
|
|
KittenTTS = _origin()._import_kittentts()
|
|
kt_config = tts_config.get("kittentts", {}) if isinstance(tts_config, dict) else {}
|
|
kt_config = kt_config or {}
|
|
model_name = kt_config.get("model", DEFAULT_KITTENTTS_MODEL)
|
|
|
|
def _load_kittentts_model():
|
|
logger.info("[KittenTTS] Loading model: %s", model_name)
|
|
m = KittenTTS(model_name)
|
|
logger.info("[KittenTTS] Model loaded successfully")
|
|
return m
|
|
|
|
model = _tts_cache_get_or_load(_kittentts_model_cache, model_name, _load_kittentts_model)
|
|
return model, kt_config
|
|
|
|
|
|
def _generate_kittentts(text: str, output_path: str, tts_config: Dict[str, Any]) -> str:
|
|
model, kt_config = _load_kittentts_model_for_config(tts_config)
|
|
audio = model.generate(
|
|
text,
|
|
voice=kt_config.get("voice", DEFAULT_KITTENTTS_VOICE),
|
|
speed=kt_config.get("speed", 1.0),
|
|
clean_text=kt_config.get("clean_text", True),
|
|
) # numpy array at 24kHz
|
|
|
|
import soundfile as sf
|
|
wav_path = _wav_sidecar_path(output_path)
|
|
sf.write(wav_path, audio, 24000)
|
|
return _finalize_wav_output(wav_path, output_path)
|