"""Download and prepare English text data for training. Downloads: - WikiText-2 (2M tokens) — small, for smoke tests - WikiText-103 (100M tokens) — medium, for real training - OpenWebText subset — large, for production training All text is English only. Tokenized with tiktoken GPT-2 BPE (50257 vocab). Cached as .pt files for fast loading. """ import os import subprocess import torch import tiktoken CACHE_DIR = "data" # Direct text URLs (no zip, no redirects) DATASETS = { "wikitext2": { # WikiText-2 raw text from the tinystories-like mirror "url": "https://raw.githubusercontent.com/pytorch/examples/main/word_language_model/data/wikitext-2/train.txt", "size_tokens": "~2M", "desc": "Wikipedia articles, 2M tokens. Good for smoke tests.", }, "tinyshakespeare": { "url": "https://raw.githubusercontent.com/karpathy/char-rnn/master/data/tinyshakespeare/input.txt", "size_tokens": "~1M", "desc": "Shakespeare text, 1M chars. Good for quick smoke tests.", }, "openwebtext_10k": { # 10k OpenWebText samples "url": "https://huggingface.co/datasets/openwebtext/resolve/main/openwebtext.txt", "size_tokens": "~10M", "desc": "OpenWebText subset, 10M tokens. Good for real training.", }, } def download_text(name: str) -> str: """Download a text dataset. Returns the text content.""" info = DATASETS[name] url = info["url"] cache_path = os.path.join(CACHE_DIR, f"{name}.txt") os.makedirs(CACHE_DIR, exist_ok=True) if os.path.exists(cache_path): with open(cache_path, "r", encoding="utf-8", errors="ignore") as f: return f.read() print(f" {name}: downloading from {url}...") subprocess.run(["curl", "-sL", "-o", cache_path, url], check=True) size_mb = os.path.getsize(cache_path) / 1e6 print(f" {name}: downloaded {size_mb:.1f} MB") with open(cache_path, "r", encoding="utf-8", errors="ignore") as f: return f.read() def prepare_data(name: str, force: bool = False) -> str: """Download, tokenize, and cache a dataset. Returns path to .pt file.""" cache_path = os.path.join(CACHE_DIR, f"{name}_tokens.pt") if os.path.exists(cache_path) and not force: tokens = torch.load(cache_path) print(f" {name}: cached {tokens.numel():,} tokens at {cache_path}") return cache_path print(f"Preparing dataset: {name}") text = download_text(name) print(f" {name}: {len(text):,} chars") enc = tiktoken.get_encoding("gpt2") tokens = torch.tensor(enc.encode_ordinary(text), dtype=torch.long) print(f" {name}: {tokens.numel():,} tokens") torch.save(tokens, cache_path) print(f" {name}: saved to {cache_path}") return cache_path def load_tokens(name: str) -> torch.Tensor: """Load pre-tokenized data. Downloads + tokenizes if needed.""" path = prepare_data(name) return torch.load(path) if __name__ == "__main__": import sys name = sys.argv[1] if len(sys.argv) > 1 else "wikitext2" prepare_data(name)