(self)
| 270 | ) |
| 271 | |
| 272 | def __iter__(self): |
| 273 | chunk_size = self.config.batch_size * self.config.seq_length |
| 274 | total_tokens = 0 |
| 275 | while True: |
| 276 | token_buffer = [] |
| 277 | loss_mask_buffer = [] |
| 278 | for index, example in enumerate(self._dataset): |
| 279 | tokens, loss_masks = self.text_processor(example) |
| 280 | token_buffer.extend(tokens) |
| 281 | loss_mask_buffer.extend(loss_masks) |
| 282 | while len(token_buffer) > chunk_size + 1: |
| 283 | total_tokens += chunk_size |
| 284 | metrics = { |
| 285 | 'dataset_example_index': index, |
| 286 | 'dataset_total_tokens': total_tokens, |
| 287 | } |
| 288 | batch = { |
| 289 | 'input_tokens': np.array(token_buffer[:chunk_size], dtype=np.int32).reshape( |
| 290 | self.config.batch_size, -1 |
| 291 | ), |
| 292 | 'target_tokens': np.array(token_buffer[1:chunk_size + 1], dtype=np.int32).reshape( |
| 293 | self.config.batch_size, -1 |
| 294 | ), |
| 295 | 'loss_masks': np.array(loss_mask_buffer[1:chunk_size + 1], dtype=np.float32).reshape( |
| 296 | self.config.batch_size, -1 |
| 297 | ), |
| 298 | } |
| 299 | if self.config.always_start_with_bos: |
| 300 | batch['input_tokens'][:, 0] = self.tokenizer.bos_token_id |
| 301 | yield batch, metrics |
| 302 | token_buffer = token_buffer[chunk_size:] |
| 303 | loss_mask_buffer = loss_mask_buffer[chunk_size:] |
| 304 | |
| 305 | def get_state_dict(self): |
| 306 | return dict(config=self.config) |
nothing calls this directly
no test coverage detected