| 10 | |
| 11 | |
| 12 | class LMDataset(torch.utils.data.Dataset): |
| 13 | def __init__(self, args, documents, tokenizer, num_original_tokens, num_tokenized_tokens): |
| 14 | self.args = args |
| 15 | self.documents = documents |
| 16 | self.max_seq_len = args.seq_length - 1 |
| 17 | self.tokenizer = tokenizer |
| 18 | self.overalapping_eval = args.overlapping_eval |
| 19 | if self.overalapping_eval is None: |
| 20 | self.overalapping_eval = self.max_seq_len |
| 21 | self.overalapping_eval = max(1, self.overalapping_eval) |
| 22 | self.num_original_tokens = num_original_tokens |
| 23 | self.num_tokenized_tokens = num_tokenized_tokens |
| 24 | # remove first sequence tokens |
| 25 | targets = [max(len(tokens) - self.max_seq_len, 0) for tokens in self.documents] |
| 26 | self.num_sequences = [max(math.ceil(target / self.overalapping_eval) + 1, 1) for target in targets] |
| 27 | self.weights = list(accumulate(self.num_sequences)) |
| 28 | self.left_weights = [0] + self.weights[:-1] |
| 29 | self.unidirectional = args.unidirectional |
| 30 | self.block_lm = args.block_lm |
| 31 | mask_token = "gMASK" if args.task_mask else 'MASK' |
| 32 | self.mask_id = self.tokenizer.get_command(mask_token).Id |
| 33 | |
| 34 | def __len__(self): |
| 35 | return sum(self.num_sequences) |
| 36 | |
| 37 | def __getitem__(self, idx): |
| 38 | document_idx = bisect_right(self.weights, idx) |
| 39 | idx = idx - self.left_weights[document_idx] |
| 40 | start_idx = idx * self.overalapping_eval |
| 41 | end_idx = start_idx + self.max_seq_len |
| 42 | tokens = self.documents[document_idx][start_idx:end_idx] |
| 43 | if self.block_lm: |
| 44 | if idx == 0 or self.unidirectional: |
| 45 | prompt, text = tokens[:1], tokens[1:] |
| 46 | else: |
| 47 | prompt_length = self.max_seq_len - self.overalapping_eval |
| 48 | prompt, text = tokens[:prompt_length], tokens[prompt_length:] |
| 49 | prompt = prompt + [self.mask_id] |
| 50 | num_special_tokens = num_special_tokens_to_add(prompt, None, text, add_cls=True, add_sep=False, |
| 51 | add_piece=True, |
| 52 | add_eos=False) |
| 53 | data = build_input_from_ids(prompt, None, text, self.max_seq_len + num_special_tokens + 1, self.tokenizer, |
| 54 | args=self.args, add_cls=True, add_sep=False, add_piece=True, add_eos=False, mask_id=self.mask_id) |
| 55 | ids, types, paddings, position_ids, sep, target_ids, loss_masks = data |
| 56 | if idx != 0 and self.unidirectional: |
| 57 | loss_masks = np.array(loss_masks, dtype=np.int64) |
| 58 | loss_masks[:-self.overalapping_eval] = 0 |
| 59 | return {'text': np.array(ids, dtype=np.int64), 'target': np.array(target_ids, dtype=np.int64), |
| 60 | 'attention_mask': np.array(sep, dtype=np.int64), 'loss_mask': np.array(loss_masks, dtype=np.int64), |
| 61 | "position_id": np.array(position_ids, dtype=np.int64)} |
| 62 | else: |
| 63 | loss_masks = [1] * len(tokens) |
| 64 | if len(tokens) < self.max_seq_len: |
| 65 | tokens = tokens + [0] * (self.max_seq_len - len(tokens)) |
| 66 | loss_masks = loss_masks + [0] * (self.max_seq_len - len(loss_masks)) |
| 67 | if idx != 0: |
| 68 | loss_masks = np.array(loss_masks, dtype=np.int64) |
| 69 | loss_masks[:-self.overalapping_eval] = 0 |
no outgoing calls
no test coverage detected