Download tokenizer/scripts/tokenizer_train.py from ViuAI/ViuMini-MoE-242M: direct link, hf CLI and curl.
- Browser
- Download file 17.6 kB
-
https://huggingface.co/ViuAI/ViuMini-MoE-242M/resolve/main/tokenizer/scripts/tokenizer_train.py
- Command line
-
hf download hf://ViuAI/ViuMini-MoE-242M/tokenizer/scripts/tokenizer_train.py
-
curl -L -o tokenizer_train.py https://huggingface.co/ViuAI/ViuMini-MoE-242M/resolve/main/tokenizer/scripts/tokenizer_train.py
17.6 kB
| """ | |
| 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 (<soch>) 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): `<mask|>` typo fixed -> `<mask>`. Count still 13. | |
| SPECIAL_TOKENS = [ | |
| "<pad>", | |
| "<bos>", | |
| "<eos>", | |
| "<unk>", | |
| "<mask>", | |
| "<|hindi|>", | |
| "<|english|>", | |
| "<|hinglish|>", | |
| "<soch>", | |
| "</soch>", | |
| "<|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|> <soch> Checking train schedule... </soch> aapki train time par hai.", | |
| # Devanagari Hindi | |
| "नमस्ते, आप कैसे हैं?", | |
| "मुझे हिंदी में कहानी सुनाओ", | |
| "भारतीय संविधान का अनुच्छेद इक्कीस जीवन के अधिकार की रक्षा करता है।", | |
| "क्षत्रिय ज्ञानी व्यक्ति त्रिशूल लेकर आया।", | |
| "<|user|> भारत की राजधानी क्या है?\n<|assistant|> <soch> नई दिल्ली भारत की राजधानी है। </soch> नई दिल्ली।", | |
| # 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="<unk>")) | |
| 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() | |