duotactic / code /leanoar /data.py
Duoia's picture
duotactic full package: checkpoints, tokenizer, config, code, docs
32c0c6c verified
Raw
History Blame Contribute Delete
2.96 kB
"""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')))