Download src/diffusion_lm/data.py from goldenfox/marimo-diffusion: direct link, hf CLI and curl.
- Browser
- Download file 11.8 kB
-
https://huggingface.co/goldenfox/marimo-diffusion/resolve/main/src/diffusion_lm/data.py
- Command line
-
hf download hf://goldenfox/marimo-diffusion/src/diffusion_lm/data.py
-
curl -L -o data.py https://huggingface.co/goldenfox/marimo-diffusion/resolve/main/src/diffusion_lm/data.py
11.8 kB
| """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() | |
| # Reopen mappings inside each DataLoader worker rather than pickling file descriptors. | |
| 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 | |