MCPcopy Create free account
hub / github.com/THUDM/GLM / BlankLMDataset

Class BlankLMDataset

tasks/seq2seq/dataset.py:651–776  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

649
650
651class 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

Callers 2

single_dataset_providerFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected