| 56 | |
| 57 | |
| 58 | class Code13BDictionary(object): |
| 59 | def __init__( |
| 60 | self, |
| 61 | dict_file: str, |
| 62 | extra_token_ids: List[str] = None, |
| 63 | pad_to_vocab_size: int = -1, |
| 64 | ): |
| 65 | self._idx = dict() |
| 66 | self._count = dict() |
| 67 | self._num_symbols = 0 |
| 68 | self._symbols = [] |
| 69 | |
| 70 | self._add_symbol("<s>", 0) |
| 71 | self._add_symbol("<pad>", 0) |
| 72 | self._add_symbol("</s>", 0) |
| 73 | self._add_symbol("<unk>", 0) |
| 74 | self._load_dict(dict_file) |
| 75 | |
| 76 | if extra_token_ids is None: |
| 77 | extra_token_ids = [ |
| 78 | str(x) for x in range(50257, 50400) |
| 79 | ] # follows GPT-J settings |
| 80 | |
| 81 | for token_id in extra_token_ids: |
| 82 | self._add_symbol(token_id, 0) |
| 83 | |
| 84 | if pad_to_vocab_size > 0: |
| 85 | self._pad_to_vocab_size(pad_to_vocab_size) |
| 86 | |
| 87 | def _pad_to_vocab_size(self, vocab_size: int): |
| 88 | num_pad = vocab_size - len(self) |
| 89 | if num_pad <= 0: |
| 90 | return |
| 91 | for i in range(1, num_pad + 1): |
| 92 | self._add_symbol("vocab_pad_token{}".format(i), 0) |
| 93 | |
| 94 | def _load_dict(self, dict_file: str): |
| 95 | with open(dict_file, "r") as f: |
| 96 | for line in f: |
| 97 | line = line.strip() |
| 98 | if line == "" or line.startswith("#"): |
| 99 | continue |
| 100 | sym, count = line.split() |
| 101 | self._add_symbol(sym, int(count)) |
| 102 | |
| 103 | def _add_symbol(self, sym: str, count: int): |
| 104 | self._idx[sym] = self._num_symbols |
| 105 | self._count[sym] = count |
| 106 | self._symbols.append(sym) |
| 107 | self._num_symbols += 1 |
| 108 | |
| 109 | def __len__(self): |
| 110 | return self._num_symbols |
| 111 | |
| 112 | def index(self, sym: str): |
| 113 | return self._idx[sym] |
| 114 | |
| 115 | def string(self, idx: int): |