"""Memory-mapped sentences, loss-preserving windows and bounded-memory shuffle. This module can prepare/check data without importing PyTorch. Each sample has already-shifted labels: the future model must NOT shift these labels again. """ import hashlib import itertools import json import random from concurrent.futures import ThreadPoolExecutor from pathlib import Path import numpy as np from vimeml.tokenizer.store import TokenStore SPLITS = ("train", "validation", "test") IGNORE_INDEX = -100 def file_sha(path): digest = hashlib.sha256() with Path(path).open("rb") as stream: for block in iter(lambda: stream.read(4 * 1024 * 1024), b""): digest.update(block) return digest.hexdigest() def write_json(path, value): path = Path(path) temporary = path.with_suffix(path.suffix + ".tmp") temporary.write_text(json.dumps(value, ensure_ascii=False, indent=2) + "\n", encoding="utf-8") temporary.replace(path) def prepare_split(token_dir, index_dir, split, context_length, manifest_sha, verification="sha256"): """Scan offsets in chunks; only long sentences require index entries.""" with TokenStore(token_dir, split) as store: count = len(store) if count == 0: raise ValueError(f"Empty split: {split}") metadata = store.manifest["splits"][split] # Verify only files consumed by training, in parallel across splits. if verification == "sha256": for suffix in ("tokens.bin", "offsets.bin", "sources.bin"): name = f"{split}.{suffix}" if file_sha(token_dir / name) != store.manifest["output_sha256"][name]: raise ValueError(f"Token data hash mismatch: {name}") if (token_dir / f"{split}.sources.bin").stat().st_size != count: raise ValueError(f"Invalid source index size: {split}") offsets = np.memmap(token_dir / f"{split}.offsets.bin", mode="r", dtype="= self.first[position]: return int(self.sentences[position]), int(index - self.first[position]) * self.context_length previous_extra = 0 if position == 0 else int(self.end[position - 1] - self.sentences[position - 1] - 1) return index - previous_extra, 0 def _open(self): if self._store is None: self._store = TokenStore(self.token_dir, self.split) source_path = self.token_dir / f"{self.split}.sources.bin" if source_path.stat().st_size != len(self._store): self.close() raise ValueError("Invalid source file size.") self._sources = np.memmap(source_path, dtype="u1", mode="r") def __getitem__(self, index): sentence, start = self.locate(index) self._open() tokens = self._store[sentence] end = min(len(tokens) - 1, start + self.context_length) source = int(self._sources[sentence]) if source not in self._store.manifest["source_ids"].values(): raise ValueError("Unknown source ID.") return {"input_ids": tokens[start:end], "labels": tokens[start + 1:end + 1], "sentence_index": sentence, "window_start": start, "source_id": source} def window_lengths(self, indices): """Vectorized lengths from offsets; no token decoding or large length table.""" indices = np.asarray(indices, dtype=np.int64) if np.any(indices < 0) or np.any(indices >= self.count): raise IndexError("Window index outside dataset.") position = np.searchsorted(self.end, indices, side="right") previous_extra = np.zeros_like(indices) has_previous = position > 0 previous = position[has_previous] - 1 previous_extra[has_previous] = self.end[previous] - self.sentences[previous] - 1 sentences = indices - previous_extra starts = np.zeros_like(indices) candidates = np.flatnonzero(position < len(self.end)) long = candidates[indices[candidates] >= self.first[position[candidates]]] sentences[long] = self.sentences[position[long]] starts[long] = (indices[long] - self.first[position[long]]) * self.context_length offsets = np.memmap(self.token_dir / f"{self.split}.offsets.bin", dtype="= self.start_batch: yield [(self.epoch, index) for index in batch] if getattr(self.dataset, "epoch_indexed", False) else batch cursor += 1 def make_loader(dataset, batch_size=128, num_workers=4, seed=42, shuffle=None, pin_memory=False, bucket_multiplier=0, indices=None, start_batch=0, worker_init_fn=None): import torch from torch.utils.data import DataLoader if batch_size < 1 or num_workers < 0 or dataset.context_length % 8: raise ValueError("Positive batch size, nonnegative workers, and context multiple of 8 required.") if shuffle is None: shuffle = dataset.split == "train" if indices is not None and shuffle: raise ValueError("Use explicit indices or shuffle, not both.") if start_batch and not bucket_multiplier: raise ValueError("Resume cursor requires bucketed batching.") if getattr(dataset, "epoch_indexed", False) and not bucket_multiplier: raise ValueError("Prefix crop requires epoch-tagged bucket batching.") sampler = indices if indices is not None else (BlockShuffleSampler(len(dataset), seed) if shuffle else None) dataset._open() pad_id = dataset._store.manifest["special_ids"]["pad"] # Main process need not keep maps open after discovering special IDs. dataset.close() options = {"multiprocessing_context": "spawn", "prefetch_factor": 2} if num_workers else {} if bucket_multiplier: batch_options = {"batch_sampler": LengthBucketBatchSampler( dataset, sampler if sampler is not None else range(len(dataset)), batch_size, bucket_multiplier, seed, start_batch)} else: batch_options = {"batch_size": batch_size, "sampler": sampler, "drop_last": False} return DataLoader(dataset, **batch_options, num_workers=num_workers, collate_fn=TorchCollator(pad_id), pin_memory=pin_memory, persistent_workers=num_workers > 0, generator=torch.Generator().manual_seed(seed), worker_init_fn=worker_init_fn, **options)