MCPcopy Create free account
hub / github.com/InternScience/SciReason / batch_encode

Method batch_encode

opencompass/models/interntrain.py:455–478  ·  view source on GitHub ↗
(self,
                     input_texts: Union[str, List[str]],
                     max_seq_len: int,
                     left_padding=False)

Source from the content-addressed store, hash-verified

453 return outputs, tokens
454
455 def batch_encode(self,
456 input_texts: Union[str, List[str]],
457 max_seq_len: int,
458 left_padding=False):
459 if isinstance(input_texts, str):
460 input_texts = [input_texts]
461 tokens = [self.tokenizer(text) for text in input_texts]
462 max_len = min(max_seq_len, max([len(t) for t in tokens]))
463 for i in range(len(tokens)):
464 cur_input = tokens[i]
465 padding_len = max_len - len(cur_input)
466 if self.mode == 'none':
467 cur_input = cur_input[:max_len]
468 elif self.mode == 'mid' and len(cur_input) > max_len:
469 mid_cut_len = max_len // 2
470 cur_input = cur_input[:mid_cut_len] + cur_input[-mid_cut_len:]
471
472 if left_padding:
473 # left padding with bos
474 tokens[i] = [self.tokenizer.bos_id] * padding_len + cur_input
475 else:
476 tokens[i] = cur_input + [self.pad_id] * padding_len
477
478 return torch.LongTensor(tokens).cuda()
479
480 def batch_decode(self,
481 outputs,

Callers 2

generateMethod · 0.95
get_logitsMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected