| 16 | |
| 17 | |
| 18 | class HuggingfaceDataset(IterableDataset): |
| 19 | |
| 20 | def __init__( |
| 21 | self, |
| 22 | dataset: Dataset, |
| 23 | tokenizer: PreTrainedTokenizer, |
| 24 | context_len: int = 2048, |
| 25 | rank: int = 0, |
| 26 | world_size: int = 1, |
| 27 | buffer_size: int = 1024 |
| 28 | ) -> HuggingfaceDataset: |
| 29 | |
| 30 | self.dataset = dataset |
| 31 | self.tokenizer = tokenizer |
| 32 | |
| 33 | self.data = dataset.shard(world_size, rank) |
| 34 | self.context_len = context_len |
| 35 | self.rank = rank |
| 36 | self.world_size = world_size |
| 37 | self.buffer_size = buffer_size |
| 38 | |
| 39 | if tokenizer.vocab_size < torch.iinfo(torch.int16).max: |
| 40 | self.dtype = torch.int16 |
| 41 | elif tokenizer.vocab_size < torch.iinfo(torch.int32).max: |
| 42 | self.dtype = torch.int32 |
| 43 | else: |
| 44 | self.dtype = torch.int64 |
| 45 | self.states = None |
| 46 | self.buffer = torch.tensor([], dtype=self.dtype) |
| 47 | self.tokens = [] |
| 48 | self.rand_id = 0 |
| 49 | self.token_id = 0 |
| 50 | self.rng_state = None |
| 51 | self._epoch = 0 |
| 52 | |
| 53 | def __iter__(self): |
| 54 | g = torch.Generator() |
| 55 | g.manual_seed(self._epoch + self.rank) |
| 56 | if self.rng_state is not None: |
| 57 | g.set_state(self.rng_state) |
| 58 | |
| 59 | rand_it = self.randint(0, self.buffer_size, g=g) |
| 60 | if self.states is not None: |
| 61 | self.data.load_state_dict(self.states) |
| 62 | |
| 63 | # max number of tokens allowed in the chunk buffer |
| 64 | n_tokens = self.buffer_size * self.context_len |
| 65 | |
| 66 | while True: |
| 67 | for sample in self.tokenize(self.data): |
| 68 | # keep appending the samples to the token buffer |
| 69 | self.tokens += sample |
| 70 | # if the token buffer is full, start sampling |
| 71 | # NOTE: we first convert the token ids to a tensor of shape [n_chunks, context_len] for efficiency |
| 72 | if len(self.buffer) == 0 and len(self.tokens) >= n_tokens: |
| 73 | self.buffer = torch.tensor(self.tokens[:n_tokens], dtype=self.dtype).view(self.buffer_size, -1) |
| 74 | self.tokens = self.tokens[n_tokens:] |
| 75 | if len(self.buffer) == self.buffer_size: |
nothing calls this directly
no outgoing calls
no test coverage detected