Views
No views yet
1import json
2import os
3import torch
4from typing import Dict, List, Optional, Sequence, Union
5
6from transformers import AutoModelForCausalLM
7from transformers.tokenization_utils import AddedToken, PreTrainedTokenizer
8
9
10class CharacterTokenizer(PreTrainedTokenizer):
11 def __init__(
12 self, characters: Sequence[str] = "", model_max_length: int = 1024, **kwargs
13 ):
14 self.characters = characters
15 self.model_max_length = model_max_length
16 cls_token = AddedToken("[CLS]", lstrip=False, rstrip=False)
17 sep_token = AddedToken("[SEP]", lstrip=False, rstrip=False)
18 bos_token = AddedToken("[BOS]", lstrip=False, rstrip=False)
19 eos_token = AddedToken("[EOS]", lstrip=False, rstrip=False)
20 mask_token = AddedToken("[MASK]", lstrip=True, rstrip=False)
21 pad_token = AddedToken("[PAD]", lstrip=False, rstrip=False)
22 unk_token = AddedToken("[UNK]", lstrip=False, rstrip=False)
23
24 self._vocab_str_to_int = {
25 "[CLS]": 0,
26 "[SEP]": 1,
27 "[BOS]": 2,
28 "[MASK]": 3,
29 "[PAD]": 4,
30 "[EOS]": 5,
31 "[UNK]": 6,
32 **{ch: i + 7 for i, ch in enumerate(characters)},
33 }
34 self._vocab_int_to_str = {v: k for k, v in self._vocab_str_to_int.items()}
35
36 super().__init__(
37 bos_token=bos_token,
38 eos_token=eos_token,
39 sep_token=sep_token,
40 cls_token=cls_token,
41 pad_token=pad_token,
42 mask_token=mask_token,
43 unk_token=unk_token,
44 add_prefix_space=False,
45 model_max_length=model_max_length,
46 **kwargs,
47 )
48
49 def vocab_size(self) -> int:
50 return len(self._vocab_str_to_int)
51
52 def get_vocab(self):
53 return self._vocab_str_to_int
54
55 def _tokenize(self, text: str) -> List[str]:
56 return list(text)
57
58 def _convert_token_to_id(self, token: str) -> int:
59 return self._vocab_str_to_int.get(token, self._vocab_str_to_int["[UNK]"])
60
61 def _convert_id_to_token(self, index: int) -> str:
62 return self._vocab_int_to_str[index]
63
64 def convert_tokens_to_string(self, tokens):
65 return "".join(tokens)
66
67 def build_inputs_with_special_tokens(
68 self, token_ids_0: List[int], token_ids_1: Optional[List[int]] = None
69 ) -> List[int]:
70 sep = [self.sep_token_id]
71 cls = [self.cls_token_id]
72 result = cls + token_ids_0 + sep
73 if token_ids_1 is not None:
74 result += token_ids_1 + sep
75 return result
76
77 def get_special_tokens_mask(
78 self,
79 token_ids_0: List[int],
80 token_ids_1: Optional[List[int]] = None,
81 already_has_special_tokens: bool = False,
82 ) -> List[int]:
83 if already_has_special_tokens:
84 return super().get_special_tokens_mask(
85 token_ids_0=token_ids_0,
86 token_ids_1=token_ids_1,
87 already_has_special_tokens=True,
88 )
89
90 result = [1] + ([0] * len(token_ids_0)) + [1]
91 if token_ids_1 is not None:
92 result += ([0] * len(token_ids_1)) + [1]
93 return result
94
95 def create_token_type_ids_from_sequences(
96 self, token_ids_0: List[int], token_ids_1: Optional[List[int]] = None
97 ) -> List[int]:
98 sep = [self.sep_token_id]
99 cls = [self.cls_token_id]
100
101 result = len(cls + token_ids_0 + sep) * [0]
102 if token_ids_1 is not None:
103 result += len(token_ids_1 + sep) * [1]
104 return result
105
106 def get_config(self) -> Dict:
107 return {
108 "char_ords": [ord(ch) for ch in self.characters],
109 "model_max_length": self.model_max_length,
110 }
111
112 @classmethod
113 def from_config(cls, config: Dict) -> "HiraganaTokenizer":
114 cfg = {}
115 cfg["characters"] = [chr(i) for i in config["char_ords"]]
116 cfg["model_max_length"] = config["model_max_length"]
117 return cls(**cfg)
118
119 def save_pretrained(self, save_directory: Union[str, os.PathLike], **kwargs):
120 cfg_file = os.path.join(save_directory, "tokenizer_config.json")
121 cfg = self.get_config()
122 with open(cfg_file, "w") as f:
123 json.dump(cfg, f, indent=4)
124
125 @classmethod
126 def _from_pretrained(
127 cls,
128 resolved_vocab_files,
129 pretrained_model_name_or_path,
130 init_configuration,
131 *init_inputs,
132 token=None,
133 cache_dir=None,
134 local_files_only=False,
135 _commit_hash=None,
136 _is_local=False,
137 trust_remote_code=False,
138 **kwargs,
139 ):
140 config_file = resolved_vocab_files["tokenizer_config_file"]
141 with open(config_file, "r", encoding="utf-8") as f:
142 config = json.load(f)
143 return cls.from_config(config)
144
145
146tokenizer = CharacterTokenizer.from_pretrained("hukuda222/hiragana-gpt2-xsmall")
147model = AutoModelForCausalLM.from_pretrained("hukuda222/hiragana-gpt2-xsmall")
148
149with torch.no_grad():
150 token_ids = tokenizer.encode(
151 "こんにちは", add_special_tokens=False, return_tensors="pt"
152 )
153 output_ids = model.generate(
154 token_ids.to(model.device),
155 max_new_tokens=50,
156 pad_token_id=tokenizer.pad_token_id,
157 eos_token_id=tokenizer.eos_token_id,
158 no_repeat_ngram_size=3,
159 )
160output = tokenizer.decode(
161 output_ids.tolist()[0][token_ids.size(1) :], skip_special_tokens=True
162)
163print(output)