ViuMini-MoE-242M / tokenizer /scripts /tokenizer_train.py
ViuAI's picture
sync local: enforce_mix sampler, tokenizer v2 (mask fix + 35/30/35), mix optional docs, repair scripts
991af26 verified
Raw History Blame Contribute Delete
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()