| 67 | return results |
| 68 | |
| 69 | def _generate(self, input: str, max_out_len: int) -> str: |
| 70 | max_num_retries = 0 |
| 71 | while max_num_retries < self.retry: |
| 72 | self.wait() |
| 73 | header = {'content-type': 'application/json'} |
| 74 | try: |
| 75 | self.logger.debug(f'input: {input}') |
| 76 | data = dict(inputs=input, parameters=self.generation_kwargs) |
| 77 | raw_response = requests.post(self.url, |
| 78 | headers=header, |
| 79 | data=json.dumps(data)) |
| 80 | except requests.ConnectionError: |
| 81 | self.logger.error('Got connection error, retrying...') |
| 82 | continue |
| 83 | try: |
| 84 | response = raw_response.json() |
| 85 | generated_text = response['generated_text'] |
| 86 | if isinstance(generated_text, list): |
| 87 | generated_text = generated_text[0] |
| 88 | self.logger.debug(f'generated_text: {generated_text}') |
| 89 | return generated_text |
| 90 | except requests.JSONDecodeError: |
| 91 | self.logger.error('JsonDecode error, got', |
| 92 | str(raw_response.content)) |
| 93 | except KeyError: |
| 94 | self.logger.error(f'KeyError. Response: {str(response)}') |
| 95 | max_num_retries += 1 |
| 96 | |
| 97 | raise RuntimeError('Calling LightllmAPI failed after retrying for ' |
| 98 | f'{max_num_retries} times. Check the logs for ' |
| 99 | 'details.') |
| 100 | |
| 101 | def get_ppl(self, inputs: List[str], max_out_len: int, |
| 102 | **kwargs) -> List[float]: |