diff --git a/rdagent/scenarios/data_science/proposal/exp_gen/proposal.py b/rdagent/scenarios/data_science/proposal/exp_gen/proposal.py index 48feb0f3..82b6d6fa 100644 --- a/rdagent/scenarios/data_science/proposal/exp_gen/proposal.py +++ b/rdagent/scenarios/data_science/proposal/exp_gen/proposal.py @@ -1,6 +1,6 @@ import json from enum import Enum -from typing import Dict, List, Optional, Tuple +from typing import Any, Dict, List, Optional, Tuple import pandas as pd from pydantic import BaseModel, Field @@ -24,6 +24,58 @@ from rdagent.utils.agent.tpl import T from rdagent.utils.repo.diff import generate_diff_from_dict from rdagent.utils.workflow import wait_retry +_COMPONENT_META: Dict[str, Dict[str, Any]] = { + "DataLoadSpec": { + "target_name": "Data loader and specification generation", + "spec_file": "spec/data_loader.md", + "output_format_key": ".prompts:output_format.data_loader", + "task_class": DataLoaderTask, + }, + "FeatureEng": { + "target_name": "Feature engineering", + "spec_file": "spec/feature.md", + "output_format_key": ".prompts:output_format.feature", + "task_class": FeatureTask, + }, + "Model": { + "target_name": "Model", + "spec_file": "spec/model.md", + "output_format_key": ".prompts:output_format.model", + "task_class": ModelTask, + }, + "Ensemble": { + "target_name": "Ensemble", + "spec_file": "spec/ensemble.md", + "output_format_key": ".prompts:output_format.ensemble", + "task_class": EnsembleTask, + }, + "Workflow": { + "target_name": "Workflow", + "spec_file": "spec/workflow.md", + "output_format_key": ".prompts:output_format.workflow", + "task_class": WorkflowTask, + }, + "Pipeline": { + "target_name": "Pipeline", + "spec_file": None, + "output_format_key": ".prompts:output_format.pipeline", + "task_class": PipelineTask, + }, +} + + +def get_component(name: str) -> Dict[str, Any]: + meta = _COMPONENT_META.get(name) + if meta is None: + raise KeyError(f"Unknown component: {name!r}") + + return { + "target_name": meta["target_name"], + "spec_file": meta["spec_file"], + "task_output_format": T(meta["output_format_key"]).r(), + "task_class": meta["task_class"], + } + class ScenarioChallengeCategory(str, Enum): DATASET_DRIVEN = "dataset-driven" @@ -233,43 +285,6 @@ def draft_exp_in_decomposition(scen: Scenario, trace: DSTrace) -> None | DSDraft class DSProposalV1ExpGen(ExpGen): - COMPONENT_TASK_MAPPING = { - "DataLoadSpec": { - "target_name": "Data loader and specification generation", - "spec_file": "spec/data_loader.md", - "task_output_format": T(".prompts:output_format.data_loader").r(), - "task_class": DataLoaderTask, - }, - "FeatureEng": { - "target_name": "Feature engineering", - "spec_file": "spec/feature.md", - "task_output_format": T(".prompts:output_format.feature").r(), - "task_class": FeatureTask, - }, - "Model": { - "target_name": "Model", - "spec_file": "spec/model.md", - "task_output_format": T(".prompts:output_format.model").r(), - "task_class": ModelTask, - }, - "Ensemble": { - "target_name": "Ensemble", - "spec_file": "spec/ensemble.md", - "task_output_format": T(".prompts:output_format.ensemble").r(), - "task_class": EnsembleTask, - }, - "Workflow": { - "target_name": "Workflow", - "spec_file": "spec/workflow.md", - "task_output_format": T(".prompts:output_format.workflow").r(), - "task_class": WorkflowTask, - }, - "Pipeline": { - "target_name": "Pipeline", - "task_output_format": T(".prompts:output_format.pipeline").r(), - "task_class": PipelineTask, - }, - } def gen(self, trace: DSTrace) -> DSExperiment: # Drafting Stage @@ -351,7 +366,7 @@ class DSProposalV1ExpGen(ExpGen): # - after we know the selected component, we can use RAG. # Step 2: Generate the rest of the hypothesis & task - component_info = self.COMPONENT_TASK_MAPPING.get(component) + component_info = get_component(component) if component_info: if DS_RD_SETTING.spec_enabled: @@ -656,9 +671,9 @@ class DSProposalV2ExpGen(ExpGen): failed_exp_feedback_list_desc: str, ) -> DSExperiment: if pipeline: - component_info = self.COMPONENT_TASK_MAPPING["Pipeline"] + component_info = get_component("Pipeline") else: - component_info = self.COMPONENT_TASK_MAPPING.get(hypothesis.component) + component_info = get_component(hypothesis.component) if pipeline: task_spec = T(f"scenarios.data_science.share:component_spec.Pipeline").r() elif DS_RD_SETTING.spec_enabled and sota_exp is not None: @@ -948,9 +963,9 @@ class DSProposalV3ExpGen(DSProposalV2ExpGen): failed_exp_feedback_list_desc: str, ) -> DSExperiment: if pipeline: - component_info = self.COMPONENT_TASK_MAPPING["Pipeline"] + component_info = get_component("Pipeline") else: - component_info = self.COMPONENT_TASK_MAPPING.get(hypotheses[0].component) + component_info = get_component(hypotheses[0].component) data_folder_info = self.scen.processed_data_folder_description sys_prompt = T(".prompts_v3:task_gen.system").r( # targets=component_info["target_name"],