Views
No views yet
1from optimum.onnxruntime import ORTModelForSeq2SeqLM
2from transformers import AutoTokenizer,pipeline
3
4model_path = 'models/src_ctx_aware_nllb_1.3B_onnx'
5
6model = ORTModelForSeq2SeqLM.from_pretrained(model_path)
7tokenizer = AutoTokenizer.from_pretrained(model_path)
8
9max_length = 100
10src_lang = 'eng_Latn'
11tgt_lang = 'deu_Latn'
12context_text = 'This is an optional context sentence.'
13sentence_text = 'Text to be translated.'
14
15# If the context is provided
16input_text = f'{context_text} {tokenizer.sep_token} {sentence_text}'
17# If no context is provided, you can use just the sentence_text as input
18# input_text = sentence_text
19
20tokenizer.src_lang = src_lang
21
22inputs = tokenizer(input_text, return_tensors='pt')
23
24input = inputs.to('cpu')
25
26forced_bos_token_id = tokenizer.lang_code_to_id[tgt_lang]
27
28output = model.generate(
29 **inputs,
30 forced_bos_token_id=forced_bos_token_id,
31 max_length=max_length
32)
33
34output_text = tokenizer.batch_decode(output, skip_special_tokens=True)[0]
35
36print(output_text)