关于编码器逻辑的问题
作者: JackxTong创建于 2024年7月20日更新于 2024年12月10日
我注意到 encode() 方法包含了一个循环,用于找到最小的合并索引:
def encode(self, text):
text_bytes = text.encode("utf-8") # raw bytes
ids = list(text_bytes) # list of integers in range 0..255
while len(ids) >= 2:
stats = get_stats(ids)
pair = min(stats, key=lambda p: self.merges.get(p, float("inf")))
if pair not in self.merges:
break # nothing else can be merged anymore
idx = self.merges[pair]
ids = merge(ids, pair, idx)
return ids我们可以将其简化为如下所示:
def encode(self, text):
tokens = text.encode("utf-8")
tokens = list(map(int, tokens))
for pair, index in self.merges.items():
tokens = merge(tokens, pair, index)
return tokens由于 merge() 会合并所有出现的情况,因此一个简单的循环似乎就足够了。是否有什么复杂的逻辑呢?我已经使用我的 tokenizer 与 basictokenizer 在一些文本数据上进行训练,并实现了完全相同的词汇表和编码器。也许我漏掉了什么。你能否澄清一下?
谢谢!
更新: 我从 我的分支仓库 中创建了一个 pytest,以便展示我的结果也是正确的: 任何想要尝试的用户都可以查看。
内容来源: karpathy/minbpe