tokenizer的encode性能优化
朴素版本
朴素版本为每次encode遍历merge,对遍历的pair从头到尾扫描pretoken判断是否存在。
优化思路
观察到朴素版本绝大部分pair对当前pretoken属于无效,因此出现了几乎等于merge_set大小*pretoken长度的无效扫描,时间复杂度爆炸。
因此引入对较短的pretoken的pairs扫描,并针对pretokens中存在且存在与merge_sets的pair进行merge。引入set进行快速存在性判断。同时注意到需要满足merge_sets在训练时的顺序,将merge_list的key,value置换以实现快速检索。
优化实现
def encode_iterable(self, iterable: Iterable[str]) -> Iterator[int]:
"""
Given an iterable of strings (e.g., a Python file handle), return a generator that lazily yields token IDs.
This is required for memory-efficient tokenization of large files that we cannot directly load into memory.
"""
for line in iterable:
splited_parts = Tokenizer._split(line, self.special_tokens)
PAT = r"""'(?:[sdmt]|ll|ve|re)| ?\p{L}+| ?\p{N}+| ?[^\s\p{L}\p{N}]+|\s+(?!\S)|\s+"""
# part-by-part merge and insert special tokens.
for part in splited_parts:
if part in self.special_tokens:
yield self.vocab_rev[part.encode('utf-8')]
else:
for pretoken in re.finditer(PAT, part):
pretoken = tuple(bytes([b]) for b in pretoken.group().encode('utf-8'))
while True:
cur_merge_idx = defaultdict(list[int])
for (idx, a),b in zip(enumerate(pretoken), pretoken[1:]):
cur_merge_idx[(a,b)].append(idx)
if not any(x in self.merge_sets for x in cur_merge_idx.keys()):
break
cur_merge_pair = min(cur_merge_idx.keys(), key = lambda x : self.merge_dict[x] if x in self.merge_sets else 1e42)
cur_merge_indices = cur_merge_idx[cur_merge_pair]
new_pretoken = []
i = 0
while i < len(pretoken):
if i in cur_merge_indices:
new_pretoken.append(cur_merge_pair[0]+cur_merge_pair[1])
i += 2
else:
new_pretoken.append(pretoken[i])
i += 1
pretoken = tuple(new_pretoken)
for vocab in pretoken:
yield self.vocab_rev[vocab]

comment 评论区
star_outline 咱快来抢个沙发吧!