Download mini_v41/data.py from Endikavi/Ines-1: direct link, hf CLI and curl.
- Browser
- Download file 16.6 kB
-
https://huggingface.co/Endikavi/Ines-1/resolve/main/mini_v41/data.py
- Command line
-
hf download hf://Endikavi/Ines-1/mini_v41/data.py
-
curl -L -o data.py https://huggingface.co/Endikavi/Ines-1/resolve/main/mini_v41/data.py
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 | |
| 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 | |