| 28 | |
| 29 | |
| 30 | class TokenExtender: |
| 31 | def __init__(self, data_path, dataset, index_file=".index.json"): |
| 32 | self.data_path = data_path |
| 33 | self.dataset = dataset |
| 34 | self.index_file = index_file |
| 35 | self.indices = None |
| 36 | self.new_tokens = None |
| 37 | |
| 38 | def _load_data(self): |
| 39 | with open(os.path.join(self.data_path, self.dataset + self.index_file), 'r') as f: |
| 40 | self.indices = json.load(f) |
| 41 | |
| 42 | def get_new_tokens(self): |
| 43 | if self.new_tokens is not None: |
| 44 | return self.new_tokens |
| 45 | |
| 46 | if self.indices is None: |
| 47 | self._load_data() |
| 48 | |
| 49 | self.new_tokens = set() |
| 50 | for index in self.indices.values(): |
| 51 | for token in index: |
| 52 | self.new_tokens.add(token) |
| 53 | self.new_tokens = sorted(list(self.new_tokens)) |
| 54 | |
| 55 | return self.new_tokens |
| 56 | |
| 57 | |
| 58 | def set_seed(seed): |