Some checks are pending
CI / ruff (push) Waiting to run
CI / mypy (push) Waiting to run
CI / pytest (push) Waiting to run
CI / package (push) Waiting to run
CI / Install smoke (${{ matrix.label }}) (linux, ubuntu-latest) (push) Blocked by required conditions
CI / Install smoke (${{ matrix.label }}) (macos-arm64, macos-14) (push) Blocked by required conditions
131 lines
4.8 KiB
Python
131 lines
4.8 KiB
Python
from __future__ import annotations
|
|
|
|
import logging
|
|
from threading import Event
|
|
from typing import Any, Iterator
|
|
|
|
import librosa
|
|
import numpy as np
|
|
import torch
|
|
from melo.api import TTS
|
|
from rich.console import Console
|
|
|
|
from speech_to_speech.baseHandler import BaseHandler
|
|
from speech_to_speech.pipeline.cancel_scope import CancelScope
|
|
from speech_to_speech.pipeline.handler_types import TTSIn, TTSOut
|
|
from speech_to_speech.pipeline.messages import AUDIO_RESPONSE_DONE, EndOfResponse
|
|
from speech_to_speech.pipeline.speculative_turns import SpeculativeTurnTracker
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
console = Console()
|
|
|
|
WHISPER_LANGUAGE_TO_MELO_LANGUAGE = {
|
|
"en": "EN",
|
|
"fr": "FR",
|
|
"es": "ES",
|
|
"zh": "ZH",
|
|
"ja": "JP",
|
|
"ko": "KR",
|
|
}
|
|
|
|
WHISPER_LANGUAGE_TO_MELO_SPEAKER = {
|
|
"en": "EN-BR",
|
|
"fr": "FR",
|
|
"es": "ES",
|
|
"zh": "ZH",
|
|
"ja": "JP",
|
|
"ko": "KR",
|
|
}
|
|
|
|
|
|
class MeloTTSHandler(BaseHandler[TTSIn, TTSOut]):
|
|
def setup(
|
|
self,
|
|
should_listen: Event,
|
|
device: str = "mps",
|
|
language: str = "en",
|
|
speaker_to_id: str = "en",
|
|
gen_kwargs: dict[str, Any] = {}, # Unused
|
|
blocksize: int = 512,
|
|
cancel_scope: CancelScope | None = None,
|
|
speculative_turns: SpeculativeTurnTracker | None = None,
|
|
) -> None:
|
|
self.should_listen = should_listen
|
|
self.cancel_scope = cancel_scope
|
|
self.speculative_turns = speculative_turns
|
|
self.device = device
|
|
self.language = language
|
|
self.model = TTS(language=WHISPER_LANGUAGE_TO_MELO_LANGUAGE[self.language], device=device)
|
|
self.speaker_id = self.model.hps.data.spk2id[WHISPER_LANGUAGE_TO_MELO_SPEAKER[speaker_to_id]]
|
|
self.blocksize = blocksize
|
|
self._initial_language = self.language
|
|
self.warmup()
|
|
|
|
def warmup(self) -> None:
|
|
logger.info(f"Warming up {self.__class__.__name__}")
|
|
_ = self.model.tts_to_file("text", self.speaker_id, quiet=True)
|
|
|
|
def process(self, tts_input: TTSIn) -> Iterator[TTSOut]:
|
|
if isinstance(tts_input, EndOfResponse):
|
|
yield AUDIO_RESPONSE_DONE
|
|
return
|
|
|
|
speculative_turns = getattr(self, "speculative_turns", None)
|
|
if speculative_turns and not speculative_turns.is_latest(
|
|
tts_input.turn_id,
|
|
tts_input.turn_revision,
|
|
):
|
|
logger.debug("Dropping stale TTS input for turn=%s rev=%s", tts_input.turn_id, tts_input.turn_revision)
|
|
return
|
|
|
|
gen = self.cancel_scope.generation if self.cancel_scope else None
|
|
language_code = tts_input.language_code
|
|
text = tts_input.text
|
|
|
|
console.print(f"[green]ASSISTANT: {text}")
|
|
|
|
if language_code is not None and self.language != language_code:
|
|
try:
|
|
self.model = TTS(
|
|
language=WHISPER_LANGUAGE_TO_MELO_LANGUAGE[language_code],
|
|
device=self.device,
|
|
)
|
|
self.speaker_id = self.model.hps.data.spk2id[WHISPER_LANGUAGE_TO_MELO_SPEAKER[language_code]]
|
|
self.language = language_code
|
|
except KeyError:
|
|
console.print(f"[red]Language {language_code} not supported by Melo. Using {self.language} instead.")
|
|
|
|
if self.device == "mps":
|
|
import time
|
|
|
|
start = time.time()
|
|
torch.mps.synchronize() # Waits for all kernels in all streams on the MPS device to complete.
|
|
torch.mps.empty_cache() # Frees all memory allocated by the MPS device.
|
|
_ = time.time() - start # Removing this line makes it fail more often. I'm looking into it.
|
|
|
|
try:
|
|
audio_chunk = self.model.tts_to_file(text, self.speaker_id, quiet=True)
|
|
except (AssertionError, RuntimeError) as e:
|
|
logger.error(f"Error in MeloTTSHandler: {e}")
|
|
audio_chunk = np.array([])
|
|
if len(audio_chunk) == 0:
|
|
return
|
|
audio_chunk = librosa.resample(audio_chunk, orig_sr=44100, target_sr=16000)
|
|
audio_chunk = (audio_chunk * 32768).astype(np.int16)
|
|
for i in range(0, len(audio_chunk), self.blocksize):
|
|
if gen is not None and self.cancel_scope is not None and self.cancel_scope.is_stale(gen):
|
|
logger.info("TTS generation cancelled (interruption)")
|
|
return
|
|
yield np.pad(
|
|
audio_chunk[i : i + self.blocksize],
|
|
(0, self.blocksize - len(audio_chunk[i : i + self.blocksize])),
|
|
)
|
|
|
|
def on_session_end(self) -> None:
|
|
if self.language != self._initial_language:
|
|
self.language = self._initial_language
|
|
self.model = TTS(language=WHISPER_LANGUAGE_TO_MELO_LANGUAGE[self.language], device=self.device)
|
|
self.speaker_id = self.model.hps.data.spk2id[WHISPER_LANGUAGE_TO_MELO_SPEAKER[self.language]]
|
|
logger.debug("Melo TTS session state reset")
|