Views
No views yet
1from transformers import NllbTokenizer, AutoModelForSeq2SeqLM, AutoConfig
2# this code is adapted from the Stopes repo of the NLLB team
3# https://github.com/facebookresearch/stopes/blob/main/stopes/pipelines/monolingual/monolingual_line_processor.py#L214
4
5import re
6import sys
7import typing as tp
8import unicodedata
9from sacremoses import MosesPunctNormalizer
10
11
12mpn = MosesPunctNormalizer(lang="en")
13mpn.substitutions = [
14 (re.compile(r), sub) for r, sub in mpn.substitutions
15]
16
17
18def get_non_printing_char_replacer(replace_by: str = " ") -> tp.Callable[[str], str]:
19 non_printable_map = {
20 ord(c): replace_by
21 for c in (chr(i) for i in range(sys.maxunicode + 1))
22 # same as \p{C} in perl
23 # see https://www.unicode.org/reports/tr44/#General_Category_Values
24 if unicodedata.category(c) in {"C", "Cc", "Cf", "Cs", "Co", "Cn"}
25 }
26
27 def replace_non_printing_char(line) -> str:
28 return line.translate(non_printable_map)
29
30 return replace_non_printing_char
31
32replace_nonprint = get_non_printing_char_replacer(" ")
33
34def preproc(text):
35 clean = mpn.normalize(text)
36 clean = replace_nonprint(clean)
37 # replace 𝓕𝔯𝔞𝔫𝔠𝔢𝔰𝔠𝔞 by Francesca
38 clean = unicodedata.normalize("NFKC", clean)
39 return clean
40
41# loading the model
42model_name = "slone/nllb-600M-azj-eng-v1"
43model = AutoModelForSeq2SeqLM.from_pretrained(model_name).cuda()
44tokenizer = NllbTokenizer.from_pretrained(model_name)
45
46def translate(text, src_lang='eng_Latn', tgt_lang='azj_Latn', a=32, b=3, max_input_length=1024, num_beams=4, **kwargs):
47 tokenizer.src_lang = src_lang
48 tokenizer.tgt_lang = tgt_lang
49 if isinstance(text, str):
50 text = [text]
51 text = [preproc(t) for t in text]
52 inputs = tokenizer(text, return_tensors='pt', padding=True, truncation=True, max_length=max_input_length)
53 result = model.generate(
54 **inputs.to(model.device),
55 forced_bos_token_id=tokenizer.convert_tokens_to_ids(tgt_lang),
56 max_new_tokens=int(a + b * inputs.input_ids.shape[1]),
57 num_beams=num_beams,
58 **kwargs
59 )
60 return tokenizer.batch_decode(result, skip_special_tokens=True)
61
62# Example of translating a couple of texts:
63texts = translate(["To be, or not to be, that is the question.", "Hello, how are you?"], src_lang='eng_Latn', tgt_lang='azj_Latn')
64print(texts)
65# ['Olmaq və ya olmamaq sualdır.', 'Salam, necə var?']1def batched_translate(texts, batch_size=16, **kwargs):
2 """Translate texts in batches of similar length"""
3 idxs, texts2 = zip(*sorted(enumerate(texts), key=lambda p: len(p[1]), reverse=True))
4 results = []
5 for i in trange(0, len(texts2), batch_size):
6 results.extend(translate(texts2[i: i+batch_size], **kwargs))
7 return [p for i, p in sorted(zip(idxs, results))]