#!/usr/bin/env python3 """Stream a mixed public data mix and write GPT-2 token .bin files.""" import argparse import json import os import pickle import shutil import time from dataclasses import dataclass from typing import Optional import numpy as np import tiktoken from datasets import load_dataset from tqdm import tqdm GPT2_VOCAB_SIZE = 50257 @dataclass(frozen=True) class SourceSpec: label: str ratio: float dataset: str name: Optional[str] split: str fields: tuple[str, ...] SOURCES = [ SourceSpec("fineweb", 0.56, "HuggingFaceFW/fineweb-edu", "sample-10BT", "train", ("text",)), SourceSpec("wikipedia", 0.18, "wikimedia/wikipedia", "20231101.en", "train", ("title", "text")), SourceSpec("science", 0.13, "ccdv/arxiv-summarization", None, "train", ("abstract", "article")), SourceSpec("books", 0.13, "common-pile/project_gutenberg", None, "train", ("text",)), ] def parse_token_count(value: Optional[str]) -> Optional[int]: if value is None: return None raw = str(value).strip().replace("_", "").lower() if raw in {"none", "all", "full", "0"}: return None multipliers = {"k": 1_000, "m": 1_000_000, "b": 1_000_000_000} suffix = raw[-1] if suffix in multipliers: return int(float(raw[:-1]) * multipliers[suffix]) return int(raw) def allocate_counts(total: int, ratios: list[float]) -> list[int]: counts = [int(total * r) for r in ratios[:-1]] counts.append(total - sum(counts)) return counts def write_tokens(handle, tokens: list[int]) -> int: if not tokens: return 0 arr = np.asarray(tokens, dtype=np.uint16) arr.tofile(handle) return int(arr.size) def doc_text(doc: dict, fields: tuple[str, ...]) -> str: parts = [] for field in fields: value = doc.get(field) if isinstance(value, str) and value.strip(): parts.append(value.strip()) return "\n\n".join(parts) def load_stream(source: SourceSpec): kwargs = {"split": source.split, "streaming": True} if source.name is None: return load_dataset(source.dataset, **kwargs) return load_dataset(source.dataset, source.name, **kwargs) def main() -> None: parser = argparse.ArgumentParser() parser.add_argument("--out_dir", default="data/mixed10b") parser.add_argument("--max_tokens", default="10B") parser.add_argument("--val_tokens", default="10M") parser.add_argument("--overwrite", action="store_true") args = parser.parse_args() max_tokens = parse_token_count(args.max_tokens) val_tokens = parse_token_count(args.val_tokens) if max_tokens is None: raise ValueError("--max_tokens must be finite for mixed data preparation") if val_tokens is None or val_tokens <= 0: raise ValueError("--val_tokens must be positive") if max_tokens <= val_tokens: raise ValueError("--max_tokens must exceed --val_tokens") os.makedirs(args.out_dir, exist_ok=True) train_path = os.path.join(args.out_dir, "train.bin") val_path = os.path.join(args.out_dir, "val.bin") meta_path = os.path.join(args.out_dir, "meta.pkl") info_path = os.path.join(args.out_dir, "data_info.json") existing = [p for p in (train_path, val_path, meta_path, info_path) if os.path.exists(p)] if existing and not args.overwrite: raise FileExistsError("output exists; pass --overwrite or choose a new --out_dir") train_part = train_path + ".part" val_part = val_path + ".part" for path in (train_part, val_part): if os.path.exists(path): if args.overwrite: os.remove(path) else: raise FileExistsError(f"partial output exists: {path}") enc = tiktoken.get_encoding("gpt2") eot = enc.eot_token ratios = [s.ratio for s in SOURCES] source_totals = allocate_counts(max_tokens, ratios) source_vals = allocate_counts(val_tokens, ratios) source_infos = [] train_written = 0 val_written = 0 start = time.time() pbar = tqdm(total=max_tokens, unit="tok", smoothing=0.05) with open(train_part, "wb") as train_f, open(val_part, "wb") as val_f: for source, source_total, source_val in zip(SOURCES, source_totals, source_vals): source_written = 0 source_train_written = 0 source_val_written = 0 docs_seen = 0 docs_used = 0 print( f"\nStreaming {source.label}: target={source_total:,} " f"val={source_val:,} dataset={source.dataset}" ) for doc in load_stream(source): docs_seen += 1 text = doc_text(doc, source.fields) if not text: continue tokens = [eot] + enc.encode_ordinary(text) remaining = source_total - source_written if remaining <= 0: break if len(tokens) > remaining: tokens = tokens[:remaining] cursor = 0 if source_val_written < source_val: take = min(source_val - source_val_written, len(tokens)) source_val_written += write_tokens(val_f, tokens[:take]) val_written += take cursor = take if cursor < len(tokens): wrote = write_tokens(train_f, tokens[cursor:]) source_train_written += wrote train_written += wrote docs_used += 1 source_written = source_train_written + source_val_written pbar.update(len(tokens)) if source_written >= source_total: break if source_written < source_total: print( f"WARNING: source {source.label} exhausted at {source_written:,} " f"of {source_total:,} tokens" ) source_infos.append({ "label": source.label, "ratio": source.ratio, "dataset": source.dataset, "name": source.name, "split": source.split, "fields": source.fields, "target_tokens": source_total, "target_val_tokens": source_val, "written_tokens": source_written, "train_tokens_written": source_train_written, "val_tokens_written": source_val_written, "docs_seen": docs_seen, "docs_used": docs_used, }) pbar.close() os.replace(train_part, train_path) os.replace(val_part, val_path) with open(meta_path, "wb") as f: pickle.dump({"vocab_size": GPT2_VOCAB_SIZE}, f) info = { "vocab_size": GPT2_VOCAB_SIZE, "tokenizer": "tiktoken:gpt2", "max_tokens": max_tokens, "val_tokens_requested": val_tokens, "train_tokens_written": train_written, "val_tokens_written": val_written, "sources": source_infos, "elapsed_sec": round(time.time() - start, 2), } with open(info_path, "w", encoding="utf-8") as f: json.dump(info, f, indent=2) total_size = os.path.getsize(train_path) + os.path.getsize(val_path) print("\nFinished mixed data preparation") print(f" train tokens: {train_written:,} -> {train_path}") print(f" val tokens: {val_written:,} -> {val_path}") print(f" disk size: {total_size / 1024**3:.2f} GiB") print(f" info: {info_path}") if shutil.disk_usage(args.out_dir).free < 5 * 1024**3: print("WARNING: less than 5 GiB free space remains in the output directory.") if __name__ == "__main__": main()