Arush kumar commited on
Commit
bf55e4e
·
1 Parent(s): 27895e8

Upload tokenizer.py

Browse files
Files changed (1) hide show
  1. tokenizer.py +80 -0
tokenizer.py ADDED
@@ -0,0 +1,80 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+
3
+ import json
4
+ from pathlib import Path
5
+ from typing import Sequence, List
6
+
7
+ import sentencepiece as spm
8
+
9
+
10
+ def train_sentencepiece(
11
+ data_files: Sequence[str],
12
+ model_prefix: str = 'tokenizer',
13
+ vocab_size: int = 32000,
14
+ model_type: str = 'bpe',
15
+ character_coverage: float = 1.0,
16
+ byte_fallback: bool = True,
17
+ pad_id: int = 1,
18
+ unk_id: int = 0,
19
+ bos_id: int = 2,
20
+ eos_id: int = 3,
21
+ ) -> str:
22
+ data_files = [str(Path(p)) for p in data_files]
23
+ if not data_files:
24
+ raise ValueError('data_files is empty')
25
+ args = [
26
+ f"--input={','.join(data_files)}",
27
+ f'--model_prefix={model_prefix}',
28
+ f'--vocab_size={int(vocab_size)}',
29
+ f'--model_type={model_type}',
30
+ f'--character_coverage={character_coverage}',
31
+ f'--pad_id={pad_id}',
32
+ f'--unk_id={unk_id}',
33
+ f'--bos_id={bos_id}',
34
+ f'--eos_id={eos_id}',
35
+ '--hard_vocab_limit=false',
36
+ '--normalization_rule_name=nmt_nfkc',
37
+ ]
38
+ if byte_fallback:
39
+ args.append('--byte_fallback=true')
40
+ spm.SentencePieceTrainer.train(' '.join(args))
41
+ return f'{model_prefix}.model'
42
+
43
+
44
+ class TokenizerWrapper:
45
+ def __init__(self, model_path: str):
46
+ model_path = str(Path(model_path))
47
+ if not Path(model_path).exists():
48
+ raise FileNotFoundError(model_path)
49
+ self.sp = spm.SentencePieceProcessor(model_file=model_path)
50
+ self.vocab_size = int(self.sp.vocab_size())
51
+ self.pad_id = self.sp.pad_id()
52
+ self.unk_id = self.sp.unk_id()
53
+ self.bos_id = self.sp.bos_id()
54
+ self.eos_id = self.sp.eos_id()
55
+ for name, val in [('pad', self.pad_id), ('unk', self.unk_id), ('bos', self.bos_id), ('eos', self.eos_id)]:
56
+ if val < 0:
57
+ raise ValueError(f'SentencePiece model missing <{name}>')
58
+
59
+ def encode(self, text: str, add_bos: bool = True, add_eos: bool = False) -> List[int]:
60
+ ids = list(self.sp.encode(text, out_type=int))
61
+ if add_bos:
62
+ ids = [self.bos_id] + ids
63
+ if add_eos:
64
+ ids = ids + [self.eos_id]
65
+ return ids
66
+
67
+ def encode_batch(self, texts: Sequence[str], add_bos: bool = True, add_eos: bool = False) -> List[List[int]]:
68
+ return [self.encode(t, add_bos=add_bos, add_eos=add_eos) for t in texts]
69
+
70
+ def decode(self, ids: Sequence[int]) -> str:
71
+ return self.sp.decode([int(i) for i in ids if int(i) != self.pad_id])
72
+
73
+ def save_config(self, path: str) -> None:
74
+ Path(path).write_text(json.dumps({
75
+ 'vocab_size': self.vocab_size,
76
+ 'pad_id': self.pad_id,
77
+ 'unk_id': self.unk_id,
78
+ 'bos_id': self.bos_id,
79
+ 'eos_id': self.eos_id,
80
+ }, indent=2), encoding='utf-8')