import time import os import spaces import torch import gradio as gr from huggingface_hub import snapshot_download # Model config HF_MODEL_ID = "dolly-vn/Vira-TTS" MODEL_PATH = "model_pretrained" # Global model variable mira_tts = None def download_model_if_needed(): """Download model from HuggingFace if not exists locally.""" if not os.path.exists(MODEL_PATH) or not os.listdir(MODEL_PATH): print(f"📥 Downloading model from HuggingFace: {HF_MODEL_ID}...") snapshot_download( repo_id=HF_MODEL_ID, local_dir=MODEL_PATH, local_dir_use_symlinks=False ) print("✅ Model downloaded!") else: print(f"✅ Model found at: {MODEL_PATH}") # Download model at startup (no GPU needed) download_model_if_needed() SAMPLE_RATE = 48000 def get_model(): """Lazy load model when GPU is available.""" global mira_tts if mira_tts is None: from mira.model import MiraTTS from mira.utils import split_text print("🔄 Loading Vira-TTS...") mira_tts = MiraTTS(MODEL_PATH) print("✅ Model loaded!") return mira_tts @spaces.GPU def generate_speech(text: str, reference_audio: str): """Generate speech from text using reference audio for voice cloning.""" from mira.utils import split_text if not text.strip(): return None, "Vui lòng nhập văn bản." if reference_audio is None: return None, "Vui lòng upload file audio tham chiếu." try: # Get model (lazy load with GPU) model = get_model() # Encode reference audio context_tokens = model.encode_audio(reference_audio) # Split text into sentences sentences = split_text(text) # Generate audio and measure time start_time = time.time() if len(sentences) == 1: audio = model.generate(sentences[0], context_tokens) else: audio = model.batch_generate(sentences, [context_tokens]) inference_time = time.time() - start_time # Calculate RTF audio_np = audio.float().cpu().numpy() audio_duration = len(audio_np) / SAMPLE_RATE rtf = inference_time / audio_duration stats = f"📝 Số câu: {len(sentences)} | ⏱️ Inference: {inference_time:.2f}s | 🎵 Audio: {audio_duration:.2f}s | 📊 RTF: {rtf:.4f}" return (SAMPLE_RATE, audio_np), stats except Exception as e: import traceback return None, f"Lỗi: {str(e)}\n{traceback.format_exc()}" # Create Gradio interface with gr.Blocks(title="Vira-TTS Vietnamese", theme=gr.themes.Soft()) as demo: gr.Markdown(""" # 🎙️ Vira-TTS Vietnamese ### Text-to-Speech với Voice Cloning Vietnamese TTS fine-tuned từ MiraTTS trên 500 giờ audio tiếng Việt. Upload một file audio tham chiếu (3-10 giây) để clone giọng nói, sau đó nhập văn bản để tạo audio. """) with gr.Row(): with gr.Column(scale=1): text_input = gr.Textbox( label="Văn bản", placeholder="Nhập văn bản tiếng Việt tại đây...", lines=5 ) reference_audio = gr.Audio( label="Audio tham chiếu (để clone giọng)", type="filepath" ) generate_btn = gr.Button("🎵 Tạo Audio", variant="primary", size="lg") with gr.Column(scale=1): output_audio = gr.Audio( label="Audio đầu ra", type="numpy" ) stats_output = gr.Textbox( label="Thống kê", interactive=False ) # Example texts gr.Markdown("### 📝 Ví dụ văn bản:") with gr.Row(): gr.Button("Xin chào").click( fn=lambda: "Xin chào, tôi là trợ lý ảo Vira-TTS.", outputs=[text_input] ) gr.Button("Thời tiết").click( fn=lambda: "Hôm nay thời tiết rất đẹp, chúng ta đi dạo nhé!", outputs=[text_input] ) gr.Button("Công nghệ").click( fn=lambda: "Công nghệ trí tuệ nhân tạo đang phát triển rất nhanh chóng.", outputs=[text_input] ) # Event handler generate_btn.click( fn=generate_speech, inputs=[text_input, reference_audio], outputs=[output_audio, stats_output] ) gr.Markdown(""" --- ### 📌 Links - [GitHub](https://github.com/iamdinhthuan/Vira-tts) - [Model](https://huggingface.co/dolly-vn/Vira-TTS) """) if __name__ == "__main__": demo.launch()