Created
March 16, 2026 05:10
-
-
Save ParagEkbote/b3877f667f84cbb9a27bdaca94ba662a to your computer and use it in GitHub Desktop.
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
| #!/usr/bin/env python3 | |
| """ | |
| Tokenizer Performance Comparison | |
| """ | |
| # ------------------------------------------------ | |
| # Dependencies | |
| # Install with: | |
| # | |
| # uv pip install \ | |
| # "datasets>=2.14.0" \ | |
| # "transformers>=4.30.0" \ | |
| # "pandas>=1.5.0" \ | |
| # "matplotlib>=3.7.0" \ | |
| # "seaborn>=0.12.0" \ | |
| # "datatrove>=0.1.0" \ | |
| # "torch" | |
| # | |
| # ------------------------------------------------ | |
| from dataclasses import dataclass | |
| from pathlib import Path | |
| from datetime import datetime | |
| import argparse | |
| import logging | |
| import pandas as pd | |
| import matplotlib.pyplot as plt | |
| import seaborn as sns | |
| from datasets import load_dataset | |
| from transformers import AutoTokenizer | |
| from datatrove.utils.word_tokenizers import load_word_tokenizer | |
| from huggingface_hub import batch_bucket_files | |
| logging.basicConfig(level=logging.INFO, format="%(message)s") | |
| logger = logging.getLogger(__name__) | |
| # ------------------------------------------------ | |
| # Run Directory | |
| # ------------------------------------------------ | |
| def create_run_dir(base="runs"): | |
| run_id = datetime.now().strftime("%Y%m%d_%H%M%S") | |
| run_dir = Path(base) / run_id | |
| run_dir.mkdir(parents=True, exist_ok=True) | |
| return run_id, run_dir | |
| # ------------------------------------------------ | |
| # Languages | |
| # ------------------------------------------------ | |
| LANGUAGE_NAMES = { | |
| "en": "English", | |
| "es": "Spanish", | |
| "fr": "French", | |
| "zh": "Chinese", | |
| } | |
| @dataclass | |
| class Language: | |
| name: str | |
| code: str | |
| # ------------------------------------------------ | |
| # Data Loading | |
| # ------------------------------------------------ | |
| class DataLoader: | |
| def __init__(self, n_chars: int = 50000): | |
| self.n_chars = n_chars | |
| def load_from_hf_dataset(self, dataset_name, lang_code, field="text"): | |
| ds = load_dataset(dataset_name, name=lang_code, split="train", streaming=True) | |
| return self._collect_text_samples(ds, field) | |
| def load_from_file(self, filepath: Path): | |
| if not filepath.exists(): | |
| return None | |
| with open(filepath, "r", encoding="utf-8") as f: | |
| return f.read()[: self.n_chars] | |
| def _collect_text_samples(self, dataset, field): | |
| text = [] | |
| total = 0 | |
| for sample in dataset: | |
| s = sample.get(field, "") | |
| text.append(s) | |
| total += len(s) | |
| if total >= self.n_chars: | |
| break | |
| return "\n".join(text) | |
| def load_all_texts(languages, data_dir=None, hf_dataset=None, content_field="text", n_chars=50000): | |
| loader = DataLoader(n_chars) | |
| texts = {} | |
| for lang in languages: | |
| logger.info(f"Loading {lang.name}") | |
| text = None | |
| if hf_dataset: | |
| try: | |
| text = loader.load_from_hf_dataset( | |
| hf_dataset, | |
| lang.code, | |
| content_field, | |
| ) | |
| except Exception as e: | |
| logger.warning(f"{lang.code} dataset load failed: {e}") | |
| if not text and data_dir: | |
| path = Path(data_dir) / f"{lang.code}.txt" | |
| text = loader.load_from_file(path) | |
| texts[lang.code] = text or "" | |
| return texts | |
| # ------------------------------------------------ | |
| # Tokenizer Analysis | |
| # ------------------------------------------------ | |
| class TokenizerAnalyzer: | |
| def compute_metrics(self, tokenizer, text, word_tokenizer): | |
| token_ids = tokenizer.encode(text, add_special_tokens=False) | |
| _ = tokenizer.convert_ids_to_tokens(token_ids) | |
| words = ( | |
| word_tokenizer.tokenize(text) | |
| if hasattr(word_tokenizer, "tokenize") | |
| else word_tokenizer(text) | |
| ) | |
| num_words = len(words) | |
| char_count = len(text.replace(" ", "")) | |
| word_fertility = len(token_ids) / num_words if num_words else 0 | |
| char_fertility = len(token_ids) / char_count if char_count else 0 | |
| return { | |
| "word_fertility": round(word_fertility, 3), | |
| "char_fertility": round(char_fertility, 3), | |
| "vocab_size": tokenizer.vocab_size, | |
| } | |
| class TokenDistributionAnalyzer: | |
| def compute_distributions(self, tokenizer, text): | |
| token_ids = tokenizer.encode(text, add_special_tokens=False) | |
| tokens = tokenizer.convert_ids_to_tokens(token_ids) | |
| lengths = [] | |
| for t in tokens: | |
| t = t.replace("▁", "").replace("Ġ", "").replace("##", "") | |
| lengths.append(len(t)) | |
| return lengths | |
| # ------------------------------------------------ | |
| # Evaluation | |
| # ------------------------------------------------ | |
| def evaluate_all_tokenizers(tokenizers, languages, texts): | |
| analyzer = TokenizerAnalyzer() | |
| dist_analyzer = TokenDistributionAnalyzer() | |
| rows = [] | |
| distribution_rows = [] | |
| for name, path in tokenizers: | |
| logger.info(f"\nEvaluating {name}") | |
| tokenizer = AutoTokenizer.from_pretrained(path, trust_remote_code=True) | |
| for lang in languages: | |
| text = texts.get(lang.code) | |
| if not text: | |
| continue | |
| try: | |
| word_tok = load_word_tokenizer(lang.code) | |
| except Exception: | |
| word_tok = str.split | |
| metrics = analyzer.compute_metrics(tokenizer, text, word_tok) | |
| rows.append( | |
| { | |
| "tokenizer": name, | |
| "language": lang.name, | |
| **metrics, | |
| } | |
| ) | |
| token_lengths = dist_analyzer.compute_distributions(tokenizer, text) | |
| for length in token_lengths: | |
| distribution_rows.append( | |
| { | |
| "tokenizer": name, | |
| "language": lang.name, | |
| "token_length": length, | |
| } | |
| ) | |
| return pd.DataFrame(rows), pd.DataFrame(distribution_rows) | |
| # ------------------------------------------------ | |
| # Visualization | |
| # ------------------------------------------------ | |
| class Visualizer: | |
| def __init__(self): | |
| sns.set_theme(style="whitegrid") | |
| def create_heatmaps(self, df, path): | |
| pivot = df.pivot( | |
| index="tokenizer", | |
| columns="language", | |
| values="word_fertility", | |
| ) | |
| plt.figure(figsize=(10,6)) | |
| sns.heatmap(pivot, annot=True, cmap="RdYlGn_r") | |
| plt.savefig(path, dpi=300, bbox_inches="tight") | |
| plt.close() | |
| def create_vocab_vs_fertility(self, df, path): | |
| summary = ( | |
| df.groupby("tokenizer") | |
| .agg( | |
| vocab_size=("vocab_size", "first"), | |
| avg_word_fertility=("word_fertility", "mean"), | |
| ) | |
| .reset_index() | |
| ) | |
| plt.figure(figsize=(8,6)) | |
| sns.scatterplot( | |
| data=summary, | |
| x="vocab_size", | |
| y="avg_word_fertility", | |
| s=200, | |
| ) | |
| for _, row in summary.iterrows(): | |
| plt.text(row["vocab_size"], row["avg_word_fertility"], row["tokenizer"]) | |
| plt.savefig(path, dpi=300, bbox_inches="tight") | |
| plt.close() | |
| def create_token_length_histograms(self, dist_df, run_dir): | |
| languages = dist_df["language"].unique() | |
| for lang in languages: | |
| subset = dist_df[dist_df["language"] == lang] | |
| plt.figure(figsize=(8,6)) | |
| sns.histplot( | |
| data=subset, | |
| x="token_length", | |
| hue="tokenizer", | |
| bins=30, | |
| element="step", | |
| stat="density", | |
| common_norm=False, | |
| ) | |
| plt.title(f"Token Length Distribution — {lang}") | |
| out = run_dir / f"token_length_distribution_{lang}.png" | |
| plt.savefig(out, dpi=300, bbox_inches="tight") | |
| plt.close() | |
| # ------------------------------------------------ | |
| # Bucket Upload | |
| # ------------------------------------------------ | |
| def upload_run_directory(bucket_id, run_id, run_dir): | |
| logger.info(f"\nUploading run → {run_id}") | |
| add_ops = [] | |
| for path in run_dir.glob("*"): | |
| if path.is_file(): | |
| add_ops.append((str(path), f"runs/{run_id}/{path.name}")) | |
| if add_ops: | |
| batch_bucket_files( | |
| bucket_id, | |
| add=add_ops, | |
| ) | |
| logger.info("Upload complete") | |
| # ------------------------------------------------ | |
| # CLI | |
| # ------------------------------------------------ | |
| def parse_arguments(): | |
| parser = argparse.ArgumentParser() | |
| parser.add_argument("--hf-dataset", type=str, default=None) | |
| parser.add_argument("--data-dir", type=str, default=None) | |
| parser.add_argument("--langs", type=str, default="en,es,fr,zh") | |
| parser.add_argument("--n-chars", type=int, default=50000) | |
| parser.add_argument("--tokenizers", type=str, required=True) | |
| parser.add_argument("--bucket", type=str, default=None) | |
| return parser.parse_args() | |
| def parse_tokenizers(spec): | |
| tokenizers = [] | |
| for item in spec.split(","): | |
| name, path = item.split(":", 1) | |
| tokenizers.append((name.strip(), path.strip())) | |
| return tokenizers | |
| def create_languages(codes): | |
| langs = [] | |
| for code in codes: | |
| name = LANGUAGE_NAMES.get(code, code.upper()) | |
| langs.append(Language(name, code)) | |
| return langs | |
| # ------------------------------------------------ | |
| # Main | |
| # ------------------------------------------------ | |
| def main(): | |
| args = parse_arguments() | |
| languages = create_languages([x.strip() for x in args.langs.split(",")]) | |
| tokenizers = parse_tokenizers(args.tokenizers) | |
| texts = load_all_texts( | |
| languages, | |
| data_dir=args.data_dir, | |
| hf_dataset=args.hf_dataset, | |
| n_chars=args.n_chars, | |
| ) | |
| df, dist_df = evaluate_all_tokenizers( | |
| tokenizers, | |
| languages, | |
| texts, | |
| ) | |
| run_id, run_dir = create_run_dir() | |
| df.to_csv(run_dir / "tokenizer_comparison.csv", index=False) | |
| dist_df.to_csv(run_dir / "token_length_distribution.csv", index=False) | |
| viz = Visualizer() | |
| viz.create_heatmaps(df, run_dir / "tokenizer_heatmaps.png") | |
| viz.create_vocab_vs_fertility(df, run_dir / "tokenizer_vocab_vs_fertility.png") | |
| viz.create_token_length_histograms(dist_df, run_dir) | |
| logger.info(f"Artifacts saved in {run_dir}") | |
| if args.bucket: | |
| upload_run_directory(args.bucket, run_id, run_dir) | |
| if __name__ == "__main__": | |
| main() |
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment