Views
No views yet
1import time
2from transformers import NllbTokenizer, AutoModelForSeq2SeqLM
3
4
5def fix_tokenizer(tokenizer, new_lang='quz_Latn'):
6 """
7 Add a new language token to the tokenizer vocabulary and update language mappings.
8 """
9 # First ensure we're working with an NLLB tokenizer
10 if not hasattr(tokenizer, 'sp_model'):
11 raise ValueError("This function expects an NLLB tokenizer")
12
13 # Add the new language token if it's not already present
14 if new_lang not in tokenizer.additional_special_tokens:
15 tokenizer.add_special_tokens({
16 'additional_special_tokens': [new_lang]
17 })
18
19 # Initialize lang_code_to_id if it doesn't exist
20 if not hasattr(tokenizer, 'lang_code_to_id'):
21 tokenizer.lang_code_to_id = {}
22
23 # Add the new language to lang_code_to_id mapping
24 if new_lang not in tokenizer.lang_code_to_id:
25 # Get the ID for the new language token
26 new_lang_id = tokenizer.convert_tokens_to_ids(new_lang)
27 tokenizer.lang_code_to_id[new_lang] = new_lang_id
28
29 # Initialize id_to_lang_code if it doesn't exist
30 if not hasattr(tokenizer, 'id_to_lang_code'):
31 tokenizer.id_to_lang_code = {}
32
33 # Update the reverse mapping
34 tokenizer.id_to_lang_code[tokenizer.lang_code_to_id[new_lang]] = new_lang
35
36 return tokenizer
37
38
39MODEL_URL = "pollitoconpapass/QnIA-translation-model"
40model = AutoModelForSeq2SeqLM.from_pretrained(MODEL_URL)
41tokenizer = NllbTokenizer.from_pretrained(MODEL_URL)
42fix_tokenizer(tokenizer)
43
44def translate(text, src_lang='spa_Latn', tgt_lang='quz_Latn', a=32, b=3, max_input_length=1024, num_beams=4, **kwargs):
45 tokenizer.src_lang = src_lang
46 tokenizer.tgt_lang = tgt_lang
47 inputs = tokenizer(text, return_tensors='pt', padding=True, truncation=True, max_length=max_input_length)
48 result = model.generate(
49 **inputs.to(model.device),
50 forced_bos_token_id=tokenizer.convert_tokens_to_ids(tgt_lang),
51 max_new_tokens=int(a + b * inputs.input_ids.shape[1]),
52 num_beams=num_beams,
53 **kwargs
54 )
55 return tokenizer.batch_decode(result, skip_special_tokens=True)
56
57
58def translate_v2(text, model, tokenizer, src_lang='spa_Latn', tgt_lang='quz_Latn',
59 max_length='auto', num_beams=4, no_repeat_ngram_size=4, n_out=None, **kwargs):
60
61 tokenizer.src_lang = src_lang
62 encoded = tokenizer(text, return_tensors="pt", truncation=True, max_length=512)
63 if max_length == 'auto':
64 max_length = int(32 + 2.0 * encoded.input_ids.shape[1])
65 model.eval()
66 generated_tokens = model.generate(
67 **encoded.to(model.device),
68 forced_bos_token_id=tokenizer.lang_code_to_id[tgt_lang],
69 max_length=max_length,
70 num_beams=num_beams,
71 no_repeat_ngram_size=no_repeat_ngram_size,
72 num_return_sequences=n_out or 1,
73 **kwargs
74 )
75 out = tokenizer.batch_decode(generated_tokens, skip_special_tokens=True)
76 if isinstance(text, str) and n_out is None:
77 return out[0]
78 return out
79
80
81# === MAIN ===
82t = '''
83Subes centelleante de labios y de ojeras!
84Por tus venas subo, como un can herido
85que busca el refugio de blandas aceras.
86
87Amor, en el mundo tú eres un pecado!
88Mi beso en la punta chispeante del cuerno
89del diablo; mi beso que es credo sagrado!
90'''
91
92start = time.time()
93result_v1 = translate(t, 'spa_Latn', 'quz_Latn')
94print(f"\n{result_v1}")
95
96end = time.time()
97print(f"\nTime for method v1: {end - start}")
98
99
100# start_v2 = time.time()
101# result_v2 = translate_v2(t, model, tokenizer)
102# print(result_v2)
103
104# end_v2 = time.time()
105# print(f"\nTime for method v1: {end_v2 - start_v2}")