MCPcopy Create free account
hub / github.com/WangXFng/RDRec / encode

Method encode

utils/utils.py:421–432  ·  view source on GitHub ↗
(self, input_list, output_list)

Source from the content-addressed store, hash-verified

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:

Callers 1

next_batchMethod · 0.95

Calls 1

compute_whole_word_idFunction · 0.85

Tested by

no test coverage detected