MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / Code13BDictionary

Class Code13BDictionary

codegeex/mindspore/src/code_tokenizer.py:58–128  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

56
57
58class 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):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected