Spaces:
Build error
Build error
| #!/usr/bin/env python3 | |
| """ | |
| VibeVoice-Hindi-1.5B Test Suite | |
| Uses the vibevoice library with ZeroGPU for actual TTS inference. | |
| """ | |
| import os | |
| import sys | |
| import time | |
| import json | |
| import tempfile | |
| import threading | |
| import numpy as np | |
| import torch | |
| import gradio as gr | |
| import soundfile as sf | |
| import librosa | |
| from pathlib import Path | |
| from datetime import datetime | |
| from huggingface_hub import hf_hub_download | |
| from vibevoice.modular.configuration_vibevoice import VibeVoiceConfig | |
| from vibevoice.modular.modeling_vibevoice_inference import VibeVoiceForConditionalGenerationInference | |
| from vibevoice.processor.vibevoice_processor import VibeVoiceProcessor | |
| from transformers import set_seed | |
| # Constants | |
| MODEL_ID = "tarun7r/vibevoice-hindi-7b" | |
| SAMPLE_RATE = 24000 | |
| TEST_CASES = { | |
| "1. Greeting": "Speaker 0: नमस्ते, आपका स्वागत है। आज आपका दिन कैसा चल रहा है?", | |
| "2. Business": "Speaker 0: हमारी कंपनी गुणवत्ता और नवाचार में विश्वास करती है। हम ग्राहक संतुष्टि को सर्वोच्च प्राथमिकता देते हैं।", | |
| "3. Numbers": "Speaker 0: साल दो हज़ार छब्बीस में हमने बीस लाख रुपये की आय प्राप्त की। यह पिछले साल से पचास प्रतिशत अधिक है।", | |
| "4. Mixed": "Speaker 0: आप यूट्यूब पर हमारा नया वीडियो देख सकते हैं। हमारी वेबसाइट पर भी जानकारी उपलब्ध है।", | |
| } | |
| # Global state | |
| _model = None | |
| _processor = None | |
| _voice_sample = None | |
| def get_voice_sample(): | |
| """Download and cache the Hindi voice sample from the model repo.""" | |
| global _voice_sample | |
| if _voice_sample is None: | |
| voice_path = hf_hub_download( | |
| repo_id=MODEL_ID, | |
| filename="hi-Priya_woman.wav", | |
| ) | |
| wav, sr = sf.read(voice_path) | |
| if len(wav.shape) > 1: | |
| wav = np.mean(wav, axis=1) | |
| if sr != SAMPLE_RATE: | |
| wav = librosa.resample(wav, orig_sr=sr, target_sr=SAMPLE_RATE) | |
| _voice_sample = wav.astype(np.float32) | |
| return _voice_sample | |
| def load_model(): | |
| """Load model and processor to GPU (T4 paid Space has CUDA always).""" | |
| global _model, _processor | |
| if _model is None: | |
| print(f"Loading VibeVoice model from {MODEL_ID}...") | |
| _processor = VibeVoiceProcessor.from_pretrained(MODEL_ID) | |
| _model = VibeVoiceForConditionalGenerationInference.from_pretrained( | |
| MODEL_ID, | |
| torch_dtype=torch.float16, | |
| device_map="cuda", | |
| attn_implementation="sdpa", | |
| ) | |
| _model.eval() | |
| # Use SDE solver | |
| _model.model.noise_scheduler = _model.model.noise_scheduler.from_config( | |
| _model.model.noise_scheduler.config, | |
| algorithm_type="sde-dpmsolver++", | |
| beta_schedule="squaredcos_cap_v2", | |
| ) | |
| _model.set_ddpm_inference_steps(num_steps=10) | |
| print("Model loaded successfully on CUDA.") | |
| return _model, _processor | |
| def convert_to_16_bit_wav(data): | |
| """Convert float audio to 16-bit PCM.""" | |
| if torch.is_tensor(data): | |
| data = data.detach().cpu().float().numpy() | |
| data = np.array(data, dtype=np.float32) | |
| if np.max(np.abs(data)) > 1.0: | |
| data = data / np.max(np.abs(data)) | |
| return (data * 32767).astype(np.int16) | |
| def generate_speech(script_text): | |
| """Generate speech from Hindi text using VibeVoice pipeline.""" | |
| if not script_text or not script_text.strip(): | |
| raise gr.Error("Please enter Hindi text.") | |
| set_seed(42) | |
| start_time = time.time() | |
| model, processor = load_model() | |
| voice_sample = get_voice_sample() | |
| device = "cuda" | |
| print(f"Running inference on: {device}") | |
| # Ensure script has Speaker format | |
| lines = script_text.strip().split("\n") | |
| formatted = [] | |
| for line in lines: | |
| line = line.strip() | |
| if not line: | |
| continue | |
| if line.startswith("Speaker ") and ":" in line: | |
| formatted.append(line) | |
| else: | |
| formatted.append(f"Speaker 0: {line}") | |
| script = "\n".join(formatted) | |
| # Process inputs | |
| inputs = processor( | |
| text=[script], | |
| padding=True, | |
| return_tensors="pt", | |
| return_attention_mask=True, | |
| voice_samples=[[voice_sample]], | |
| ) | |
| for k, v in inputs.items(): | |
| if torch.is_tensor(v): | |
| inputs[k] = v.to(device) | |
| # Generate | |
| from vibevoice.modular.streamer import AudioStreamer | |
| audio_streamer = AudioStreamer(batch_size=1, stop_signal=None, timeout=None) | |
| def _run_generate(): | |
| try: | |
| model.generate( | |
| **inputs, | |
| max_new_tokens=None, | |
| cfg_scale=1.3, | |
| tokenizer=processor.tokenizer, | |
| generation_config={"do_sample": False}, | |
| audio_streamer=audio_streamer, | |
| verbose=False, | |
| refresh_negative=True, | |
| is_prefill=True, | |
| ) | |
| except Exception as e: | |
| print(f"Generation error: {e}") | |
| audio_streamer.end() | |
| gen_thread = threading.Thread(target=_run_generate) | |
| gen_thread.start() | |
| # Collect audio chunks | |
| all_chunks = [] | |
| audio_stream = audio_streamer.get_stream(0) | |
| for chunk in audio_stream: | |
| if torch.is_tensor(chunk): | |
| if chunk.dtype == torch.bfloat16: | |
| chunk = chunk.float() | |
| chunk_np = chunk.cpu().numpy().astype(np.float32) | |
| else: | |
| chunk_np = np.array(chunk, dtype=np.float32) | |
| if len(chunk_np.shape) > 1: | |
| chunk_np = chunk_np.squeeze() | |
| all_chunks.append(convert_to_16_bit_wav(chunk_np)) | |
| gen_thread.join(timeout=10) | |
| elapsed = time.time() - start_time | |
| if not all_chunks: | |
| raise gr.Error("No audio generated. Model may have failed.") | |
| complete_audio = np.concatenate(all_chunks) | |
| duration = len(complete_audio) / SAMPLE_RATE | |
| # Save to temp file | |
| tmp = tempfile.NamedTemporaryFile(suffix=".wav", delete=False) | |
| sf.write(tmp.name, complete_audio, SAMPLE_RATE, subtype="PCM_16") | |
| status = f"Generated {duration:.1f}s audio in {elapsed:.1f}s" | |
| return tmp.name, status | |
| def process_test_case(test_name): | |
| """Run a preset test case.""" | |
| if not test_name: | |
| return "", None, "" | |
| text = TEST_CASES.get(test_name, "") | |
| audio_path, status = generate_speech(text) | |
| return text, audio_path, status | |
| def generate_custom(text): | |
| """Generate from custom text input.""" | |
| if not text or not text.strip(): | |
| raise gr.Error("Please enter Hindi text.") | |
| audio_path, status = generate_speech(text) | |
| return audio_path, status | |
| def save_evaluation(test_name, d1, d2, d3, d4, d5, notes): | |
| """Save evaluation results.""" | |
| checks = [d1, d2, d3, d4, d5] | |
| passed = sum(checks) | |
| result = { | |
| "timestamp": datetime.now().isoformat(), | |
| "test_case": test_name, | |
| "dimensions_passed": passed, | |
| "dimensions_total": 5, | |
| "notes": notes, | |
| } | |
| log_file = Path("evaluation_log.jsonl") | |
| with open(log_file, "a") as f: | |
| f.write(json.dumps(result) + "\n") | |
| labels = { | |
| 5: "Production Ready", | |
| 4: "Good, Minor Issues", | |
| 3: "Acceptable, Needs Work", | |
| 2: "Poor Quality", | |
| 1: "Archive This Model", | |
| 0: "Complete Failure", | |
| } | |
| return f"Saved! Score: {passed}/5 — {labels.get(passed, 'N/A')}" | |
| def load_results(): | |
| log_file = Path("evaluation_log.jsonl") | |
| if log_file.exists(): | |
| lines = log_file.read_text().strip().split("\n") | |
| summary = f"Total Evaluations: {len(lines)}\n\n" | |
| for i, line in enumerate(lines[-10:], 1): | |
| data = json.loads(line) | |
| summary += f"{i}. {data['test_case']} | Score: {data['dimensions_passed']}/5 | {data['timestamp']}\n" | |
| return summary | |
| return "No evaluations yet." | |
| # === GRADIO UI === | |
| DIMS = [ | |
| "Understand (Can you follow what is being said?)", | |
| "Natural (Does it sound like a real person?)", | |
| "Speed (Was inference latency acceptable?)", | |
| "Clarity (Can you hear each word clearly?)", | |
| "Consistency (Same pronunciation throughout?)", | |
| ] | |
| with gr.Blocks(title="VibeVoice-Hindi-7B Test") as demo: | |
| gr.Markdown("# VibeVoice-Hindi-7B — TTS Evaluation") | |
| gr.Markdown("Hindi fine-tuned TTS on Qwen2.5-7B backbone. A100 80GB GPU.") | |
| with gr.Tabs(): | |
| with gr.Tab("Preset Tests"): | |
| with gr.Row(): | |
| test_selector = gr.Dropdown( | |
| choices=list(TEST_CASES.keys()), | |
| label="Select Test Case", | |
| ) | |
| run_btn = gr.Button("Generate Speech", variant="primary") | |
| test_text = gr.Textbox(label="Script", interactive=False, lines=3) | |
| audio_out = gr.Audio(label="Generated Speech", type="filepath") | |
| status_out = gr.Textbox(label="Status", interactive=False) | |
| run_btn.click( | |
| fn=process_test_case, | |
| inputs=test_selector, | |
| outputs=[test_text, audio_out, status_out], | |
| ) | |
| with gr.Tab("Custom Text"): | |
| custom_text = gr.Textbox( | |
| label="Enter Hindi Text", | |
| placeholder="नमस्ते, यह एक परीक्षण है।", | |
| lines=4, | |
| ) | |
| custom_btn = gr.Button("Generate", variant="primary") | |
| custom_audio = gr.Audio(label="Generated Speech", type="filepath") | |
| custom_status = gr.Textbox(label="Status", interactive=False) | |
| custom_btn.click( | |
| fn=generate_custom, | |
| inputs=custom_text, | |
| outputs=[custom_audio, custom_status], | |
| ) | |
| with gr.Tab("Evaluate"): | |
| eval_test = gr.Dropdown( | |
| choices=list(TEST_CASES.keys()) + ["Custom"], | |
| label="Test Case Evaluated", | |
| ) | |
| checkboxes = [gr.Checkbox(label=d, value=False) for d in DIMS] | |
| eval_notes = gr.Textbox(label="Notes", placeholder="Optional", lines=2) | |
| save_btn = gr.Button("Save Evaluation", variant="primary") | |
| eval_result = gr.Textbox(label="Result", interactive=False) | |
| save_btn.click( | |
| fn=save_evaluation, | |
| inputs=[eval_test] + checkboxes + [eval_notes], | |
| outputs=eval_result, | |
| ) | |
| with gr.Tab("Results"): | |
| results_display = gr.Textbox(label="Evaluations", interactive=False, lines=12) | |
| refresh_btn = gr.Button("Refresh") | |
| refresh_btn.click(fn=load_results, outputs=results_display) | |
| demo.load(fn=load_results, outputs=results_display) | |
| if __name__ == "__main__": | |
| demo.launch() | |