spec100m / code /data_pipeline.py
Akahsizrr's picture
squash to reclaim LFS quota
a8f07a3
Raw History Blame Contribute Delete
12.4 kB
"""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()