| 71 | |
| 72 | |
| 73 | class LambadaDataset(torch.utils.data.Dataset): |
| 74 | def __init__(self, args, tokenizer, strict=True): |
| 75 | data_path = args.valid_data[0] |
| 76 | print_rank_0('> building lambada dataset from {} ...'.format(data_path)) |
| 77 | self.args = args |
| 78 | self.max_seq_length = args.seq_length |
| 79 | self.tokenizer = tokenizer |
| 80 | self.pad_idx = tokenizer.get_command('pad').Id |
| 81 | self.strict = strict |
| 82 | self.block_lm = args.block_lm |
| 83 | self.unidirectional = args.unidirectional |
| 84 | mask_token = "gMASK" if args.task_mask else 'MASK' |
| 85 | self.mask_id = self.tokenizer.get_command(mask_token).Id |
| 86 | |
| 87 | self.tokens = [] |
| 88 | self.labels = [] |
| 89 | with open(data_path, 'r') as f: |
| 90 | for line in f.readlines(): |
| 91 | text = json.loads(line)['text'] |
| 92 | tokens, labels = self.get_tokens(text) |
| 93 | self.tokens.append(tokens) |
| 94 | self.labels.append(labels) |
| 95 | |
| 96 | def get_tokens(self, text): |
| 97 | if not self.strict: |
| 98 | tokens = self.tokenizer.EncodeAsIds(text).tokenization |
| 99 | return tokens[:-1], [tokens[-1]] |
| 100 | last_token = text.split()[-1] |
| 101 | start_idx = text.rfind(last_token) |
| 102 | beginning_tokens = self.tokenizer.EncodeAsIds(text[:start_idx].strip()).tokenization |
| 103 | last_token = self.tokenizer.EncodeAsIds(' ' + last_token).tokenization |
| 104 | return beginning_tokens, last_token |
| 105 | |
| 106 | def __len__(self): |
| 107 | return len(self.tokens) |
| 108 | |
| 109 | def __getitem__(self, idx): |
| 110 | tokens, answer = self.tokens[idx], self.labels[idx] |
| 111 | if self.block_lm: |
| 112 | if self.unidirectional: |
| 113 | tokens, answer_tokens = tokens[:1], tokens[1:] + answer |
| 114 | else: |
| 115 | answer_tokens = answer |
| 116 | tokens = tokens + [self.mask_id] |
| 117 | num_special_tokens = num_special_tokens_to_add(tokens, None, answer_tokens, add_cls=True, add_sep=False, |
| 118 | add_piece=True) |
| 119 | left_shift = len(tokens) + len(answer_tokens) + num_special_tokens - self.max_seq_length |
| 120 | if left_shift > 0: |
| 121 | tokens = tokens[left_shift:] |
| 122 | data = build_input_from_ids(tokens, None, answer_tokens, self.max_seq_length, self.tokenizer, |
| 123 | args=self.args, add_cls=True, add_sep=False, add_piece=True, |
| 124 | mask_id=self.mask_id) |
| 125 | ids, types, paddings, position_ids, sep, target_ids, loss_masks = data |
| 126 | if self.unidirectional: |
| 127 | loss_masks = np.array(loss_masks, dtype=np.int64) |
| 128 | last_index = len(loss_masks) |
| 129 | while loss_masks[last_index - 1] == 0: |
| 130 | last_index -= 1 |
no outgoing calls
no test coverage detected