| 499 | |
| 500 | |
| 501 | class XLDataset(data.Dataset): |
| 502 | def __init__(self, ds, tokenizer, max_seq_len=1024, mem_len=None, sample_across_doc=True, **kwargs): |
| 503 | self.ds = ds |
| 504 | self.tokenizer = tokenizer |
| 505 | self.max_seq_len = max_seq_len |
| 506 | if mem_len is None: |
| 507 | mem_len = max_seq_len |
| 508 | self.mem_len = mem_len |
| 509 | self.sample_across_doc = sample_across_doc |
| 510 | self.indices, self.num_samples = None, None |
| 511 | if hasattr(self.ds, 'is_lazy') and self.ds.is_lazy: |
| 512 | self.is_lazy = True |
| 513 | self.init_indices() |
| 514 | |
| 515 | def init_indices(self): |
| 516 | if self.is_lazy: |
| 517 | lens = np.array([self.ds.get_text_len(idx) for idx in range(len(self.ds))]) |
| 518 | else: |
| 519 | lens = np.array([len(d['prompt']) + len(d['text']) if isinstance(d, dict) else len(d) for d in self.ds]) |
| 520 | self.indices = list(accumulate(lens)) |
| 521 | print_rank_0(f"Dataset document count {len(lens)}, token count {self.indices[-1]}") |
| 522 | self.num_samples = self.indices[-1] // self.max_seq_len + 1 |
| 523 | |
| 524 | def __len__(self): |
| 525 | return self.num_samples |
| 526 | |
| 527 | def __getitem__(self, idx): |
| 528 | tokens, targets, loss_mask, attention_mask = self.getidx(idx) |
| 529 | tokens = self.pad_seq(tokens) |
| 530 | targets = self.pad_seq(targets) |
| 531 | loss_mask = self.pad_seq(loss_mask, pad_id=0) |
| 532 | return {'text': np.array(tokens), "target": np.array(targets), "loss_mask": np.array(loss_mask), |
| 533 | "attention_mask": np.array(attention_mask)} |
| 534 | |
| 535 | def getidx(self, idx): |
| 536 | tokens, targets, loss_masks = [], [], [] |
| 537 | attention_mask = np.concatenate((np.zeros((self.max_seq_len, self.mem_len), dtype=np.long), |
| 538 | np.ones((self.max_seq_len, self.max_seq_len), dtype=np.long)), axis=1) |
| 539 | sample_idx = bisect_right(self.indices, idx * self.max_seq_len) |
| 540 | last_end = 0 if sample_idx == 0 else self.indices[sample_idx - 1] |
| 541 | token_offset = idx * self.max_seq_len - last_end |
| 542 | if token_offset != 0: |
| 543 | history = min(self.mem_len, token_offset) |
| 544 | attention_mask[:, -self.max_seq_len - history:-self.max_seq_len] = 1 |
| 545 | count = 0 |
| 546 | while len(tokens) < self.max_seq_len and sample_idx < len(self.ds): |
| 547 | item = self.ds[sample_idx] |
| 548 | text, masks = item['tokens'], item['loss_masks'] |
| 549 | text = text + [self.tokenizer.get_command('eos').Id] |
| 550 | end = min(len(text) - 1, token_offset + self.max_seq_len - len(tokens)) |
| 551 | masks = masks + [1] |
| 552 | if count > 0: |
| 553 | current = len(tokens) |
| 554 | attention_mask[current:, :current + self.mem_len] = 0 |
| 555 | tokens += text[token_offset: end] |
| 556 | targets += text[token_offset + 1: end + 1] |
| 557 | loss_masks += masks[token_offset + 1: end + 1] |
| 558 | count += 1 |