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

Class Seq2SeqDataset

tasks/seq2seq/dataset.py:420–550  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

418
419
420class Seq2SeqDataset(torch.utils.data.Dataset):
421 def __init__(self, args, split, tokenizer):
422 self.args = args
423 self.task, self.data_dir = args.task.lower(), args.data_dir
424 self.max_src_length, self.max_tgt_length = args.src_seq_length, args.tgt_seq_length
425 self.split = split
426 self.tokenizer = tokenizer
427 self.dataset_name = split
428 if self.task in ["gigaword", "cnn_dm", "cnn_dm_original"]:
429 self.processor = SummmaryProcessor(self.task, self.data_dir, tokenizer)
430 elif self.task in ["xsum"]:
431 self.processor = XSumProcessor(self.data_dir, tokenizer)
432 elif self.task in ["squad_generation"]:
433 self.processor = SQuADGenerationProcessor(self.data_dir, tokenizer)
434 elif self.task in ["squad", "squad_v1"]:
435 self.processor = SQuADProcessor(self.data_dir, tokenizer, self.max_src_length, args)
436 elif self.task in ['cmrc']:
437 self.processor = CMRCProcessor(self.data_dir, tokenizer)
438 else:
439 raise NotImplementedError(self.task)
440 example_list = self.processor.create_examples(split)
441 self.example_list = example_list
442 self.examples = {example.guid: example for example in example_list}
443
444 print_rank_0(f"Return {len(self.examples)} {split} examples")
445
446 def __len__(self):
447 return len(self.example_list)
448
449 def __getitem__(self, idx):
450 example = self.example_list[idx]
451 cls_id = self.tokenizer.get_command('ENC').Id
452 mask_token = 'sMASK' if self.args.task_mask else 'MASK'
453 mask_id = self.tokenizer.get_command(mask_token).Id
454 pad_id = self.tokenizer.get_command('pad').Id
455 sop_id = self.tokenizer.get_command('sop').Id
456 eop_id = self.tokenizer.get_command('eop').Id
457 if self.task in ["gigaword", "cnn_dm", "cnn_dm_original", "xsum"]:
458 source_text, target_text = example.text_a, example.text_b
459 source_tokens = self.tokenizer.EncodeAsIds(" " + source_text).tokenization
460 prompt = [cls_id, mask_id] + self.tokenizer.EncodeAsIds(" Content:").tokenization
461 if len(source_tokens) > self.max_src_length - len(prompt):
462 source_tokens = source_tokens[:self.max_src_length - len(prompt)]
463 source_tokens = prompt + source_tokens
464 elif self.task == "squad_generation":
465 source_text = example.text_a
466 target_text, answer = example.meta["question"], example.meta["answer"]
467 source_tokens = self.tokenizer.EncodeAsIds(source_text.rstrip() + " Question:").tokenization
468 answer_tokens = self.tokenizer.EncodeAsIds(" Answer: " + answer).tokenization
469 if len(source_tokens) > self.max_src_length - len(answer_tokens) - 2:
470 max_src_length = self.max_src_length - len(answer_tokens) - 2
471 answer_pattern = self.tokenizer.EncodeAsIds(" " + answer).tokenization
472
473 def sub_finder(mylist, pattern):
474 matches = []
475 for i in range(len(mylist)):
476 if mylist[i] == pattern[0] and mylist[i:i + len(pattern)] == pattern:
477 matches.append(i)

Callers 2

single_dataset_providerFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected