Skip to content

Instantly share code, notes, and snippets.

@ParagEkbote
Created March 16, 2026 05:10
Show Gist options
  • Select an option

  • Save ParagEkbote/b3877f667f84cbb9a27bdaca94ba662a to your computer and use it in GitHub Desktop.

Select an option

Save ParagEkbote/b3877f667f84cbb9a27bdaca94ba662a to your computer and use it in GitHub Desktop.
#!/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