File size: 17,593 Bytes
140f9f6
991af26
 
 
 
 
 
 
140f9f6
991af26
140f9f6
 
991af26
140f9f6
 
 
 
991af26
140f9f6
 
 
 
991af26
 
140f9f6
 
991af26
140f9f6
991af26
 
 
 
 
 
 
 
 
 
 
140f9f6
 
 
991af26
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
140f9f6
991af26
140f9f6
991af26
140f9f6
 
 
 
991af26
 
140f9f6
 
991af26
 
 
 
 
 
 
 
140f9f6
 
 
991af26
 
140f9f6
 
 
 
 
 
 
 
 
 
 
991af26
 
140f9f6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
991af26
 
 
140f9f6
 
 
 
991af26
 
 
140f9f6
991af26
140f9f6
 
 
 
991af26
 
140f9f6
 
 
991af26
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
140f9f6
 
991af26
140f9f6
991af26
140f9f6
991af26
 
 
 
 
 
140f9f6
 
 
 
 
 
 
991af26
140f9f6
 
 
 
 
 
 
 
991af26
140f9f6
 
991af26
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
140f9f6
 
 
991af26
 
 
 
140f9f6
 
 
 
 
 
 
 
 
 
991af26
 
 
 
 
 
 
 
 
 
 
 
 
 
 
140f9f6
991af26
140f9f6
 
 
 
991af26
140f9f6
991af26
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
140f9f6
 
 
991af26
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
"""
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()