Views
No views yet
1from transformers import(
2 EncoderDecoderModel,
3 PreTrainedTokenizerFast,
4 BertJapaneseTokenizer,
5)
6
7import torch
8
9encoder_model_name = "cl-tohoku/bert-base-japanese-v2"
10decoder_model_name = "skt/kogpt2-base-v2"
11
12src_tokenizer = BertJapaneseTokenizer.from_pretrained(encoder_model_name)
13trg_tokenizer = PreTrainedTokenizerFast.from_pretrained(decoder_model_name)
14
15# You should change following `./best_model` to the path of model **directory**
16model = EncoderDecoderModel.from_pretrained("./best_model")
17
18text = "ギルガメッシュ討伐戦"
19# text = "ギルガメッシュ討伐戦に行ってきます。一緒に行きましょうか?"
20
21def translate(text_src):
22 embeddings = src_tokenizer(text_src, return_attention_mask=False, return_token_type_ids=False, return_tensors='pt')
23 embeddings = {k: v for k, v in embeddings.items()}
24 output = model.generate(**embeddings, max_length=500)[0, 1:-1]
25 text_trg = trg_tokenizer.decode(output.cpu())
26 return text_trg
27
28print(translate(text))1from transformers import BertJapaneseTokenizer,PreTrainedTokenizerFast
2from optimum.onnxruntime import ORTModelForSeq2SeqLM
3from onnxruntime import SessionOptions
4import torch
5
6encoder_model_name = "cl-tohoku/bert-base-japanese-v2"
7decoder_model_name = "skt/kogpt2-base-v2"
8
9src_tokenizer = BertJapaneseTokenizer.from_pretrained(encoder_model_name)
10trg_tokenizer = PreTrainedTokenizerFast.from_pretrained(decoder_model_name)
11
12sess_options = SessionOptions()
13sess_options.log_severity_level = 3 # mute warnings including CleanUnusedInitializersAndNodeArgs
14# change subfolder to "onnxq" if you want to use the quantized model
15model = ORTModelForSeq2SeqLM.from_pretrained("sappho192/ffxiv-ja-ko-translator",
16 sess_options=sess_options, subfolder="onnx")
17
18texts = [
19 "逃げろ!", # Should be "도망쳐!"
20 "初めまして.", # "반가워요"
21 "よろしくお願いします.", # "잘 부탁드립니다."
22 "ギルガメッシュ討伐戦", # "길가메쉬 토벌전"
23 "ギルガメッシュ討伐戦に行ってきます。一緒に行きましょうか?", # "길가메쉬 토벌전에 갑니다. 같이 가실래요?"
24 "夜になりました", # "밤이 되었습니다"
25 "ご飯を食べましょう." # "음, 이제 식사도 해볼까요"
26 ]
27
28
29def translate(text_src):
30 embeddings = src_tokenizer(text_src, return_attention_mask=False, return_token_type_ids=False, return_tensors='pt')
31 print(f'Src tokens: {embeddings.data["input_ids"]}')
32 embeddings = {k: v for k, v in embeddings.items()}
33
34 output = model.generate(**embeddings, max_length=500)[0, 1:-1]
35 print(f'Trg tokens: {output}')
36 text_trg = trg_tokenizer.decode(output.cpu())
37 return text_trg
38
39
40for text in texts:
41 print(translate(text))
42 print()