tolivert's picture
deploy: financial_rag streamlit app
4e316d6
Raw
History Blame Contribute Delete
7.58 kB
"""
Memmap-backed token dataset for next-token prediction pretraining.
Why memmap?
===========
A pretraining corpus is one long sequence of token ids. The naive
approach (load the whole thing into a Python list and slice) wastes
memory and is slow to copy onto the GPU. The standard pattern, used
by nanoGPT and basically every educational pretraining repo, is:
1. Tokenize once at preparation time and write the result to a binary
file as raw uint16 little-endian (for vocab <= 65535).
2. At training time, ``np.memmap`` the file. The OS handles paging,
so the dataset behaves as if fully in RAM but only the slices you
actually touch are read from disk.
3. For each training step, sample ``batch_size`` random offsets into
the memmap and read ``block_size + 1`` tokens at each one. Split
into ``x = tokens[:-1]`` (input) and ``y = tokens[1:]`` (target)
for next-token prediction.
This is fast (sequential reads), trivial to implement (~30 lines),
scales to corpora that don't fit in RAM, and has zero
preprocessing-at-train-time overhead.
uint16 vs int32
===============
uint16 covers vocab sizes up to 65,535, which is ample for our 8k
custom BPE. If you ever need a larger vocab, switch to uint32 — but
the file size doubles, and you almost certainly don't need a larger
vocab for from-scratch tiny pretraining (the whole point is keeping
the embedding small).
"""
from pathlib import Path
from typing import Protocol, Tuple, Union, runtime_checkable
import numpy as np
import torch
@runtime_checkable
class TokenSource(Protocol):
"""Anything the trainer can pull next-token-prediction batches from.
Implementing this is the *only* thing required to plug a new corpus
into the pretraining loop. Concretely, ``get_batch`` must return a
pair of ``torch.long`` tensors of shape ``(batch_size, block_size)``
such that ``y[:, t] == x[:, t+1]`` for some implicit underlying token
stream — i.e. the standard shifted-by-one next-token target.
The trainer treats sources as opaque: no iteration order, no epoch
boundaries, no shuffling state. It just calls ``get_batch`` once per
train step and once per eval step. If you need a streaming HF
dataset, a multi-file shard, or anything fancier, hide it behind
this method.
The optional ``generator`` argument is used to make train batch
sampling reproducible across runs; eval calls pass ``None`` and
rely on the global RNG, which is fine because eval is a noisy
average over many batches anyway.
"""
def get_batch(
self,
batch_size: int,
device: Union[str, torch.device] = "cpu",
generator: Union[torch.Generator, None] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
...
# Token dtype for the on-disk binary. Must match what
# scripts/prepare_tinystories.py writes.
TOKEN_DTYPE = np.uint16
class TokenDataset:
"""
Random-access view over a memmapped token binary.
Not a ``torch.utils.data.Dataset`` — for from-scratch pretraining
we don't need iteration order, multi-worker shuffling, or any of
the DataLoader machinery. ``get_batch`` directly samples random
offsets, which is simpler and faster.
Args:
bin_path: Path to a uint16 binary written by the preparation
script.
block_size: Number of *input* tokens per sample. Each sample
actually reads ``block_size + 1`` tokens so the target can
be the same window shifted by one.
"""
def __init__(self, bin_path: Union[str, Path], block_size: int):
self.bin_path = Path(bin_path)
self.block_size = block_size
if not self.bin_path.exists():
raise FileNotFoundError(
f"Token binary not found at {self.bin_path}. Run "
f"scripts/prepare_tinystories.py first."
)
# mode='r' = read-only memmap. The OS will lazily page in
# whatever bytes we touch.
self.tokens = np.memmap(self.bin_path, dtype=TOKEN_DTYPE, mode="r")
if len(self.tokens) < block_size + 1:
raise ValueError(
f"Token binary at {self.bin_path} has only {len(self.tokens)} "
f"tokens, but block_size+1 = {block_size + 1} are required "
f"for a single sample."
)
def __len__(self) -> int:
"""Number of valid starting positions for a (block_size+1)-token window."""
return len(self.tokens) - self.block_size
def get_batch(
self,
batch_size: int,
device: Union[str, torch.device] = "cpu",
generator: Union[torch.Generator, None] = None,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Sample a random batch of (input, target) sequences.
Args:
batch_size: Number of independent sequences in the batch.
device: Device to copy the tensors to. For MPS / CUDA the
copy is asynchronous and overlaps with the next forward.
generator: Optional ``torch.Generator`` for reproducible
sampling. If ``None``, uses the default RNG.
Returns:
``(x, y)`` where each is shape ``(batch_size, block_size)``
of dtype ``torch.long``. ``y`` is ``x`` shifted by one
position to the right (the standard next-token target).
"""
max_start = len(self) - 1 # inclusive max for randint
# torch.randint is faster than np.random.randint for our needs
# and integrates with torch.Generator for reproducibility.
starts = torch.randint(
low=0,
high=max_start + 1,
size=(batch_size,),
generator=generator,
)
# Build the batch on CPU as int64 (the dtype embeddings expect),
# then ship to the target device in one transfer.
x = torch.empty((batch_size, self.block_size), dtype=torch.long)
y = torch.empty((batch_size, self.block_size), dtype=torch.long)
for i, start in enumerate(starts.tolist()):
window = self.tokens[start : start + self.block_size + 1]
# .astype(np.int64) realizes the slice into RAM (the memmap
# slice is a uint16 view); from_numpy then shares the buffer.
window = window.astype(np.int64, copy=False)
x[i] = torch.from_numpy(window[:-1])
y[i] = torch.from_numpy(window[1:])
if str(device) != "cpu":
x = x.to(device, non_blocking=True)
y = y.to(device, non_blocking=True)
return x, y
def write_tokens_bin(tokens: np.ndarray, output_path: Union[str, Path]) -> None:
"""
Persist a 1-D token id array to disk in the format ``TokenDataset``
expects.
Args:
tokens: 1-D numpy array of token ids. Will be cast to
``TOKEN_DTYPE`` — caller is responsible for ensuring all
ids fit (i.e. ``tokens.max() < 65536``).
output_path: Destination ``.bin`` file. Parent directory is
created if needed.
"""
output_path = Path(output_path)
output_path.parent.mkdir(parents=True, exist_ok=True)
if tokens.ndim != 1:
raise ValueError(f"Expected 1-D token array, got shape {tokens.shape}")
if tokens.size > 0 and tokens.max() >= np.iinfo(TOKEN_DTYPE).max + 1:
raise ValueError(
f"Token id {int(tokens.max())} does not fit in {TOKEN_DTYPE}. "
f"Use a smaller vocabulary or switch to uint32."
)
tokens.astype(TOKEN_DTYPE, copy=False).tofile(output_path)