MCPcopy Create free account
hub / github.com/SalesforceAIResearch/perfcodegen / infer_many_chats

Method infer_many_chats

src/model.py:286–309  ·  view source on GitHub ↗
(self, dialogs, n = 10)

Source from the content-addressed store, hash-verified

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)

Callers 1

batch_infer_local_modelFunction · 0.95

Calls 2

format_chatsMethod · 0.95
generateMethod · 0.80

Tested by

no test coverage detected