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

Class LMDataset

tasks/language_model/dataset.py:12–70  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

10
11
12class 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

Callers 2

build_lm_datasetFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected