Anupam007 commited on
Commit
0faa3aa
Β·
verified Β·
1 Parent(s): 1bfc659

Update app.py

Browse files
Files changed (1) hide show
  1. app.py +104 -68
app.py CHANGED
@@ -1,72 +1,108 @@
1
  import gradio as gr
2
  import numpy as np
3
  import torch
4
- from datasets import load_dataset
5
-
6
- from transformers import SpeechT5ForTextToSpeech, SpeechT5HifiGan, SpeechT5Processor, pipeline
7
-
8
-
 
 
 
 
 
 
 
 
 
 
 
9
  device = "cuda:0" if torch.cuda.is_available() else "cpu"
10
-
11
- # load speech translation checkpoint
12
- asr_pipe = pipeline("automatic-speech-recognition", model="openai/whisper-base", device=device)
13
-
14
- # load text-to-speech checkpoint and speaker embeddings
15
- processor = SpeechT5Processor.from_pretrained("microsoft/speecht5_tts")
16
-
17
- model = SpeechT5ForTextToSpeech.from_pretrained("microsoft/speecht5_tts").to(device)
18
- vocoder = SpeechT5HifiGan.from_pretrained("microsoft/speecht5_hifigan").to(device)
19
-
20
- embeddings_dataset = load_dataset("Matthijs/cmu-arctic-xvectors", split="validation")
21
- speaker_embeddings = torch.tensor(embeddings_dataset[7306]["xvector"]).unsqueeze(0)
22
-
23
-
24
- def translate(audio):
25
- outputs = asr_pipe(audio, max_new_tokens=256, generate_kwargs={"task": "translate"})
26
- return outputs["text"]
27
-
28
-
29
- def synthesise(text):
30
- inputs = processor(text=text, return_tensors="pt")
31
- speech = model.generate_speech(inputs["input_ids"].to(device), speaker_embeddings.to(device), vocoder=vocoder)
32
- return speech.cpu()
33
-
34
-
35
- def speech_to_speech_translation(audio):
36
- translated_text = translate(audio)
37
- synthesised_speech = synthesise(translated_text)
38
- synthesised_speech = (synthesised_speech.numpy() * 32767).astype(np.int16)
39
- return 16000, synthesised_speech
40
-
41
-
42
- title = "Cascaded STST"
43
- description = """
44
- Demo for cascaded speech-to-speech translation (STST), mapping from source speech in any language to target speech in English. Demo uses OpenAI's [Whisper Base](https://huggingface.co/openai/whisper-base) model for speech translation, and Microsoft's
45
- [SpeechT5 TTS](https://huggingface.co/microsoft/speecht5_tts) model for text-to-speech:
46
-
47
- ![Cascaded STST](https://huggingface.co/datasets/huggingface-course/audio-course-images/resolve/main/s2st_cascaded.png "Diagram of cascaded speech to speech translation")
48
- """
49
-
50
- demo = gr.Blocks()
51
-
52
- mic_translate = gr.Interface(
53
- fn=speech_to_speech_translation,
54
- inputs=gr.Audio(source="microphone", type="filepath"),
55
- outputs=gr.Audio(label="Generated Speech", type="numpy"),
56
- title=title,
57
- description=description,
58
- )
59
-
60
- file_translate = gr.Interface(
61
- fn=speech_to_speech_translation,
62
- inputs=gr.Audio(source="upload", type="filepath"),
63
- outputs=gr.Audio(label="Generated Speech", type="numpy"),
64
- examples=[["./example.wav"]],
65
- title=title,
66
- description=description,
67
- )
68
-
69
- with demo:
70
- gr.TabbedInterface([mic_translate, file_translate], ["Microphone", "Audio File"])
71
-
72
- demo.launch()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  import gradio as gr
2
  import numpy as np
3
  import torch
4
+ import os
5
+ import tempfile
6
+ import logging
7
+ import asyncio
8
+ from typing import Optional
9
+ from transformers import AutoProcessor, AutoModelForSpeechSeq2Seq, AutoModelForTextToSpectrogram
10
+ from deep_translator import GoogleTranslator
11
+ import time
12
+ import threading
13
+ import torchaudio
14
+
15
+ # Set up logging
16
+ logging.basicConfig(level=logging.INFO, format="%(asctime)s - %(levelname)s - %(message)s")
17
+ logger = logging.getLogger(__name__)
18
+
19
+ # Device setup
20
  device = "cuda:0" if torch.cuda.is_available() else "cpu"
21
+ logger.info(f"Device set to use {device}")
22
+
23
+ # Supported languages
24
+ SUPPORTED_LANGUAGES = {
25
+ "English": "en",
26
+ "Hindi": "hi",
27
+ }
28
+
29
+ # Load ASR (Speech Recognition) Model
30
+ asr_processor = AutoProcessor.from_pretrained("openai/whisper-base")
31
+ asr_model = AutoModelForSpeechSeq2Seq.from_pretrained("openai/whisper-base").to(device)
32
+
33
+ # Load TTS (Text-to-Speech) Model
34
+ tts_processor = AutoProcessor.from_pretrained("microsoft/speecht5_tts")
35
+ tts_model = AutoModelForTextToSpectrogram.from_pretrained("microsoft/speecht5_tts").to(device)
36
+
37
+ def speech_to_text(audio_data: np.ndarray, sample_rate: int) -> str:
38
+ try:
39
+ if len(audio_data.shape) > 1:
40
+ audio_data = np.mean(audio_data, axis=1)
41
+ audio_data = torch.tensor(audio_data, dtype=torch.float32).to(device)
42
+
43
+ inputs = asr_processor(audio_data, sampling_rate=sample_rate, return_tensors="pt").to(device)
44
+ outputs = asr_model.generate(**inputs)
45
+ text = asr_processor.batch_decode(outputs, skip_special_tokens=True)[0]
46
+ logger.info(f"Transcribed text: {text}")
47
+ return text if text else ""
48
+ except Exception as e:
49
+ logger.error(f"STT failed: {str(e)}")
50
+ return ""
51
+
52
+ def text_to_speech(text: str, lang: str) -> Optional[str]:
53
+ try:
54
+ if not text or not text.strip():
55
+ logger.warning("Empty text received for TTS.")
56
+ return None
57
+ target_lang_code = SUPPORTED_LANGUAGES.get(lang, "en")
58
+ translated_text = GoogleTranslator(source="auto", target=target_lang_code).translate(text)
59
+
60
+ inputs = tts_processor(translated_text, return_tensors="pt").to(device)
61
+ spectrogram = tts_model.generate(**inputs)
62
+
63
+ with tempfile.NamedTemporaryFile(delete=False, suffix=".wav") as temp_file:
64
+ torchaudio.save(temp_file.name, spectrogram.cpu(), 16000)
65
+ return temp_file.name
66
+ except Exception as e:
67
+ logger.error(f"TTS failed: {str(e)}")
68
+ return None
69
+
70
+ async def process_audio_stream(audio_stream, target_lang: str):
71
+ async for sample_rate, audio_chunk in audio_stream:
72
+ try:
73
+ if audio_chunk is None or not isinstance(audio_chunk, np.ndarray):
74
+ yield "", None
75
+ continue
76
+ text = speech_to_text(audio_chunk, sample_rate)
77
+ dubbed_audio = text_to_speech(text, target_lang) if text and text.strip() else None
78
+ yield text, dubbed_audio
79
+ except Exception as e:
80
+ logger.error(f"Stream processing failed: {str(e)}")
81
+ yield "", None
82
+
83
+ def run_cleanup():
84
+ while True:
85
+ for root, dirs, files in os.walk(tempfile.gettempdir()):
86
+ for file in files:
87
+ if file.endswith(".wav"):
88
+ try:
89
+ os.remove(os.path.join(root, file))
90
+ except:
91
+ pass
92
+ time.sleep(300)
93
+
94
+ with gr.Blocks(title="Real-Time Multilingual Dubbing") as demo:
95
+ gr.Markdown("<h1>Real-Time Multilingual Dubbing</h1>")
96
+ with gr.Row():
97
+ audio_input = gr.Audio(sources=["microphone"], type="numpy", streaming=True, label="Speak")
98
+ lang_dropdown = gr.Dropdown(choices=list(SUPPORTED_LANGUAGES.keys()), label="Target Language", value="English")
99
+ with gr.Row():
100
+ stt_output = gr.Textbox(label="Transcription")
101
+ dub_output = gr.Audio(label="Dubbed Audio", autoplay=True)
102
+
103
+ audio_input.stream(fn=process_audio_stream, inputs=[audio_input, lang_dropdown], outputs=[stt_output, dub_output])
104
+
105
+ if __name__ == "__main__":
106
+ cleanup_thread = threading.Thread(target=run_cleanup, daemon=True)
107
+ cleanup_thread.start()
108
+ demo.launch()