From 75e496d46131a7a7a1c989bd4de79734f269706c Mon Sep 17 00:00:00 2001 From: Xu Yang Date: Tue, 18 Feb 2025 15:39:34 +0800 Subject: [PATCH] fix restart bug (#609) --- rdagent/app/data_science/loop.py | 9 +++++++-- .../components/coder/data_science/raw_data_loader/exp.py | 2 -- 2 files changed, 7 insertions(+), 4 deletions(-) diff --git a/rdagent/app/data_science/loop.py b/rdagent/app/data_science/loop.py index 3d538b40..5a5e5084 100644 --- a/rdagent/app/data_science/loop.py +++ b/rdagent/app/data_science/loop.py @@ -125,9 +125,14 @@ class DataScienceRDLoop(RDLoop): ) if self.trace.sota_experiment() is None and len(self.trace.hist) >= DS_RD_SETTING.consecutive_errors: trace_exp_next_component_list = [ - exp.next_component_required() for exp, _ in self.trace.hist[-DS_RD_SETTING.consecutive_errors :] + type(exp.pending_tasks_list[0][0]) + for exp, _ in self.trace.hist[-DS_RD_SETTING.consecutive_errors :] ] - if None not in trace_exp_next_component_list and len(set(trace_exp_next_component_list)) == 1: + last_successful_exp = self.trace.last_successful_exp() + if ( + last_successful_exp not in [exp for exp, _ in self.trace.hist[-DS_RD_SETTING.consecutive_errors :]] + and len(set(trace_exp_next_component_list)) == 1 + ): logger.error("Consecutive errors reached the limit. Dumping trace.") logger.log_object(self.trace, tag="trace before restart") self.trace = DSTrace(scen=self.trace.scen, knowledge_base=self.trace.knowledge_base) diff --git a/rdagent/components/coder/data_science/raw_data_loader/exp.py b/rdagent/components/coder/data_science/raw_data_loader/exp.py index 98a775c2..9263e209 100644 --- a/rdagent/components/coder/data_science/raw_data_loader/exp.py +++ b/rdagent/components/coder/data_science/raw_data_loader/exp.py @@ -11,8 +11,6 @@ from rdagent.oai.llm_utils import md5_hash from rdagent.utils.agent.tpl import T from rdagent.utils.env import DockerEnv, DSDockerConf -DataLoaderTask = CoSTEERTask - # Because we use isinstance to distinguish between different types of tasks, we need to use sub classes to represent different types of tasks class DataLoaderTask(CoSTEERTask):