nllb-test / app.py
cidjeu's picture
Create app.py
3d33eb8 verified
Raw
History Blame Contribute Delete
2.95 kB
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
import gradio as gr
model_name = "facebook/nllb-200-distilled-600M"
tokenizer = AutoTokenizer.from_pretrained(model_name)
model = AutoModelForSeq2SeqLM.from_pretrained(model_name)
lang_readable = {
"aka_Latn": "Akan", "bam_Latn": "Bambara", "dyu_Latn": "Dioula",
"ewe_Latn": "Éwé", "fon_Latn": "Fon", "fuv_Latn": "Peul (Nigéria)",
"hau_Latn": "Haoussa", "ibo_Latn": "Igbo", "kab_Latn": "Kabyle",
"knc_Arab": "Kanouri (arabe)", "knc_Latn": "Kanouri (latin)",
"mos_Latn": "Mooré", "som_Latn": "Somali", "twi_Latn": "Twi",
"wol_Latn": "Wolof", "yor_Latn": "Yoruba", "kbp_Latn": "Kabiye",
"kon_Latn": "Kikongo", "lin_Latn": "Lingala", "lua_Latn": "Luba-Lulua",
"lug_Latn": "Luganda", "luo_Latn": "Luo", "run_Latn": "Kirundi",
"sag_Latn": "Sango", "amh_Ethi": "Amharique", "gaz_Latn": "Oromo",
"kam_Latn": "Kamba", "kik_Latn": "Kikuyu", "kin_Latn": "Kinyarwanda",
"nus_Latn": "Nuer", "swh_Latn": "Swahili", "tir_Ethi": "Tigrigna",
"afr_Latn": "Afrikaans", "bem_Latn": "Bemba", "cjk_Latn": "Chokwe",
"dik_Latn": "Dinka du Sud-Ouest", "kmb_Latn": "Kimbundu",
"nso_Latn": "Sepedi", "nya_Latn": "Chichewa", "sna_Latn": "Shona",
"sot_Latn": "Sesotho", "ssw_Latn": "Swati", "tsn_Latn": "Tswana",
"tso_Latn": "Tsonga", "tum_Latn": "Tumbuka", "umb_Latn": "Umbundu",
"xho_Latn": "Xhosa", "zul_Latn": "Zoulou", "aeb_Arab": "Arabe tunisien",
"ary_Arab": "Arabe marocain", "arz_Arab": "Arabe égyptien",
"arb_Arab": "Arabe standard moderne", "tzm_Tfng": "Tamazight (Tifinagh)",
"taq_Latn": "Tamasheq (latin)", "taq_Tfng": "Tamasheq (Tifinagh)",
"kea_Latn": "Créole capverdien", "plt_Latn": "Malgache (Plateau)",
"fra_Latn": "Français", "eng_Latn": "Anglais",
"por_Latn": "Portugais", "spa_Latn": "Espagnol", "ara_Arab": "Arabe"
}
lang_names = list(lang_readable.values())
def translate(text, src_lang_name, tgt_lang_name):
src_lang = [k for k, v in lang_readable.items() if v == src_lang_name][0]
tgt_lang = [k for k, v in lang_readable.items() if v == tgt_lang_name][0]
# ✅ Définir la langue source dans le tokenizer
tokenizer.src_lang = src_lang
inputs = tokenizer(text, return_tensors="pt")
# ✅ Récupérer l'ID de la langue cible correctement
tgt_lang_id = tokenizer.convert_tokens_to_ids(tgt_lang)
translated_tokens = model.generate(
**inputs,
forced_bos_token_id=tgt_lang_id
)
return tokenizer.decode(translated_tokens[0], skip_special_tokens=True)
demo = gr.Interface(
fn=translate,
inputs=[
gr.Textbox(label="Texte à traduire"),
gr.Dropdown(choices=lang_names, label="Langue source", filterable=True),
gr.Dropdown(choices=lang_names, label="Langue cible", filterable=True)
],
outputs="text",
title="NLLB Translator",
description="Traduction multilingue avec NLLB"
)
demo.launch()