"""Data plumbing for stage-1 training. * StepDataset : uint16 [N,768] ids + mask_start + len, length-bucketed so each batch pads only to the longest member (fixed 768 padding wastes 77% of the compute). * LMDataset : random ctx windows out of the packed mathlib token stream. """ from __future__ import annotations import json import os import random import numpy as np import torch from torch.utils.data import Dataset class StepDataset(Dataset): def __init__(self, root: str, split: str, ctx: int = 768, pad_id: int = 0, bucket: int = 512, seed: int = 0): self.ids = np.load(f'{root}/{split}_ids.npy', mmap_mode='r') self.mask_start = np.load(f'{root}/{split}_mask_start.npy') self.lens = np.load(f'{root}/{split}_len.npy') self.ctx = ctx self.pad_id = pad_id self.bucket = bucket self.seed = seed self.epoch = 0 self.order = None self._reorder() def _reorder(self): """Sort by length into buckets, shuffle bucket order + inside buckets.""" rng = random.Random(self.seed + self.epoch) idx = np.argsort(self.lens, kind='stable') buckets = [idx[i:i + self.bucket] for i in range(0, len(idx), self.bucket)] for b in buckets: rng.shuffle(b.tolist()) rng.shuffle(buckets) self.order = np.concatenate(buckets) if buckets else idx def set_epoch(self, epoch: int): self.epoch = epoch self._reorder() def __len__(self): return len(self.ids) def __getitem__(self, i): j = int(self.order[i]) n = int(self.lens[j]) ids = np.asarray(self.ids[j, :n], dtype=np.int64) return ids, int(self.mask_start[j]) def collate_steps(batch, pad_id: int = 0, pad_multiple: int = 8): lens = [len(b[0]) for b in batch] m = max(lens) m = min(((m + pad_multiple - 1) // pad_multiple) * pad_multiple, max(lens)) m = ((m + pad_multiple - 1) // pad_multiple) * pad_multiple ids = np.full((len(batch), m), pad_id, dtype=np.int64) mask_start = np.zeros(len(batch), dtype=np.int64) for k, (seq, ms) in enumerate(batch): ids[k, :len(seq)] = seq mask_start[k] = min(ms, m) return (torch.from_numpy(ids), torch.from_numpy(mask_start)) class LMDataset(Dataset): def __init__(self, path: str, ctx: int = 768, length: int = 100_000, seed: int = 0): self.tokens = np.load(path, mmap_mode='r') self.ctx = ctx self.length = length self.seed = seed def __len__(self): return self.length def __getitem__(self, i): rng = random.Random(self.seed * 1_000_003 + i) s = rng.randrange(0, len(self.tokens) - self.ctx - 1) seq = np.asarray(self.tokens[s:s + self.ctx + 1], dtype=np.int64) return seq[:-1], seq[1:] def load_specials(tok_dir: str): return json.load(open(os.path.join(tok_dir, 'special_tokens.json')))