"""Download and prepare English text data for training. Datasets (via HF datasets): - wikitext2: ~2M tokens (smoke test) - wikitext103: ~103M tokens (Wikipedia articles) - mixture250m: ~255M tokens — FineWeb-Edu + Cosmopedia + WikiText-103 All text is English only. Tokenized with tiktoken GPT-2 BPE (50257 vocab). Cached as .pt files for fast loading. """ import os import torch import tiktoken CACHE_DIR = "data" def _tokenize_stream(rows, target, enc, name): """Tokenize text rows until target tokens collected. Returns [N] tensor.""" chunks = [] total = 0 batch = [] for t in rows: if not t or not t.strip(): continue batch.append(t) if len(batch) >= 256: ids = enc.encode_ordinary("\n".join(batch)) chunks.append(torch.tensor(ids, dtype=torch.long)) total += len(ids) if total % 20_000_000 < 3_000_000: print(f" {name}: {total/1e6:.0f}M tokens...", flush=True) if total >= target: break batch = [] if batch and total < target: ids = enc.encode_ordinary("\n".join(batch)) chunks.append(torch.tensor(ids, dtype=torch.long)) total += len(ids) print(f" {name}: done — {total:,} tokens") return torch.cat(chunks)[:target] if chunks else torch.empty(0, dtype=torch.long) def prepare_wikitext103(target_tokens: int = 80_000_000) -> str: """Download wikitext-103 via HF datasets, tokenize, cache. Returns .pt path.""" cache_path = os.path.join(CACHE_DIR, "wikitext103_tokens.pt") if os.path.exists(cache_path): tokens = torch.load(cache_path) print(f"wikitext103: cached {tokens.numel():,} tokens") return cache_path from datasets import load_dataset print("wikitext103: downloading via HF datasets...") ds = load_dataset("Salesforce/wikitext", "wikitext-103-raw-v1", split="train") print(f" {len(ds):,} rows") enc = tiktoken.get_encoding("gpt2") os.makedirs(CACHE_DIR, exist_ok=True) tokens = _tokenize_stream((r["text"] for r in ds), target_tokens, enc, "wikitext103") torch.save(tokens, cache_path) print(f" saved to {cache_path}") return cache_path def prepare_wikitext2() -> str: """Download wikitext-2 (small smoke test corpus).""" cache_path = os.path.join(CACHE_DIR, "wikitext2_tokens.pt") if os.path.exists(cache_path): tokens = torch.load(cache_path) print(f"wikitext2: cached {tokens.numel():,} tokens") return cache_path import subprocess os.makedirs(CACHE_DIR, exist_ok=True) txt_path = os.path.join(CACHE_DIR, "wikitext2.txt") if not os.path.exists(txt_path): url = ("https://raw.githubusercontent.com/pytorch/examples/main/" "word_language_model/data/wikitext-2/train.txt") subprocess.run(["curl", "-sL", "-o", txt_path, url], check=True) with open(txt_path, "r", encoding="utf-8", errors="ignore") as f: text = f.read() enc = tiktoken.get_encoding("gpt2") tokens = torch.tensor(enc.encode_ordinary(text), dtype=torch.long) print(f"wikitext2: {tokens.numel():,} tokens") torch.save(tokens, cache_path) return cache_path def _stream_hf(repo, config, split, text_field, target, enc, name, skip_docs=0): """Stream an HF dataset, return token tensor (empty on failure). skip_docs: drop the first N documents (fresh data past what the model already trained on).""" from datasets import load_dataset try: print(f"{name}: streaming {repo} / {config} (skip {skip_docs:,})...") ds = load_dataset(repo, config, split=split, streaming=True) it = (r.get(text_field, "") for r in ds) if skip_docs: import itertools it = itertools.islice(it, skip_docs, None) return _tokenize_stream(it, target, enc, name) except Exception as e: print(f" {name} FAILED: {type(e).__name__}: {e}") return torch.empty(0, dtype=torch.long) def prepare_mixture250m() -> str: """~255M token English mixture: FineWeb-Edu + Cosmopedia + WikiText-103. FineWeb-Edu: educationally filtered web text — highest quality per token. Cosmopedia: synthetic textbooks — clean structured English. WikiText-103: Wikipedia articles (reuses existing cache). """ cache_path = os.path.join(CACHE_DIR, "mixture250m_tokens.pt") if os.path.exists(cache_path): tokens = torch.load(cache_path) print(f"mixture250m: cached {tokens.numel():,} tokens") return cache_path enc = tiktoken.get_encoding("gpt2") os.makedirs(CACHE_DIR, exist_ok=True) parts = [] # 1. FineWeb-Edu — educational web text (primary source) t = _stream_hf("HuggingFaceTB/fineweb-edu", "sample-10BT", "train", "text", 110_000_000, enc, "fineweb-edu") if t.numel() > 0: parts.append(t) # 2. Cosmopedia — synthetic textbooks (multiple configs for diversity) cosmo_budget = 65_000_000 for cfg in ["openstax", "stanford", "wikihow", "stories", "khanacademy"]: if cosmo_budget <= 0: break t = _stream_hf("HuggingFaceTB/cosmopedia", cfg, "train", "text", cosmo_budget, enc, f"cosmopedia-{cfg}") if t.numel() > 0: parts.append(t) cosmo_budget -= t.numel() # 3. WikiText-103 — Wikipedia (reuse cache if present) wt_path = os.path.join(CACHE_DIR, "wikitext103_tokens.pt") if os.path.exists(wt_path): t = torch.load(wt_path) print(f"wikitext103: reusing cache — {t.numel():,} tokens") else: t = _stream_hf("Salesforce/wikitext", "wikitext-103-raw-v1", "train", "text", 80_000_000, enc, "wikitext103") if t.numel() > 0: parts.append(t) # Fallback: if fineweb/cosmopedia both failed, top up with openwebtext total = sum(t.numel() for t in parts) if total < 200_000_000: need = 255_000_000 - total print(f"mixture short ({total/1e6:.0f}M) — topping up with openwebtext ({need/1e6:.0f}M)...") t = _stream_hf("Skylion007/openwebtext", None, "train", "text", need, enc, "openwebtext") if t.numel() > 0: parts.append(t) tokens = torch.cat(parts) print(f"mixture250m: {tokens.numel():,} tokens total " f"({', '.join(f'{t.numel()//1_000_000}M' for t in parts)})") torch.save(tokens, cache_path) print(f" saved to {cache_path}") return cache_path def prepare_mixture500m() -> str: """~510M token English mixture — SmolLM2-style recipe for max quality jump. fineweb-edu-dedup: 200M — educationally-filtered web (accuracy) cosmopedia-v2: 100M — synthetic textbooks (coherent exposition) openwebtext: 100M — diverse general web (narrative English) wikitext103: 80M — Wikipedia (encyclopedic, cached) finepdfs: 30M — long-form books/papers (topic coherence) """ cache_path = os.path.join(CACHE_DIR, "mixture500m_tokens.pt") if os.path.exists(cache_path): tokens = torch.load(cache_path) print(f"mixture500m: cached {tokens.numel():,} tokens") return cache_path enc = tiktoken.get_encoding("gpt2") os.makedirs(CACHE_DIR, exist_ok=True) parts = [] # 1. FineWeb-Edu dedup (via cosmopedia-v2 repo — same data, not gated) t = _stream_hf("HuggingFaceTB/cosmopedia-v2", "fineweb-edu-dedup", "train", "text", 200_000_000, enc, "fineweb-edu-dedup") if t.numel() > 0: parts.append(t) if t.numel() < 100_000_000: # fallback: smollm-corpus mirror t2 = _stream_hf("HuggingFaceTB/smollm-corpus", "fineweb-edu-dedup", "train", "text", 200_000_000 - t.numel(), enc, "fineweb-edu-dedup-b") if t2.numel() > 0: parts.append(t2) # 2. Cosmopedia v2 — synthetic textbooks t = _stream_hf("HuggingFaceTB/cosmopedia-v2", "cosmopedia-v2", "train", "text", 100_000_000, enc, "cosmopedia-v2") if t.numel() > 0: parts.append(t) # 3. OpenWebText — diverse general web t = _stream_hf("Skylion007/openwebtext", None, "train", "text", 100_000_000, enc, "openwebtext") if t.numel() > 0: parts.append(t) # 4. WikiText-103 — cached Wikipedia wt_path = os.path.join(CACHE_DIR, "wikitext103_tokens.pt") if os.path.exists(wt_path): t = torch.load(wt_path) print(f"wikitext103: reusing cache — {t.numel():,} tokens") else: t = _stream_hf("Salesforce/wikitext", "wikitext-103-raw-v1", "train", "text", 80_000_000, enc, "wikitext103") if t.numel() > 0: parts.append(t) # 5. FinePDFs — long-form books/papers (small slice, fights topic drift) t = _stream_hf("HuggingFaceFW/finepdfs", "eng_Latn", "train", "text", 30_000_000, enc, "finepdfs") if t.numel() > 0: parts.append(t) # Top-up with c4 if any source failed badly total = sum(t.numel() for t in parts) if total < 450_000_000: need = 510_000_000 - total print(f"mixture short ({total/1e6:.0f}M) — topping up with c4 ({need/1e6:.0f}M)...") t = _stream_hf("allenai/c4", "en", "train", "text", need, enc, "c4") if t.numel() > 0: parts.append(t) tokens = torch.cat(parts) print(f"mixture500m: {tokens.numel():,} tokens total " f"({', '.join(f'{t.numel()//1_000_000}M' for t in parts)})") torch.save(tokens, cache_path) print(f" saved to {cache_path}") return cache_path def prepare_mixture1b() -> str: """~1.3B FRESH tokens — skips past docs the 500m mixture already used. fineweb-edu: 700M — primary (skip ~800K docs) cosmopedia-v2: 250M — textbooks/lessons (skip ~150K docs) c4: 250M — common-crawl diversity (unused before) openwebtext: 150M — skip ~50K docs """ cache_path = os.path.join(CACHE_DIR, "mixture1b_tokens.pt") if os.path.exists(cache_path): tokens = torch.load(cache_path) print(f"mixture1b: cached {tokens.numel():,} tokens") return cache_path enc = tiktoken.get_encoding("gpt2") parts = [] total = 0 t = _stream_hf("HuggingFaceTB/fineweb-edu", "sample-10BT", "train", "text", 700_000_000, enc, "fineweb-edu", skip_docs=800_000) if t.numel() < 400_000_000: t2 = _stream_hf("HuggingFaceTB/smollm-corpus", "fineweb-edu-dedup", "train", "text", 700_000_000 - t.numel(), enc, "smollm-fineweb", skip_docs=500_000) t = torch.cat([t, t2]) parts.append(t); total += t.numel() t = _stream_hf("HuggingFaceTB/cosmopedia-v2", "cosmopedia-v2", "train", "text", 250_000_000, enc, "cosmopedia-v2", skip_docs=150_000) parts.append(t); total += t.numel() t = _stream_hf("allenai/c4", "en", "train", "text", 250_000_000, enc, "c4") parts.append(t); total += t.numel() t = _stream_hf("Skylion007/openwebtext", None, "train", "text", 150_000_000, enc, "openwebtext", skip_docs=50_000) parts.append(t); total += t.numel() tokens = torch.cat(parts) print(f"mixture1b: {tokens.numel():,} tokens total " f"({total/1e6:.0f}M collected)") torch.save(tokens, cache_path) print(f" saved to {cache_path}") return cache_path def load_tokens(name: str) -> torch.Tensor: if name == "mixture1b": return torch.load(prepare_mixture1b()) if name == "mixture500m": return torch.load(prepare_mixture500m()) if name == "mixture250m": return torch.load(prepare_mixture250m()) if name == "wikitext103": return torch.load(prepare_wikitext103()) return torch.load(prepare_wikitext2()) if __name__ == "__main__": import sys name = sys.argv[1] if len(sys.argv) > 1 else "wikitext2" if name == "mixture1b": prepare_mixture1b() elif name == "mixture500m": prepare_mixture500m() elif name == "mixture250m": prepare_mixture250m() elif name == "wikitext103": prepare_wikitext103() else: prepare_wikitext2()