Felipe97/llama-cpp-compiled
01.1k
1# Test libllama tokenizer == AutoTokenizer.2# Brute force random words/text generation.3#4# Sample usage:5#6# python3 tests/test-tokenizer-random.py ./models/ggml-vocab-llama-bpe.gguf ./models/tokenizers/llama-bpe7#8 9from __future__ import annotations10 11import time12import logging13import argparse14import subprocess15import random16import unicodedata17 18from pathlib import Path19from typing import Any, Iterator20 21import cffi22from transformers import AutoTokenizer, PreTrainedTokenizer23 24 25logger = logging.getLogger("test-tokenizer-random")26 27 28class LibLlama:29 30 DEFAULT_PATH_LLAMA_H = "./include/llama.h"31 DEFAULT_PATH_INCLUDES = ["./ggml/include/", "./include/"]32 DEFAULT_PATH_LIBLLAMA = "./build/src/libllama.so" # CMakeLists.txt: BUILD_SHARED_LIBS ON33 34 def __init__(self, path_llama_h: str | None = None, path_includes: list[str] = [], path_libllama: str | None = None):35 path_llama_h = path_llama_h or self.DEFAULT_PATH_LLAMA_H36 path_includes = path_includes or self.DEFAULT_PATH_INCLUDES37 path_libllama = path_libllama or self.DEFAULT_PATH_LIBLLAMA38 (self.ffi, self.lib) = self._load_libllama_cffi(path_llama_h, path_includes, path_libllama)39 self.lib.llama_backend_init()40 41 def _load_libllama_cffi(self, path_llama_h: str, path_includes: list[str], path_libllama: str) -> tuple[cffi.FFI, Any]:42 cmd = ["gcc", "-O0", "-E", "-P", "-D__restrict=", "-D__attribute__(x)=", "-D__asm__(x)="]43 cmd += ["-I" + path for path in path_includes] + [path_llama_h]44 res = subprocess.run(cmd, stdout=subprocess.PIPE)45 assert (res.returncode == 0)46 source = res.stdout.decode()47 ffi = cffi.FFI()48 if True: # workarounds for pycparser49 source = "typedef struct { } __builtin_va_list;" + "\n" + source50 source = source.replace("sizeof (int)", str(ffi.sizeof("int")))51 source = source.replace("sizeof (void *)", str(ffi.sizeof("void*")))52 source = source.replace("sizeof (size_t)", str(ffi.sizeof("size_t")))53 source = source.replace("sizeof(int32_t)", str(ffi.sizeof("int32_t")))54 ffi.cdef(source, override=True)55 lib = ffi.dlopen(path_libllama)56 return (ffi, lib)57 58 def model_default_params(self, **kwargs):59 mparams = self.lib.llama_model_default_params()60 for k, v in kwargs.items():61 setattr(mparams, k, v)62 return mparams63 64 def context_default_params(self, **kwargs):65 cparams = self.lib.llama_context_default_params()66 for k, v in kwargs.items():67 setattr(cparams, k, v)68 return cparams69 70 71class LibLlamaModel:72 73 def __init__(self, libllama: LibLlama, path_model: str, mparams={}, cparams={}):74 self.lib: Any = libllama.lib75 self.ffi = libllama.ffi76 if isinstance(mparams, dict):77 mparams = libllama.model_default_params(**mparams)78 self.model = self.lib.llama_model_load_from_file(path_model.encode(), mparams)79 if not self.model:80 raise RuntimeError("error: failed to load model '%s'" % path_model)81 if isinstance(cparams, dict):82 cparams = libllama.context_default_params(**cparams)83 self.ctx = self.lib.llama_new_context_with_model(self.model, cparams)84 if not self.ctx:85 raise RuntimeError("error: failed to create context for model '%s'" % path_model)86 n_tokens_max = self.lib.llama_n_ctx(self.ctx)87 self.token_ids = self.ffi.new("llama_token[]", n_tokens_max)88 self.text_buff = self.ffi.new("uint8_t[]", 1024)89 90 def free(self):91 if self.ctx:92 self.lib.llama_free(self.ctx)93 if self.model:94 self.lib.llama_model_free(self.model)95 self.ctx = None96 self.model = None97 self.lib = None98 99 def tokenize(self, text: str, add_special: bool = False, parse_special: bool = False) -> list[int]:100 encoded_text: bytes = text.encode("utf-8")101 num = self.lib.llama_tokenize(self.model, encoded_text, len(encoded_text), self.token_ids, len(self.token_ids), add_special, parse_special)102 while num < 0 and len(self.token_ids) < (16 << 20):103 self.token_ids = self.ffi.new("llama_token[]", -2 * num)104 num = self.lib.llama_tokenize(self.model, encoded_text, len(encoded_text), self.token_ids, len(self.token_ids), add_special, parse_special)105 return list(self.token_ids[0:num])106 107 def detokenize(self, ids: list[int], remove_special: bool = False, unparse_special: bool = False) -> str:108 if len(self.token_ids) < len(ids):109 self.token_ids = self.ffi.new("llama_token[]", 2 * len(ids))110 for i, id in enumerate(ids):111 self.token_ids[i] = id112 num = self.lib.llama_detokenize(self.model, self.token_ids, len(ids), self.text_buff, len(self.text_buff), remove_special, unparse_special)113 while num < 0 and len(self.text_buff) < (16 << 20):114 self.text_buff = self.ffi.new("uint8_t[]", -2 * num)115 num = self.lib.llama_detokenize(self.model, self.token_ids, len(ids), self.text_buff, len(self.text_buff), remove_special, unparse_special)116 return str(self.ffi.buffer(self.text_buff, num), encoding="utf-8", errors="replace") # replace errors with '\uFFFD' # pyright: ignore[reportArgumentType]117 118 119class Tokenizer:120 121 def encode(self, text: str) -> list[int]:122 raise NotImplementedError123 124 def decode(self, ids: list[int]) -> str:125 raise NotImplementedError126 127 128class TokenizerGroundtruth (Tokenizer):129 130 def __init__(self, dir_tokenizer: str):131 self.model: PreTrainedTokenizer = AutoTokenizer.from_pretrained(dir_tokenizer) # ty: ignore[invalid-assignment]132 # guess BOS and EOS133 ids = self.encode("a")134 assert 1 <= len(ids) <= 3135 add_bos_token = len(ids) > 1 and self.model.bos_token_id == ids[0]136 add_eos_token = len(ids) > 1 and self.model.eos_token_id == ids[-1]137 self.add_bos_token = getattr(self.model, "add_bos_token", add_bos_token)138 self.add_eos_token = getattr(self.model, "add_eos_token", add_eos_token)139 # build vocab140 tokens = list(self.model.get_vocab().values())141 self.vocab = self.model.batch_decode(tokens, skip_special_tokens=True)142 self.vocab = list(sorted(self.vocab))143 # tokens and lists144 self.special_tokens = list(self.model.all_special_tokens)145 self.added_tokens = self.model.batch_decode(list(self.model.added_tokens_encoder.values()), skip_special_tokens=False)146 self.bos_token = self.model.bos_token147 self.eos_token = self.model.eos_token148 149 def encode(self, text: str) -> list[int]:150 return self.model.encode(text, add_special_tokens=True)151 152 def decode(self, ids: list[int]) -> str:153 return self.model.decode(ids, skip_special_tokens=False) # ty: ignore[invalid-return-type]154 155 156class TokenizerLlamaCpp (Tokenizer):157 158 libllama: LibLlama | None = None159 160 def __init__(self, vocab_file: str):161 if not self.libllama:162 self.libllama = LibLlama()163 self.model = LibLlamaModel(self.libllama, vocab_file, mparams=dict(vocab_only=True), cparams=dict(n_ctx=4096))164 165 def encode(self, text: str) -> list[int]:166 return self.model.tokenize(text, add_special=True, parse_special=True)167 168 def decode(self, ids: list[int]) -> str:169 return self.model.detokenize(ids, remove_special=False, unparse_special=True)170 171 172def generator_custom_text() -> Iterator[str]:173 """General tests"""174 yield from [175 "",176 " ",177 " ",178 " ",179 "\t",180 "\n",181 "\n\n",182 "\n\n\n",183 "\t\n",184 "Hello world",185 " Hello world",186 "Hello World",187 " Hello World",188 " Hello World!",189 "Hello, world!",190 " Hello, world!",191 " this is 🦙.cpp",192 "w048 7tuijk dsdfhu",193 "нещо на Български",194 "កាន់តែពិសេសអាចខលចេញ",195 "🚀 (normal) 😶🌫️ (multiple emojis concatenated) ✅ (only emoji that has its own token)",196 "Hello",197 " Hello",198 " Hello",199 " Hello",200 " Hello",201 " Hello\n Hello",202 " (",203 "\n =",204 "' era",205 "Hello, y'all! How are you 😁 ?我想在apple工作1314151天~",206 "3",207 "33",208 "333",209 "3333",210 "33333",211 "333333",212 "3333333",213 "33333333",214 "333333333",215 ]216 217 218def generator_custom_text_edge_cases() -> Iterator[str]:219 """Edge cases found while debugging"""220 yield from [221 '\x1f-a', # unicode_ranges_control, {0x00001C, 0x00001F}222 '¼-a', # unicode_ranges_digit, 0x00BC223 '½-a', # unicode_ranges_digit, 0x00BD224 '¾-a', # unicode_ranges_digit, 0x00BE225 'a 〇b', # unicode_ranges_digit, 0x3007226 'Ⅵ-a', # unicode_ranges_digit, {0x00002150, 0x0000218F} // Number Forms227 '\uFEFF//', # unicode_ranges_control, 0xFEFF (BOM)228 'Cửa Việt', # llama-3, ignore_merges = true229 '<s>a', # Phi-3 fail230 '<unk><|endoftext|><s>', # Phi-3 fail231 'a\na', # bert fail232 '"`', # falcon233 ' \u2e4e', # falcon234 '\n\x0b ', # falcon235 'a\xa0\xa0\x00b', # jina-v2-es236 'one <mask>', # jina-v2-es <mask> lstrip=true237 'a </s> b', # rstrip phi-3238 'a <mask> b', # lstrip jina-v2239 '\xa0aC', # deepseek240 '\u2029 \uA3E4', # deepseek-llm241 "a ?",242 'å', # mpt243 '\U000ac517', # utf-8 encode error, falcon244 '\U000522f4', # utf-8 encode error, starcoder245 "<s><s><unk><s>a<s>b<s>c<unk>d<unk></s>",246 "<s> <s> <unk><s>a<s>b<s>c<unk>d<unk></s>",247 ]248 249 250def generator_vocab_words(tokenizer: TokenizerGroundtruth) -> Iterator[str]:251 """Brute force check all vocab words"""252 yield from tokenizer.vocab253 254 255def generator_ascii_lr_strip() -> Iterator[str]:256 WHITESPACES = ["", " ", " "]257 CHARACTERS = list(chr(i) for i in range(1, 0x80)) + [""]258 for char1 in CHARACTERS:259 for char2 in CHARACTERS:260 for lstrip in WHITESPACES:261 for rstrip in WHITESPACES:262 yield lstrip + char1 + char2 + rstrip263 yield lstrip + char1 + rstrip + char2264 yield char1 + lstrip + char2 + rstrip265 266 267def generator_apostrophe() -> Iterator[str]:268 WHITESPACES = ["", " ", " "]269 CHARACTERS = list(chr(i) for i in range(1, 0x80)) + [""]270 for char1 in CHARACTERS:271 for char2 in CHARACTERS:272 for lstrip in WHITESPACES:273 for rstrip in WHITESPACES:274 yield char1 + lstrip + "'" + rstrip + char2275 yield char1 + char2 + lstrip + "'" + rstrip + "z"276 yield "a" + lstrip + "'" + rstrip + char1 + char2277 278 279def generator_added_lr_strip(tokenizer: TokenizerGroundtruth) -> Iterator[str]:280 WHITESPACES = ["", " ", " ", "\n", "\r\n", "\n\n", "\t", "\t\t"]281 all_tokens = list(sorted(set(tokenizer.special_tokens + tokenizer.added_tokens)))282 for token in all_tokens:283 for lstrip in WHITESPACES:284 for rstrip in WHITESPACES:285 yield lstrip + token + rstrip286 yield "a" + lstrip + token + rstrip287 yield lstrip + token + rstrip + "z"288 yield "a" + lstrip + token + rstrip + "z"289 290 291def generator_random_added_tokens(tokenizer: TokenizerGroundtruth, iterations=100) -> Iterator[str]:292 separations = [" ", "\n", "\t", "-", "!", "one", "1", "<s>", "</s>"]293 all_tokens = list(sorted(set(tokenizer.special_tokens + tokenizer.added_tokens + separations)))294 rand = random.Random()295 for m in range(iterations):296 rand.seed(m)297 words = rand.choices(all_tokens, k=500)298 if words and words[0] == tokenizer.bos_token: # skip spam warning of double BOS299 while len(words) > 1 and words[1] == tokenizer.bos_token: # leave one starting BOS300 words.pop(0)301 if tokenizer.add_bos_token: # drop all starting BOS302 words.pop(0)303 if words and words[-1] == tokenizer.eos_token: # skip spam warning of double EOS304 while len(words) > 1 and words[-2] == tokenizer.eos_token: # leave one trailing EOS305 words.pop(-1)306 if tokenizer.add_bos_token: # drop all trailing EOS307 words.pop(-1)308 yield "".join(words)309 310 311def generator_random_chars(iterations=100) -> Iterator[str]:312 """Brute force random text with simple characters"""313 314 NUM_WORDS = 400315 WHITESPACES = list(" " * 20 + "\n" * 5 + "\r\n" * 5 + "\t" * 5)316 CHARS = list(sorted(set("""317 ABCDEFGHIJKLMNOPQRSTUVWXYZ318 abcdefghijklmnopqrstuvwxyz319 ÁÉÍÓÚÀÈÌÒÙÂÊÎÔÛÄËÏÖÜ320 áéíóúàèìòùâêîôûäëïöü321 .-,*/-+ª!"·$%&/()=?¿[]{}<>\\|@#~½¬~;:_322 """)))323 324 rand = random.Random()325 for m in range(iterations):326 rand.seed(m)327 text = []328 for _ in range(NUM_WORDS):329 k = rand.randint(1, 7)330 word = rand.choices(CHARS, k=k)331 word.append(rand.choice(WHITESPACES))332 text.append("".join(word))333 yield "".join(text)334 335 336def generator_unicodes() -> Iterator[str]:337 """Iterate unicode characters"""338 339 MAX_CODEPOINTS = 0x30000 # 0x110000340 341 def _valid(cpt):342 if cpt >= 0x30000: # unassigned and supplementary343 return False344 # if cpt == 0x2029: # deepseek-llm345 # return False346 if unicodedata.category(chr(cpt)) in ("Cn", "Cs", "Co"): # undefined, surrogates, private347 return False348 return True349 350 characters = [chr(cpt) for cpt in range(0, MAX_CODEPOINTS) if _valid(cpt)]351 352 yield from characters353 354 355def generator_random_unicodes(iterations=100) -> Iterator[str]:356 """Brute force random text with unicode characters"""357 358 NUM_WORDS = 200359 WHITESPACES = list(" " * 20 + "\n" * 5 + "\r\n" * 5 + "\t" * 5)360 361 characters = list(generator_unicodes())362 363 rand = random.Random()364 for m in range(iterations):365 rand.seed(m)366 text = []367 for _ in range(NUM_WORDS):368 k = rand.randint(1, 7)369 word = rand.choices(characters, k=k)370 word.append(rand.choice(WHITESPACES))371 text.append("".join(word))372 yield "".join(text)373 374 375def generator_random_vocab_chars(tokenizer: TokenizerGroundtruth, iterations=100) -> Iterator[str]:376 """Brute force random text with vocab characters"""377 378 vocab_chars = set()379 for word in tokenizer.vocab:380 vocab_chars.update(word)381 vocab_chars = list(sorted(vocab_chars))382 383 rand = random.Random()384 for m in range(iterations):385 rand.seed(m)386 text = rand.choices(vocab_chars, k=1024)387 yield "".join(text)388 389 390def generator_random_vocab_words(tokenizer: TokenizerGroundtruth, iterations=100) -> Iterator[str]:391 """Brute force random text from vocab words"""392 393 vocab = [w.strip() for w in tokenizer.vocab]394 yield from vocab395 396 rand = random.Random()397 for m in range(iterations):398 rand.seed(m)399 text = []400 num_words = rand.randint(300, 400)401 for i in range(num_words):402 k = rand.randint(1, 3)403 words = rand.choices(vocab, k=k)404 sep = rand.choice(" \n\r\t")405 text.append("".join(words) + sep)406 yield "".join(text)407 408 409def compare_tokenizers(tokenizer1: TokenizerGroundtruth, tokenizer2: TokenizerLlamaCpp, generator: Iterator[str]):410 411 def find_first_mismatch(ids1: list[int] | str, ids2: list[int] | str):412 for i, (a, b) in enumerate(zip(ids1, ids2)):413 if a != b:414 return i415 if len(ids1) == len(ids2):416 return -1417 return min(len(ids1), len(ids2))418 419 def check_detokenizer(text: str, text1: str, text2: str) -> bool:420 if text1 == text2: # equal to TokenizerGroundtruth?421 return True422 # equal to source text?423 if tokenizer1.add_bos_token and tokenizer1.bos_token and isinstance(tokenizer1.bos_token, str): # remove BOS424 if text2.startswith(tokenizer1.bos_token):425 text2 = text2[len(tokenizer1.bos_token):]426 if tokenizer1.add_eos_token and tokenizer1.eos_token and isinstance(tokenizer1.eos_token, str): # remove EOS427 if text2.endswith(tokenizer1.eos_token):428 text2 = text2[:-len(tokenizer1.eos_token)]429 return text == text2430 431 t_encode1 = 0432 t_encode2 = 0433 t_decode1 = 0434 t_decode2 = 0435 t_start = time.perf_counter()436 encode_errors = 0437 decode_errors = 0438 MAX_ERRORS = 10439 440 logger.info("%s: %s" % (getattr(generator, "__qualname__", ""), "ini"))441 for text in generator:442 # print(repr(text), text.encode())443 # print(repr(text), hex(ord(text[0])), text.encode())444 t0 = time.perf_counter()445 ids1 = tokenizer1.encode(text)446 t1 = time.perf_counter()447 ids2 = tokenizer2.encode(text)448 t2 = time.perf_counter()449 text1 = tokenizer1.decode(ids1)450 t3 = time.perf_counter()451 text2 = tokenizer2.decode(ids1)452 t4 = time.perf_counter()453 t_encode1 += t1 - t0454 t_encode2 += t2 - t1455 t_decode1 += t3 - t2456 t_decode2 += t4 - t3457 if encode_errors < MAX_ERRORS and ids1 != ids2:458 i = find_first_mismatch(ids1, ids2)459 ids1 = list(ids1)[max(0, i - 2) : i + 5 + 1]460 ids2 = list(ids2)[max(0, i - 2) : i + 5 + 1]461 logger.error(" Expected: " + str(ids1))462 logger.error(" Result: " + str(ids2))463 encode_errors += 1464 logger.error(f" {encode_errors=}")465 if decode_errors < MAX_ERRORS and not check_detokenizer(text, text1, text2):466 i = find_first_mismatch(text1, text2)467 text1 = list(text1[max(0, i - 2) : i + 5 + 1])468 text2 = list(text2[max(0, i - 2) : i + 5 + 1])469 logger.error(" Expected: " + " ".join(hex(ord(x)) for x in text1))470 logger.error(" Result: " + " ".join(hex(ord(x)) for x in text2))471 decode_errors += 1472 logger.error(f" {decode_errors=}")473 if encode_errors >= MAX_ERRORS and decode_errors >= MAX_ERRORS:474 logger.error(f" EXIT: {encode_errors=} {decode_errors=}")475 # raise Exception()476 break477 478 t_total = time.perf_counter() - t_start479 logger.info(f"{getattr(generator, '__qualname__', '')}: end, {t_encode1=:.3f} {t_encode2=:.3f} {t_decode1=:.3f} {t_decode2=:.3f} {t_total=:.3f}")480 481 482def main(argv: list[str] | None = None):483 parser = argparse.ArgumentParser()484 parser.add_argument("vocab_file", type=str, help="path to vocab 'gguf' file")485 parser.add_argument("dir_tokenizer", type=str, help="directory containing 'tokenizer.model' file")486 parser.add_argument("--verbose", action="store_true", help="increase output verbosity")487 args = parser.parse_args(argv)488 489 logging.basicConfig(level = logging.DEBUG if args.verbose else logging.INFO)490 logger.info(f"VOCABFILE: '{args.vocab_file}'")491 492 tokenizer1 = TokenizerGroundtruth(args.dir_tokenizer)493 tokenizer2 = TokenizerLlamaCpp(args.vocab_file)494 495 # compare_tokenizers(tokenizer1, tokenizer2, generator_custom_text())496 # compare_tokenizers(tokenizer1, tokenizer2, generator_custom_text_edge_cases())497 compare_tokenizers(tokenizer1, tokenizer2, generator_ascii_lr_strip())498 compare_tokenizers(tokenizer1, tokenizer2, generator_apostrophe())499 compare_tokenizers(tokenizer1, tokenizer2, generator_unicodes())500 compare_tokenizers(tokenizer1, tokenizer2, generator_vocab_words(tokenizer1))501 compare_tokenizers(tokenizer1, tokenizer2, generator_added_lr_strip(tokenizer1))502 # compare_tokenizers(tokenizer1, tokenizer2, generator_random_added_tokens(tokenizer1, 10_000))503 # compare_tokenizers(tokenizer1, tokenizer2, generator_random_chars(10_000))504 # compare_tokenizers(tokenizer1, tokenizer2, generator_random_unicodes(10_000))505 # compare_tokenizers(tokenizer1, tokenizer2, generator_random_vocab_chars(tokenizer1, 10_000))506 # compare_tokenizers(tokenizer1, tokenizer2, generator_random_vocab_words(tokenizer1, 5_000))507 508 tokenizer2.model.free()509 510 511if __name__ == "__main__":512 # main()513 514 if True:515 logging.basicConfig(516 level = logging.DEBUG,517 format = "%(asctime)s.%(msecs)03d %(name)s %(levelname)s %(message)s",518 datefmt = "%Y-%m-%d %H:%M:%S",519 filename = logger.name + ".log",520 filemode = "a"521 )522 logging.basicConfig(523 level = logging.DEBUG,524 format = "%(levelname)s %(message)s",525 )526 527 path_tokenizers = Path("./models/tokenizers/")528 path_vocab_format = "./models/ggml-vocab-%s.gguf"529 530 tokenizers = [531 "llama-spm", # SPM532 "phi-3", # SPM533 "gemma", # SPM534 "gemma-2", # SPM535 "baichuan", # SPM536 "bert-bge", # WPM537 "jina-v2-en", # WPM538 "llama-bpe", # BPE539 "phi-2", # BPE540 "deepseek-llm", # BPE541 "deepseek-coder", # BPE542 "falcon", # BPE543 "mpt", # BPE544 "starcoder", # BPE545 "gpt-2", # BPE546 "stablelm2", # BPE547 "refact", # BPE548 "qwen2", # BPE549 "olmo", # BPE550 "jina-v2-es", # BPE551 "jina-v2-de", # BPE552 "smaug-bpe", # BPE553 "poro-chat", # BPE554 "jina-v2-code", # BPE555 "viking", # BPE556 "jais", # BPE557 ]558 559 logger.info("=" * 50)560 for tokenizer in tokenizers:561 logger.info("-" * 50)562 logger.info(f"TOKENIZER: '{tokenizer}'")563 vocab_file = Path(path_vocab_format % tokenizer)564 dir_tokenizer = path_tokenizers / tokenizer565 main([str(vocab_file), str(dir_tokenizer), "--verbose"])566 