(self)
| 297 | can_repeat: bool = True |
| 298 | |
| 299 | def run(self) -> List[Experience]: |
| 300 | # TODO: Optimize the generate function |
| 301 | messages = self.format_messages() |
| 302 | |
| 303 | self.logger.debug("start chat") |
| 304 | responses = self.model.chat(messages, **self.rollout_args) |
| 305 | for i, response in enumerate(responses): |
| 306 | reward_dict = self.reward_fn( # type: ignore [misc] |
| 307 | response=response.response_text, # type: ignore [arg-type] |
| 308 | truth=self.truth, |
| 309 | ) |
| 310 | |
| 311 | if response.metrics is None: |
| 312 | response.metrics = {} |
| 313 | response.metrics.update(reward_dict) |
| 314 | reward = sum(reward_dict.values()) |
| 315 | response.reward = reward |
| 316 | response.eid.run = i + self.run_id_base |
| 317 | |
| 318 | self.logger.debug( |
| 319 | f"self.task_desc: {self.task_desc}, messages: {messages}, response: {response.response_text}, reward: {reward}" |
| 320 | ) |
| 321 | return responses |
| 322 | |
| 323 | |
| 324 | class AsyncSimpleWorkflow(BaseSimpleWorkflow): |
nothing calls this directly
no test coverage detected