services(riva): make FastpitchTTSService asyncio
This commit is contained in:
committed by
vipyne
parent
7e407e5548
commit
4b55c73fbe
@@ -60,8 +60,6 @@ class FastpitchTTSService(TTSService):
|
|||||||
self.voice_id = voice_id
|
self.voice_id = voice_id
|
||||||
self.sample_rate_hz = sample_rate_hz
|
self.sample_rate_hz = sample_rate_hz
|
||||||
self.language_code = params.language
|
self.language_code = params.language
|
||||||
self.nchannels = 1
|
|
||||||
self.sampwidth = 2
|
|
||||||
self.quality = None
|
self.quality = None
|
||||||
|
|
||||||
metadata = [
|
metadata = [
|
||||||
@@ -73,35 +71,39 @@ class FastpitchTTSService(TTSService):
|
|||||||
self.service = riva.client.SpeechSynthesisService(auth)
|
self.service = riva.client.SpeechSynthesisService(auth)
|
||||||
|
|
||||||
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
async def run_tts(self, text: str) -> AsyncGenerator[Frame, None]:
|
||||||
|
def read_audio_responses():
|
||||||
|
try:
|
||||||
|
custom_dictionary_input = {}
|
||||||
|
responses = self.service.synthesize_online(
|
||||||
|
text,
|
||||||
|
self.voice_id,
|
||||||
|
self.language_code,
|
||||||
|
sample_rate_hz=self.sample_rate_hz,
|
||||||
|
audio_prompt_file=None,
|
||||||
|
quality=20 if self.quality is None else self.quality,
|
||||||
|
custom_dictionary=custom_dictionary_input,
|
||||||
|
)
|
||||||
|
return responses
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"{self} exception: {e}")
|
||||||
|
return []
|
||||||
|
|
||||||
logger.debug(f"Generating TTS: [{text}]")
|
logger.debug(f"Generating TTS: [{text}]")
|
||||||
|
|
||||||
await self.start_ttfb_metrics()
|
await self.start_ttfb_metrics()
|
||||||
yield TTSStartedFrame()
|
yield TTSStartedFrame()
|
||||||
|
|
||||||
try:
|
responses = await asyncio.to_thread(read_audio_responses)
|
||||||
custom_dictionary_input = {}
|
|
||||||
responses = self.service.synthesize_online(
|
for resp in responses:
|
||||||
text,
|
await self.stop_ttfb_metrics()
|
||||||
self.voice_id,
|
|
||||||
self.language_code,
|
frame = TTSAudioRawFrame(
|
||||||
sample_rate_hz=self.sample_rate_hz,
|
audio=resp.audio,
|
||||||
audio_prompt_file=None,
|
sample_rate=self.sample_rate_hz,
|
||||||
quality=20 if self.quality is None else self.quality,
|
num_channels=1,
|
||||||
custom_dictionary=custom_dictionary_input,
|
|
||||||
)
|
)
|
||||||
|
yield frame
|
||||||
for resp in responses:
|
|
||||||
await self.stop_ttfb_metrics()
|
|
||||||
|
|
||||||
frame = TTSAudioRawFrame(
|
|
||||||
audio=resp.audio,
|
|
||||||
sample_rate=self.sample_rate_hz,
|
|
||||||
num_channels=self.nchannels,
|
|
||||||
)
|
|
||||||
yield frame
|
|
||||||
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"{self} exception: {e}")
|
|
||||||
|
|
||||||
await self.start_tts_usage_metrics(text)
|
await self.start_tts_usage_metrics(text)
|
||||||
yield TTSStoppedFrame()
|
yield TTSStoppedFrame()
|
||||||
|
|||||||
Reference in New Issue
Block a user