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)
| 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], |
no test coverage detected