From bd8122e0cdba2bedf94f555c41072a65328dc7f3 Mon Sep 17 00:00:00 2001 From: XianBW <36835909+XianBW@users.noreply.github.com> Date: Thu, 23 Jan 2025 16:12:22 +0800 Subject: [PATCH] feat: base data science scenario UI (#525) * base data science ui * fix bug * fix mle grade * not cache when mle prepare * fix * fix grade sample * fix a small bug * fix * cache mle score * fix * add gen_mle_score script * update for mle score * simple debug show * small change * summary folder * add evo loop tag * add loop id * add comment * fix CI * add enable_cache for docker conf * CI * use setting data path * fix ui bug --------- Co-authored-by: yuanteli <1957922024@qq.com> --- rdagent/app/data_science/loop.py | 8 +- rdagent/core/evolving_agent.py | 82 ++++----- rdagent/log/mle_summary.py | 137 +++++++++++++++ rdagent/log/ui/app.py | 8 + rdagent/log/ui/dsapp.py | 285 +++++++++++++++++++++++++++++++ rdagent/log/ui/llm_st.py | 63 ++++++- rdagent/utils/env.py | 8 +- rdagent/utils/workflow.py | 44 ++--- 8 files changed, 565 insertions(+), 70 deletions(-) create mode 100644 rdagent/log/mle_summary.py create mode 100644 rdagent/log/ui/dsapp.py diff --git a/rdagent/app/data_science/loop.py b/rdagent/app/data_science/loop.py index 1a9045ed..3d1f8449 100644 --- a/rdagent/app/data_science/loop.py +++ b/rdagent/app/data_science/loop.py @@ -60,7 +60,7 @@ class DataScienceRDLoop(RDLoop): def direct_exp_gen(self, prev_out: dict[str, Any]): exp = self.exp_gen.gen(self.trace) - logger.log_object(exp, tag="direct_exp_gen") + logger.log_object(exp) # FIXME: this is for LLM debug webapp, remove this when the debugging is done. logger.log_object(exp, tag="debug_exp_gen") @@ -83,14 +83,14 @@ class DataScienceRDLoop(RDLoop): else: raise NotImplementedError(f"Unsupported component in DataScienceRDLoop: {exp.hypothesis.component}") exp.sub_tasks = [] - logger.log_object(exp, tag="coding") + logger.log_object(exp) return exp def running(self, prev_out: dict[str, Any]): exp: DSExperiment = prev_out["coding"] if exp.next_component_required() is None: new_exp = self.runner.develop(exp) - logger.log_object(new_exp, tag="running") + logger.log_object(new_exp) return new_exp else: return exp @@ -104,7 +104,7 @@ class DataScienceRDLoop(RDLoop): reason=f"{exp.hypothesis.component} is completed.", decision=True, ) - logger.log_object(feedback, tag="feedback") + logger.log_object(feedback) return feedback def record(self, prev_out: dict[str, Any]): diff --git a/rdagent/core/evolving_agent.py b/rdagent/core/evolving_agent.py index eded60b3..4d7e8e5c 100644 --- a/rdagent/core/evolving_agent.py +++ b/rdagent/core/evolving_agent.py @@ -58,51 +58,51 @@ class RAGEvoAgent(EvoAgent): eva: Evaluator | Feedback, filter_final_evo: bool = False, ) -> EvolvableSubjects: - for _ in tqdm(range(self.max_loop), "Implementing"): - # with logger.tag(f"evo_loop_{evo_loop_id}"): - # 1. knowledge self-evolving - if self.knowledge_self_gen and self.rag is not None: - self.rag.generate_knowledge(self.evolving_trace) - # 2. RAG - queried_knowledge = None - if self.with_knowledge and self.rag is not None: - # TODO: Putting the evolving trace in here doesn't actually work - queried_knowledge = self.rag.query(evo, self.evolving_trace) + for evo_loop_id in tqdm(range(self.max_loop), "Implementing"): + with logger.tag(f"evo_loop_{evo_loop_id}"): + # 1. knowledge self-evolving + if self.knowledge_self_gen and self.rag is not None: + self.rag.generate_knowledge(self.evolving_trace) + # 2. RAG + queried_knowledge = None + if self.with_knowledge and self.rag is not None: + # TODO: Putting the evolving trace in here doesn't actually work + queried_knowledge = self.rag.query(evo, self.evolving_trace) - # 3. evolve - evo = self.evolving_strategy.evolve( - evo=evo, - evolving_trace=self.evolving_trace, - queried_knowledge=queried_knowledge, - ) - # TODO: Due to design issues, we have chosen to ignore this mypy error. - logger.log_object(evo.sub_workspace_list, tag="evolving code") # type: ignore[attr-defined] - for sw in evo.sub_workspace_list: # type: ignore[attr-defined] - logger.info(f"evolving code workspace: {sw}") - - # 4. Pack evolve results - es = EvoStep(evo, queried_knowledge) - - # 5. Evaluation - if self.with_feedback: - es.feedback = ( - # TODO: Due to the irregular design of rdagent.core.evaluation.Evaluator, - # it fails mypy's test here, so we'll ignore this error for now. - eva - if isinstance(eva, Feedback) - else eva.evaluate(evo, queried_knowledge=queried_knowledge) # type: ignore[arg-type, call-arg] + # 3. evolve + evo = self.evolving_strategy.evolve( + evo=evo, + evolving_trace=self.evolving_trace, + queried_knowledge=queried_knowledge, ) - logger.log_object(es.feedback, tag="evolving feedback") + # TODO: Due to design issues, we have chosen to ignore this mypy error. + logger.log_object(evo.sub_workspace_list, tag="evolving code") # type: ignore[attr-defined] + for sw in evo.sub_workspace_list: # type: ignore[attr-defined] + logger.info(f"evolving code workspace: {sw}") - # 6. update trace - self.evolving_trace.append(es) + # 4. Pack evolve results + es = EvoStep(evo, queried_knowledge) - # 7. check if all tasks are completed - if self.with_feedback: - all_completed = all(es.feedback) if isinstance(es.feedback, list) else es.feedback - if all_completed: - logger.info("All tasks in evolving subject have been completed.") - break + # 5. Evaluation + if self.with_feedback: + es.feedback = ( + # TODO: Due to the irregular design of rdagent.core.evaluation.Evaluator, + # it fails mypy's test here, so we'll ignore this error for now. + eva + if isinstance(eva, Feedback) + else eva.evaluate(evo, queried_knowledge=queried_knowledge) # type: ignore[arg-type, call-arg] + ) + logger.log_object(es.feedback, tag="evolving feedback") + + # 6. update trace + self.evolving_trace.append(es) + + # 7. check if all tasks are completed + if self.with_feedback: + all_completed = all(es.feedback) if isinstance(es.feedback, list) else es.feedback + if all_completed: + logger.info("All tasks in evolving subject have been completed.") + break if self.with_feedback and filter_final_evo: evo = self.filter_evolvable_subjects_by_feedback(evo, self.evolving_trace[-1].feedback) diff --git a/rdagent/log/mle_summary.py b/rdagent/log/mle_summary.py new file mode 100644 index 00000000..383c4393 --- /dev/null +++ b/rdagent/log/mle_summary.py @@ -0,0 +1,137 @@ +import json +import re +from collections import defaultdict +from pathlib import Path + +import fire +import pandas as pd + +from rdagent.app.data_science.conf import DS_RD_SETTING +from rdagent.log.storage import FileStorage +from rdagent.utils.env import DockerEnv, MLEBDockerConf + +mle_de_conf = MLEBDockerConf() +mle_de_conf.extra_volumes = { + f"{DS_RD_SETTING.local_data_path}/zip_files": "/mle/data", +} +de = DockerEnv(conf=mle_de_conf) +de.prepare() + + +def extract_mle_json(log_content): + match = re.search(r"\{.*\}", log_content, re.DOTALL) + if match: + return json.loads(match.group(0)) + return None + + +def save_grade_info(log_trace_path: Path): + for msg in FileStorage(log_trace_path).iter_msg(): + if "competition" in msg.tag: + competition = msg.content + + if "running" in msg.tag: + msg.content.experiment_workspace.execute( + env=de, + entry=f"bash -c 'mlebench grade-sample submission.csv {competition} --data-dir /mle/data > mle_score.txt 2>&1'", + ) + msg.content.experiment_workspace.execute(env=de, entry="chmod 777 mle_score.txt") + + +def save_all_grade_info(log_folder): + for log_trace_path in log_folder.iterdir(): + save_grade_info(log_trace_path) + + +def summarize_folder(log_folder: Path): + stat = defaultdict(dict) + for log_trace_path in log_folder.iterdir(): # One log trace + if not log_trace_path.is_dir(): + continue + loop_num = 0 + made_submission_num = 0 + test_scores = {} + valid_scores = {} + medal = "None" + success_loop_num = 0 + + for msg in FileStorage(log_trace_path).iter_msg(): # messages in log trace + if "competition" in msg.tag: + stat[log_trace_path.name]["competition"] = msg.content + + if "direct_exp_gen" in msg.tag: + loop_num += 1 + if "running" in msg.tag: + + submission_path = msg.content.experiment_workspace.workspace_path / "submission.csv" + if submission_path.exists(): + made_submission_num += 1 + scores_path = msg.content.experiment_workspace.workspace_path / "scores.csv" + valid_scores[loop_num - 1] = pd.read_csv(scores_path, index_col=0) + grade_output_path = msg.content.experiment_workspace.workspace_path / "mle_score.txt" + if not grade_output_path.exists(): + raise FileNotFoundError(f"mle_score.txt in {grade_output_path} not found, genarate it first!") + grade_output = extract_mle_json(grade_output_path.read_text()) + if grade_output["score"] is not None: + test_scores[loop_num - 1] = grade_output["score"] + if grade_output["any_medal"]: + medal = ( + "gold" + if grade_output["gold_medal"] + else "silver" if grade_output["silver_medal"] else "bronze" + ) + + if "feedback" in msg.tag and "evolving" not in msg.tag: + if bool(msg.content): + success_loop_num += 1 + + stat[log_trace_path.name].update( + { + "loop_num": loop_num, + "made_submission_num": made_submission_num, + "test_scores": test_scores, + "valid_scores": valid_scores, + "medal": medal, + "success_loop_num": success_loop_num, + } + ) + if (log_folder / "summary.pkl").exists(): + (log_folder / "summary.pkl").unlink() + print("Old summary file removed.") + pd.to_pickle(stat, log_folder / "summary.pkl") + + +# { +# "competition_id": "stanford-covid-vaccine", +# "score": null, +# "gold_threshold": 0.34728, +# "silver_threshold": 0.35175, +# "bronze_threshold": 0.3534, +# "median_threshold": 0.363095, +# "any_medal": false, +# "gold_medal": false, +# "silver_medal": false, +# "bronze_medal": false, +# "above_median": false, +# "submission_exists": true, +# "valid_submission": false, +# "is_lower_better": true, +# "created_at": "2025-01-21T11:59:33.788201", +# "submission_path": "submission.csv" +# } + + +def grade_summary(log_folder): + log_folder = Path(log_folder) + save_all_grade_info(log_folder) + summarize_folder(log_folder) + + +if __name__ == "__main__": + fire.Fire( + { + "grade": save_all_grade_info, + "summary": summarize_folder, + "grade_summary": grade_summary, + } + ) diff --git a/rdagent/log/ui/app.py b/rdagent/log/ui/app.py index 04f49ed2..39bfe0e7 100644 --- a/rdagent/log/ui/app.py +++ b/rdagent/log/ui/app.py @@ -1,4 +1,5 @@ import argparse +import re import textwrap from collections import defaultdict from datetime import datetime, timezone @@ -143,6 +144,13 @@ def get_msgs_until(end_func: Callable[[Message], bool] = lambda _: True): while True: try: msg = next(state.fs) + + # new scenario gen this tags, old version UI not have these tags. + msg.tag = re.sub(r"\.evo_loop_\d+", "", msg.tag) + msg.tag = re.sub(r"Loop_\d+\.[^.]+", "", msg.tag) + msg.tag = re.sub(r"\.\.", ".", msg.tag) + msg.tag = msg.tag.strip(".") + if should_display(msg): tags = msg.tag.split(".") if "r" not in state.current_tags and "r" in tags: diff --git a/rdagent/log/ui/dsapp.py b/rdagent/log/ui/dsapp.py new file mode 100644 index 00000000..5786e4a9 --- /dev/null +++ b/rdagent/log/ui/dsapp.py @@ -0,0 +1,285 @@ +import json +import re +from collections import defaultdict +from pathlib import Path + +import pandas as pd +import plotly.express as px +import plotly.graph_objects as go +import streamlit as st +from plotly.subplots import make_subplots +from streamlit import session_state as state + +from rdagent.log.mle_summary import extract_mle_json +from rdagent.log.storage import FileStorage + +st.set_page_config(layout="wide", page_title="RD-Agent", page_icon="🎓", initial_sidebar_state="expanded") + +# 设置主日志路径 +if "log_folder" not in state: + state.log_folder = Path("./log") +if "log_path" not in state: + state.log_path = None + + +# @st.cache_data +def load_data(log_path): + data = defaultdict(lambda: defaultdict(dict)) + li = -1 # loop id + ei = -1 # evo id + for msg in FileStorage(state.log_folder / log_path).iter_msg(): + if msg.tag and "llm" not in msg.tag and "session" not in msg.tag: + if msg.tag == "competition": + data["competition"] = msg.content + continue + if msg.tag == "direct_exp_gen": + li += 1 + ei = -1 + + if "evolving " in msg.tag: + if "evolving code" in msg.tag: + ei += 1 + data[li][ei][msg.tag] = msg.content + else: + data[li][msg.tag] = msg.content + + return data + + +@st.cache_data +def get_folders_sorted(log_path): + """缓存并返回排序后的文件夹列表,并加入进度打印""" + with st.spinner("正在加载文件夹列表..."): + folders = sorted( + (folder for folder in log_path.iterdir() if folder.is_dir() and list(folder.iterdir())), + key=lambda folder: folder.stat().st_mtime, + reverse=True, + ) + st.write(f"找到 {len(folders)} 个文件夹") + return [folder.name for folder in folders] + + +# UI - Sidebar +with st.sidebar: + state.log_folder = Path(st.text_input("**Log Folder**", placeholder=state.log_folder, value=state.log_folder)) + if not state.log_folder.exists(): + st.warning(f"Path {state.log_folder} does not exist!") + folders = get_folders_sorted(state.log_folder) + st.selectbox(f"Select from :blue[**{state.log_folder.absolute()}**]", folders, key="log_path") + + if st.button("Refresh Data"): + if state.log_path is None: + st.toast("Please select a log path first!", type="error") + st.stop() + + state.data = load_data(state.log_path) + st.rerun() + + show_all_summary = st.toggle("One Trace / Log Folder Summary", value=True) + + +# UI windows +def task_win(data): + with st.container(border=True): + st.markdown(f"**:violet[{data.name}]**") + st.markdown(data.description) + if hasattr(data, "architecture"): # model task + st.markdown( + f""" + | Model_type | Architecture | hyperparameters | + |------------|--------------|-----------------| + | {data.model_type} | {data.architecture} | {data.hyperparameters} | + """ + ) + + +def workspace_win(data): + show_files = {k: v for k, v in data.file_dict.items() if not "test" in k} + if len(show_files) > 0: + with st.expander(f"Files in :blue[{data.workspace_path}]"): + code_tabs = st.tabs(show_files.keys()) + for ct, codename in zip(code_tabs, show_files.keys()): + with ct: + st.code( + show_files[codename], + language=("python" if codename.endswith(".py") else "markdown"), + wrap_lines=True, + ) + else: + st.markdown("No files in the workspace") + + +def exp_gen_win(data): + st.header("Exp Gen", divider="blue") + st.subheader("Hypothesis") + st.markdown(data.hypothesis) + + st.subheader("pending_tasks") + for tasks in data.pending_tasks_list: + task_win(tasks[0]) + st.subheader("Exp Workspace", anchor="exp-workspace") + workspace_win(data.experiment_workspace) + + +def evolving_win(data): + st.header("Code Evolving", divider="green") + if len(data) > 1: + evo_id = st.slider("Evolving", 0, len(data) - 1, 0) + else: + evo_id = 0 + + if evo_id in data: + st.subheader("codes") + workspace_win(data[evo_id]["evolving code"][0]) + fb = data[evo_id]["evolving feedback"][0] + st.subheader("evolving feedback" + ("✅" if bool(fb) else "❌"), anchor="c_feedback") + f1, f2, f3 = st.tabs(["execution", "return_checking", "code"]) + f1.code(fb.execution, wrap_lines=True) + f2.code(fb.return_checking, wrap_lines=True) + f3.code(fb.code, wrap_lines=True) + else: + st.markdown("No evolving.") + + +def exp_after_coding_win(data): + st.header("Exp After Coding", divider="blue") + st.subheader("Exp Workspace", anchor="eac-exp-workspace") + workspace_win(data.experiment_workspace) + + +def exp_after_running_win(data, mle_score): + st.header("Exp After Running", divider="blue") + st.subheader("Exp Workspace", anchor="ear-exp-workspace") + workspace_win(data.experiment_workspace) + st.subheader("Result") + st.write(data.result) + st.subheader("MLE Submission Score") + st.json(mle_score) + + +def feedback_win(data): + st.header("Feedback" + ("✅" if bool(data) else "❌"), divider="orange") + st.code(data, wrap_lines=True) + if data.exception is not None: + st.markdown(f"**:red[Exception]**: {data.exception}") + + +def sota_win(data): + st.header("SOTA Experiment", divider="rainbow") + if data: + st.subheader("Exp Workspace", anchor="sota-exp-workspace") + workspace_win(data.experiment_workspace) + else: + st.markdown("No SOTA experiment.") + + +def main_win(data): + exp_gen_win(data["direct_exp_gen"]) + evo_data = {k: v for k, v in data.items() if isinstance(k, int)} + evolving_win(evo_data) + if "coding" in data: + exp_after_coding_win(data["coding"]) + if "running" in data: + exp_after_running_win(data["running"], data["mle_score"]) + if "feedback" in data: + feedback_win(data["feedback"]) + sota_win(data["SOTA experiment"]) + + with st.sidebar: + st.markdown( + f""" +- [Exp Gen](#exp-gen) + - [Hypothesis](#hypothesis) + - [pending_tasks](#pending-tasks) + - [Exp Workspace](#exp-workspace) +- [Code Evolving ({len(evo_data)})](#code-evolving) + - [codes](#codes) + - [evolving feedback](#c_feedback) +{"- [Exp After Coding](#exp-after-coding)" if "coding" in data else ""} +{"- [Exp After Running](#exp-after-running)" if "running" in data else ""} +{"- [Feedback](#feedback)" if "feedback" in data else ""} +- [SOTA Experiment](#sota-experiment) +""" + ) + + +def summarize_data(): + st.header("Summary", divider="rainbow") + df = pd.DataFrame(columns=["Component", "Running Score", "Feedback"], index=range(len(state.data) - 1)) + + for loop in range(len(state.data) - 1): + loop_data = state.data[loop] + df.loc[loop, "Component"] = loop_data["direct_exp_gen"].hypothesis.component + + if "running" in loop_data: + if "mle_score" not in state.data[loop]: + mle_score_path = loop_data["running"].experiment_workspace.workspace_path / "mle_score.txt" + try: + state.data[loop]["mle_score"] = extract_mle_json(mle_score_path.read_text()) + df.loc[loop, "Running Score"] = str(state.data[loop]["mle_score"]["score"]) + except Exception as e: + state.data[loop]["mle_score"] = str(e) + df.loc[loop, "Running Score"] = "❌" + else: + df.loc[loop, "Running Score"] = "N/A" + + if "feedback" in loop_data: + df.loc[loop, "Feedback"] = "✅" if bool(loop_data["feedback"]) else "❌" + else: + df.loc[loop, "Feedback"] = "N/A" + st.dataframe(df) + + +def all_summarize_win(): + if not (state.log_folder / "summary.pkl").exists(): + st.warning( + f"No summary file found in {state.log_folder}\nRun:`dotenv run -- python rdagent/log/mle_summary.py grade_summary --log_folder=`" + ) + return + summary = pd.read_pickle(state.log_folder / "summary.pkl") + base_df = pd.DataFrame( + columns=["Competition", "Total Loops", "Made Submission", "Successful Final Decision", "Medal"], + index=summary.keys(), + ) + for k, v in summary.items(): + loop_num = v["loop_num"] + base_df.loc[k, "Competition"] = v["competition"] + base_df.loc[k, "Total Loops"] = loop_num + base_df.loc[k, "Made Submission"] = ( + f"{v['made_submission_num']} ({round(v['made_submission_num'] / loop_num * 100, 2)}%)" + ) + base_df.loc[k, "Successful Final Decision"] = ( + f"{v['success_loop_num']} ({round(v['success_loop_num'] / loop_num * 100, 2)}%)" + ) + base_df.loc[k, "Medal"] = v["medal"] + st.dataframe(base_df) + # write curve + for k, v in summary.items(): + with st.container(border=True): + st.markdown(f"**:blue[{k}] - :violet[{v['competition']}]**") + vscores = {k: v.iloc[:, 0] for k, v in v["valid_scores"].items()} + if len(vscores) > 0: + metric_name = list(vscores.values())[0].name + else: + metric_name = "None" + + fc1, fc2 = st.columns(2) + vdf = pd.DataFrame(vscores) + vdf.columns = [f"loop {i}" for i in vdf.columns] + f1 = px.line(vdf.T, markers=True, title=f"Valid scores (metric: {metric_name})") + fc1.plotly_chart(f1, key=f"{k}_v") + + tscores = {f"loop {k}": v for k, v in v["test_scores"].items()} + tdf = pd.Series(tscores, name="score") + f2 = px.line(tdf, markers=True, title="Test scores") + fc2.plotly_chart(f2, key=k) + + +# UI - Main +if show_all_summary: + all_summarize_win() +elif "data" in state: + st.title(state.data["competition"]) + summarize_data() + loop_id = st.slider("Loop", 0, len(state.data) - 2, 0) + main_win(state.data[loop_id]) diff --git a/rdagent/log/ui/llm_st.py b/rdagent/log/ui/llm_st.py index 9c05441c..e0c270ba 100644 --- a/rdagent/log/ui/llm_st.py +++ b/rdagent/log/ui/llm_st.py @@ -72,8 +72,6 @@ with st.sidebar: load_data() st.rerun() - expand_all = st.toggle("Expand All", key="expand_all") - # Helper functions def show_text(text, lang=None): @@ -126,6 +124,67 @@ sorted_loop_ids = sorted(loop_groups.keys(), key=int) # 假设 Loop ID 是数 total_loops = len(sorted_loop_ids) total_pages = total_loops # 每页展示一个 Loop + +# simple display +# FIXME: Delete this simple UI if trace have tag(evo_id & loop_id) +# with st.sidebar: +# start = int(st.text_input("start", 0)) +# end = int(st.text_input("end", 100)) +# for m in session_state.data[start:end]: +# if "tpl" in m["tag"]: +# obj = m["obj"] +# uri = obj["uri"] +# tpl = obj["template"] +# cxt = obj["context"] +# rd = obj["rendered"] +# with st.expander(highlight_prompts_uri(uri), expanded=False, icon="⚙️"): +# t1, t2, t3 = st.tabs([":green[**Rendered**]", ":blue[**Template**]", ":orange[**Context**]"]) +# with t1: +# show_text(rd) +# with t2: +# show_text(tpl, lang="django") +# with t3: +# st.json(cxt) +# if "llm" in m["tag"]: +# obj = m["obj"] +# system = obj.get("system", None) +# user = obj["user"] +# resp = obj["resp"] +# with st.expander(f"**LLM**", expanded=False, icon="🤖"): +# t1, t2, t3 = st.tabs([":green[**Response**]", ":blue[**User**]", ":orange[**System**]"]) +# with t1: +# try: +# rdict = json.loads(resp) +# if "code" in rdict: +# code = rdict["code"] +# st.markdown(":red[**Code in response dict:**]") +# st.code(code, language="python", wrap_lines=True, line_numbers=True) +# rdict.pop("code") +# elif "spec" in rdict: +# spec = rdict["spec"] +# st.markdown(":red[**Spec in response dict:**]") +# st.markdown(spec) +# rdict.pop("spec") +# else: +# # show model codes +# showed_keys = [] +# for k, v in rdict.items(): +# if k.startswith("model_") and k.endswith(".py"): +# st.markdown(f":red[**{k}**]") +# st.code(v, language="python", wrap_lines=True, line_numbers=True) +# showed_keys.append(k) +# for k in showed_keys: +# rdict.pop(k) +# st.write(":red[**Other parts (except for the code or spec) in response dict:**]") +# st.json(rdict) +# except: +# st.json(resp) +# with t2: +# show_text(user) +# with t3: +# show_text(system or "No system prompt available") + + if total_pages: # 初始化 current_loop if "current_loop" not in st.session_state: diff --git a/rdagent/utils/env.py b/rdagent/utils/env.py index 811aef8a..23d89f9d 100644 --- a/rdagent/utils/env.py +++ b/rdagent/utils/env.py @@ -144,6 +144,8 @@ class DockerConf(ExtendedBaseSettings): running_timeout_period: int = 3600 # 1 hour + enable_cache: bool = True # enable the cache mechanism + class QlibDockerConf(DockerConf): model_config = ExtendedSettingsConfigDict(env_prefix="QLIB_DOCKER_") @@ -227,6 +229,7 @@ class MLEBDockerConf(DockerConf): mem_limit: str | None = ( "48g" # Add memory limit attribute # new-york-city-taxi-fare-prediction may need more memory ) + enable_cache: bool = False # physionet.org/files/mimic-eicu-fiddle-feature/1.0.0/FIDDLE_mimic3 @@ -458,7 +461,10 @@ class DockerEnv(Env[DockerConf]): ) start = time.time() - out = self.cached_run(entry_add_timeout, local_path, env, running_extra_volume) + if self.conf.enable_cache: + out = self.cached_run(entry_add_timeout, local_path, env, running_extra_volume) + else: + out = self.__run(entry, local_path, env, running_extra_volume, remove_timestamp=False) end = time.time() if end - start + 1 >= self.conf.running_timeout_period: diff --git a/rdagent/utils/workflow.py b/rdagent/utils/workflow.py index e2e83b7a..11305916 100644 --- a/rdagent/utils/workflow.py +++ b/rdagent/utils/workflow.py @@ -111,29 +111,29 @@ class LoopBase: li, si = self.loop_idx, self.step_idx name = self.steps[si] - # with logger.tag(f"Loop_{li}.{name}"): - start = datetime.datetime.now(datetime.timezone.utc) - func = getattr(self, name) - try: - self.loop_prev_out[name] = func(self.loop_prev_out) - # TODO: Fix the error logger.exception(f"Skip loop {li} due to {e}") - except self.skip_loop_error as e: - # FIXME: This does not support previous demo (due to their last step is not for recording) - logger.warning(f"Skip loop {li} due to {e}") - # NOTE: strong assumption! The last step is responsible for recording information - self.step_idx = len(self.steps) - 1 # directly jump to the last step. - self.loop_prev_out[self.EXCEPTION_KEY] = e - continue - finally: - # make sure failure steps are displayed correclty - end = datetime.datetime.now(datetime.timezone.utc) - self.loop_trace[li].append(LoopTrace(start, end, step_idx=si)) + with logger.tag(f"Loop_{li}.{name}"): + start = datetime.datetime.now(datetime.timezone.utc) + func = getattr(self, name) + try: + self.loop_prev_out[name] = func(self.loop_prev_out) + # TODO: Fix the error logger.exception(f"Skip loop {li} due to {e}") + except self.skip_loop_error as e: + # FIXME: This does not support previous demo (due to their last step is not for recording) + logger.warning(f"Skip loop {li} due to {e}") + # NOTE: strong assumption! The last step is responsible for recording information + self.step_idx = len(self.steps) - 1 # directly jump to the last step. + self.loop_prev_out[self.EXCEPTION_KEY] = e + continue + finally: + # make sure failure steps are displayed correclty + end = datetime.datetime.now(datetime.timezone.utc) + self.loop_trace[li].append(LoopTrace(start, end, step_idx=si)) - # Update tqdm progress bar directly to step_idx - pbar.n = si + 1 - pbar.set_postfix( - loop_index=li, step_index=si + 1, step_name=name - ) # step_name indicate last finished step_name + # Update tqdm progress bar directly to step_idx + pbar.n = si + 1 + pbar.set_postfix( + loop_index=li, step_index=si + 1, step_name=name + ) # step_name indicate last finished step_name # index increase and save session self.step_idx = (self.step_idx + 1) % len(self.steps)