dwolf/seamless_m4t_d_wolf
0
1from __future__ import annotations2 3import os4 5import gradio as gr6import numpy as np7import torch8import torchaudio9from seamless_communication.models.inference.translator import Translator10 11from lang_list import (12 LANGUAGE_NAME_TO_CODE,13 S2ST_TARGET_LANGUAGE_NAMES,14 S2TT_TARGET_LANGUAGE_NAMES,15 T2TT_TARGET_LANGUAGE_NAMES,16 TEXT_SOURCE_LANGUAGE_NAMES,17)18 19DESCRIPTION = """# SeamlessM4T20 21[SeamlessM4T](https://github.com/facebookresearch/seamless_communication) is designed to provide high-quality22translation, allowing people from different linguistic communities to communicate effortlessly through speech and text.23 24This unified model enables multiple tasks like Speech-to-Speech (S2ST), Speech-to-Text (S2TT), Text-to-Speech (T2ST)25translation and more, without relying on multiple separate models.26"""27 28CACHE_EXAMPLES = os.getenv("CACHE_EXAMPLES") == "1"29 30TASK_NAMES = [31 "S2ST (Speech to Speech translation)",32 "S2TT (Speech to Text translation)",33 "T2ST (Text to Speech translation)",34 "T2TT (Text to Text translation)",35 "ASR (Automatic Speech Recognition)",36]37AUDIO_SAMPLE_RATE = 16000.038MAX_INPUT_AUDIO_LENGTH = 60 # in seconds39DEFAULT_TARGET_LANGUAGE = "French"40 41device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")42translator = Translator(43 model_name_or_card="seamlessM4T_large",44 vocoder_name_or_card="vocoder_36langs",45 device=device,46 sample_rate=AUDIO_SAMPLE_RATE,47)48 49 50def predict(51 task_name: str,52 audio_source: str,53 input_audio_mic: str | None,54 input_audio_file: str | None,55 input_text: str | None,56 source_language: str | None,57 target_language: str,58) -> tuple[tuple[int, np.ndarray] | None, str]:59 task_name = task_name.split()[0]60 source_language_code = LANGUAGE_NAME_TO_CODE[source_language] if source_language else None61 target_language_code = LANGUAGE_NAME_TO_CODE[target_language]62 63 if task_name in ["S2ST", "S2TT", "ASR"]:64 if audio_source == "microphone":65 input_data = input_audio_mic66 else:67 input_data = input_audio_file68 69 arr, org_sr = torchaudio.load(input_data)70 new_arr = torchaudio.functional.resample(arr, orig_freq=org_sr, new_freq=AUDIO_SAMPLE_RATE)71 max_length = int(MAX_INPUT_AUDIO_LENGTH * AUDIO_SAMPLE_RATE)72 if new_arr.shape[1] > max_length:73 new_arr = new_arr[:, :max_length]74 gr.Warning(f"Input audio is too long. Only the first {MAX_INPUT_AUDIO_LENGTH} seconds is used.")75 torchaudio.save(input_data, new_arr, sample_rate=int(AUDIO_SAMPLE_RATE))76 else:77 input_data = input_text78 text_out, wav, sr = translator.predict(79 input=input_data,80 task_str=task_name,81 tgt_lang=target_language_code,82 src_lang=source_language_code,83 ngram_filtering=True,84 )85 if task_name in ["S2ST", "T2ST"]:86 return (sr, wav.cpu().detach().numpy()), text_out87 else:88 return None, text_out89 90 91def process_s2st_example(input_audio_file: str, target_language: str) -> tuple[tuple[int, np.ndarray] | None, str]:92 return predict(93 task_name="S2ST",94 audio_source="file",95 input_audio_mic=None,96 input_audio_file=input_audio_file,97 input_text=None,98 source_language=None,99 target_language=target_language,100 )101 102 103def process_s2tt_example(input_audio_file: str, target_language: str) -> tuple[tuple[int, np.ndarray] | None, str]:104 return predict(105 task_name="S2TT",106 audio_source="file",107 input_audio_mic=None,108 input_audio_file=input_audio_file,109 input_text=None,110 source_language=None,111 target_language=target_language,112 )113 114 115def process_t2st_example(116 input_text: str, source_language: str, target_language: str117) -> tuple[tuple[int, np.ndarray] | None, str]:118 return predict(119 task_name="T2ST",120 audio_source="",121 input_audio_mic=None,122 input_audio_file=None,123 input_text=input_text,124 source_language=source_language,125 target_language=target_language,126 )127 128 129def process_t2tt_example(130 input_text: str, source_language: str, target_language: str131) -> tuple[tuple[int, np.ndarray] | None, str]:132 return predict(133 task_name="T2TT",134 audio_source="",135 input_audio_mic=None,136 input_audio_file=None,137 input_text=input_text,138 source_language=source_language,139 target_language=target_language,140 )141 142 143def process_asr_example(input_audio_file: str, target_language: str) -> tuple[tuple[int, np.ndarray] | None, str]:144 return predict(145 task_name="ASR",146 audio_source="file",147 input_audio_mic=None,148 input_audio_file=input_audio_file,149 input_text=None,150 source_language=None,151 target_language=target_language,152 )153 154 155def update_audio_ui(audio_source: str) -> tuple[dict, dict]:156 mic = audio_source == "microphone"157 return (158 gr.update(visible=mic, value=None), # input_audio_mic159 gr.update(visible=not mic, value=None), # input_audio_file160 )161 162 163def update_input_ui(task_name: str) -> tuple[dict, dict, dict, dict]:164 task_name = task_name.split()[0]165 if task_name == "S2ST":166 return (167 gr.update(visible=True), # audio_box168 gr.update(visible=False), # input_text169 gr.update(visible=False), # source_language170 gr.update(171 visible=True, choices=S2ST_TARGET_LANGUAGE_NAMES, value=DEFAULT_TARGET_LANGUAGE172 ), # target_language173 )174 elif task_name == "S2TT":175 return (176 gr.update(visible=True), # audio_box177 gr.update(visible=False), # input_text178 gr.update(visible=False), # source_language179 gr.update(180 visible=True, choices=S2TT_TARGET_LANGUAGE_NAMES, value=DEFAULT_TARGET_LANGUAGE181 ), # target_language182 )183 elif task_name == "T2ST":184 return (185 gr.update(visible=False), # audio_box186 gr.update(visible=True), # input_text187 gr.update(visible=True), # source_language188 gr.update(189 visible=True, choices=S2ST_TARGET_LANGUAGE_NAMES, value=DEFAULT_TARGET_LANGUAGE190 ), # target_language191 )192 elif task_name == "T2TT":193 return (194 gr.update(visible=False), # audio_box195 gr.update(visible=True), # input_text196 gr.update(visible=True), # source_language197 gr.update(198 visible=True, choices=T2TT_TARGET_LANGUAGE_NAMES, value=DEFAULT_TARGET_LANGUAGE199 ), # target_language200 )201 elif task_name == "ASR":202 return (203 gr.update(visible=True), # audio_box204 gr.update(visible=False), # input_text205 gr.update(visible=False), # source_language206 gr.update(207 visible=True, choices=S2TT_TARGET_LANGUAGE_NAMES, value=DEFAULT_TARGET_LANGUAGE208 ), # target_language209 )210 else:211 raise ValueError(f"Unknown task: {task_name}")212 213 214def update_output_ui(task_name: str) -> tuple[dict, dict]:215 task_name = task_name.split()[0]216 if task_name in ["S2ST", "T2ST"]:217 return (218 gr.update(visible=True, value=None), # output_audio219 gr.update(value=None), # output_text220 )221 elif task_name in ["S2TT", "T2TT", "ASR"]:222 return (223 gr.update(visible=False, value=None), # output_audio224 gr.update(value=None), # output_text225 )226 else:227 raise ValueError(f"Unknown task: {task_name}")228 229 230def update_example_ui(task_name: str) -> tuple[dict, dict, dict, dict, dict]:231 task_name = task_name.split()[0]232 return (233 gr.update(visible=task_name == "S2ST"), # s2st_example_row234 gr.update(visible=task_name == "S2TT"), # s2tt_example_row235 gr.update(visible=task_name == "T2ST"), # t2st_example_row236 gr.update(visible=task_name == "T2TT"), # t2tt_example_row237 gr.update(visible=task_name == "ASR"), # asr_example_row238 )239 240 241with gr.Blocks(css="style.css") as demo:242 gr.Markdown(DESCRIPTION)243 gr.DuplicateButton(244 value="Duplicate Space for private use",245 elem_id="duplicate-button",246 visible=os.getenv("SHOW_DUPLICATE_BUTTON") == "1",247 )248 with gr.Group():249 task_name = gr.Dropdown(250 label="Task",251 choices=TASK_NAMES,252 value=TASK_NAMES[0],253 )254 with gr.Row():255 source_language = gr.Dropdown(256 label="Source language",257 choices=TEXT_SOURCE_LANGUAGE_NAMES,258 value="English",259 visible=False,260 )261 target_language = gr.Dropdown(262 label="Target language",263 choices=S2ST_TARGET_LANGUAGE_NAMES,264 value=DEFAULT_TARGET_LANGUAGE,265 )266 with gr.Row() as audio_box:267 audio_source = gr.Radio(268 label="Audio source",269 choices=["file", "microphone"],270 value="file",271 )272 input_audio_mic = gr.Audio(273 label="Input speech",274 type="filepath",275 source="microphone",276 visible=False,277 )278 input_audio_file = gr.Audio(279 label="Input speech",280 type="filepath",281 source="upload",282 visible=True,283 )284 input_text = gr.Textbox(label="Input text", visible=False)285 btn = gr.Button("Translate")286 with gr.Column():287 output_audio = gr.Audio(288 label="Translated speech",289 autoplay=False,290 streaming=False,291 type="numpy",292 )293 output_text = gr.Textbox(label="Translated text")294 295 with gr.Row(visible=True) as s2st_example_row:296 s2st_examples = gr.Examples(297 examples=[298 ["assets/sample_input.mp3", "French"],299 ["assets/sample_input.mp3", "Mandarin Chinese"],300 ["assets/sample_input_2.mp3", "Hindi"],301 ["assets/sample_input_2.mp3", "Spanish"],302 ],303 inputs=[input_audio_file, target_language],304 outputs=[output_audio, output_text],305 fn=process_s2st_example,306 cache_examples=CACHE_EXAMPLES,307 )308 with gr.Row(visible=False) as s2tt_example_row:309 s2tt_examples = gr.Examples(310 examples=[311 ["assets/sample_input.mp3", "French"],312 ["assets/sample_input.mp3", "Mandarin Chinese"],313 ["assets/sample_input_2.mp3", "Hindi"],314 ["assets/sample_input_2.mp3", "Spanish"],315 ],316 inputs=[input_audio_file, target_language],317 outputs=[output_audio, output_text],318 fn=process_s2tt_example,319 cache_examples=CACHE_EXAMPLES,320 )321 with gr.Row(visible=False) as t2st_example_row:322 t2st_examples = gr.Examples(323 examples=[324 ["My favorite animal is the elephant.", "English", "French"],325 ["My favorite animal is the elephant.", "English", "Mandarin Chinese"],326 [327 "Meta AI's Seamless M4T model is democratising spoken communication across language barriers",328 "English",329 "Hindi",330 ],331 [332 "Meta AI's Seamless M4T model is democratising spoken communication across language barriers",333 "English",334 "Spanish",335 ],336 ],337 inputs=[input_text, source_language, target_language],338 outputs=[output_audio, output_text],339 fn=process_t2st_example,340 cache_examples=CACHE_EXAMPLES,341 )342 with gr.Row(visible=False) as t2tt_example_row:343 t2tt_examples = gr.Examples(344 examples=[345 ["My favorite animal is the elephant.", "English", "French"],346 ["My favorite animal is the elephant.", "English", "Mandarin Chinese"],347 [348 "Meta AI's Seamless M4T model is democratising spoken communication across language barriers",349 "English",350 "Hindi",351 ],352 [353 "Meta AI's Seamless M4T model is democratising spoken communication across language barriers",354 "English",355 "Spanish",356 ],357 ],358 inputs=[input_text, source_language, target_language],359 outputs=[output_audio, output_text],360 fn=process_t2tt_example,361 cache_examples=CACHE_EXAMPLES,362 )363 with gr.Row(visible=False) as asr_example_row:364 asr_examples = gr.Examples(365 examples=[366 ["assets/sample_input.mp3", "English"],367 ["assets/sample_input_2.mp3", "English"],368 ],369 inputs=[input_audio_file, target_language],370 outputs=[output_audio, output_text],371 fn=process_asr_example,372 cache_examples=CACHE_EXAMPLES,373 )374 375 audio_source.change(376 fn=update_audio_ui,377 inputs=audio_source,378 outputs=[379 input_audio_mic,380 input_audio_file,381 ],382 queue=False,383 api_name=False,384 )385 task_name.change(386 fn=update_input_ui,387 inputs=task_name,388 outputs=[389 audio_box,390 input_text,391 source_language,392 target_language,393 ],394 queue=False,395 api_name=False,396 ).then(397 fn=update_output_ui,398 inputs=task_name,399 outputs=[output_audio, output_text],400 queue=False,401 api_name=False,402 ).then(403 fn=update_example_ui,404 inputs=task_name,405 outputs=[406 s2st_example_row,407 s2tt_example_row,408 t2st_example_row,409 t2tt_example_row,410 asr_example_row,411 ],412 queue=False,413 api_name=False,414 )415 416 btn.click(417 fn=predict,418 inputs=[419 task_name,420 audio_source,421 input_audio_mic,422 input_audio_file,423 input_text,424 source_language,425 target_language,426 ],427 outputs=[output_audio, output_text],428 api_name="run",429 )430demo.queue(max_size=50).launch()431 432# Linking models to the space433# 'facebook/seamless-m4t-large'434# 'facebook/SONAR'435 