Views
No views yet
1from transformers import EncoderDecoderModel
2from importlib.machinery import SourceFileLoader
3from transformers.file_utils import cached_path, hf_bucket_url
4import torch
5import os
6
7## Load model & tokenizer
8cache_dir='./cache'
9model_name='nguyenvulebinh/spelling-oov'
10
11def download_tokenizer_files():
12 resources = ['envibert_tokenizer.py', 'dict.txt', 'sentencepiece.bpe.model']
13 for item in resources:
14 if not os.path.exists(os.path.join(cache_dir, item)):
15 tmp_file = hf_bucket_url(model_name, filename=item)
16 tmp_file = cached_path(tmp_file,cache_dir=cache_dir)
17 os.rename(tmp_file, os.path.join(cache_dir, item))
18
19download_tokenizer_files()
20spell_tokenizer = SourceFileLoader("envibert.tokenizer",os.path.join(cache_dir,'envibert_tokenizer.py')).load_module().RobertaTokenizer(cache_dir)
21spell_model = EncoderDecoderModel.from_pretrained(model_name)
22
23def oov_spelling(word, num_candidate=1):
24 result = []
25 inputs = spell_tokenizer([word.lower()])
26 input_ids = inputs['input_ids']
27 attention_mask = inputs['attention_mask']
28 inputs = {
29 "input_ids": torch.tensor(input_ids),
30 "attention_mask": torch.tensor(attention_mask)
31 }
32 outputs = spell_model.generate(**inputs, num_return_sequences=num_candidate)
33 for output in outputs.cpu().detach().numpy().tolist():
34 result.append(spell_tokenizer.sp_model.DecodePieces(spell_tokenizer.decode(output, skip_special_tokens=True).split()))
35 return result
36
37oov_spelling('spacespeaker')
38# output: ['x pây x pếch cơ']
39