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

Class XLDataset

data_utils/datasets.py:501–567  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

499
500
501class XLDataset(data.Dataset):
502 def __init__(self, ds, tokenizer, max_seq_len=1024, mem_len=None, sample_across_doc=True, **kwargs):
503 self.ds = ds
504 self.tokenizer = tokenizer
505 self.max_seq_len = max_seq_len
506 if mem_len is None:
507 mem_len = max_seq_len
508 self.mem_len = mem_len
509 self.sample_across_doc = sample_across_doc
510 self.indices, self.num_samples = None, None
511 if hasattr(self.ds, 'is_lazy') and self.ds.is_lazy:
512 self.is_lazy = True
513 self.init_indices()
514
515 def init_indices(self):
516 if self.is_lazy:
517 lens = np.array([self.ds.get_text_len(idx) for idx in range(len(self.ds))])
518 else:
519 lens = np.array([len(d['prompt']) + len(d['text']) if isinstance(d, dict) else len(d) for d in self.ds])
520 self.indices = list(accumulate(lens))
521 print_rank_0(f"Dataset document count {len(lens)}, token count {self.indices[-1]}")
522 self.num_samples = self.indices[-1] // self.max_seq_len + 1
523
524 def __len__(self):
525 return self.num_samples
526
527 def __getitem__(self, idx):
528 tokens, targets, loss_mask, attention_mask = self.getidx(idx)
529 tokens = self.pad_seq(tokens)
530 targets = self.pad_seq(targets)
531 loss_mask = self.pad_seq(loss_mask, pad_id=0)
532 return {'text': np.array(tokens), "target": np.array(targets), "loss_mask": np.array(loss_mask),
533 "attention_mask": np.array(attention_mask)}
534
535 def getidx(self, idx):
536 tokens, targets, loss_masks = [], [], []
537 attention_mask = np.concatenate((np.zeros((self.max_seq_len, self.mem_len), dtype=np.long),
538 np.ones((self.max_seq_len, self.max_seq_len), dtype=np.long)), axis=1)
539 sample_idx = bisect_right(self.indices, idx * self.max_seq_len)
540 last_end = 0 if sample_idx == 0 else self.indices[sample_idx - 1]
541 token_offset = idx * self.max_seq_len - last_end
542 if token_offset != 0:
543 history = min(self.mem_len, token_offset)
544 attention_mask[:, -self.max_seq_len - history:-self.max_seq_len] = 1
545 count = 0
546 while len(tokens) < self.max_seq_len and sample_idx < len(self.ds):
547 item = self.ds[sample_idx]
548 text, masks = item['tokens'], item['loss_masks']
549 text = text + [self.tokenizer.get_command('eos').Id]
550 end = min(len(text) - 1, token_offset + self.max_seq_len - len(tokens))
551 masks = masks + [1]
552 if count > 0:
553 current = len(tokens)
554 attention_mask[current:, :current + self.mem_len] = 0
555 tokens += text[token_offset: end]
556 targets += text[token_offset + 1: end + 1]
557 loss_masks += masks[token_offset + 1: end + 1]
558 count += 1

Callers 1

wrap_datasetFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected