shabdham_backend / tts_engine.py
akma01's picture
Update tts_engine.py
389d4d2 verified
Raw
History Blame Contribute Delete
3.53 kB
import os
import sys
# Mock spaces if running locally to support @spaces.GPU decorator
try:
import spaces
except ImportError:
class MockSpaces:
def GPU(self, func):
return func
sys.modules["spaces"] = MockSpaces()
import spaces
import wave
try:
import pyaudio
except ImportError:
pyaudio = None
# Ensure custom f5_tts from sooktam2 is resolvable
repo_dir = os.path.join(os.path.dirname(os.path.abspath(__file__)), "sooktam2")
sys.path.append(repo_dir)
sys.path.append(os.path.join(repo_dir, "src"))
from transformers import AutoModel
# Load the model globally so it doesn't reload on every request
MODEL_ID = "bharatgenai/sooktam2"
print(f"Loading TTS Model: {MODEL_ID}...")
try:
import torch
device = "cuda" if torch.cuda.is_available() else ("mps" if torch.backends.mps.is_available() else "cpu")
model_tts = AutoModel.from_pretrained(
MODEL_ID,
trust_remote_code=True,
).to(device)
print(f"TTS Model loaded successfully and moved to {device}.")
except Exception as e:
print(f"Error loading TTS model: {e}")
model_tts = None
# We need a reference audio and text for Sooktam-2
# Using the same ones from gen_voice.py
REF_AUDIO = "ishita_tts_audio.mp3 (1) (1).wav" # Ensure this exists locally or update path
REF_TEXT = "வணக்கம், ஜென் லேப் ல இருந்து பேசுறேன். டெஸ்ட் புக் பண்ணணுமா இல்ல அப்பாயின்ட்மென்ட் ஷெட்யூல் பண்ணணுமா?"
OUT_DIR = "outputs"
os.makedirs(OUT_DIR, exist_ok=True)
@spaces.GPU
def synthesize_speech(text: str, output_filename: str = "response.wav"):
if not model_tts:
print("TTS model is not loaded. Cannot synthesize speech.")
return None
out_wav_path = os.path.join(OUT_DIR, output_filename)
# Check if reference audio exists
if not os.path.exists(REF_AUDIO):
print(f"Warning: Reference audio {REF_AUDIO} not found.")
# Try to find a fallback or just pass it anyway, inference might fail
print(f"Synthesizing speech: '{text}'")
try:
wav, sr, _ = model_tts.infer(
ref_file=REF_AUDIO,
ref_text=REF_TEXT,
gen_text=text,
tokenizer="cls",
cls_language="tamil", # Can be made dynamic if needed
file_wave=out_wav_path,
)
print(f"Speech saved to: {out_wav_path}")
return out_wav_path
except Exception as e:
print(f"Error during TTS synthesis: {e}")
return None
def play_audio(file_path: str):
if not pyaudio:
print("pyaudio is not installed. Cannot play audio locally.")
return
if not file_path or not os.path.exists(file_path):
print(f"Cannot play audio: file {file_path} not found.")
return
try:
wf = wave.open(file_path, 'rb')
p = pyaudio.PyAudio()
stream = p.open(format=p.get_format_from_width(wf.getsampwidth()),
channels=wf.getnchannels(),
rate=wf.getframerate(),
output=True)
data = wf.readframes(1024)
while len(data) > 0:
stream.write(data)
data = wf.readframes(1024)
stream.stop_stream()
stream.close()
p.terminate()
except Exception as e:
print(f"Error playing audio: {e}")