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>
This commit is contained in:
XianBW
2025-01-23 16:12:22 +08:00
committed by GitHub
parent 3717cd3bcd
commit bd8122e0cd
8 changed files with 565 additions and 70 deletions
+4 -4
View File
@@ -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]):
+41 -41
View File
@@ -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)
+137
View File
@@ -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,
}
)
+8
View File
@@ -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:
+285
View File
@@ -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=<your trace 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])
+61 -2
View File
@@ -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:
+7 -1
View File
@@ -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:
+22 -22
View File
@@ -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)