(self, dialogs, n = 10)
| 284 | return str(e) |
| 285 | |
| 286 | def infer_many_chats(self, dialogs, n = 10): |
| 287 | prompt_tokens = self.format_chats(dialogs, pad = True) |
| 288 | predictions = {} |
| 289 | try: |
| 290 | for index, prompt_token in enumerate(prompt_tokens): |
| 291 | if n <= 10: |
| 292 | outputs = self.model.generate(prompt_token_ids = prompt_token, sampling_params=self.sampling_params) |
| 293 | preds = [] |
| 294 | for o in outputs[0].outputs: |
| 295 | preds.append(o.text) |
| 296 | predictions[dialogs[index][0]["content"]] = [preds, True] |
| 297 | else: |
| 298 | iters = math.ceil(n/10) |
| 299 | preds = [] |
| 300 | for i in range(0, iters): |
| 301 | outputs = self.model.generate(prompt_token_ids = prompt_token, sampling_params=self.sampling_params) |
| 302 | for o in outputs[0].outputs: |
| 303 | preds.append(o.text) |
| 304 | predictions[dialogs[index][0]["content"]] = [preds, True] |
| 305 | return predictions |
| 306 | except Exception as e: |
| 307 | logger.error("Error occurred with reason {} when processing multiple prompts".format(str(e))) |
| 308 | print(traceback.print_exc()) |
| 309 | return str(e) |
no test coverage detected