""" Tokenizer Training Pipeline - ViuMini-MoE-242M ============================================== - Vocabulary Size: 48,000 (Byte-Level BPE) - v2 Distribution (OPTIONAL rule, --mix/--no-mix): 35% Hinglish + 30% Hindi + 35% English (English-boosted; v1 Hub was 35/45/20) - Specialized Tokens: Native Chain-of-Thought () and Chat Turn Delimiters - Library: Hugging Face Tokenizers (Rust-accelerated) Usage: python tokenizer/scripts/tokenizer_train.py --hf_dataset ViuAI/viu-mini-raw-pretrain --max_lines 500000 --out tokenizer/outputs/tokenizer.json """ import argparse import os import random import sys from pathlib import Path # Fix console encoding on Windows for ByteLevel characters try: sys.stdout.reconfigure(encoding="utf-8") except Exception: pass from tokenizers import Tokenizer, Regex from tokenizers.models import BPE from tokenizers.trainers import BpeTrainer from tokenizers.pre_tokenizers import ByteLevel, Split, Sequence from tokenizers.decoders import ByteLevel as ByteLevelDecoder from tokenizers.normalizers import NFC, Sequence as NormSequence # Indic-optimized Llama-3 regex pattern with \p{M} for combining marks (matras, halants) INDIC_LLAMA3_PATTERN = ( r"(?i:'s|'t|'re|'ve|'m|'ll|'d)|" r"[^\r\n\p{L}\p{N}]?[\p{L}\p{M}]+|" r"\p{N}{1,3}|" r" ?[^\s\p{L}\p{N}\p{M}]+[\r\n]*|" r"\s*[\r\n]+|" r"\s+(?!\S)|\s+" ) VOCAB_SIZE = 48000 # v2 (2026-09-24): `` typo fixed -> ``. Count still 13. SPECIAL_TOKENS = [ "", "", "", "", "", "<|hindi|>", "<|english|>", "<|hinglish|>", "", "", "<|user|>", "<|assistant|>", "<|system|>", ] # Benchmark sample sentences for fertility and coverage assessment COVERAGE_SAMPLES = [ # Hinglish (Romanized Hindi) "mai tumse bahut pyaar karta hu", "bhai kal party me kya scene hai?", "yaar ye phone ka network bahut slow hai", "tumne khana khaya kya abhi tak?", "<|user|> mujhe train ka status batao\n<|assistant|> Checking train schedule... aapki train time par hai.", # Devanagari Hindi "नमस्ते, आप कैसे हैं?", "मुझे हिंदी में कहानी सुनाओ", "भारतीय संविधान का अनुच्छेद इक्कीस जीवन के अधिकार की रक्षा करता है।", "क्षत्रिय ज्ञानी व्यक्ति त्रिशूल लेकर आया।", "<|user|> भारत की राजधानी क्या है?\n<|assistant|> नई दिल्ली भारत की राजधानी है। नई दिल्ली।", # English "The quick brown fox jumps over the lazy dog.", "Artificial intelligence is transforming real-world problem solving.", "Can you explain quantum computing in simple terms?", "DeepSeek-R1 utilizes reinforcement learning for reasoning verification.", ] def collect_local_files(data_dir: Path): """Scan and group local text files by language category.""" hinglish = list(data_dir.rglob("*hinglish*.txt")) hindi = [p for p in data_dir.rglob("*hindi*.txt") if "hinglish" not in p.name.lower()] english = list(data_dir.rglob("*english*.txt")) if hinglish or hindi or english: return {"hinglish": hinglish, "hindi": hindi, "english": english} all_txt = sorted(data_dir.rglob("*.txt")) return {"all": all_txt} def iter_file_lines(files, limit=None): """Yield non-empty lines from a list of text files up to limit.""" count = 0 for fp in files: try: with open(fp, "r", encoding="utf-8") as f: for line in f: line = line.strip() if not line: continue yield line count += 1 if limit and count >= limit: return except FileNotFoundError: continue def build_local_mixed_iterator(grouped, total_lines=500_000, seed=42, mix=None, enforce_mix=True): """Interleave lines per mix weights. Mix rule OPTIONAL: enforce_mix=False = natural file order. mix = (hinglish, hindi, english) weights, default v2 (0.35, 0.30, 0.35).""" random.seed(seed) if "all" in grouped: files = grouped["all"] if not files: raise FileNotFoundError("No text files found in specified directory.") def gen_all(): emitted = 0 for line in iter_file_lines(files): yield line emitted += 1 if emitted >= total_lines: return return gen_all() h, hi, e = grouped.get("hinglish", []), grouped.get("hindi", []), grouped.get("english", []) if not (h or hi or e): raise FileNotFoundError("No language-specific text files found in data directory.") if not enforce_mix: # OPTIONAL rule OFF: natural file order, no ratio enforcement. def gen_natural(): emitted = 0 for pool in (h, hi, e): for line in iter_file_lines(pool): yield line emitted += 1 if emitted >= total_lines: return return gen_natural() w_h, w_hi, w_e = (list(mix) + [0.35, 0.30, 0.35])[:3] if mix else (0.35, 0.30, 0.35) def gen_interleaved(): h_lines = list(iter_file_lines(h, limit=int(total_lines * w_h) or None)) if h else [] hi_lines = list(iter_file_lines(hi, limit=int(total_lines * w_hi) or None)) if hi else [] e_lines = list(iter_file_lines(e, limit=int(total_lines * w_e) or None)) if e else [] pools = [] if h_lines: pools.append(("hinglish", h_lines)) if hi_lines: pools.append(("hindi", hi_lines)) if e_lines: pools.append(("english", e_lines)) order = ["hinglish"] * max(int(round(w_h * 10)), 1) + ["hindi"] * max(int(round(w_hi * 10)), 1) + ["english"] * max(int(round(w_e * 10)), 1) lookup = {name: lines for name, lines in pools} idx = {name: 0 for name, _ in pools} emitted = 0 while emitted < total_lines: progressed = False for key in order: if key not in lookup or not lookup[key]: continue lines = lookup[key] yield lines[idx[key] % len(lines)] idx[key] += 1 emitted += 1 progressed = True if emitted >= total_lines: return if not progressed: return return gen_interleaved() def build_hf_streaming_iterator(repo_id: str, max_lines: int = 500_000, token: str = None, mix=None, enforce_mix=True): """Stream lines from Hub. Mix rule OPTIONAL: enforce_mix=False = sequential (hindi, hinglish, english). v2 default mix (hinglish 0.35, hindi 0.30, english 0.35): English-boosted vs v1 (0.35/0.45/0.20) taaki English fertility ~1.15 aaye; Hindi 1.00 floor par hai isliye uska cut safe hai.""" from datasets import load_dataset from huggingface_hub import HfApi print(f"[info] Connecting to Hugging Face dataset: {repo_id}") api = HfApi(token=token) repo_files = api.list_repo_files(repo_id, repo_type="dataset") hinglish_files = [f for f in repo_files if f.startswith("hinglish/") and (f.endswith(".parquet") or f.endswith(".jsonl"))] hindi_files = [f for f in repo_files if (f.startswith("hindi/") or f.startswith("hindi_fixed/")) and f.endswith(".parquet")] english_files = [f for f in repo_files if (f.startswith("distilled/") or f.startswith("english_fixed/")) and f.endswith(".parquet")] print(f"[info] Discovered files - Hinglish: {len(hinglish_files)}, Hindi: {len(hindi_files)}, English/Reasoning: {len(english_files)}") w_h, w_hi, w_e = (list(mix) + [0.35, 0.30, 0.35])[:3] if mix else (0.35, 0.30, 0.35) tot = (w_h + w_hi + w_e) or 1.0 h_target = int(max_lines * w_h / tot) hi_target = int(max_lines * w_hi / tot) e_target = max_lines - h_target - hi_target def fetch_stream(file_list, target_count, category_name): emitted = 0 for fpath in file_list: if emitted >= target_count: break try: ds = load_dataset(repo_id, data_files=fpath, split="train", streaming=True, token=token) for row in ds: text = row.get("text", "") or "" text = text.strip() if len(text) >= 15: yield text emitted += 1 if emitted >= target_count: break except Exception as exc: print(f"[warning] Skipping {fpath} due to error: {exc}") continue print(f"[info] Collected {emitted:,} lines for {category_name}") def stream_gen(): if not enforce_mix: # OPTIONAL rule OFF: sequential streams capped at max_lines total. emitted = 0 for fl, cat in ((hindi_files, "Hindi"), (hinglish_files, "Hinglish"), (english_files, "English/Reasoning")): for line in fetch_stream(fl, max_lines - emitted, cat): yield line emitted += 1 if emitted >= max_lines: return return hi_gen = fetch_stream(hindi_files, hi_target, "Hindi") h_gen = fetch_stream(hinglish_files, h_target, "Hinglish") e_gen = fetch_stream(english_files, e_target, "English/Reasoning") generators = {"hindi": hi_gen, "hinglish": h_gen, "english": e_gen} pattern = ["hindi"] * 5 + ["hinglish"] * 3 + ["english"] * 2 active = {"hindi": True, "hinglish": True, "english": True} total_emitted = 0 while total_emitted < max_lines and any(active.values()): progress = False for cat in pattern: if not active[cat]: continue try: line = next(generators[cat]) yield line total_emitted += 1 progress = True if total_emitted >= max_lines: return except StopIteration: active[cat] = False if not progress: break return stream_gen() def evaluate_tokenizer(tok: Tokenizer): """Run comprehensive fertility, coverage, and special-token tests.""" print("\n" + "=" * 65) print("TOKENIZER EVALUATION & QUALITY AUDIT") print("=" * 65) # 1. Special token isolation check print("[test 1] Special Tokens Encoding Isolation:") all_special_isolated = True for st in SPECIAL_TOKENS: enc = tok.encode(st) is_single = len(enc.tokens) == 1 and enc.tokens[0] == st if not is_single: all_special_isolated = False print(f" [FAIL] '{st}' split into {enc.tokens} (IDs: {enc.ids})") else: print(f" [PASS] '{st}' -> ID: {enc.ids[0]}") if all_special_isolated: print(" -> All 13 special tokens correctly isolated as atomic single IDs.") # 2. Benchmark fertility and out-of-vocabulary check print("\n[test 2] Multi-Lingual Fertility & Out-Of-Vocabulary Audit:") category_metrics = {"Hinglish": [], "Hindi": [], "English": []} for s in COVERAGE_SAMPLES: enc = tok.encode(s) words = s.split() fertility = len(enc.tokens) / max(len(words), 1) # Categorize if any(ord(c) >= 0x0900 and ord(c) <= 0x097F for c in s): category_metrics["Hindi"].append(fertility) elif "mai" in s or "bhai" in s or "yaar" in s or "train" in s: category_metrics["Hinglish"].append(fertility) else: category_metrics["English"].append(fertility) safe_in = s[:50] + ("..." if len(s) > 50 else "") print(f" Text: {safe_in:55} | Tokens: {len(enc.tokens):2d} | Words: {len(words):2d} | Ratio: {fertility:.2f}") print("\nFertility Summary (Tokens per Word):") for cat, vals in category_metrics.items(): if vals: avg_f = sum(vals) / len(vals) target = 1.5 if cat == "Hinglish" else (1.8 if cat == "Hindi" else 1.3) status = "PASS" if avg_f <= (target + 0.3) else "WARN" print(f" {cat:10}: {avg_f:.2f} tokens/word (Target <= {target:.1f}) [{status}]") # 3. Round-trip fidelity check print("\n[test 3] Round-Trip Encoding Fidelity Check:") fidelity_pass = True for s in COVERAGE_SAMPLES: enc = tok.encode(s) decoded = tok.decode(enc.ids, skip_special_tokens=False) if decoded.strip() != s.strip(): fidelity_pass = False print(f" [FAIL] Mismatch:\n Orig: {s}\n Dec : {decoded}") if fidelity_pass: print(" [PASS] 100% Round-trip fidelity verified across all evaluation samples.") print("=" * 65 + "\n") def _parse_mix(s): """'40,30,30' -> (0.4, 0.3, 0.3) order (hinglish, hindi, english). None = defaults.""" if not s: return None try: parts = [float(x.strip()) for x in str(s).split(",")] if len(parts) != 3 or sum(parts) <= 0: raise ValueError tot = sum(parts) return (parts[0] / tot, parts[1] / tot, parts[2] / tot) except ValueError: raise ValueError("--mix format: 'hinglish,hindi,english' e.g. '40,30,30'") def train_tokenizer(args): """Main tokenizer training routine.""" out_path = Path(args.out) tok = Tokenizer(BPE(unk_token="")) tok.normalizer = Sequence([NFC()]) tok.pre_tokenizer = Sequence([ Split(pattern=Regex(INDIC_LLAMA3_PATTERN), behavior="isolated"), ByteLevel(add_prefix_space=False, use_regex=False), ]) tok.decoder = ByteLevelDecoder() trainer = BpeTrainer( vocab_size=args.vocab_size, min_frequency=2, show_progress=True, special_tokens=SPECIAL_TOKENS, initial_alphabet=ByteLevel.alphabet(), ) if args.hf_dataset: print(f"[info] Preparing training iterator from Hugging Face: {args.hf_dataset}") token = os.environ.get("HF_TOKEN") mix = _parse_mix(args.mix) iterator = build_hf_streaming_iterator(args.hf_dataset, max_lines=args.max_lines, token=token, mix=mix, enforce_mix=not args.no_mix) else: data_dir = Path(args.data_dir) if not data_dir.exists(): raise FileNotFoundError(f"Local data directory not found: {data_dir}") grouped = collect_local_files(data_dir) print(f"[info] Preparing local training iterator from {data_dir}") mix = _parse_mix(args.mix) iterator = build_local_mixed_iterator(grouped, total_lines=args.max_lines, mix=mix, enforce_mix=not args.no_mix) print(f"[info] Training Byte-Level BPE model (Target Vocab: {args.vocab_size:,}) ...") tok.train_from_iterator(iterator, trainer=trainer, length=args.max_lines) out_path.parent.mkdir(parents=True, exist_ok=True) tok.save(str(out_path)) print(f"[success] Production tokenizer successfully saved to: {out_path}") # Run comprehensive quality audit evaluate_tokenizer(tok) if args.push_to_hub: repo_target = args.push_to_hub token = os.environ.get("HF_TOKEN") print(f"[info] Uploading trained tokenizer to Hugging Face Hub: {repo_target}") from huggingface_hub import HfApi api = HfApi(token=token) repo_type = "model" if "/" in repo_target and not repo_target.endswith("pretrain") else "dataset" api.upload_file( path_or_fileobj=str(out_path), path_in_repo="tokenizer/tokenizer.json", repo_id=repo_target, repo_type=repo_type, ) print(f"[success] Uploaded tokenizer/tokenizer.json to {repo_target} ({repo_type})") def main(): parser = argparse.ArgumentParser(description="Train Byte-Level BPE Tokenizer for ViuMini-MoE-242M") parser.add_argument("--data_dir", default="data/raw", help="Path to directory containing raw text files") parser.add_argument("--hf_dataset", default="ViuAI/viu-mini-raw-pretrain", help="Hugging Face dataset repository for streaming") parser.add_argument("--out", default="tokenizer/outputs/tokenizer.json", help="Path where trained tokenizer will be saved") parser.add_argument("--vocab_size", type=int, default=VOCAB_SIZE, help="Target vocabulary size") parser.add_argument("--max_lines", type=int, default=500_000, help="Maximum number of lines to sample for training") parser.add_argument("--mix", default=None, help="Mix weights 'hinglish,hindi,english' e.g. '35,30,35' (v2 default both paths: 35,30,35). Rule OPTIONAL.") parser.add_argument("--no-mix", action="store_true", help="Mix rule OFF: natural file/stream order, no ratio enforcement") parser.add_argument("--push_to_hub", default=None, help="Optional Hugging Face repo ID to upload trained tokenizer") args = parser.parse_args() train_tokenizer(args) if __name__ == "__main__": main()