(self, input_list, output_list)
| 419 | self.step = 0 |
| 420 | |
| 421 | def encode(self, input_list, output_list): |
| 422 | sample_num = len(input_list) |
| 423 | encoded_source = self.tokenizer(input_list, padding=True, return_tensors='pt') |
| 424 | source_seq = encoded_source['input_ids'].contiguous() |
| 425 | source_mask = encoded_source['attention_mask'].contiguous() |
| 426 | max_len = source_seq.size(1) |
| 427 | whole_word_ids = compute_whole_word_id(input_list, self.tokenizer, max_len) |
| 428 | whole_word = torch.tensor(whole_word_ids, dtype=torch.int64).contiguous() |
| 429 | encoded_target = self.tokenizer(output_list, padding=True, return_tensors='pt') |
| 430 | target_seq = encoded_target['input_ids'] |
| 431 | task = torch.ones((sample_num,), dtype=torch.int64) * self.task_id |
| 432 | return task, source_seq, source_mask, whole_word, target_seq |
| 433 | |
| 434 | def next_batch(self, valid=True): |
| 435 | if self.step == self.total_step: |
no test coverage detected