"""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 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 `/val_.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