Formatter for task data. Example Input: .. code-block:: python { "input": "Hello", "output": "Hi" }
| 18 | |
| 19 | |
| 20 | class TaskFormatter: |
| 21 | """Formatter for task data. |
| 22 | |
| 23 | Example Input: |
| 24 | |
| 25 | .. code-block:: python |
| 26 | |
| 27 | { |
| 28 | "input": "Hello", |
| 29 | "output": "Hi" |
| 30 | } |
| 31 | """ |
| 32 | |
| 33 | def __init__(self, config: StorageConfig): |
| 34 | self.config = config |
| 35 | self.default_workflow_cls = WORKFLOWS.get(config.default_workflow_type) # type: ignore |
| 36 | self.default_reward_fn_cls = REWARD_FUNCTIONS.get(config.default_reward_fn_type) # type: ignore |
| 37 | self.workflow_key = config.format.workflow_key |
| 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): |