Align factor coder into new framework (#47)

* use CoSTEER as component name

* rename factorimplementation to avoid confusion

* rename modelimplementation

* align benchmark and evolving evaluators

* add scenario to evaluator init function

* rename all factorimplementationknowledge in CoSTEER

* remove all scenario related information in component

* remove useless code

---------

Co-authored-by: xuyang1 <xuyang1@microsoft.com>
This commit is contained in:
Xu Yang
2024-07-05 17:42:00 +08:00
committed by GitHub
parent f61453fbb1
commit 1d9b4cd2ec
57 changed files with 974 additions and 1093 deletions
+24 -28
View File
@@ -4,22 +4,18 @@ from typing import List, Tuple, Union
from tqdm import tqdm
from rdagent.components.task_implementation.factor_implementation.config import (
FACTOR_IMPLEMENT_SETTINGS,
)
from rdagent.components.task_implementation.factor_implementation.evolving.evaluators import (
FactorImplementationCorrelationEvaluator,
FactorImplementationEvaluator,
FactorImplementationIndexEvaluator,
FactorImplementationIndexFormatEvaluator,
FactorImplementationMissingValuesEvaluator,
FactorImplementationRowCountEvaluator,
FactorImplementationSingleColumnEvaluator,
FactorImplementationValuesEvaluator,
)
from rdagent.components.task_implementation.factor_implementation.factor import (
FileBasedFactorImplementation,
from rdagent.components.coder.factor_coder.config import FACTOR_IMPLEMENT_SETTINGS
from rdagent.components.coder.factor_coder.CoSTEER.evaluators import (
FactorCorrelationEvaluator,
FactorEqualValueCountEvaluator,
FactorEvaluator,
FactorIndexEvaluator,
FactorMissingValuesEvaluator,
FactorOutputFormatEvaluator,
FactorRowCountEvaluator,
FactorSingleColumnEvaluator,
)
from rdagent.components.coder.factor_coder.factor import FileBasedFactorImplementation
from rdagent.core.exception import ImplementRunException
from rdagent.core.experiment import Implementation, Task
from rdagent.core.task_generator import TaskGenerator
@@ -43,7 +39,7 @@ class BaseEval:
def __init__(
self,
evaluator_l: List[FactorImplementationEvaluator],
evaluator_l: List[FactorEvaluator],
test_cases: List[TestCase],
generate_method: TaskGenerator,
catch_eval_except: bool = True,
@@ -52,7 +48,7 @@ class BaseEval:
----------
test_cases : List[TestCase]
cases to be evaluated, ground truth are included in the test cases.
evaluator_l : List[FactorImplementationEvaluator]
evaluator_l : List[FactorEvaluator]
A list of evaluators to evaluate the generated code.
catch_eval_except : bool
If we want to debug the evaluators, we recommend to set the this parameter to True.
@@ -81,7 +77,7 @@ class BaseEval:
self,
case_gt: Implementation,
case_gen: Implementation,
) -> List[Union[Tuple[FactorImplementationEvaluator, object], Exception]]:
) -> List[Union[Tuple[FactorEvaluator, object], Exception]]:
"""Parameters
----------
case_gt : FactorImplementation
@@ -91,14 +87,14 @@ class BaseEval:
Returns
-------
List[Union[Tuple[FactorImplementationEvaluator, object],Exception]]
List[Union[Tuple[FactorEvaluator, object],Exception]]
for each item
If the evaluation run successfully, return the evaluate results. Otherwise, return the exception.
"""
eval_res = []
for ev in self.evaluator_l:
try:
eval_res.append((ev, ev.evaluate(case_gt, case_gen)))
eval_res.append((ev, ev.evaluate(implementation=case_gen, gt_implementation=case_gt)))
# if the corr ev is successfully evaluated and achieve the best performance, then break
except ImplementRunException as e:
return e
@@ -116,18 +112,18 @@ class FactorImplementEval(BaseEval):
self,
test_cases: TestCase,
method: TaskGenerator,
test_round: int = 10,
*args,
test_round: int = 10,
**kwargs,
):
online_evaluator_l = [
FactorImplementationSingleColumnEvaluator(),
FactorImplementationIndexFormatEvaluator(),
FactorImplementationRowCountEvaluator(),
FactorImplementationIndexEvaluator(),
FactorImplementationMissingValuesEvaluator(),
FactorImplementationValuesEvaluator(),
FactorImplementationCorrelationEvaluator(hard_check=False),
FactorSingleColumnEvaluator(),
FactorOutputFormatEvaluator(),
FactorRowCountEvaluator(),
FactorIndexEvaluator(),
FactorMissingValuesEvaluator(),
FactorEqualValueCountEvaluator(),
FactorCorrelationEvaluator(hard_check=False),
]
super().__init__(online_evaluator_l, test_cases, method, *args, **kwargs)
self.test_round = test_round