File size: 2,949 Bytes
3d33eb8
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
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()