CoolFace
Apppublic

dwolf/seamless_m4t_d_wolf

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
app.py435 linesDownload Raw Back to root
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