Ines-1 / mini_v41 /data.py
Endikavi's picture
Ines-1 RC1 (private staging; release commit b7f5644)
61b6fb9 verified
Raw History Blame Contribute Delete
16.6 kB
"""Data pipeline: local files -> tokenized shards -> packed sequences.
Design goals, in order: reproducibility, resumability, then throughput.
A run is reproducible because (a) shards are content-addressed by a manifest
containing the tokenizer hash and packing parameters, and (b) the sampler order
is a deterministic function of (seed, epoch) with the position recorded in the
checkpoint, so a resume replays the identical token stream.
"""
from __future__ import annotations
import hashlib
import json
from dataclasses import asdict, dataclass
from pathlib import Path
import numpy as np
import torch
from .config import DataConfig
from .tokenizer import Tokenizer
__all__ = ["build_shards", "PackedDataset", "ShardedSampler", "make_dataloaders"]
def _iter_texts(source: Path, text_key: str):
"""Yield documents from .txt (whole file) and .jsonl (one field per line)."""
paths = sorted(source.rglob("*")) if source.is_dir() else [source]
for path in paths:
if path.suffix == ".txt":
yield path.read_text(encoding="utf-8", errors="replace")
elif path.suffix in (".jsonl", ".ndjson"):
with path.open(encoding="utf-8", errors="replace") as fh:
for line in fh:
line = line.strip()
if not line:
continue
try:
obj = json.loads(line)
except json.JSONDecodeError:
continue
text = obj.get(text_key) if isinstance(obj, dict) else obj
if isinstance(text, str) and text:
yield text
# Bump when the packing scheme changes, so stale shards are rebuilt rather than
# silently reused (the manifest otherwise matches on sequence_length alone).
PACKING_VERSION = 2
@dataclass
class ShardManifest:
packing_version: int
tokenizer_sha256: str
vocab_size: int
sequence_length: int
dtype: str
n_train_sequences: int
n_val_sequences: int
n_tokens: int
n_documents: int
val_fraction: float
seed: int
source_fingerprint: str
def matches(self, other: ShardManifest) -> bool:
"""Everything except the realised counts must agree for a cache hit."""
keys = (
"packing_version", "tokenizer_sha256", "vocab_size", "sequence_length",
"dtype", "val_fraction", "seed", "source_fingerprint",
)
return all(getattr(self, k) == getattr(other, k) for k in keys)
def _source_fingerprint(source: Path) -> str:
"""Hash of (relative path, size, mtime) over the corpus. Cheap; catches
edits without re-reading gigabytes."""
h = hashlib.sha256()
paths = sorted(source.rglob("*")) if source.is_dir() else [source]
for p in paths:
if p.is_file():
st = p.stat()
h.update(str(p.relative_to(source) if source.is_dir() else p.name).encode())
h.update(f"{st.st_size}:{int(st.st_mtime)}".encode())
return h.hexdigest()[:16]
def build_shards(cfg: DataConfig, tokenizer: Tokenizer, verbose: bool = True) -> ShardManifest:
"""Tokenize + pack the corpus into `train.bin` / `val.bin`.
Packing is the standard concatenate-with-EOS-then-chunk scheme: documents are
joined by <eos> into one stream and cut into windows of exactly
`sequence_length` tokens.
A window is used as both input and labels; the model shifts internally
(`logits[:, :-1]` against `labels[:, 1:]`), so one window of L tokens yields
L-1 next-token targets. We deliberately do NOT cut `L+1`-token windows to
recover that last target: it would make the model process L+1 positions,
so `sequence_length: 2048` with `max_sequence_length: 2048` would overrun
RoPE by one. Losing 1 target in 2048 is cheaper than that off-by-one trap.
"""
source = Path(cfg.source)
cache = Path(cfg.cache_dir)
cache.mkdir(parents=True, exist_ok=True)
if not source.exists():
raise FileNotFoundError(f"data.source {source} does not exist")
dtype = "uint16" if tokenizer.vocab_size < 2**16 else "uint32"
manifest_path = cache / "manifest.json"
want = ShardManifest(
packing_version=PACKING_VERSION,
tokenizer_sha256=tokenizer.identity()["sha256"],
vocab_size=tokenizer.vocab_size,
sequence_length=cfg.sequence_length,
dtype=dtype,
n_train_sequences=0,
n_val_sequences=0,
n_tokens=0,
n_documents=0,
val_fraction=cfg.val_fraction,
seed=cfg.seed,
source_fingerprint=_source_fingerprint(source),
)
if manifest_path.exists() and not cfg.force_retokenize:
try:
have = ShardManifest(**json.loads(manifest_path.read_text()))
except (TypeError, ValueError, json.JSONDecodeError) as exc:
# A manifest written by an older packing scheme: treat it as a cache
# miss and rebuild rather than crashing. Never reuse shards we cannot
# fully validate -- a silently mismatched packing is worse than the
# minutes it costs to retokenize.
if verbose:
print(f"[data] unreadable manifest ({exc}); rebuilding shards")
else:
if have.matches(want) and (cache / "train.bin").exists():
if verbose:
print(f"[data] cache hit: {have.n_train_sequences} train / "
f"{have.n_val_sequences} val sequences")
return have
if verbose:
print(f"[data] tokenizing {source} ...")
np_dtype = np.uint16 if dtype == "uint16" else np.uint32
stream: list[int] = []
n_docs = 0
for text in _iter_texts(source, cfg.text_key):
ids = tokenizer.encode(text)
if not ids:
continue
stream.extend(ids)
stream.append(tokenizer.eos_id)
n_docs += 1
if not stream:
raise ValueError(f"no text found under {source}")
window = cfg.sequence_length
n_seq = len(stream) // window
if n_seq < 2:
raise ValueError(
f"corpus yields only {n_seq} sequences of length {window}; "
f"lower data.sequence_length or add more text"
)
arr = np.asarray(stream[: n_seq * window], dtype=np_dtype).reshape(n_seq, window)
# Deterministic split on a dedicated RNG, so changing train seed does not
# reshuffle which sequences are held out.
rng = np.random.default_rng(cfg.seed)
perm = rng.permutation(n_seq)
n_val = max(1, int(round(n_seq * cfg.val_fraction))) if cfg.val_fraction > 0 else 0
n_val = min(n_val, n_seq - 1)
val_idx, train_idx = perm[:n_val], perm[n_val:]
arr[train_idx].tofile(cache / "train.bin")
if n_val:
arr[val_idx].tofile(cache / "val.bin")
want.n_train_sequences = int(len(train_idx))
want.n_val_sequences = int(n_val)
want.n_tokens = int(len(stream))
want.n_documents = n_docs
manifest_path.write_text(json.dumps(asdict(want), indent=2), encoding="utf-8")
if verbose:
print(f"[data] {n_docs} docs -> {len(stream)} tokens -> "
f"{want.n_train_sequences} train / {want.n_val_sequences} val sequences")
return want
class PackedDataset(torch.utils.data.Dataset):
"""Memory-mapped view over a packed shard. Item i is [sequence_length] int64."""
def __init__(self, path: str | Path, sequence_length: int, dtype: str = "uint16"):
self.path = Path(path)
self.window = sequence_length
np_dtype = np.uint16 if dtype == "uint16" else np.uint32
self.data = np.memmap(self.path, dtype=np_dtype, mode="r")
if self.data.size % self.window != 0:
raise ValueError(f"{path} size {self.data.size} is not a multiple of {self.window}")
self.data = self.data.reshape(-1, self.window)
def __len__(self) -> int:
return self.data.shape[0]
def __getitem__(self, idx: int) -> torch.Tensor:
return torch.from_numpy(self.data[idx].astype(np.int64))
class ShardedSampler:
"""Deterministic, resumable index stream.
`state_dict()` / `load_state_dict()` capture (epoch, offset) so a resumed run
continues from the exact sequence it was about to consume -- the
"dataset/shard position" the checkpoint spec asks for.
"""
TERMINAL_POLICIES = ("partial", "wrap_v1")
def __init__(self, n_items: int, batch_size: int, seed: int = 0,
shuffle: bool = True, terminal_policy: str = "partial"):
"""What happens at the end of an epoch is a POLICY, and it is recorded.
`partial` (default): when fewer than `batch_size` items remain, hand
back the short remainder and only then start the next epoch. Every
sequence is consumed exactly once per epoch, and no global batch ever
mixes two epochs.
`wrap_v1`: the original behaviour, kept only so `1p6b-pretrain-v1` can
be replayed. When fewer than `batch_size` items remain it skips them
and wraps to a fresh permutation mid-batch. With an odd `n_items` and
`batch_size` 2 that silently drops ONE sequence per epoch and starts
replaying already-seen sequences to fill the batch. In the 1B run the
dropped sequence was id 441274 and 20 sequences were replayed -- see
docs/TERMINAL_TOKEN_AUDIT.md. Do not choose this for new runs.
"""
if terminal_policy not in self.TERMINAL_POLICIES:
raise ValueError(
f"unknown terminal_policy {terminal_policy!r}; "
f"expected one of {self.TERMINAL_POLICIES}"
)
self.n_items = n_items
self.batch_size = batch_size
self.seed = seed
self.shuffle = shuffle
self.terminal_policy = terminal_policy
self.epoch = 0
self.offset = 0
self._order = self._make_order()
def _make_order(self) -> np.ndarray:
if not self.shuffle:
return np.arange(self.n_items)
# Order depends only on (seed, epoch): replayable from a checkpoint
# without storing the permutation itself.
return np.random.default_rng([self.seed, self.epoch]).permutation(self.n_items)
def sequences_per_epoch(self) -> int:
"""How many sequences one epoch actually delivers under this policy.
Not always `n_items`. A target derived from the dataset rather than
from this number can be unreachable: the 1B run set its terminal
milestone at the content of all 490,765 sequences while `wrap_v1` could
only ever deliver 490,764, so the run had to cross into a second epoch
to satisfy it.
"""
if self.terminal_policy == "wrap_v1":
return self.n_items - (self.n_items % self.batch_size)
return self.n_items
def epoch_order(self) -> np.ndarray:
"""The sequence ids this epoch will deliver, in order."""
return self._order[: self.sequences_per_epoch()]
def next_batch_indices(self) -> np.ndarray:
if self.terminal_policy == "wrap_v1":
if self.offset + self.batch_size > self.n_items:
self.epoch += 1
self.offset = 0
self._order = self._make_order()
idx = self._order[self.offset : self.offset + self.batch_size]
self.offset += self.batch_size
return idx
# "partial": never skip the tail, never mix epochs inside one batch.
if self.offset >= self.n_items:
self.epoch += 1
self.offset = 0
self._order = self._make_order()
idx = self._order[self.offset : self.offset + self.batch_size]
self.offset += len(idx)
return idx
def rewind(self) -> None:
"""Return to the start of the current order.
Used before each validation pass so every evaluation covers the same
sequences. Without it a validation number is not comparable to the
previous one, because it was measured on different data.
"""
self.offset = 0
self._order = self._make_order()
def state_dict(self) -> dict:
return {"epoch": self.epoch, "offset": self.offset, "seed": self.seed,
"n_items": self.n_items, "batch_size": self.batch_size,
"terminal_policy": self.terminal_policy}
def load_state_dict(self, state: dict) -> None:
if state.get("n_items") != self.n_items:
raise ValueError(
f"dataset changed size ({state.get('n_items')} -> {self.n_items}); "
f"a resume would replay a different token stream"
)
# Checkpoints written before the policy existed carry no key. They were
# all produced by the wrapping sampler, so that is what they mean --
# assuming the current default would resume them onto a different
# stream while reporting nothing.
saved_policy = state.get("terminal_policy", "wrap_v1")
if saved_policy != self.terminal_policy:
raise ValueError(
f"terminal_policy changed ({saved_policy} -> "
f"{self.terminal_policy}); the end of every epoch would behave "
f"differently, so this is a different token stream, not a resume"
)
saved_bs = state.get("batch_size", self.batch_size)
if saved_bs != self.batch_size:
raise ValueError(
f"micro_batch_size changed ({saved_bs} -> {self.batch_size}); "
f"the same sequence order would be regrouped into different "
f"optimizer steps, which is a different run, not a resume"
)
self.epoch = state["epoch"]
self.offset = state["offset"]
self.seed = state["seed"]
self._order = self._make_order()
class BatchStream:
"""Pairs a PackedDataset with a ShardedSampler and yields (inputs, labels)."""
def __init__(self, dataset: PackedDataset, sampler: ShardedSampler, device: str = "cpu"):
self.dataset = dataset
self.sampler = sampler
self.device = device
def next(self) -> tuple[torch.Tensor, torch.Tensor]:
idx = self.sampler.next_batch_indices()
rows = np.stack([self.dataset.data[i] for i in idx]).astype(np.int64)
batch = torch.from_numpy(rows)
# Both are the full window: the model shifts internally (logits[:, :-1]
# vs labels[:, 1:]), so one tensor of seq_len+1 serves as both.
batch = batch.to(self.device, non_blocking=True)
return batch, batch
def state_dict(self) -> dict:
return self.sampler.state_dict()
def load_state_dict(self, state: dict) -> None:
self.sampler.load_state_dict(state)
def make_domain_val_streams(
cfg: DataConfig, manifest: ShardManifest, batch_size: int, device: str
) -> dict:
"""One validation stream per domain, in addition to the global one.
Domains come from `data.val_domains` when the data agent supplies explicit
paths; otherwise `<cache_dir>/val_<domain>.bin` is auto-discovered, which is
the convention the packing stage writes. Non-shuffled, so a domain's
validation number is comparable across checkpoints.
"""
cache = Path(cfg.cache_dir)
found: dict[str, Path] = {}
for domain, path in (cfg.val_domains or {}).items():
found[domain] = Path(path)
if not found:
for candidate in sorted(cache.glob("val_*.bin")):
found[candidate.stem[len("val_"):]] = candidate
streams = {}
for domain, path in found.items():
if not path.exists():
print(f"[data] domain validation shard missing, skipping: {path}")
continue
ds = PackedDataset(path, cfg.sequence_length, manifest.dtype)
streams[domain] = BatchStream(
ds, ShardedSampler(len(ds), batch_size, cfg.seed, shuffle=False), device
)
return streams
def make_dataloaders(cfg: DataConfig, manifest: ShardManifest, batch_size: int, device: str):
cache = Path(cfg.cache_dir)
train_ds = PackedDataset(cache / "train.bin", cfg.sequence_length, manifest.dtype)
train = BatchStream(train_ds, ShardedSampler(len(train_ds), batch_size, cfg.seed), device)
val = None
if manifest.n_val_sequences and (cache / "val.bin").exists():
val_ds = PackedDataset(cache / "val.bin", cfg.sequence_length, manifest.dtype)
val = BatchStream(
val_ds, ShardedSampler(len(val_ds), batch_size, cfg.seed, shuffle=False), device
)
return train, val