File size: 2,959 Bytes
32c0c6c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
"""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')))