diff --git a/rdagent/components/coder/CoSTEER/evaluators.py b/rdagent/components/coder/CoSTEER/evaluators.py index 612b6895..4c4f4dae 100644 --- a/rdagent/components/coder/CoSTEER/evaluators.py +++ b/rdagent/components/coder/CoSTEER/evaluators.py @@ -43,6 +43,34 @@ class CoSTEERSingleFeedback(Feedback): code: str final_decision: bool + @staticmethod + def val_and_update_init_dict(data: dict) -> dict: + # TODO: (bowen) use a more general method to validate and update the data dictionary before init, like pydantic + """ + Validates and converts the 'final_decision' field in the given data dictionary. + + Args: + data (dict): The data dictionary containing the 'final_decision' field. + + Returns: + dict: The updated data dictionary with 'final_decision' as a boolean. + + Raises: + ValueError: If 'final_decision' is not present or not a boolean. + """ + if "final_decision" not in data: + raise ValueError("'final_decision' is required") + + if isinstance(data["final_decision"], str): + if data["final_decision"] == "false" or data["final_decision"] == "False": + data["final_decision"] = False + elif data["final_decision"] == "true" or data["final_decision"] == "True": + data["final_decision"] = True + + if not isinstance(data["final_decision"], bool): + raise ValueError(f"'final_decision' must be a boolean, not {type(data['final_decision'])}") + return data + def __str__(self) -> str: return f"""------------------Execution------------------ {self.execution} diff --git a/rdagent/components/coder/data_science/ensemble/eval.py b/rdagent/components/coder/data_science/ensemble/eval.py index e6a4b72f..d84e0758 100644 --- a/rdagent/components/coder/data_science/ensemble/eval.py +++ b/rdagent/components/coder/data_science/ensemble/eval.py @@ -81,6 +81,11 @@ class EnsembleCoSTEEREvaluator(CoSTEEREvaluator): stdout=stdout, workflow_stdout=workflow_stdout, ) - efb = build_cls_from_json_with_retry(EnsembleEvalFeedback, system_prompt=system_prompt, user_prompt=user_prompt) + efb = build_cls_from_json_with_retry( + EnsembleEvalFeedback, + system_prompt=system_prompt, + user_prompt=user_prompt, + init_kwargs_update_func=EnsembleEvalFeedback.val_and_update_init_dict, + ) efb.final_decision = efb.final_decision and ret_code == 0 return efb diff --git a/rdagent/components/coder/data_science/feature/eval.py b/rdagent/components/coder/data_science/feature/eval.py index f09f8eda..148a511c 100644 --- a/rdagent/components/coder/data_science/feature/eval.py +++ b/rdagent/components/coder/data_science/feature/eval.py @@ -72,4 +72,9 @@ class FeatureCoSTEEREvaluator(CoSTEEREvaluator): workflow_stdout=workflow_stdout, ) - return build_cls_from_json_with_retry(FeatureEvalFeedback, system_prompt=system_prompt, user_prompt=user_prompt) + return build_cls_from_json_with_retry( + FeatureEvalFeedback, + system_prompt=system_prompt, + user_prompt=user_prompt, + init_kwargs_update_func=FeatureEvalFeedback.val_and_update_init_dict, + ) diff --git a/rdagent/components/coder/data_science/model/eval.py b/rdagent/components/coder/data_science/model/eval.py index a1b03fb2..4b145896 100644 --- a/rdagent/components/coder/data_science/model/eval.py +++ b/rdagent/components/coder/data_science/model/eval.py @@ -90,4 +90,9 @@ class ModelGeneralCaseSpecEvaluator(CoSTEEREvaluator): stdout=stdout, workflow_stdout=workflow_stdout, ) - return build_cls_from_json_with_retry(ModelSingleFeedback, system_prompt=system_prompt, user_prompt=user_prompt) + return build_cls_from_json_with_retry( + ModelSingleFeedback, + system_prompt=system_prompt, + user_prompt=user_prompt, + init_kwargs_update_func=ModelSingleFeedback.val_and_update_init_dict, + ) diff --git a/rdagent/components/coder/data_science/raw_data_loader/eval.py b/rdagent/components/coder/data_science/raw_data_loader/eval.py index 58d84327..8a6f6972 100644 --- a/rdagent/components/coder/data_science/raw_data_loader/eval.py +++ b/rdagent/components/coder/data_science/raw_data_loader/eval.py @@ -81,5 +81,8 @@ class DataLoaderCoSTEEREvaluator(CoSTEEREvaluator): ) return build_cls_from_json_with_retry( - DataLoaderEvalFeedback, system_prompt=system_prompt, user_prompt=user_prompt + DataLoaderEvalFeedback, + system_prompt=system_prompt, + user_prompt=user_prompt, + init_kwargs_update_func=DataLoaderEvalFeedback.val_and_update_init_dict, ) diff --git a/rdagent/components/coder/data_science/workflow/eval.py b/rdagent/components/coder/data_science/workflow/eval.py index 9df1a07c..5d25777c 100644 --- a/rdagent/components/coder/data_science/workflow/eval.py +++ b/rdagent/components/coder/data_science/workflow/eval.py @@ -121,7 +121,10 @@ class WorkflowGeneralCaseSpecEvaluator(CoSTEEREvaluator): code=implementation.file_dict["main.py"], ) wfb = build_cls_from_json_with_retry( - WorkflowSingleFeedback, system_prompt=system_prompt, user_prompt=user_prompt + WorkflowSingleFeedback, + system_prompt=system_prompt, + user_prompt=user_prompt, + init_kwargs_update_func=WorkflowSingleFeedback.val_and_update_init_dict, ) if score_ret_code != 0: wfb.final_decision = False diff --git a/rdagent/scenarios/data_science/dev/runner/eval.py b/rdagent/scenarios/data_science/dev/runner/eval.py index 1b2e6cc0..35e2507f 100644 --- a/rdagent/scenarios/data_science/dev/runner/eval.py +++ b/rdagent/scenarios/data_science/dev/runner/eval.py @@ -85,7 +85,10 @@ class DSCoSTEERCoSTEEREvaluator(CoSTEEREvaluator): ) feedback = build_cls_from_json_with_retry( - DSCoSTEEREvalFeedback, system_prompt=system_prompt, user_prompt=user_prompt + DSCoSTEEREvalFeedback, + system_prompt=system_prompt, + user_prompt=user_prompt, + init_kwargs_update_func=DSCoSTEEREvalFeedback.val_and_update_init_dict, ) if feedback: