| """ |
| Data Pipeline: Streaming tokenizer, packing, and dataset caching |
| - GPT-2 Tokenizer via Tokenizer / tiktoken |
| - Continuous Token Packing (concatenates documents with <|endoftext|> and chunks into 512 tokens) |
| - Dataloader constructor build_dataloaders() |
| """ |
|
|
| import os |
| import json |
| from typing import Optional |
| import torch |
| import numpy as np |
| from torch.utils.data import Dataset, DataLoader |
| from datasets import load_dataset |
| from tokenizer import Tokenizer, get_tokenizer |
|
|
|
|
| class PackedDataset(Dataset): |
| """Memory-mapped array of packed token sequences.""" |
| def __init__(self, data_path: str, context_length: int = 512): |
| self.context_length = context_length |
| if data_path.endswith(".npy"): |
| self.data = np.load(data_path, mmap_mode='r') |
| else: |
| self.data = np.memmap(data_path, dtype=np.uint16, mode='r') |
| self.num_samples = len(self.data) // context_length |
|
|
| def __len__(self): |
| return self.num_samples |
|
|
| def __getitem__(self, idx): |
| start = idx * self.context_length |
| end = start + self.context_length + 1 |
| chunk = self.data[start:end] |
| |
| if len(chunk) < self.context_length + 1: |
| chunk = self.data[start:start + self.context_length] |
| x = torch.from_numpy(np.array(chunk, dtype=np.int64)) |
| y = torch.roll(x, -1) |
| y[-1] = -1 |
| return x, y |
|
|
| x = torch.from_numpy(np.array(chunk[:-1], dtype=np.int64)) |
| y = torch.from_numpy(np.array(chunk[1:], dtype=np.int64)) |
| return x, y |
|
|
|
|
| def iter_local_texts(corpus_dir: str): |
| """ |
| Yields text documents from a local corpus directory. |
| Supported formats: .txt, .md, .json, .jsonl |
| """ |
| if not corpus_dir or not os.path.isdir(corpus_dir): |
| raise ValueError(f"Local corpus directory not found: {corpus_dir}") |
|
|
| supported_suffixes = {".txt", ".md", ".json", ".jsonl"} |
| for root, _, files in os.walk(corpus_dir): |
| for name in sorted(files): |
| _, ext = os.path.splitext(name) |
| if ext.lower() not in supported_suffixes: |
| continue |
|
|
| path = os.path.join(root, name) |
| try: |
| if ext.lower() in {".txt", ".md"}: |
| with open(path, "r", encoding="utf-8") as f: |
| text = f.read().strip() |
| if text: |
| yield text |
| elif ext.lower() == ".json": |
| with open(path, "r", encoding="utf-8") as f: |
| payload = json.load(f) |
|
|
| if isinstance(payload, dict): |
| for key in ("text", "content", "body"): |
| value = payload.get(key) |
| if isinstance(value, str) and value.strip(): |
| yield value.strip() |
| break |
| elif isinstance(payload, list): |
| for item in payload: |
| if isinstance(item, str) and item.strip(): |
| yield item.strip() |
| elif isinstance(item, dict): |
| for key in ("text", "content", "body"): |
| value = item.get(key) |
| if isinstance(value, str) and value.strip(): |
| yield value.strip() |
| break |
| else: |
| with open(path, "r", encoding="utf-8") as f: |
| for line in f: |
| line = line.strip() |
| if not line: |
| continue |
| try: |
| record = json.loads(line) |
| except json.JSONDecodeError: |
| yield line |
| continue |
|
|
| if isinstance(record, str) and record.strip(): |
| yield record.strip() |
| elif isinstance(record, dict): |
| for key in ("text", "content", "body"): |
| value = record.get(key) |
| if isinstance(value, str) and value.strip(): |
| yield value.strip() |
| break |
| except (OSError, UnicodeDecodeError, json.JSONDecodeError) as exc: |
| print(f"[!] Skipping {path}: {exc}") |
|
|
|
|
| def get_cache_manifest_path(cache_dir: str) -> str: |
| return os.path.join(cache_dir, "dataset_manifest.json") |
|
|
|
|
| def load_cache_manifest(cache_dir: str) -> Optional[dict]: |
| manifest_path = get_cache_manifest_path(cache_dir) |
| if not os.path.exists(manifest_path): |
| return None |
| with open(manifest_path, "r", encoding="utf-8") as f: |
| return json.load(f) |
|
|
|
|
| def save_cache_manifest(cache_dir: str, manifest: dict): |
| manifest_path = get_cache_manifest_path(cache_dir) |
| with open(manifest_path, "w", encoding="utf-8") as f: |
| json.dump(manifest, f, indent=2) |
|
|
|
|
| def build_and_cache_dataset( |
| cache_dir: str = "/tmp/data", |
| max_books: int = None, |
| val_ratio: float = 0.02, |
| dataset_name: str = "wikitext", |
| corpus_dir: str = None, |
| target_tokens: int = None, |
| ): |
| """ |
| Builds tokenized train/val arrays from either a Hugging Face dataset or a local corpus directory. |
| Supported datasets: "wikitext" (wikitext-103-raw-v1), "openwebtext", |
| "tiny_shakespeare", "c4_en", "slimpajama", "local_files" |
| """ |
| os.makedirs(cache_dir, exist_ok=True) |
| train_path = os.path.join(cache_dir, "train.bin") |
| val_path = os.path.join(cache_dir, "val.bin") |
| expected_manifest = { |
| "dataset_name": dataset_name, |
| "corpus_dir": corpus_dir, |
| "max_books": max_books, |
| "target_tokens": target_tokens, |
| "val_ratio": val_ratio, |
| } |
|
|
| if os.path.exists(train_path) and os.path.exists(val_path): |
| cached_manifest = load_cache_manifest(cache_dir) |
| if cached_manifest == expected_manifest: |
| print(f"[*] Found cached dataset in {cache_dir}") |
| return train_path, val_path |
| print("[*] Cached dataset settings changed; rebuilding token cache...") |
|
|
| tok = get_tokenizer() |
| eot_id = tok.eot_token |
|
|
| |
| if dataset_name == "wikitext": |
| print("[*] Loading WikiText-103...") |
| ds = load_dataset("wikitext", "wikitext-103-raw-v1", split="train", streaming=True) |
| text_key = "text" |
| elif dataset_name == "openwebtext": |
| print("[*] Loading OpenWebText...") |
| ds = load_dataset("openwebtext", split="train", streaming=True) |
| text_key = "text" |
| elif dataset_name == "tiny_shakespeare": |
| print("[*] Loading Tiny Shakespeare...") |
| ds = load_dataset("tiny_shakespeare", split="train", streaming=True) |
| text_key = "text" |
| elif dataset_name == "c4_en": |
| print("[*] Loading C4 English...") |
| ds = load_dataset("allenai/c4", "en", split="train", streaming=True) |
| text_key = "text" |
| elif dataset_name == "slimpajama": |
| print("[*] Loading SlimPajama...") |
| ds = load_dataset("cerebras/SlimPajama-627B", split="train", streaming=True) |
| text_key = "text" |
| elif dataset_name == "local_files": |
| print(f"[*] Loading local corpus from {corpus_dir}...") |
| ds = iter_local_texts(corpus_dir) |
| text_key = None |
| else: |
| raise ValueError( |
| f"Unknown dataset: {dataset_name}. Use 'wikitext', 'openwebtext', " |
| f"'tiny_shakespeare', 'c4_en', 'slimpajama', or 'local_files'" |
| ) |
|
|
| val_interval = int(1.0 / val_ratio) if val_ratio > 0 else 0 |
| total_tokens = 0 |
| train_tokens_count = 0 |
| val_tokens_count = 0 |
| doc_count = 0 |
|
|
| print(f"[*] Streaming & tokenizing directly to binary files in {cache_dir}...") |
| with open(train_path, "wb") as f_train, open(val_path, "wb") as f_val: |
| for item in ds: |
| text = item if text_key is None else item.get(text_key, "") |
| if not text.strip(): |
| continue |
| |
| tokens = tok.encode(text) |
| tokens.append(eot_id) |
| token_arr = np.array(tokens, dtype=np.uint16) |
| token_bytes = token_arr.tobytes() |
|
|
| if val_interval > 0 and (doc_count % val_interval == 0): |
| f_val.write(token_bytes) |
| val_tokens_count += len(tokens) |
| else: |
| f_train.write(token_bytes) |
| train_tokens_count += len(tokens) |
|
|
| total_tokens += len(tokens) |
| doc_count += 1 |
|
|
| if doc_count % 1000 == 0: |
| print(f" Processed {doc_count} documents ({total_tokens:,} tokens)...") |
|
|
| if target_tokens and total_tokens >= target_tokens: |
| print(f"[*] Reached target token budget: {total_tokens:,} / {target_tokens:,}") |
| break |
|
|
| if max_books and doc_count >= max_books: |
| break |
|
|
| print(f"[*] Finished processing {doc_count} documents.") |
| print(f"[*] Total tokens: {total_tokens:,} (Train: {train_tokens_count:,}, Val: {val_tokens_count:,})") |
| print(f"[*] Saved train tokens -> {train_path}") |
| print(f"[*] Saved val tokens -> {val_path}") |
| save_cache_manifest(cache_dir, expected_manifest) |
| print("[*] Dataset cache complete.") |
|
|
| return train_path, val_path |
|
|
|
|
| def build_dataloaders( |
| cache_dir: str = "/tmp/data", |
| context_length: int = 512, |
| batch_size: int = 32, |
| num_workers: int = 2, |
| dataset_name: str = "wikitext", |
| corpus_dir: str = None, |
| target_tokens: int = None, |
| ): |
| """ |
| Returns (train_loader, val_loader, tokenizer) |
| """ |
| train_path, val_path = build_and_cache_dataset( |
| cache_dir=cache_dir, |
| dataset_name=dataset_name, |
| corpus_dir=corpus_dir, |
| target_tokens=target_tokens, |
| ) |
| train_ds = PackedDataset(train_path, context_length=context_length) |
| val_ds = PackedDataset(val_path, context_length=context_length) |
| |
| train_loader = DataLoader(train_ds, batch_size=batch_size, shuffle=True, num_workers=num_workers, pin_memory=True) |
| val_loader = DataLoader(val_ds, batch_size=batch_size, shuffle=False, num_workers=num_workers, pin_memory=True) |
| tok = get_tokenizer() |
| |
| return train_loader, val_loader, tok |
|
|