Format a raw sample dict into a Task.
(self, sample: Dict)
| 38 | self.reward_fn_key = config.format.reward_fn_key |
| 39 | |
| 40 | def format(self, sample: Dict) -> Task: |
| 41 | """Format a raw sample dict into a Task.""" |
| 42 | |
| 43 | workflow_name = sample.get(self.workflow_key, None) if self.workflow_key else None |
| 44 | reward_fn_name = sample.get(self.reward_fn_key, None) if self.reward_fn_key else None |
| 45 | |
| 46 | workflow_cls = ( |
| 47 | WORKFLOWS.get(workflow_name) if workflow_name else None |
| 48 | ) or self.default_workflow_cls |
| 49 | reward_fn_cls = ( |
| 50 | REWARD_FUNCTIONS.get(reward_fn_name) if reward_fn_name else None |
| 51 | ) or self.default_reward_fn_cls |
| 52 | assert workflow_cls is not None, "`default_workflow_type` or `workflow_key` is required" |
| 53 | return Task( |
| 54 | workflow=workflow_cls, |
| 55 | reward_fn=reward_fn_cls, |
| 56 | format_args=self.config.format, |
| 57 | repeat_times=self.config.repeat_times, |
| 58 | rollout_args=self.config.rollout_args, |
| 59 | workflow_args=self.config.workflow_args, |
| 60 | reward_fn_args=self.config.reward_fn_args, |
| 61 | is_eval=self.config.is_eval, |
| 62 | raw_task=sample, |
| 63 | ) |
| 64 | |
| 65 | |
| 66 | class CPTFormatter(ExperienceFormatter): |