File size: 4,339 Bytes
29f25be
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
"""Deterministic prefix crops over the unchanged sentence window token store."""

import math

import numpy as np

from vimeml.training.data import SentenceWindowDataset

MASK64 = (1 << 64) - 1
GOLDEN64 = 0x9E3779B97F4A7C15
EPOCH64 = 0xD1B54A32D192ED03


def mix64(value):
    value = ((value ^ (value >> 30)) * 0xBF58476D1CE4E5B9) & MASK64
    value = ((value ^ (value >> 27)) * 0x94D049BB133111EB) & MASK64
    return value ^ (value >> 31)


def mix64_array(values):
    values = (values ^ (values >> np.uint64(30))) * np.uint64(0xBF58476D1CE4E5B9)
    values = (values ^ (values >> np.uint64(27))) * np.uint64(0x94D049BB133111EB)
    return values ^ (values >> np.uint64(31))


class PrefixCropWindowDataset(SentenceWindowDataset):
    # Bucket samplers pass (epoch, index) to workers, including prefetched work.
    epoch_indexed = True

    def __init__(
        self,
        token_dir,
        index_dir,
        split="train",
        *,
        seed=42,
        probability=0.30,
        min_remaining_tokens=8,
    ):
        if split != "train":
            raise ValueError("Prefix crops are only defined for the V2 training split.")
        if not math.isfinite(probability) or not 0 <= probability <= 1:
            raise ValueError("Crop probability must be in [0, 1].")
        if type(min_remaining_tokens) is not int or min_remaining_tokens < 1:
            raise ValueError("Minimum remaining tokens must be a positive integer.")
        if type(seed) is not int or not 0 <= seed <= MASK64:
            raise ValueError("Crop seed must be a uint64 integer.")
        super().__init__(token_dir, index_dir, split)
        self.seed = seed
        self.probability = probability
        self.threshold = int(probability * (1 << 53))
        self.min_remaining_tokens = min_remaining_tokens
        self.epoch = 0

    def set_epoch(self, epoch):
        if type(epoch) is not int or epoch < 0:
            raise ValueError("Crop epoch must be a nonnegative integer.")
        self.epoch = epoch

    def crop_offset(self, index, length, original_start, epoch):
        maximum = length - self.min_remaining_tokens
        if original_start != 0 or maximum < 1 or self.threshold == 0:
            return 0
        bits = mix64((self.seed + epoch * EPOCH64 + index + GOLDEN64) & MASK64)
        if (bits >> 11) >= self.threshold:
            return 0
        return 1 + mix64((bits + GOLDEN64) & MASK64) % maximum

    def __getitem__(self, key):
        epoch, index = key if isinstance(key, tuple) else (self.epoch, key)
        index = int(index)
        if index < 0:
            index += len(self)
        if type(epoch) is not int or epoch < 0:
            raise ValueError("Crop epoch must be a nonnegative integer.")
        sample = super().__getitem__(index)
        original_start = sample["window_start"]
        length = len(sample["input_ids"])
        offset = self.crop_offset(index, length, original_start, epoch)
        return {
            **sample,
            "input_ids": sample["input_ids"][offset:],
            "labels": sample["labels"][offset:],
            "window_start": original_start + offset,
            "original_window_start": original_start,
            "crop_offset": offset,
            "uncropped_length": length,
            "sample_index": index,
            "sample_epoch": epoch,
        }

    def window_lengths(self, indices):
        indices = np.asarray(indices, dtype=np.int64)
        lengths = super().window_lengths(indices)
        if not self.threshold:
            return lengths
        position = np.searchsorted(self.end, indices, side="right")
        first_window = np.ones(len(indices), dtype=np.bool_)
        candidates = np.flatnonzero(position < len(self.end))
        first_window[candidates] = indices[candidates] <= self.first[position[candidates]]
        maximum = lengths - self.min_remaining_tokens
        base = (self.seed + self.epoch * EPOCH64 + GOLDEN64) & MASK64
        bits = mix64_array(indices.astype(np.uint64) + np.uint64(base))
        apply = first_window & (maximum > 0) & ((bits >> np.uint64(11)) < np.uint64(self.threshold))
        choices = mix64_array(bits + np.uint64(GOLDEN64)) % np.maximum(maximum, 1).astype(
            np.uint64
        ) + np.uint64(1)
        return lengths - np.where(apply, choices, 0).astype(np.int64)