| 649 | |
| 650 | |
| 651 | class BlankLMDataset(torch.utils.data.Dataset): |
| 652 | def __init__(self, args, split, tokenizer): |
| 653 | self.args = args |
| 654 | task, data_dir = args.task.lower(), args.data_dir |
| 655 | self.max_src_length, self.max_tgt_length = args.src_seq_length, args.tgt_seq_length |
| 656 | self.split = split |
| 657 | assert args.tokenizer_type == "BertWordPieceTokenizer" |
| 658 | self.tokenizer = tokenizer |
| 659 | if split == "train": |
| 660 | filename = "train" |
| 661 | elif split == "dev": |
| 662 | filename = "valid" |
| 663 | elif split == "test": |
| 664 | filename = "test" |
| 665 | else: |
| 666 | raise NotImplementedError(split) |
| 667 | print_rank_0(f"Creating {task}-{split} dataset from {data_dir}") |
| 668 | self.dataset_name = split |
| 669 | detokenizer = blanklm_detokenize |
| 670 | source_texts, target_texts = [], [] |
| 671 | with open(os.path.join(data_dir, f"{filename}.txt"), encoding='utf-8') as file: |
| 672 | for line in file: |
| 673 | line = line.strip() |
| 674 | line = detokenizer(line) if detokenizer else line |
| 675 | target_texts.append(line) |
| 676 | if split == 'test': |
| 677 | with open(os.path.join(data_dir, f"blank/test.maskratio{args.blank_maskratio:.1f}.blank"), |
| 678 | encoding='utf-8') as file: |
| 679 | for line in file: |
| 680 | line = line.strip() |
| 681 | line = detokenizer(line) if detokenizer else line |
| 682 | source_texts.append(line) |
| 683 | else: |
| 684 | source_texts = target_texts |
| 685 | self.examples, self.example_list = {}, [] |
| 686 | for idx, (source_text, target_text) in enumerate(zip(source_texts, target_texts)): |
| 687 | # if idx > 10000: |
| 688 | # break |
| 689 | if (idx + 1) % 20000 == 0: |
| 690 | print_rank_0(f"Complete {idx + 1} examples") |
| 691 | guid = "%s-%s" % (split, idx) |
| 692 | meta = {"ref": target_text} |
| 693 | example = InputExample(guid=guid, text_a=source_text, text_b=target_text, meta=meta) |
| 694 | self.examples[guid] = example |
| 695 | self.example_list.append(example) |
| 696 | print_rank_0(f"Return {len(self.examples)} {split} examples") |
| 697 | self.random = random.Random(args.seed) |
| 698 | |
| 699 | def __len__(self): |
| 700 | return len(self.example_list) |
| 701 | |
| 702 | def __getitem__(self, idx): |
| 703 | example = self.example_list[idx] |
| 704 | source_text, target_text = example.text_a, example.text_b |
| 705 | mask_token = 'gMASK' if self.args.task_mask else 'MASK' |
| 706 | mask_id = self.tokenizer.get_command(mask_token).Id |
| 707 | sop_id = self.tokenizer.get_command('sop').Id |
| 708 | eop_id = self.tokenizer.get_command('eop').Id |
no outgoing calls
no test coverage detected