diff --git a/rdagent/scenarios/data_science/proposal/exp_gen/draft/draft.py b/rdagent/scenarios/data_science/proposal/exp_gen/draft/draft.py index e368f634..ff42a022 100644 --- a/rdagent/scenarios/data_science/proposal/exp_gen/draft/draft.py +++ b/rdagent/scenarios/data_science/proposal/exp_gen/draft/draft.py @@ -125,10 +125,10 @@ class DSDraftExpGen(ExpGen): return exp -class DSDraftExpGenV2(ExpGen): +class DSDraftV2ExpGen(ExpGen): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - self.support_function_calling = APIBackend().support_function_calling() + self.supports_response_schema = APIBackend().supports_response_schema() def tag_gen(self, scenario_desc: str) -> str: sys_prompt = T(".prompts_draft:tag_gen.system").r(tag_desc=T(".prompts_draft:description.tag_description").r()) @@ -192,7 +192,7 @@ class DSDraftExpGenV2(ExpGen): component_info = get_component(hypothesis.component) data_folder_info = self.scen.processed_data_folder_description sys_prompt = T(".prompts_draft:task_gen.system").r( - task_output_format=component_info["task_output_format"] if not self.support_function_calling else None, + task_output_format=component_info["task_output_format"] if not self.supports_response_schema else None, component_desc=component_desc, workflow_check=not pipeline and hypothesis.component != "Workflow", ) @@ -206,12 +206,12 @@ class DSDraftExpGenV2(ExpGen): response = APIBackend().build_messages_and_create_chat_completion( user_prompt=user_prompt, system_prompt=sys_prompt, - response_format=CodingSketch if self.support_function_calling else {"type": "json_object"}, - json_target_type=Dict[str, str | Dict[str, str]] if not self.support_function_calling else None, + response_format=CodingSketch if self.supports_response_schema else {"type": "json_object"}, + json_target_type=Dict[str, str | Dict[str, str]] if not self.supports_response_schema else None, ) task_dict = json.loads(response) task_design = ( - task_dict.get("task_design", {}) if not self.support_function_calling else task_dict.get("sketch", {}) + task_dict.get("task_design", {}) if not self.supports_response_schema else task_dict.get("sketch", {}) ) logger.info(f"Task design:\n{task_design}") task_name = hypothesis.component diff --git a/rdagent/scenarios/data_science/proposal/exp_gen/router/__init__.py b/rdagent/scenarios/data_science/proposal/exp_gen/router/__init__.py index e7a445b2..4621dd81 100644 --- a/rdagent/scenarios/data_science/proposal/exp_gen/router/__init__.py +++ b/rdagent/scenarios/data_science/proposal/exp_gen/router/__init__.py @@ -2,7 +2,7 @@ from rdagent.app.data_science.conf import DS_RD_SETTING from rdagent.core.proposal import ExpGen from rdagent.scenarios.data_science.experiment.experiment import DSExperiment from rdagent.scenarios.data_science.proposal.exp_gen.base import DSTrace -from rdagent.scenarios.data_science.proposal.exp_gen.draft.draft import DSDraftExpGenV2 +from rdagent.scenarios.data_science.proposal.exp_gen.draft.draft import DSDraftV2ExpGen from rdagent.scenarios.data_science.proposal.exp_gen.proposal import DSProposalV2ExpGen @@ -14,7 +14,7 @@ class DraftRouterExpGen(ExpGen): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) - self.draft_exp_gen = DSDraftExpGenV2(self.scen) + self.draft_exp_gen = DSDraftV2ExpGen(self.scen) self.base_exp_gen = DSProposalV2ExpGen(self.scen) def gen(self, trace: DSTrace) -> DSExperiment: