diff --git a/benchmarks/bench_chinese.py b/benchmarks/bench_chinese.py index 608c043..0c4a19d 100644 --- a/benchmarks/bench_chinese.py +++ b/benchmarks/bench_chinese.py @@ -1,12 +1,14 @@ """Benchmark Chinese BPE training at various scales.""" +import hashlib import time from itertools import islice +import complex_tokenization_fast from datasets import load_dataset -from complex_tokenization import BPETokenizer -from complex_tokenization.graphs.units import register_script -from complex_tokenization.languages.chinese.graph import chinese_character_to_graph +import complex_tokenization + +IMPLS = {"reference (Python)": complex_tokenization, "fast (Rust)": complex_tokenization_fast} def load_texts(n): @@ -19,21 +21,28 @@ def load_texts(n): return [row["text"][:500] for row in islice(ds, n) if row["text"]] -def bench(texts, num_merges=10, label=""): +def bench(module, impl, texts, num_merges=10): + # Each implementation has its own script registry and BPETokenizer. + from importlib import import_module + pkg = module.__name__ + units = import_module(f"{pkg}.graphs.units") + chinese = import_module(f"{pkg}.languages.chinese.graph") + units.register_script("Han", chinese.chinese_character_to_graph) + t0 = time.perf_counter() - tok = BPETokenizer() + tok = module.BPETokenizer() merges = tok.train(texts, num_merges=num_merges) elapsed = time.perf_counter() - t0 - per_merge = elapsed / num_merges if num_merges else 0 - print(f" {label:30s} {elapsed:7.3f}s ({per_merge:.4f}s/merge, {len(merges)} merges)") + digest = hashlib.md5(repr(list(merges)).encode()).hexdigest()[:10] + print(f" {impl:22s} {elapsed:7.3f}s ({elapsed / num_merges:.4f}s/merge, " + f"{len(merges)} merges, digest={digest})") return elapsed if __name__ == "__main__": - register_script("Han", chinese_character_to_graph) - for n in [10, 50, 100]: texts = load_texts(n) - print(f"\n--- {len(texts)} docs ---") - bench(texts, num_merges=10, label=f"{len(texts)} docs, 10 merges") - bench(texts, num_merges=50, label=f"{len(texts)} docs, 50 merges") + for num_merges in [10, 50]: + print(f"\n--- {len(texts)} docs, {num_merges} merges ---") + for impl, module in IMPLS.items(): + bench(module, impl, texts, num_merges=num_merges)