| """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'))) |
|
|