thefinalboss commited on
Commit
64d122d
·
verified ·
1 Parent(s): 6b3dc3b

Upload cognet_data_prep.py with huggingface_hub

Browse files
Files changed (1) hide show
  1. cognet_data_prep.py +1252 -0
cognet_data_prep.py ADDED
@@ -0,0 +1,1252 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ CogNet 1B Data Preparation Script
4
+ ===================================
5
+ Downloads, preprocesses, and tokenizes datasets from HuggingFace
6
+ for training CogNet-1B on syntax, math, and code.
7
+
8
+ Target: Character-level tokenizer (136 vocab: ASCII + French accents)
9
+ Output: Pre-tokenized .pt files ready for training
10
+
11
+ Usage:
12
+ python cognet_data_prep.py --output_dir /root/CogNet/data_1b [--max_gb 50] [--dry_run]
13
+
14
+ Datasets covered:
15
+ CODE: the-stack-smol, codeparrot-clean, CodeAlpaca, CodeSearchNet, python_code_instructions
16
+ MATH: MathPile, OpenMathInstruct-1, MetaMathQA, GSM8K, HendrycksMath
17
+ SYNTAX: WikiText-103, C4 subset, Penn Treebank, Universal Dependencies
18
+ GENERAL: The Pile subset (for code+math+prose mix)
19
+ """
20
+
21
+ import argparse
22
+ import json
23
+ import os
24
+ import sys
25
+ import time
26
+ import unicodedata
27
+ from pathlib import Path
28
+ from typing import Dict, List, Optional, Tuple
29
+
30
+ # ─── Character-level Tokenizer ───────────────────────────────────────────────
31
+
32
+ class CharTokenizer:
33
+ """Character-level tokenizer: printable ASCII + French accents + newline/tab."""
34
+
35
+ def __init__(self):
36
+ self.chars = sorted(set(
37
+ [chr(i) for i in range(32, 127)]
38
+ + list('àâäéèêëïîôùûüÿçœæÀÂÄÉÈÊËÏÎÔÙÛÜŸÇŒÆ')
39
+ + list('ëßñ¿«»')
40
+ + ['\t', '\n']
41
+ ))
42
+ self.char_to_id = {c: i for i, c in enumerate(self.chars)}
43
+ self.id_to_char = {i: c for i, c in enumerate(self.chars)}
44
+ self.vocab_size = len(self.chars)
45
+ self._allowed = set(self.chars)
46
+
47
+ def is_compatible(self, text: str) -> float:
48
+ """Return fraction of characters that are in vocab (1.0 = all compatible)."""
49
+ if not text:
50
+ return 1.0
51
+ compatible = sum(1 for c in text if c in self._allowed)
52
+ return compatible / len(text)
53
+
54
+ def clean_text(self, text: str) -> str:
55
+ """Clean text to be compatible with our 136-char vocab.
56
+
57
+ Strategy:
58
+ - Map common Unicode to ASCII equivalents
59
+ - Replace smart quotes, em-dashes, etc.
60
+ - Strip characters that can't be mapped
61
+ - Preserve newlines and tabs
62
+ """
63
+ if not text:
64
+ return text
65
+
66
+ # Phase 1: Common Unicode replacements
67
+ replacements = {
68
+ '\u2018': "'", # left single quote
69
+ '\u2019': "'", # right single quote
70
+ '\u201c': '"', # left double quote
71
+ '\u201d': '"', # right double quote
72
+ '\u2013': '-', # en dash
73
+ '\u2014': '--', # em dash
74
+ '\u2026': '...', # ellipsis
75
+ '\u00a0': ' ', # non-breaking space
76
+ '\u2028': '\n', # line separator
77
+ '\u2029': '\n', # paragraph separator
78
+ '\u200b': '', # zero-width space
79
+ '\u200c': '', # zero-width non-joiner
80
+ '\u200d': '', # zero-width joiner
81
+ '\ufeff': '', # BOM
82
+ '\u2192': '->', # right arrow
83
+ '\u2190': '<-', # left arrow
84
+ '\u2194': '<->', # left-right arrow
85
+ '\u2264': '<=', # less than or equal
86
+ '\u2265': '>=', # greater than or equal
87
+ '\u2260': '!=', # not equal
88
+ '\u00d7': '*', # multiplication sign
89
+ '\u00f7': '/', # division sign
90
+ '\u00b1': '+-', # plus-minus
91
+ '\u2212': '-', # minus sign
92
+ '\u2248': '~=', # approximately equal
93
+ '\u221e': 'inf', # infinity
94
+ '\u03c0': 'pi', # pi
95
+ '\u03b1': 'alpha', # alpha
96
+ '\u03b2': 'beta', # beta
97
+ '\u03b3': 'gamma', # gamma
98
+ '\u03b4': 'delta', # delta
99
+ '\u03b5': 'epsilon', # epsilon
100
+ '\u03b8': 'theta', # theta
101
+ '\u03bb': 'lambda', # lambda
102
+ '\u03c3': 'sigma', # sigma
103
+ '\u03c9': 'omega', # omega
104
+ '\u2211': 'sum', # summation
105
+ '\u220f': 'prod', # product
106
+ '\u222b': 'int', # integral
107
+ '\u221a': 'sqrt', # square root
108
+ '\u2202': 'partial', # partial derivative
109
+ '\u2208': 'in', # element of
110
+ '\u2282': 'subset', # subset
111
+ '\u2229': 'intersect', # intersection
112
+ '\u222a': 'union', # union
113
+ '\u00b2': '^2', # superscript 2
114
+ '\u00b3': '^3', # superscript 3
115
+ '\u2082': '_2', # subscript 2
116
+ }
117
+
118
+ for old, new in replacements.items():
119
+ text = text.replace(old, new)
120
+
121
+ # Phase 2: NFKC normalization for remaining Unicode
122
+ # This handles accented chars decomposition, etc.
123
+ normalized = []
124
+ for ch in text:
125
+ if ch in self._allowed:
126
+ normalized.append(ch)
127
+ else:
128
+ # Try NFKC normalization
129
+ nfkc = unicodedata.normalize('NFKC', ch)
130
+ if all(c in self._allowed for c in nfkc):
131
+ normalized.append(nfkc)
132
+ else:
133
+ # Try stripping diacritics
134
+ stripped = unicodedata.normalize('NFD', ch)
135
+ stripped = ''.join(
136
+ c for c in stripped
137
+ if unicodedata.category(c) != 'Mn'
138
+ )
139
+ if all(c in self._allowed for c in stripped):
140
+ normalized.append(stripped)
141
+ # else: skip this character entirely
142
+
143
+ return ''.join(normalized)
144
+
145
+ def encode(self, text: str) -> List[int]:
146
+ return [self.char_to_id.get(c, self.char_to_id.get(' ', 0)) for c in text]
147
+
148
+ def decode(self, ids: List[int]) -> str:
149
+ return ''.join(self.id_to_char.get(i, ' ') for i in ids)
150
+
151
+ def save(self, path: str):
152
+ with open(path, 'w', encoding='utf-8') as f:
153
+ json.dump({
154
+ 'chars': self.chars,
155
+ 'vocab_size': self.vocab_size,
156
+ }, f, ensure_ascii=False, indent=2)
157
+
158
+ @classmethod
159
+ def load(cls, path: str) -> 'CharTokenizer':
160
+ tok = cls.__new__(cls)
161
+ with open(path, 'r', encoding='utf-8') as f:
162
+ data = json.load(f)
163
+ tok.chars = data['chars']
164
+ tok.char_to_id = {c: i for i, c in enumerate(tok.chars)}
165
+ tok.id_to_char = {i: c for i, c in enumerate(tok.chars)}
166
+ tok.vocab_size = data['vocab_size']
167
+ tok._allowed = set(tok.chars)
168
+ return tok
169
+
170
+
171
+ # ─── Dataset Processors ──────────────────────────────────────────────────────
172
+
173
+ class DatasetProcessor:
174
+ """Base class for dataset processors."""
175
+
176
+ def __init__(self, name: str, category: str, tokenizer: CharTokenizer,
177
+ output_dir: str, max_gb: float = 50):
178
+ self.name = name
179
+ self.category = category
180
+ self.tokenizer = tokenizer
181
+ self.output_dir = output_dir
182
+ self.max_gb = max_gb
183
+ self.stats = {
184
+ 'name': name,
185
+ 'category': category,
186
+ 'raw_chars': 0,
187
+ 'clean_chars': 0,
188
+ 'tokens': 0,
189
+ 'files': 0,
190
+ 'skipped_chars': 0,
191
+ }
192
+
193
+ def _should_stop(self, total_bytes: int) -> bool:
194
+ """Check if we've exceeded our storage budget."""
195
+ gb = total_bytes / (1024**3)
196
+ return gb >= self.max_gb
197
+
198
+ def _save_tokens(self, token_ids: List[int], split_name: str = 'train'):
199
+ """Save token IDs to a .pt file."""
200
+ import torch
201
+ out_path = os.path.join(self.output_dir, f'{self.name}_{split_name}.pt')
202
+ os.makedirs(self.output_dir, exist_ok=True)
203
+ torch.save(torch.tensor(token_ids, dtype=torch.long), out_path)
204
+ self.stats['files'] += 1
205
+ print(f" Saved {len(token_ids):,} tokens to {out_path} "
206
+ f"({len(token_ids) * 8 / 1024**2:.1f} MB)")
207
+ return out_path
208
+
209
+ def process(self, dry_run: bool = False) -> Dict:
210
+ """Process the dataset. Override in subclasses."""
211
+ raise NotImplementedError
212
+
213
+
214
+ class TheStackSmolProcessor(DatasetProcessor):
215
+ """bigcode/the-stack-smol — Multi-language code, manageable size."""
216
+
217
+ def __init__(self, tokenizer, output_dir, max_gb):
218
+ super().__init__('the_stack_smol', 'code', tokenizer, output_dir, max_gb)
219
+ self.languages = ['python', 'javascript', 'c', 'cpp', 'java', 'rust', 'go', 'typescript']
220
+
221
+ def process(self, dry_run=False):
222
+ from datasets import load_dataset
223
+
224
+ print(f"\n{'='*60}")
225
+ print(f"Processing: {self.name} (CODE)")
226
+ print(f"{'='*60}")
227
+
228
+ all_tokens = []
229
+ total_bytes = 0
230
+
231
+ for lang in self.languages:
232
+ print(f" Loading language: {lang}")
233
+ try:
234
+ ds = load_dataset(
235
+ "bigcode/the-stack-smol",
236
+ data_dir=f"data/{lang}",
237
+ split="train",
238
+ streaming=True,
239
+ trust_remote_code=True,
240
+ )
241
+
242
+ for i, item in enumerate(ds):
243
+ content = item.get('content', '')
244
+ if not content or len(content) < 10:
245
+ continue
246
+
247
+ # Clean for char-level
248
+ clean = self.tokenizer.clean_text(content)
249
+ self.stats['raw_chars'] += len(content)
250
+ self.stats['skipped_chars'] += len(content) - len(clean)
251
+
252
+ # Add separator between files
253
+ text = clean + '\n\n'
254
+ self.stats['clean_chars'] += len(text)
255
+
256
+ # Encode
257
+ tokens = self.tokenizer.encode(text)
258
+ all_tokens.extend(tokens)
259
+ self.stats['tokens'] += len(tokens)
260
+ total_bytes += len(text)
261
+
262
+ if i % 5000 == 0 and i > 0:
263
+ print(f" {lang}: {i:,} files, {self.stats['tokens']:,} tokens total")
264
+
265
+ if self._should_stop(total_bytes):
266
+ print(f" Storage limit reached at {total_bytes/1024**3:.1f} GB")
267
+ break
268
+
269
+ if i >= 8000: # Cap per language for smol
270
+ break
271
+
272
+ except Exception as e:
273
+ print(f" Error loading {lang}: {e}")
274
+ continue
275
+
276
+ if self._should_stop(total_bytes):
277
+ break
278
+
279
+ if not dry_run and all_tokens:
280
+ self._save_tokens(all_tokens)
281
+
282
+ return self.stats
283
+
284
+
285
+ class CodeParrotProcessor(DatasetProcessor):
286
+ """codeparrot/codeparrot-clean — Clean Python code."""
287
+
288
+ def __init__(self, tokenizer, output_dir, max_gb):
289
+ super().__init__('codeparrot_clean', 'code', tokenizer, output_dir, max_gb)
290
+
291
+ def process(self, dry_run=False):
292
+ from datasets import load_dataset
293
+
294
+ print(f"\n{'='*60}")
295
+ print(f"Processing: {self.name} (CODE - Python)")
296
+ print(f"{'='*60}")
297
+
298
+ all_tokens = []
299
+ total_bytes = 0
300
+
301
+ ds = load_dataset("codeparrot/codeparrot-clean", split="train", streaming=True)
302
+
303
+ for i, item in enumerate(ds):
304
+ content = item.get('content', '')
305
+ if not content or len(content) < 20:
306
+ continue
307
+
308
+ clean = self.tokenizer.clean_text(content)
309
+ self.stats['raw_chars'] += len(content)
310
+ self.stats['skipped_chars'] += len(content) - len(clean)
311
+
312
+ text = clean + '\n\n'
313
+ self.stats['clean_chars'] += len(text)
314
+
315
+ tokens = self.tokenizer.encode(text)
316
+ all_tokens.extend(tokens)
317
+ self.stats['tokens'] += len(tokens)
318
+ total_bytes += len(text)
319
+
320
+ if i % 10000 == 0 and i > 0:
321
+ print(f" {i:,} files, {self.stats['tokens']:,} tokens, "
322
+ f"{total_bytes/1024**3:.1f} GB")
323
+
324
+ if self._should_stop(total_bytes) or i >= 200000:
325
+ break
326
+
327
+ if not dry_run and all_tokens:
328
+ self._save_tokens(all_tokens)
329
+
330
+ return self.stats
331
+
332
+
333
+ class CodeAlpacaProcessor(DatasetProcessor):
334
+ """sahil2801/CodeAlpaca-20k — Instruction-code pairs."""
335
+
336
+ def __init__(self, tokenizer, output_dir, max_gb):
337
+ super().__init__('code_alpaca', 'code', tokenizer, output_dir, max_gb)
338
+
339
+ def process(self, dry_run=False):
340
+ from datasets import load_dataset
341
+
342
+ print(f"\n{'='*60}")
343
+ print(f"Processing: {self.name} (CODE - Instruction)")
344
+ print(f"{'='*60}")
345
+
346
+ ds = load_dataset("sahil2801/CodeAlpaca-20k", split="train")
347
+
348
+ all_tokens = []
349
+ total_bytes = 0
350
+
351
+ for i, item in enumerate(ds):
352
+ instruction = item.get('instruction', '')
353
+ inp = item.get('input', '')
354
+ output = item.get('output', '')
355
+
356
+ # Flatten into linear text
357
+ parts = [f"### Instruction:\n{instruction}"]
358
+ if inp:
359
+ parts.append(f"### Input:\n{inp}")
360
+ parts.append(f"### Output:\n{output}")
361
+ text = '\n\n'.join(parts) + '\n\n'
362
+
363
+ clean = self.tokenizer.clean_text(text)
364
+ self.stats['raw_chars'] += len(text)
365
+ self.stats['clean_chars'] += len(clean)
366
+ self.stats['skipped_chars'] += len(text) - len(clean)
367
+
368
+ tokens = self.tokenizer.encode(clean)
369
+ all_tokens.extend(tokens)
370
+ self.stats['tokens'] += len(tokens)
371
+ total_bytes += len(clean)
372
+
373
+ print(f" {len(ds):,} samples, {self.stats['tokens']:,} tokens")
374
+
375
+ if not dry_run and all_tokens:
376
+ self._save_tokens(all_tokens)
377
+
378
+ return self.stats
379
+
380
+
381
+ class CodeSearchNetProcessor(DatasetProcessor):
382
+ """code_search_net — Code + documentation in 6 languages."""
383
+
384
+ def __init__(self, tokenizer, output_dir, max_gb):
385
+ super().__init__('codesearchnet', 'code', tokenizer, output_dir, max_gb)
386
+ self.languages = ['python', 'javascript', 'java', 'go', 'ruby', 'php']
387
+
388
+ def process(self, dry_run=False):
389
+ from datasets import load_dataset
390
+
391
+ print(f"\n{'='*60}")
392
+ print(f"Processing: {self.name} (CODE - Multi-language)")
393
+ print(f"{'='*60}")
394
+
395
+ all_tokens = []
396
+ total_bytes = 0
397
+
398
+ for lang in self.languages:
399
+ print(f" Loading language: {lang}")
400
+ try:
401
+ ds = load_dataset("code_search_net", languages=[lang],
402
+ split="train", streaming=True, trust_remote_code=True)
403
+
404
+ for i, item in enumerate(ds):
405
+ code = item.get('func_code_string', '')
406
+ doc = item.get('func_documentation_string', '')
407
+
408
+ text = ''
409
+ if doc:
410
+ text += f"# {doc}\n"
411
+ text += code + '\n\n'
412
+
413
+ clean = self.tokenizer.clean_text(text)
414
+ self.stats['raw_chars'] += len(text)
415
+ self.stats['clean_chars'] += len(clean)
416
+ self.stats['skipped_chars'] += len(text) - len(clean)
417
+
418
+ tokens = self.tokenizer.encode(clean)
419
+ all_tokens.extend(tokens)
420
+ self.stats['tokens'] += len(tokens)
421
+ total_bytes += len(clean)
422
+
423
+ if i % 10000 == 0 and i > 0:
424
+ print(f" {lang}: {i:,} funcs, {self.stats['tokens']:,} tokens")
425
+
426
+ if self._should_stop(total_bytes) or i >= 50000:
427
+ break
428
+ except Exception as e:
429
+ print(f" Error with {lang}: {e}")
430
+ continue
431
+
432
+ if self._should_stop(total_bytes):
433
+ break
434
+
435
+ if not dry_run and all_tokens:
436
+ self._save_tokens(all_tokens)
437
+
438
+ return self.stats
439
+
440
+
441
+ class PythonCodeInstructionsProcessor(DatasetProcessor):
442
+ """iamtarun/python_code_instructions_18k_alpaca — Python instruction pairs."""
443
+
444
+ def __init__(self, tokenizer, output_dir, max_gb):
445
+ super().__init__('python_code_instructions', 'code', tokenizer, output_dir, max_gb)
446
+
447
+ def process(self, dry_run=False):
448
+ from datasets import load_dataset
449
+
450
+ print(f"\n{'='*60}")
451
+ print(f"Processing: {self.name} (CODE - Python Instructions)")
452
+ print(f"{'='*60}")
453
+
454
+ ds = load_dataset("iamtarun/python_code_instructions_18k_alpaca", split="train")
455
+
456
+ all_tokens = []
457
+ total_bytes = 0
458
+
459
+ for i, item in enumerate(ds):
460
+ instruction = item.get('instruction', '')
461
+ inp = item.get('input', '')
462
+ output_text = item.get('output', '')
463
+
464
+ parts = [f"### Instruction:\n{instruction}"]
465
+ if inp:
466
+ parts.append(f"### Input:\n{inp}")
467
+ parts.append(f"### Output:\n{output_text}")
468
+ text = '\n\n'.join(parts) + '\n\n'
469
+
470
+ clean = self.tokenizer.clean_text(text)
471
+ self.stats['raw_chars'] += len(text)
472
+ self.stats['clean_chars'] += len(clean)
473
+ self.stats['skipped_chars'] += len(text) - len(clean)
474
+
475
+ tokens = self.tokenizer.encode(clean)
476
+ all_tokens.extend(tokens)
477
+ self.stats['tokens'] += len(tokens)
478
+ total_bytes += len(clean)
479
+
480
+ print(f" {len(ds):,} samples, {self.stats['tokens']:,} tokens")
481
+
482
+ if not dry_run and all_tokens:
483
+ self._save_tokens(all_tokens)
484
+
485
+ return self.stats
486
+
487
+
488
+ class MathPileProcessor(DatasetProcessor):
489
+ """GAIR/MathPile — Large-scale math pretraining corpus."""
490
+
491
+ def __init__(self, tokenizer, output_dir, max_gb):
492
+ super().__init__('mathpile', 'math', tokenizer, output_dir, max_gb)
493
+
494
+ def process(self, dry_run=False):
495
+ from datasets import load_dataset
496
+
497
+ print(f"\n{'='*60}")
498
+ print(f"Processing: {self.name} (MATH)")
499
+ print(f"{'='*60}")
500
+
501
+ all_tokens = []
502
+ total_bytes = 0
503
+
504
+ try:
505
+ ds = load_dataset("zwhe99/mathpile-text", split="train", streaming=True)
506
+ except Exception:
507
+ try:
508
+ ds = load_dataset("GAIR/MathPile", split="train", streaming=True)
509
+ except Exception as e:
510
+ print(f" Could not load MathPile: {e}")
511
+ return self.stats
512
+
513
+ for i, item in enumerate(ds):
514
+ text = item.get('text', '') or item.get('content', '')
515
+ if not text or len(text) < 20:
516
+ continue
517
+
518
+ clean = self.tokenizer.clean_text(text)
519
+ self.stats['raw_chars'] += len(text)
520
+ self.stats['clean_chars'] += len(clean)
521
+ self.stats['skipped_chars'] += len(text) - len(clean)
522
+
523
+ text_out = clean + '\n\n'
524
+ tokens = self.tokenizer.encode(text_out)
525
+ all_tokens.extend(tokens)
526
+ self.stats['tokens'] += len(tokens)
527
+ total_bytes += len(text_out)
528
+
529
+ if i % 5000 == 0 and i > 0:
530
+ print(f" {i:,} docs, {self.stats['tokens']:,} tokens, "
531
+ f"{total_bytes/1024**3:.1f} GB")
532
+
533
+ if self._should_stop(total_bytes) or i >= 100000:
534
+ break
535
+
536
+ if not dry_run and all_tokens:
537
+ self._save_tokens(all_tokens)
538
+
539
+ return self.stats
540
+
541
+
542
+ class OpenMathInstructProcessor(DatasetProcessor):
543
+ """nvidia/OpenMathInstruct-1 — Math reasoning with step-by-step solutions."""
544
+
545
+ def __init__(self, tokenizer, output_dir, max_gb):
546
+ super().__init__('openmath_instruct1', 'math', tokenizer, output_dir, max_gb)
547
+
548
+ def process(self, dry_run=False):
549
+ from datasets import load_dataset
550
+
551
+ print(f"\n{'='*60}")
552
+ print(f"Processing: {self.name} (MATH - Reasoning)")
553
+ print(f"{'='*60}")
554
+
555
+ all_tokens = []
556
+ total_bytes = 0
557
+
558
+ try:
559
+ ds = load_dataset("nvidia/OpenMathInstruct-1", split="train", streaming=True)
560
+ except Exception as e:
561
+ print(f" Could not load OpenMathInstruct-1: {e}")
562
+ return self.stats
563
+
564
+ for i, item in enumerate(ds):
565
+ problem = item.get('problem', '')
566
+ solution = item.get('generated_solution', '')
567
+
568
+ text = f"Problem:\n{problem}\n\nSolution:\n{solution}\n\n"
569
+
570
+ clean = self.tokenizer.clean_text(text)
571
+ self.stats['raw_chars'] += len(text)
572
+ self.stats['clean_chars'] += len(clean)
573
+ self.stats['skipped_chars'] += len(text) - len(clean)
574
+
575
+ tokens = self.tokenizer.encode(clean)
576
+ all_tokens.extend(tokens)
577
+ self.stats['tokens'] += len(tokens)
578
+ total_bytes += len(clean)
579
+
580
+ if i % 50000 == 0 and i > 0:
581
+ print(f" {i:,} problems, {self.stats['tokens']:,} tokens, "
582
+ f"{total_bytes/1024**3:.1f} GB")
583
+
584
+ if self._should_stop(total_bytes) or i >= 500000:
585
+ break
586
+
587
+ if not dry_run and all_tokens:
588
+ self._save_tokens(all_tokens)
589
+
590
+ return self.stats
591
+
592
+
593
+ class MetaMathQAProcessor(DatasetProcessor):
594
+ """meta-math/MetaMathQA — Augmented math Q&A."""
595
+
596
+ def __init__(self, tokenizer, output_dir, max_gb):
597
+ super().__init__('metamath_qa', 'math', tokenizer, output_dir, max_gb)
598
+
599
+ def process(self, dry_run=False):
600
+ from datasets import load_dataset
601
+
602
+ print(f"\n{'='*60}")
603
+ print(f"Processing: {self.name} (MATH - Q&A)")
604
+ print(f"{'='*60}")
605
+
606
+ ds = load_dataset("meta-math/MetaMathQA", split="train")
607
+
608
+ all_tokens = []
609
+ total_bytes = 0
610
+
611
+ for i, item in enumerate(ds):
612
+ query = item.get('query', '')
613
+ response = item.get('response', '')
614
+
615
+ text = f"Question:\n{query}\n\nAnswer:\n{response}\n\n"
616
+
617
+ clean = self.tokenizer.clean_text(text)
618
+ self.stats['raw_chars'] += len(text)
619
+ self.stats['clean_chars'] += len(clean)
620
+ self.stats['skipped_chars'] += len(text) - len(clean)
621
+
622
+ tokens = self.tokenizer.encode(clean)
623
+ all_tokens.extend(tokens)
624
+ self.stats['tokens'] += len(tokens)
625
+ total_bytes += len(clean)
626
+
627
+ print(f" {len(ds):,} samples, {self.stats['tokens']:,} tokens")
628
+
629
+ if not dry_run and all_tokens:
630
+ self._save_tokens(all_tokens)
631
+
632
+ return self.stats
633
+
634
+
635
+ class GSM8KProcessor(DatasetProcessor):
636
+ """openai/gsm8k — Grade-school math word problems."""
637
+
638
+ def __init__(self, tokenizer, output_dir, max_gb):
639
+ super().__init__('gsm8k', 'math', tokenizer, output_dir, max_gb)
640
+
641
+ def process(self, dry_run=False):
642
+ from datasets import load_dataset
643
+
644
+ print(f"\n{'='*60}")
645
+ print(f"Processing: {self.name} (MATH - Grade School)")
646
+ print(f"{'='*60}")
647
+
648
+ ds = load_dataset("openai/gsm8k", "main", split="train")
649
+
650
+ all_tokens = []
651
+ total_bytes = 0
652
+
653
+ for i, item in enumerate(ds):
654
+ question = item.get('question', '')
655
+ answer = item.get('answer', '')
656
+
657
+ text = f"Question:\n{question}\n\nAnswer:\n{answer}\n\n"
658
+
659
+ clean = self.tokenizer.clean_text(text)
660
+ self.stats['raw_chars'] += len(text)
661
+ self.stats['clean_chars'] += len(clean)
662
+ self.stats['skipped_chars'] += len(text) - len(clean)
663
+
664
+ tokens = self.tokenizer.encode(clean)
665
+ all_tokens.extend(tokens)
666
+ self.stats['tokens'] += len(tokens)
667
+ total_bytes += len(clean)
668
+
669
+ print(f" {len(ds):,} problems, {self.stats['tokens']:,} tokens")
670
+
671
+ if not dry_run and all_tokens:
672
+ self._save_tokens(all_tokens)
673
+
674
+ return self.stats
675
+
676
+
677
+ class HendrycksMathProcessor(DatasetProcessor):
678
+ """EleutherAI/hendrycks_math — Competition-level math with LaTeX."""
679
+
680
+ def __init__(self, tokenizer, output_dir, max_gb):
681
+ super().__init__('hendrycks_math', 'math', tokenizer, output_dir, max_gb)
682
+ self.subjects = [
683
+ 'algebra', 'counting_and_probability', 'geometry',
684
+ 'intermediate_algebra', 'number_theory', 'prealgebra', 'precalculus'
685
+ ]
686
+
687
+ def process(self, dry_run=False):
688
+ from datasets import load_dataset
689
+
690
+ print(f"\n{'='*60}")
691
+ print(f"Processing: {self.name} (MATH - Competition)")
692
+ print(f"{'='*60}")
693
+
694
+ all_tokens = []
695
+ total_bytes = 0
696
+
697
+ for subject in self.subjects:
698
+ print(f" Loading subject: {subject}")
699
+ try:
700
+ ds = load_dataset("EleutherAI/hendrycks_math", subject, split="train")
701
+
702
+ for i, item in enumerate(ds):
703
+ problem = item.get('problem', '')
704
+ solution = item.get('solution', '')
705
+
706
+ # LaTeX is ASCII — great for char-level!
707
+ text = f"Problem:\n{problem}\n\nSolution:\n{solution}\n\n"
708
+
709
+ clean = self.tokenizer.clean_text(text)
710
+ self.stats['raw_chars'] += len(text)
711
+ self.stats['clean_chars'] += len(clean)
712
+ self.stats['skipped_chars'] += len(text) - len(clean)
713
+
714
+ tokens = self.tokenizer.encode(clean)
715
+ all_tokens.extend(tokens)
716
+ self.stats['tokens'] += len(tokens)
717
+ total_bytes += len(clean)
718
+ except Exception as e:
719
+ print(f" Error with {subject}: {e}")
720
+ continue
721
+
722
+ print(f" {self.stats['tokens']:,} tokens total")
723
+
724
+ if not dry_run and all_tokens:
725
+ self._save_tokens(all_tokens)
726
+
727
+ return self.stats
728
+
729
+
730
+ class WikiText103Processor(DatasetProcessor):
731
+ """Salesforce/wikitext-103-raw-v1 — High-quality English prose."""
732
+
733
+ def __init__(self, tokenizer, output_dir, max_gb):
734
+ super().__init__('wikitext103', 'syntax', tokenizer, output_dir, max_gb)
735
+
736
+ def process(self, dry_run=False):
737
+ from datasets import load_dataset
738
+
739
+ print(f"\n{'='*60}")
740
+ print(f"Processing: {self.name} (SYNTAX - English Prose)")
741
+ print(f"{'='*60}")
742
+
743
+ ds = load_dataset("Salesforce/wikitext", "wikitext-103-raw-v1")
744
+
745
+ for split_name in ['train', 'validation', 'test']:
746
+ split_ds = ds[split_name]
747
+ all_tokens = []
748
+
749
+ for item in split_ds:
750
+ text = item.get('text', '')
751
+ if not text or text.strip() == '':
752
+ continue
753
+
754
+ clean = self.tokenizer.clean_text(text)
755
+ self.stats['raw_chars'] += len(text)
756
+ self.stats['clean_chars'] += len(clean)
757
+ self.stats['skipped_chars'] += len(text) - len(clean)
758
+
759
+ # Add newline between paragraphs
760
+ text_out = clean + '\n'
761
+ tokens = self.tokenizer.encode(text_out)
762
+ all_tokens.extend(tokens)
763
+ self.stats['tokens'] += len(tokens)
764
+
765
+ if not dry_run and all_tokens:
766
+ self._save_tokens(all_tokens, split_name)
767
+ print(f" {split_name}: {len(all_tokens):,} tokens")
768
+
769
+ return self.stats
770
+
771
+
772
+ class C4SubsetProcessor(DatasetProcessor):
773
+ """allenai/c4 — Massive English web text (subset for syntax)."""
774
+
775
+ def __init__(self, tokenizer, output_dir, max_gb):
776
+ super().__init__('c4_subset', 'syntax', tokenizer, output_dir, max_gb)
777
+
778
+ def process(self, dry_run=False):
779
+ from datasets import load_dataset
780
+
781
+ print(f"\n{'='*60}")
782
+ print(f"Processing: {self.name} (SYNTAX - Web Text)")
783
+ print(f"{'='*60}")
784
+
785
+ all_tokens = []
786
+ total_bytes = 0
787
+
788
+ ds = load_dataset("allenai/c4", "en", split="train", streaming=True)
789
+
790
+ for i, item in enumerate(ds):
791
+ text = item.get('text', '')
792
+ if not text or len(text) < 100:
793
+ continue
794
+
795
+ clean = self.tokenizer.clean_text(text)
796
+ self.stats['raw_chars'] += len(text)
797
+ self.stats['clean_chars'] += len(clean)
798
+ self.stats['skipped_chars'] += len(text) - len(clean)
799
+
800
+ text_out = clean + '\n\n'
801
+ tokens = self.tokenizer.encode(text_out)
802
+ all_tokens.extend(tokens)
803
+ self.stats['tokens'] += len(tokens)
804
+ total_bytes += len(text_out)
805
+
806
+ if i % 10000 == 0 and i > 0:
807
+ print(f" {i:,} docs, {self.stats['tokens']:,} tokens, "
808
+ f"{total_bytes/1024**3:.1f} GB")
809
+
810
+ # Cap at ~5GB of raw text for C4
811
+ if self._should_stop(total_bytes) or i >= 100000:
812
+ break
813
+
814
+ if not dry_run and all_tokens:
815
+ self._save_tokens(all_tokens)
816
+
817
+ return self.stats
818
+
819
+
820
+ class PennTreebankProcessor(DatasetProcessor):
821
+ """ptb_text_only — Gold-standard syntactic English."""
822
+
823
+ def __init__(self, tokenizer, output_dir, max_gb):
824
+ super().__init__('ptb', 'syntax', tokenizer, output_dir, max_gb)
825
+
826
+ def process(self, dry_run=False):
827
+ from datasets import load_dataset
828
+
829
+ print(f"\n{'='*60}")
830
+ print(f"Processing: {self.name} (SYNTAX - Penn Treebank)")
831
+ print(f"{'='*60}")
832
+
833
+ ds = load_dataset("ptb_text_only")
834
+
835
+ for split_name in ['train', 'validation', 'test']:
836
+ if split_name not in ds:
837
+ continue
838
+ split_ds = ds[split_name]
839
+ all_tokens = []
840
+
841
+ for item in split_ds:
842
+ text = item.get('sentence', '')
843
+ if not text:
844
+ continue
845
+
846
+ clean = self.tokenizer.clean_text(text)
847
+ self.stats['raw_chars'] += len(text)
848
+ self.stats['clean_chars'] += len(clean)
849
+ self.stats['skipped_chars'] += len(text) - len(clean)
850
+
851
+ text_out = clean + '\n'
852
+ tokens = self.tokenizer.encode(text_out)
853
+ all_tokens.extend(tokens)
854
+ self.stats['tokens'] += len(tokens)
855
+
856
+ if not dry_run and all_tokens:
857
+ self._save_tokens(all_tokens, split_name)
858
+ print(f" {split_name}: {len(all_tokens):,} tokens")
859
+
860
+ return self.stats
861
+
862
+
863
+ class UniversalDependenciesProcessor(DatasetProcessor):
864
+ """universal-dependencies — Syntax annotations (CoNLL-U format)."""
865
+
866
+ def __init__(self, tokenizer, output_dir, max_gb):
867
+ super().__init__('universal_deps', 'syntax', tokenizer, output_dir, max_gb)
868
+ # English + French treebanks for accent coverage
869
+ self.treebanks = ['en_gum', 'en_ewt', 'fr_gsd', 'fr_sequoia']
870
+
871
+ def process(self, dry_run=False):
872
+ from datasets import load_dataset
873
+
874
+ print(f"\n{'='*60}")
875
+ print(f"Processing: {self.name} (SYNTAX - UD)")
876
+ print(f"{'='*60}")
877
+
878
+ all_tokens = []
879
+ total_bytes = 0
880
+
881
+ for tb in self.treebanks:
882
+ print(f" Loading treebank: {tb}")
883
+ try:
884
+ ds = load_dataset(
885
+ "universal-dependencies/universal_dependencies", tb,
886
+ split="train", trust_remote_code=True
887
+ )
888
+
889
+ for i, item in enumerate(ds):
890
+ # Build CoNLL-U style text from tokens
891
+ tokens_list = item.get('tokens', [])
892
+ lemmas = item.get('lemmas', [])
893
+ upos = item.get('upos_tags', [])
894
+
895
+ # Create a linear text: word/lemma/UPOS per line
896
+ lines = []
897
+ for j, (tok, lem, pos) in enumerate(zip(tokens_list, lemmas, upos)):
898
+ lines.append(f"{tok}\t{lem}\t{pos}")
899
+
900
+ text = '\n'.join(lines) + '\n\n'
901
+
902
+ clean = self.tokenizer.clean_text(text)
903
+ self.stats['raw_chars'] += len(text)
904
+ self.stats['clean_chars'] += len(clean)
905
+ self.stats['skipped_chars'] += len(text) - len(clean)
906
+
907
+ tokens = self.tokenizer.encode(clean)
908
+ all_tokens.extend(tokens)
909
+ self.stats['tokens'] += len(tokens)
910
+ total_bytes += len(clean)
911
+
912
+ if self._should_stop(total_bytes):
913
+ break
914
+ except Exception as e:
915
+ print(f" Error with {tb}: {e}")
916
+ continue
917
+
918
+ if self._should_stop(total_bytes):
919
+ break
920
+
921
+ if not dry_run and all_tokens:
922
+ self._save_tokens(all_tokens)
923
+
924
+ return self.stats
925
+
926
+
927
+ class ThePileSubsetProcessor(DatasetProcessor):
928
+ """EleutherAI/pile — All-in-one: code + math + prose (subset)."""
929
+
930
+ def __init__(self, tokenizer, output_dir, max_gb):
931
+ super().__init__('pile_subset', 'general', tokenizer, output_dir, max_gb)
932
+
933
+ def process(self, dry_run=False):
934
+ from datasets import load_dataset
935
+
936
+ print(f"\n{'='*60}")
937
+ print(f"Processing: {self.name} (GENERAL - The Pile)")
938
+ print(f"{'='*60}")
939
+
940
+ all_tokens = []
941
+ total_bytes = 0
942
+
943
+ try:
944
+ ds = load_dataset("EleutherAI/pile", split="train", streaming=True)
945
+ except Exception as e:
946
+ print(f" Could not load The Pile: {e}")
947
+ return self.stats
948
+
949
+ for i, item in enumerate(ds):
950
+ text = item.get('text', '')
951
+ if not text or len(text) < 50:
952
+ continue
953
+
954
+ clean = self.tokenizer.clean_text(text)
955
+ self.stats['raw_chars'] += len(text)
956
+ self.stats['clean_chars'] += len(clean)
957
+ self.stats['skipped_chars'] += len(text) - len(clean)
958
+
959
+ text_out = clean + '\n\n'
960
+ tokens = self.tokenizer.encode(text_out)
961
+ all_tokens.extend(tokens)
962
+ self.stats['tokens'] += len(tokens)
963
+ total_bytes += len(text_out)
964
+
965
+ if i % 10000 == 0 and i > 0:
966
+ print(f" {i:,} docs, {self.stats['tokens']:,} tokens, "
967
+ f"{total_bytes/1024**3:.1f} GB")
968
+
969
+ # Cap at ~10GB for Pile subset
970
+ if self._should_stop(total_bytes) or i >= 200000:
971
+ break
972
+
973
+ if not dry_run and all_tokens:
974
+ self._save_tokens(all_tokens)
975
+
976
+ return self.stats
977
+
978
+
979
+ class CustomCodeProcessor(DatasetProcessor):
980
+ """Process user-provided code files from a directory."""
981
+
982
+ def __init__(self, tokenizer, output_dir, max_gb, code_dir: str):
983
+ super().__init__('custom_code', 'code', tokenizer, output_dir, max_gb)
984
+ self.code_dir = code_dir
985
+
986
+ def process(self, dry_run=False):
987
+ print(f"\n{'='*60}")
988
+ print(f"Processing: {self.name} (CODE - Custom Files)")
989
+ print(f"{'='*60}")
990
+
991
+ if not os.path.exists(self.code_dir):
992
+ print(f" Directory not found: {self.code_dir}")
993
+ return self.stats
994
+
995
+ all_tokens = []
996
+ total_bytes = 0
997
+ supported_exts = {'.py', '.js', '.ts', '.c', '.cpp', '.h', '.hpp',
998
+ '.java', '.rs', '.go', '.rb', '.php', '.sh', '.sql',
999
+ '.html', '.css', '.json', '.yaml', '.yml', '.toml',
1000
+ '.md', '.txt', '.tex', '.cfg', '.ini'}
1001
+
1002
+ for root, dirs, files in os.walk(self.code_dir):
1003
+ for fname in files:
1004
+ ext = Path(fname).suffix.lower()
1005
+ if ext not in supported_exts:
1006
+ continue
1007
+
1008
+ fpath = os.path.join(root, fname)
1009
+ try:
1010
+ with open(fpath, 'r', encoding='utf-8', errors='replace') as f:
1011
+ content = f.read()
1012
+ except Exception:
1013
+ continue
1014
+
1015
+ if not content or len(content) < 10:
1016
+ continue
1017
+
1018
+ clean = self.tokenizer.clean_text(content)
1019
+ self.stats['raw_chars'] += len(content)
1020
+ self.stats['clean_chars'] += len(clean)
1021
+ self.stats['skipped_chars'] += len(content) - len(clean)
1022
+
1023
+ text = clean + '\n\n'
1024
+ tokens = self.tokenizer.encode(text)
1025
+ all_tokens.extend(tokens)
1026
+ self.stats['tokens'] += len(tokens)
1027
+ self.stats['files'] += 1
1028
+ total_bytes += len(text)
1029
+
1030
+ if self.stats['files'] % 100 == 0:
1031
+ print(f" {self.stats['files']} files, {self.stats['tokens']:,} tokens")
1032
+
1033
+ print(f" Total: {self.stats['files']} files, {self.stats['tokens']:,} tokens")
1034
+
1035
+ if not dry_run and all_tokens:
1036
+ self._save_tokens(all_tokens)
1037
+
1038
+ return self.stats
1039
+
1040
+
1041
+ # ─── Merge All Tokenized Data ────────────────────────────────────────────────
1042
+
1043
+ def merge_all_data(output_dir: str, tokenizer_path: str, val_ratio: float = 0.02):
1044
+ """Merge all .pt token files into unified train/val splits."""
1045
+ import torch
1046
+
1047
+ print(f"\n{'='*60}")
1048
+ print(f"Merging all tokenized data")
1049
+ print(f"{'='*60}")
1050
+
1051
+ # Find all .pt files
1052
+ pt_files = sorted(Path(output_dir).glob('*.pt'))
1053
+ print(f"Found {len(pt_files)} tokenized files:")
1054
+ for f in pt_files:
1055
+ size_mb = f.stat().st_size / 1024**2
1056
+ print(f" {f.name} ({size_mb:.1f} MB)")
1057
+
1058
+ if not pt_files:
1059
+ print("No tokenized files found!")
1060
+ return
1061
+
1062
+ # Load and concatenate
1063
+ all_tokens = []
1064
+ for f in pt_files:
1065
+ print(f" Loading {f.name}...")
1066
+ t = torch.load(f, weights_only=True)
1067
+ all_tokens.append(t)
1068
+
1069
+ all_tokens = torch.cat(all_tokens, dim=0)
1070
+ total_tokens = len(all_tokens)
1071
+ print(f"\nTotal tokens: {total_tokens:,}")
1072
+ print(f"Total size: {total_tokens * 8 / 1024**3:.2f} GB (as int64)")
1073
+
1074
+ # Convert to int16 to save space (vocab_size=136 fits in uint8 but int16 is safer)
1075
+ all_tokens = all_tokens.to(torch.int16)
1076
+ print(f"Compressed size: {total_tokens * 2 / 1024**3:.2f} GB (as int16)")
1077
+
1078
+ # Shuffle
1079
+ print("Shuffling tokens...")
1080
+ perm = torch.randperm(total_tokens)
1081
+ all_tokens = all_tokens[perm]
1082
+
1083
+ # Split train/val
1084
+ val_size = int(total_tokens * val_ratio)
1085
+ train_tokens = all_tokens[val_size:]
1086
+ val_tokens = all_tokens[:val_size]
1087
+
1088
+ print(f"Train tokens: {len(train_tokens):,}")
1089
+ print(f"Val tokens: {len(val_tokens):,}")
1090
+
1091
+ # Save
1092
+ train_path = os.path.join(output_dir, 'train_merged.pt')
1093
+ val_path = os.path.join(output_dir, 'val_merged.pt')
1094
+
1095
+ torch.save(train_tokens, train_path)
1096
+ torch.save(val_tokens, val_path)
1097
+
1098
+ print(f"\nSaved: {train_path} ({len(train_tokens)*2/1024**3:.2f} GB)")
1099
+ print(f"Saved: {val_path} ({len(val_tokens)*2/1024**3:.2f} GB)")
1100
+
1101
+ # Save dataset manifest
1102
+ manifest = {
1103
+ 'total_tokens': total_tokens,
1104
+ 'train_tokens': len(train_tokens),
1105
+ 'val_tokens': len(val_tokens),
1106
+ 'vocab_size': 136,
1107
+ 'tokenizer': tokenizer_path,
1108
+ 'source_files': [f.name for f in pt_files],
1109
+ }
1110
+ manifest_path = os.path.join(output_dir, 'manifest.json')
1111
+ with open(manifest_path, 'w') as f:
1112
+ json.dump(manifest, f, indent=2)
1113
+ print(f"Saved manifest: {manifest_path}")
1114
+
1115
+
1116
+ # ─── Main ────────────────────────────────────────────────────────────────────
1117
+
1118
+ def main():
1119
+ parser = argparse.ArgumentParser(description='CogNet 1B Data Preparation')
1120
+ parser.add_argument('--output_dir', type=str, default='/root/CogNet/data_1b',
1121
+ help='Output directory for tokenized data')
1122
+ parser.add_argument('--max_gb', type=float, default=50,
1123
+ help='Maximum GB of text data to download (default: 50)')
1124
+ parser.add_argument('--dry_run', action='store_true',
1125
+ help='Only show what would be downloaded, dont save')
1126
+ parser.add_argument('--tokenizer', type=str, default=None,
1127
+ help='Path to tokenizer JSON (default: create new)')
1128
+ parser.add_argument('--skip_merge', action='store_true',
1129
+ help='Skip the final merge step')
1130
+ parser.add_argument('--only', type=str, nargs='*', default=None,
1131
+ help='Only process these datasets (e.g., --only code_alpaca gsm8k)')
1132
+ parser.add_argument('--custom_code_dir', type=str, default=None,
1133
+ help='Directory of custom code files to include')
1134
+ args = parser.parse_args()
1135
+
1136
+ # Tokenizer
1137
+ if args.tokenizer and os.path.exists(args.tokenizer):
1138
+ print(f"Loading tokenizer from {args.tokenizer}")
1139
+ tokenizer = CharTokenizer.load(args.tokenizer)
1140
+ else:
1141
+ tokenizer = CharTokenizer()
1142
+ print(f"Tokenizer: vocab_size={tokenizer.vocab_size}")
1143
+
1144
+ # Save tokenizer
1145
+ os.makedirs(args.output_dir, exist_ok=True)
1146
+ tok_path = os.path.join(args.output_dir, 'tokenizer_v3.json')
1147
+ tokenizer.save(tok_path)
1148
+
1149
+ # Dataset processors (ordered by priority)
1150
+ all_processors = [
1151
+ # ── CODE ──
1152
+ ('the_stack_smol', TheStackSmolProcessor),
1153
+ ('codeparrot_clean', CodeParrotProcessor),
1154
+ ('code_alpaca', CodeAlpacaProcessor),
1155
+ ('codesearchnet', CodeSearchNetProcessor),
1156
+ ('python_code_instructions', PythonCodeInstructionsProcessor),
1157
+ # ── MATH ──
1158
+ ('mathpile', MathPileProcessor),
1159
+ ('openmath_instruct1', OpenMathInstructProcessor),
1160
+ ('metamath_qa', MetaMathQAProcessor),
1161
+ ('gsm8k', GSM8KProcessor),
1162
+ ('hendrycks_math', HendrycksMathProcessor),
1163
+ # ── SYNTAX ──
1164
+ ('wikitext103', WikiText103Processor),
1165
+ ('c4_subset', C4SubsetProcessor),
1166
+ ('ptb', PennTreebankProcessor),
1167
+ ('universal_deps', UniversalDependenciesProcessor),
1168
+ # ── GENERAL ──
1169
+ ('pile_subset', ThePileSubsetProcessor),
1170
+ ]
1171
+
1172
+ # Filter if --only specified
1173
+ if args.only:
1174
+ all_processors = [(n, p) for n, p in all_processors if n in args.only]
1175
+ print(f"Processing only: {args.only}")
1176
+
1177
+ # Add custom code if specified
1178
+ if args.custom_code_dir:
1179
+ all_processors.append((
1180
+ 'custom_code',
1181
+ lambda t, o, m: CustomCodeProcessor(t, o, m, args.custom_code_dir)
1182
+ ))
1183
+
1184
+ # Process all datasets
1185
+ all_stats = []
1186
+ start_time = time.time()
1187
+
1188
+ for name, processor_cls in all_processors:
1189
+ try:
1190
+ processor = processor_cls(tokenizer, args.output_dir, args.max_gb)
1191
+ stats = processor.process(dry_run=args.dry_run)
1192
+ all_stats.append(stats)
1193
+ except Exception as e:
1194
+ print(f"\n ERROR processing {name}: {e}")
1195
+ import traceback
1196
+ traceback.print_exc()
1197
+ continue
1198
+
1199
+ elapsed = time.time() - start_time
1200
+
1201
+ # Summary
1202
+ print(f"\n{'='*60}")
1203
+ print(f"DATA PREPARATION SUMMARY")
1204
+ print(f"{'='*60}")
1205
+ print(f"Elapsed: {elapsed/60:.1f} minutes")
1206
+ print(f"")
1207
+
1208
+ total_tokens = 0
1209
+ total_chars = 0
1210
+ by_category = {}
1211
+
1212
+ for s in all_stats:
1213
+ total_tokens += s['tokens']
1214
+ total_chars += s['clean_chars']
1215
+ cat = s['category']
1216
+ if cat not in by_category:
1217
+ by_category[cat] = {'tokens': 0, 'chars': 0, 'datasets': 0}
1218
+ by_category[cat]['tokens'] += s['tokens']
1219
+ by_category[cat]['chars'] += s['clean_chars']
1220
+ by_category[cat]['datasets'] += 1
1221
+
1222
+ print(f" {s['name']:30s} [{s['category']:8s}] "
1223
+ f"{s['tokens']:>12,} tokens "
1224
+ f"{s['clean_chars']:>12,} chars "
1225
+ f"skipped: {s['skipped_chars']:,}")
1226
+
1227
+ print(f"\n {'TOTAL':30s} {'':8s} {total_tokens:>12,} tokens {total_chars:>12,} chars")
1228
+ print(f"\n By category:")
1229
+ for cat, info in by_category.items():
1230
+ print(f" {cat:10s}: {info['tokens']:>12,} tokens ({info['datasets']} datasets)")
1231
+
1232
+ # Save stats
1233
+ stats_path = os.path.join(args.output_dir, 'prep_stats.json')
1234
+ with open(stats_path, 'w') as f:
1235
+ json.dump({
1236
+ 'elapsed_seconds': elapsed,
1237
+ 'total_tokens': total_tokens,
1238
+ 'total_chars': total_chars,
1239
+ 'by_category': by_category,
1240
+ 'datasets': all_stats,
1241
+ }, f, indent=2)
1242
+ print(f"\nStats saved to: {stats_path}")
1243
+
1244
+ # Merge
1245
+ if not args.skip_merge and not args.dry_run:
1246
+ merge_all_data(args.output_dir, tok_path)
1247
+
1248
+ print(f"\nDone! Tokenized data ready in: {args.output_dir}")
1249
+
1250
+
1251
+ if __name__ == '__main__':
1252
+ main()