mirror of
https://github.com/NicolasBohn/NexQuant.git
synced 2026-08-08 04:27:44 +00:00
fix: refactor Bench (#302)
* refactor for better bench * autolint * add cmd * lint
This commit is contained in:
@@ -20,7 +20,7 @@ from rdagent.components.coder.factor_coder.factor import FactorFBWorkspace
|
||||
from rdagent.core.conf import RD_AGENT_SETTINGS
|
||||
from rdagent.core.developer import Developer
|
||||
from rdagent.core.exception import CoderError
|
||||
from rdagent.core.experiment import Task, Workspace
|
||||
from rdagent.core.experiment import Experiment, Task, Workspace
|
||||
from rdagent.core.scenario import Scenario
|
||||
from rdagent.core.utils import multiprocessing_wrapper
|
||||
|
||||
@@ -33,11 +33,34 @@ EVAL_RES = Dict[
|
||||
class TestCase:
|
||||
def __init__(
|
||||
self,
|
||||
target_task: list[Task] = [],
|
||||
ground_truth: list[Workspace] = [],
|
||||
target_task: Task,
|
||||
ground_truth: Workspace,
|
||||
):
|
||||
self.ground_truth = ground_truth
|
||||
self.target_task = target_task
|
||||
self.ground_truth = ground_truth
|
||||
|
||||
|
||||
class TestCases:
|
||||
def __init__(self, test_case_l: list[TestCase] = []):
|
||||
# self.test_case_l = [TestCase(task, gt) for task, gt in zip(target_task, ground_truth)]
|
||||
self.test_case_l = test_case_l
|
||||
|
||||
def __getitem__(self, item):
|
||||
return self.test_case_l[item]
|
||||
|
||||
def __len__(self):
|
||||
return len(self.test_case_l)
|
||||
|
||||
def get_exp(self):
|
||||
return Experiment([case.target_task for case in self.test_case_l])
|
||||
|
||||
@property
|
||||
def target_task(self):
|
||||
return [case.target_task for case in self.test_case_l]
|
||||
|
||||
@property
|
||||
def ground_truth(self):
|
||||
return [case.ground_truth for case in self.test_case_l]
|
||||
|
||||
|
||||
class BaseEval:
|
||||
@@ -48,13 +71,13 @@ class BaseEval:
|
||||
def __init__(
|
||||
self,
|
||||
evaluator_l: List[FactorEvaluator],
|
||||
test_cases: List[TestCase],
|
||||
test_cases: TestCases,
|
||||
generate_method: Developer,
|
||||
catch_eval_except: bool = True,
|
||||
):
|
||||
"""Parameters
|
||||
----------
|
||||
test_cases : List[TestCase]
|
||||
test_cases : TestCases
|
||||
cases to be evaluated, ground truth are included in the test cases.
|
||||
evaluator_l : List[FactorEvaluator]
|
||||
A list of evaluators to evaluate the generated code.
|
||||
@@ -105,6 +128,7 @@ class BaseEval:
|
||||
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 (CoderError, AttributeError) as e:
|
||||
# TODO: remove AttributeError.
|
||||
return e
|
||||
except Exception as e:
|
||||
# exception when evaluation
|
||||
@@ -118,7 +142,7 @@ class BaseEval:
|
||||
class FactorImplementEval(BaseEval):
|
||||
def __init__(
|
||||
self,
|
||||
test_cases: TestCase,
|
||||
test_cases: TestCases,
|
||||
method: Developer,
|
||||
*args,
|
||||
scen: Scenario,
|
||||
@@ -146,7 +170,7 @@ class FactorImplementEval(BaseEval):
|
||||
print(f"Eval {_}-th times...")
|
||||
print("========================================================\n")
|
||||
try:
|
||||
gen_factor_l = self.generate_method.develop(self.test_cases.target_task)
|
||||
gen_factor_l = self.generate_method.develop(self.test_cases.get_exp())
|
||||
except KeyboardInterrupt:
|
||||
# TODO: Why still need to save result after KeyboardInterrupt?
|
||||
print("Manually interrupted the evaluation. Saving existing results")
|
||||
|
||||
Reference in New Issue
Block a user