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

Method generate

opencompass/models/vllm.py:64–117  ·  view source on GitHub ↗

Generate results given a list of inputs. Args: inputs (List[str]): A list of strings. max_out_len (int): The maximum length of the output. Returns: List[str]: A list of generated strings.

(self,
                 inputs: List[str],
                 max_out_len: int,
                 stopping_criteria: List[str] = [],
                 **kwargs)

Source from the content-addressed store, hash-verified

62 self.model = LLM(path, **model_kwargs)
63
64 def generate(self,
65 inputs: List[str],
66 max_out_len: int,
67 stopping_criteria: List[str] = [],
68 **kwargs) -> List[str]:
69 """Generate results given a list of inputs.
70
71 Args:
72 inputs (List[str]): A list of strings.
73 max_out_len (int): The maximum length of the output.
74
75 Returns:
76 List[str]: A list of generated strings.
77 """
78
79 if self.mode == 'mid':
80 input_ids = self.tokenizer(inputs, truncation=False)['input_ids']
81 inputs = []
82 for input_id in input_ids:
83 if len(input_id) > self.max_seq_len - max_out_len:
84 half = int((self.max_seq_len - max_out_len) / 2)
85 inputs.append(
86 self.tokenizer.decode(input_id[:half],
87 skip_special_tokens=True) +
88 self.tokenizer.decode(input_id[-half:],
89 skip_special_tokens=True))
90 else:
91 inputs.append(
92 self.tokenizer.decode(input_id,
93 skip_special_tokens=True))
94
95 generation_kwargs = kwargs.copy()
96 generation_kwargs.update(self.generation_kwargs)
97 generation_kwargs.update({'max_tokens': max_out_len})
98 _stop = list(set(self.stop_words + stopping_criteria))
99 generation_kwargs.update({'stop': _stop})
100 sampling_kwargs = SamplingParams(**generation_kwargs)
101 if not self.lora_path:
102 outputs = self.model.generate(inputs, sampling_kwargs)
103 else:
104 outputs = self.model.generate(inputs,
105 sampling_kwargs,
106 lora_request=LoRARequest(
107 'sql_adapter', 1,
108 self.lora_path))
109
110 prompt_list, output_strs = [], []
111 for output in outputs:
112 prompt = output.prompt
113 generated_text = output.outputs[0].text
114 prompt_list.append(prompt)
115 output_strs.append(generated_text)
116
117 return output_strs
118
119 def get_ppl(self,
120 inputs: List[str],

Callers 4

_evaluate_datasetMethod · 0.45
api_scoreMethod · 0.45
thread_workerMethod · 0.45
get_pplMethod · 0.45

Calls 2

decodeMethod · 0.80
updateMethod · 0.80

Tested by

no test coverage detected