fix draft bugs (#1048)

Co-authored-by: Xu <v-xuminrui@microsoft.com>
This commit is contained in:
Roland Minrui
2025-07-10 11:47:23 +08:00
committed by GitHub
parent 53d7c3b6da
commit c998ec9ccf
2 changed files with 8 additions and 8 deletions
@@ -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
@@ -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: