| """Memory-mapped packed-token datasets and deterministic sampling.""" |
|
|
| from __future__ import annotations |
|
|
| import json |
| import math |
| from bisect import bisect_right |
| from pathlib import Path |
| from typing import Any |
|
|
| import numpy as np |
| import torch |
| from torch import Tensor |
| from torch.utils.data import Dataset, Sampler |
|
|
| from diffusion_lm.tokenizer import token_metadata_path |
|
|
|
|
| PACKED_TOKEN_FORMAT = "mini-diffusion-lm-packed-tokens-v1" |
| PACKED_MANIFEST_FORMAT = "mini-diffusion-lm-packed-manifest-v1" |
|
|
|
|
| def _read_json(path: Path) -> dict[str, Any]: |
| try: |
| with path.open("r", encoding="utf-8") as handle: |
| value = json.load(handle) |
| except json.JSONDecodeError as exc: |
| raise ValueError(f"invalid JSON in {path}: {exc}") from exc |
| if not isinstance(value, dict): |
| raise ValueError(f"expected a JSON object in {path}") |
| return value |
|
|
|
|
| def load_token_metadata(path: str | Path) -> dict[str, Any]: |
| metadata_path = token_metadata_path(path) |
| if not metadata_path.is_file(): |
| raise FileNotFoundError( |
| f"token metadata not found: {metadata_path}; encode data with mini-mdlm-tokenizer" |
| ) |
| metadata = _read_json(metadata_path) |
| if metadata.get("format") != PACKED_TOKEN_FORMAT: |
| raise ValueError(f"unsupported token file metadata in {metadata_path}") |
| return metadata |
|
|
|
|
| def load_packed_manifest(path: str | Path) -> dict[str, Any]: |
| manifest_path = Path(path) |
| if not manifest_path.is_file(): |
| raise FileNotFoundError(f"packed-token manifest not found: {manifest_path}") |
| manifest = _read_json(manifest_path) |
| if manifest.get("format") != PACKED_MANIFEST_FORMAT: |
| raise ValueError(f"unsupported packed-token manifest in {manifest_path}") |
| shards = manifest.get("shards") |
| if not isinstance(shards, list) or not shards: |
| raise ValueError(f"packed-token manifest has no shards: {manifest_path}") |
| return manifest |
|
|
|
|
| def _validate_dtype(name: object, *, context: Path) -> np.dtype[Any]: |
| try: |
| dtype = np.dtype(name) |
| except TypeError as exc: |
| raise ValueError(f"invalid token dtype {name!r} in {context}") from exc |
| if dtype not in (np.dtype("uint16"), np.dtype("uint32")): |
| raise ValueError(f"unsupported token dtype {dtype.name!r} in {context}") |
| return dtype |
|
|
|
|
| class PackedTokenDataset(Dataset[Tensor]): |
| """Expose deterministic non-overlapping blocks from one packed token file.""" |
|
|
| def __init__(self, path: str | Path, sequence_length: int) -> None: |
| if sequence_length <= 0: |
| raise ValueError("sequence_length must be positive") |
| self.path = Path(path) |
| if not self.path.is_file(): |
| raise FileNotFoundError(f"packed token file not found: {self.path}") |
| self.metadata = load_token_metadata(self.path) |
| self.sequence_length = sequence_length |
| self._dtype = _validate_dtype(self.metadata.get("dtype"), context=self.path) |
| self._tokens = np.memmap(self.path, mode="r", dtype=self._dtype) |
| expected_count = int(self.metadata["token_count"]) |
| if self._tokens.size != expected_count: |
| raise ValueError( |
| f"metadata says {expected_count} tokens but {self.path} contains " |
| f"{self._tokens.size}" |
| ) |
| self._blocks = self._tokens.size // sequence_length |
| if self._blocks == 0: |
| raise ValueError( |
| f"dataset has {self._tokens.size} tokens, fewer than one " |
| f"{sequence_length}-token block" |
| ) |
|
|
| def __len__(self) -> int: |
| return self._blocks |
|
|
| def __getitem__(self, index: int) -> Tensor: |
| if index < 0: |
| index += self._blocks |
| if not 0 <= index < self._blocks: |
| raise IndexError(index) |
| start = index * self.sequence_length |
| block = np.asarray(self._tokens[start : start + self.sequence_length]).astype( |
| np.int64, copy=True |
| ) |
| return torch.from_numpy(block) |
|
|
|
|
| class ManifestPackedTokenDataset(Dataset[Tensor]): |
| """Expose one logical dataset backed by independently memory-mapped token shards. |
| |
| Blocks never cross shard boundaries. At most ``sequence_length - 1`` trailing tokens per shard |
| are ignored, keeping source shards independently replaceable and resumable. |
| """ |
|
|
| def __init__(self, path: str | Path, sequence_length: int) -> None: |
| if sequence_length <= 0: |
| raise ValueError("sequence_length must be positive") |
| self.path = Path(path) |
| self.metadata = load_packed_manifest(self.path) |
| self.sequence_length = sequence_length |
| self._dtype = _validate_dtype(self.metadata.get("dtype"), context=self.path) |
|
|
| declared_token_count = int(self.metadata.get("token_count", -1)) |
| declared_document_count = int(self.metadata.get("document_count", -1)) |
| shard_token_count = 0 |
| shard_document_count = 0 |
| self._shard_paths: list[Path] = [] |
| self._shard_token_counts: list[int] = [] |
| self._block_ends: list[int] = [] |
| total_blocks = 0 |
|
|
| for position, raw_shard in enumerate(self.metadata["shards"]): |
| if not isinstance(raw_shard, dict): |
| raise ValueError(f"shard {position} in {self.path} is not an object") |
| raw_path = raw_shard.get("path") |
| if not isinstance(raw_path, str) or not raw_path: |
| raise ValueError(f"shard {position} in {self.path} has no path") |
| shard_path = Path(raw_path) |
| if not shard_path.is_absolute(): |
| shard_path = self.path.parent / shard_path |
| if not shard_path.is_file(): |
| raise FileNotFoundError(f"packed token shard not found: {shard_path}") |
|
|
| token_count = int(raw_shard.get("token_count", -1)) |
| document_count = int(raw_shard.get("document_count", -1)) |
| if token_count < 0 or document_count < 0: |
| raise ValueError(f"invalid counts for shard {shard_path}") |
| expected_bytes = token_count * self._dtype.itemsize |
| if shard_path.stat().st_size != expected_bytes: |
| raise ValueError( |
| f"manifest says {token_count} tokens but {shard_path} has " |
| f"{shard_path.stat().st_size} bytes" |
| ) |
|
|
| shard_token_count += token_count |
| shard_document_count += document_count |
| blocks = token_count // sequence_length |
| if blocks: |
| self._shard_paths.append(shard_path) |
| self._shard_token_counts.append(token_count) |
| total_blocks += blocks |
| self._block_ends.append(total_blocks) |
|
|
| if shard_token_count != declared_token_count: |
| raise ValueError( |
| f"manifest token_count is {declared_token_count}, shard total is " |
| f"{shard_token_count}" |
| ) |
| if shard_document_count != declared_document_count: |
| raise ValueError( |
| f"manifest document_count is {declared_document_count}, shard total is " |
| f"{shard_document_count}" |
| ) |
| if total_blocks == 0: |
| raise ValueError( |
| f"dataset has no shard containing a full {sequence_length}-token block" |
| ) |
| self._blocks = total_blocks |
| self._maps: list[np.memmap[Any, Any] | None] = [None] * len(self._shard_paths) |
|
|
| def __len__(self) -> int: |
| return self._blocks |
|
|
| def _map(self, shard_index: int) -> np.memmap[Any, Any]: |
| tokens = self._maps[shard_index] |
| if tokens is None: |
| tokens = np.memmap(self._shard_paths[shard_index], mode="r", dtype=self._dtype) |
| self._maps[shard_index] = tokens |
| return tokens |
|
|
| def __getitem__(self, index: int) -> Tensor: |
| if index < 0: |
| index += self._blocks |
| if not 0 <= index < self._blocks: |
| raise IndexError(index) |
| shard_index = bisect_right(self._block_ends, index) |
| previous_end = 0 if shard_index == 0 else self._block_ends[shard_index - 1] |
| local_block = index - previous_end |
| start = local_block * self.sequence_length |
| tokens = self._map(shard_index) |
| block = np.asarray(tokens[start : start + self.sequence_length]).astype( |
| np.int64, copy=True |
| ) |
| return torch.from_numpy(block) |
|
|
| def __getstate__(self) -> dict[str, Any]: |
| state = self.__dict__.copy() |
| |
| state["_maps"] = [None] * len(self._shard_paths) |
| return state |
|
|
|
|
| PackedDataset = PackedTokenDataset | ManifestPackedTokenDataset |
|
|
|
|
| def load_packed_dataset(path: str | Path, sequence_length: int) -> PackedDataset: |
| """Load a legacy single-file dataset or a manifest-backed sharded dataset.""" |
|
|
| candidate = Path(path) |
| if candidate.suffix == ".json" and candidate.is_file(): |
| value = _read_json(candidate) |
| if value.get("format") == PACKED_MANIFEST_FORMAT: |
| return ManifestPackedTokenDataset(candidate, sequence_length) |
| return PackedTokenDataset(candidate, sequence_length) |
|
|
|
|
| _UINT64_MASK = (1 << 64) - 1 |
|
|
|
|
| def _splitmix64(value: int) -> int: |
| value = (value + 0x9E3779B97F4A7C15) & _UINT64_MASK |
| value = ((value ^ (value >> 30)) * 0xBF58476D1CE4E5B9) & _UINT64_MASK |
| value = ((value ^ (value >> 27)) * 0x94D049BB133111EB) & _UINT64_MASK |
| return value ^ (value >> 31) |
|
|
|
|
| def _affine_permutation_parameters(size: int, seed: int, epoch: int) -> tuple[int, int]: |
| """Return ``a, b`` for the bijection ``(a*x+b) mod size``.""" |
|
|
| if size == 1: |
| return 0, 0 |
| mixed = _splitmix64((seed & _UINT64_MASK) ^ _splitmix64(epoch & _UINT64_MASK)) |
| offset = mixed % size |
| multiplier = _splitmix64(mixed) % size |
| if multiplier == 0: |
| multiplier = 1 |
| while math.gcd(multiplier, size) != 1: |
| multiplier += 1 |
| if multiplier == size: |
| multiplier = 1 |
| return multiplier, offset |
|
|
|
|
| class DeterministicBatchSampler(Sampler[list[int]]): |
| """Infinite epoch permutations with exact batch-cursor resume and O(1) memory. |
| |
| Each epoch uses a seeded affine bijection over dataset indexes. Unlike ``randperm().tolist()``, |
| memory use is independent of corpus size and a resumed batch can be calculated directly. |
| """ |
|
|
| def __init__( |
| self, |
| dataset_size: int, |
| batch_size: int, |
| *, |
| seed: int, |
| start_batch: int = 0, |
| ) -> None: |
| if dataset_size <= 0 or batch_size <= 0: |
| raise ValueError("dataset_size and batch_size must be positive") |
| if start_batch < 0: |
| raise ValueError("start_batch must be non-negative") |
| self.dataset_size = dataset_size |
| self.batch_size = batch_size |
| self.seed = seed |
| self.start_batch = start_batch |
| self.batches_per_epoch = (dataset_size + batch_size - 1) // batch_size |
|
|
| def __iter__(self): |
| epoch = self.start_batch // self.batches_per_epoch |
| batch_in_epoch = self.start_batch % self.batches_per_epoch |
| while True: |
| multiplier, offset = _affine_permutation_parameters( |
| self.dataset_size, self.seed, epoch |
| ) |
| for batch_index in range(batch_in_epoch, self.batches_per_epoch): |
| start = batch_index * self.batch_size |
| stop = min(start + self.batch_size, self.dataset_size) |
| yield [ |
| (multiplier * position + offset) % self.dataset_size |
| for position in range(start, stop) |
| ] |
| epoch += 1 |
| batch_in_epoch = 0 |
|
|