mirror of
https://github.com/NicolasBohn/NexQuant.git
synced 2026-07-27 23:47:46 +00:00
feat: streamlit webapp demo for different scenarios (#135)
add streamlit webapp demo & docs
This commit is contained in:
Vendored
BIN
Binary file not shown.
|
After Width: | Height: | Size: 123 KiB |
@@ -5,14 +5,39 @@ Demo and Introduction
|
||||
Introduction
|
||||
============
|
||||
|
||||
TODO: copy the content in the README.md
|
||||
RD-Agent will generate some logs during the R&D process. These logs are very useful for debugging and understanding the R&D process. However, just viewing the terminal log is not intuitive enough. RD-Agent provides a web app to visualize the R&D process. You can easily view the R&D process and understand the R&D process better.
|
||||
|
||||
A Quick Demo
|
||||
============
|
||||
|
||||
TODO:
|
||||
- copy the demo content in the README.md
|
||||
- Quick start for the demo.
|
||||
Start Web App
|
||||
-------------
|
||||
|
||||
In `RD-Agent/` folder, run:
|
||||
|
||||
TODO: link to more demos
|
||||
.. code-block:: bash
|
||||
|
||||
streamlit run rdagent/log/ui/app.py --server.port <port> -- --log_dir <log_dir>
|
||||
|
||||
This will start a web app on `http://localhost:<port>`.
|
||||
|
||||
**NOTE**: The log_dir parameter is not required. You can manually enter the log_path in the web app. If you set the log_dir parameter, you can easily select a different log_path in the web app.
|
||||
|
||||
Use Web App
|
||||
-----------
|
||||
|
||||
1. Open the sidebar.
|
||||
|
||||
2. Select the scenario you want to show. There are some pre-defined scenarios:
|
||||
- Qlib Model
|
||||
- Qlib Factor
|
||||
- Data Mining
|
||||
- Model from Paper
|
||||
|
||||
3. Click the `Config⚙️` button and input the log path (if you set the log_dir parameter, you can select a log_path in the dropdown list).
|
||||
|
||||
4. Click the buttons below Config⚙️ to show the scenario execution process. Buttons are:
|
||||
- All Loops: Show complete scenario execution process.
|
||||
- Next Loop: Show one success **R&D Loop**.
|
||||
- One Evolving: Show one **evolving** step of **development** part.
|
||||
- refresh logs: clear shown logs.
|
||||
@@ -31,10 +31,10 @@ class GeneralModelScenario(Scenario):
|
||||
@property
|
||||
def rich_style_description(self) -> str:
|
||||
return """
|
||||
# [Model Research & Development Co-Pilot] (#_scenario)
|
||||
|
||||
## [Overview](#_summary)
|
||||
|
||||
### [Model Research & Development Co-Pilot](#_scenario)
|
||||
|
||||
#### [Overview](#_summary)
|
||||
|
||||
This demo automates the extraction and development of PyTorch models from academic papers. It supports various model types through two main components: Reader and Coder.
|
||||
|
||||
#### [Workflow Components](#_rdloops)
|
||||
|
||||
@@ -1,36 +0,0 @@
|
||||
# start web app
|
||||
|
||||
in `RD-Agent/` folder, run `streamlit run rdagent/log/ui/app.py --server.port 14000`
|
||||
|
||||
## custom windows
|
||||
|
||||
in `rdagent.log.ui.web.py`
|
||||
|
||||
StWindow is the base window, can show common log messages like logger in terminal.
|
||||
|
||||
Windows can be nested.
|
||||
|
||||
Simply, just write a window for one log object type, and use `isinstance()` to display messages in different windows.
|
||||
|
||||
### some base windows
|
||||
|
||||
- StWindow
|
||||
- LLMWindow
|
||||
- CodeWindow
|
||||
|
||||
### multi tabs window
|
||||
|
||||
More convenient to use nested windows
|
||||
|
||||
- ProgressTabsWindow
|
||||
- ObjectsTabsWindow
|
||||
|
||||
### main trace window
|
||||
|
||||
- QlibFactorTraceWindow
|
||||
|
||||
## TODOS
|
||||
|
||||
- make it like a living trace
|
||||
- display real living trace
|
||||
- Window Styles
|
||||
+482
-262
@@ -1,6 +1,8 @@
|
||||
import time
|
||||
import argparse
|
||||
import textwrap
|
||||
from collections import defaultdict
|
||||
from datetime import datetime, timezone
|
||||
from pathlib import Path
|
||||
from typing import Callable, Type
|
||||
|
||||
import pandas as pd
|
||||
@@ -8,9 +10,11 @@ import plotly.express as px
|
||||
import plotly.graph_objects as go
|
||||
import streamlit as st
|
||||
from plotly.subplots import make_subplots
|
||||
from st_btn_select import st_btn_select
|
||||
from streamlit import session_state as state
|
||||
from streamlit.delta_generator import DeltaGenerator
|
||||
|
||||
from rdagent.app.model_extraction_and_code.GeneralModel import GeneralModelScenario
|
||||
from rdagent.components.coder.factor_coder.CoSTEER.evaluators import (
|
||||
FactorSingleFeedback,
|
||||
)
|
||||
@@ -21,20 +25,47 @@ from rdagent.core.proposal import Hypothesis, HypothesisFeedback
|
||||
from rdagent.log.base import Message
|
||||
from rdagent.log.storage import FileStorage
|
||||
from rdagent.log.ui.qlib_report_figure import report_figure
|
||||
from rdagent.scenarios.qlib.experiment.factor_experiment import QlibFactorExperiment
|
||||
from rdagent.scenarios.data_mining.experiment.model_experiment import DMModelScenario
|
||||
from rdagent.scenarios.qlib.experiment.factor_experiment import (
|
||||
QlibFactorExperiment,
|
||||
QlibFactorScenario,
|
||||
)
|
||||
from rdagent.scenarios.qlib.experiment.model_experiment import (
|
||||
QlibModelExperiment,
|
||||
QlibModelScenario,
|
||||
)
|
||||
|
||||
st.set_page_config(layout="wide")
|
||||
st.set_page_config(layout="wide", page_title="RD-Agent", page_icon="🎓", initial_sidebar_state="expanded")
|
||||
|
||||
|
||||
if "log_path" not in state:
|
||||
state.log_path = ""
|
||||
# 获取log_path参数
|
||||
parser = argparse.ArgumentParser(description="RD-Agent Streamlit App")
|
||||
parser.add_argument("--log_dir", type=str, help="Path to the log directory")
|
||||
args = parser.parse_args()
|
||||
if args.log_dir:
|
||||
main_log_path = Path(args.log_dir)
|
||||
if not main_log_path.exists():
|
||||
st.error(f"Log dir `{main_log_path}` does not exist!")
|
||||
st.stop()
|
||||
else:
|
||||
main_log_path = None
|
||||
|
||||
|
||||
SELECTED_METRICS = [
|
||||
"IC",
|
||||
"1day.excess_return_without_cost.annualized_return",
|
||||
"1day.excess_return_without_cost.information_ratio",
|
||||
"1day.excess_return_without_cost.max_drawdown",
|
||||
]
|
||||
|
||||
if "log_type" not in state:
|
||||
state.log_type = "qlib_model"
|
||||
state.log_type = "Qlib Model"
|
||||
|
||||
if "log_path" not in state:
|
||||
if main_log_path:
|
||||
state.log_path = next(main_log_path.iterdir()).relative_to(main_log_path)
|
||||
else:
|
||||
state.log_path = ""
|
||||
|
||||
if "fs" not in state:
|
||||
state.fs = None
|
||||
@@ -54,6 +85,9 @@ if "lround" not in state:
|
||||
if "erounds" not in state:
|
||||
state.erounds = defaultdict(int) # Evolving Rounds in each RD Loop
|
||||
|
||||
if "e_decisions" not in state:
|
||||
state.e_decisions = defaultdict(lambda: defaultdict(tuple))
|
||||
|
||||
# Summary Info
|
||||
if "hypotheses" not in state:
|
||||
# Hypotheses in each RD Loop
|
||||
@@ -65,17 +99,26 @@ if "h_decisions" not in state:
|
||||
if "metric_series" not in state:
|
||||
state.metric_series = []
|
||||
|
||||
# Factor Task Baseline
|
||||
if "alpha158_metrics" not in state:
|
||||
state.alpha158_metrics = None
|
||||
|
||||
|
||||
def refresh():
|
||||
state.fs = FileStorage(state.log_path).iter_msg()
|
||||
if main_log_path:
|
||||
state.fs = FileStorage(main_log_path / state.log_path).iter_msg()
|
||||
else:
|
||||
state.fs = FileStorage(state.log_path).iter_msg()
|
||||
state.msgs = defaultdict(lambda: defaultdict(list))
|
||||
state.lround = 0
|
||||
state.erounds = defaultdict(int)
|
||||
state.e_decisions = defaultdict(lambda: defaultdict(tuple))
|
||||
state.hypotheses = defaultdict(None)
|
||||
state.h_decisions = defaultdict(bool)
|
||||
state.metric_series = []
|
||||
state.last_msg = None
|
||||
state.current_tags = []
|
||||
state.alpha158_metrics = None
|
||||
|
||||
|
||||
def should_display(msg: Message):
|
||||
@@ -103,54 +146,341 @@ def get_msgs_until(end_func: Callable[[Message], bool] = lambda _: True):
|
||||
|
||||
state.current_tags = tags
|
||||
state.last_msg = msg
|
||||
state.msgs[state.lround][msg.tag].append(msg)
|
||||
|
||||
# Update Summary Info
|
||||
if "model runner result" in tags or "factor runner result" in tags or "runner result" in tags:
|
||||
# factor baseline exp metrics
|
||||
if state.log_type == "Qlib Factor" and state.alpha158_metrics is None:
|
||||
sms = msg.content.based_experiments[0].result.loc[SELECTED_METRICS]
|
||||
sms.name = "alpha158"
|
||||
state.alpha158_metrics = sms
|
||||
|
||||
# common metrics
|
||||
if msg.content.result is None:
|
||||
state.metric_series.append(pd.Series([None], index=["AUROC"]))
|
||||
state.metric_series.append(pd.Series([None], index=["AUROC"], name=f"Round {state.lround}"))
|
||||
else:
|
||||
if msg.content.result.name == "AUROC":
|
||||
if len(msg.content.result) < 4:
|
||||
ps = msg.content.result
|
||||
ps.index = ["AUROC"]
|
||||
ps.name = f"Round {state.lround}"
|
||||
state.metric_series.append(ps)
|
||||
else:
|
||||
state.metric_series.append(
|
||||
msg.content.result.loc[
|
||||
[
|
||||
"IC",
|
||||
"1day.excess_return_without_cost.annualized_return",
|
||||
"1day.excess_return_without_cost.information_ratio",
|
||||
"1day.excess_return_without_cost.max_drawdown",
|
||||
]
|
||||
]
|
||||
)
|
||||
sms = msg.content.result.loc[SELECTED_METRICS]
|
||||
sms.name = f"Round {state.lround}"
|
||||
state.metric_series.append(sms)
|
||||
elif "hypothesis generation" in tags:
|
||||
state.hypotheses[state.lround] = msg.content
|
||||
elif "ef" in tags and "feedback" in tags:
|
||||
state.h_decisions[state.lround] = msg.content.decision
|
||||
elif "d" in tags:
|
||||
if "evolving code" in tags:
|
||||
msg.content = [i for i in msg.content if i]
|
||||
if "evolving feedback" in tags:
|
||||
msg.content = [i for i in msg.content if i]
|
||||
if len(msg.content) != len(state.msgs[state.lround]["d.evolving code"][-1].content):
|
||||
st.toast(":red[**Evolving Feedback Length Error!**]", icon="‼️")
|
||||
right_num = 0
|
||||
for wsf in msg.content:
|
||||
if wsf.final_decision:
|
||||
right_num += 1
|
||||
wrong_num = len(msg.content) - right_num
|
||||
state.e_decisions[state.lround][state.erounds[state.lround]] = (right_num, wrong_num)
|
||||
|
||||
state.msgs[state.lround][msg.tag].append(msg)
|
||||
# Stop Getting Logs
|
||||
if end_func(msg):
|
||||
break
|
||||
except StopIteration:
|
||||
st.toast(":red[**No More Logs to Show!**]", icon="🛑")
|
||||
break
|
||||
|
||||
|
||||
def evolving_feedback_window(wsf: FactorSingleFeedback | ModelCoderFeedback):
|
||||
if isinstance(wsf, FactorSingleFeedback):
|
||||
ffc, efc, cfc, vfc = st.tabs(
|
||||
["**Final Feedback🏁**", "Execution Feedback🖥️", "Code Feedback📄", "Value Feedback🔢"]
|
||||
)
|
||||
with ffc:
|
||||
st.markdown(wsf.final_feedback)
|
||||
with efc:
|
||||
st.code(wsf.execution_feedback, language="log")
|
||||
with cfc:
|
||||
st.markdown(wsf.code_feedback)
|
||||
with vfc:
|
||||
st.markdown(wsf.factor_value_feedback)
|
||||
elif isinstance(wsf, ModelCoderFeedback):
|
||||
ffc, efc, cfc, msfc, vfc = st.tabs(
|
||||
[
|
||||
"**Final Feedback🏁**",
|
||||
"Execution Feedback🖥️",
|
||||
"Code Feedback📄",
|
||||
"Model Shape Feedback📐",
|
||||
"Value Feedback🔢",
|
||||
]
|
||||
)
|
||||
with ffc:
|
||||
st.markdown(wsf.final_feedback)
|
||||
with efc:
|
||||
st.code(wsf.execution_feedback, language="log")
|
||||
with cfc:
|
||||
st.markdown(wsf.code_feedback)
|
||||
with msfc:
|
||||
st.markdown(wsf.shape_feedback)
|
||||
with vfc:
|
||||
st.markdown(wsf.value_feedback)
|
||||
|
||||
|
||||
def display_hypotheses(hypotheses: dict[int, Hypothesis], decisions: dict[int, bool], success_only: bool = False):
|
||||
if success_only:
|
||||
shd = {k: v.__dict__ for k, v in hypotheses.items() if decisions[k]}
|
||||
else:
|
||||
shd = {k: v.__dict__ for k, v in hypotheses.items()}
|
||||
df = pd.DataFrame(shd).T
|
||||
if "reason" in df.columns:
|
||||
df.drop(["reason"], axis=1, inplace=True)
|
||||
df.columns = df.columns.map(lambda x: x.replace("_", " ").capitalize())
|
||||
|
||||
def style_rows(row):
|
||||
if decisions[row.name]:
|
||||
return ["color: green;"] * len(row)
|
||||
return [""] * len(row)
|
||||
|
||||
def style_columns(col):
|
||||
if col.name != "Hypothesis":
|
||||
return ["font-style: italic;"] * len(col)
|
||||
return ["font-weight: bold;"] * len(col)
|
||||
|
||||
# st.dataframe(df.style.apply(style_rows, axis=1).apply(style_columns, axis=0))
|
||||
st.markdown(df.style.apply(style_rows, axis=1).apply(style_columns, axis=0).to_html(), unsafe_allow_html=True)
|
||||
|
||||
|
||||
def metrics_window(df: pd.DataFrame, R: int, C: int, *, height: int = 300, colors: list[str] = None):
|
||||
fig = make_subplots(rows=R, cols=C, subplot_titles=df.columns)
|
||||
|
||||
def hypothesis_hover_text(h: Hypothesis, d: bool = False):
|
||||
color = "green" if d else "black"
|
||||
text = h.hypothesis
|
||||
lines = textwrap.wrap(text, width=60)
|
||||
return f"<span style='color: {color};'>{'<br>'.join(lines)}</span>"
|
||||
|
||||
hover_texts = [
|
||||
hypothesis_hover_text(state.hypotheses[int(i[6:])], state.h_decisions[int(i[6:])])
|
||||
for i in df.index
|
||||
if i != "alpha158"
|
||||
]
|
||||
if state.alpha158_metrics is not None:
|
||||
hover_texts = ["Baseline: alpha158"] + hover_texts
|
||||
for ci, col in enumerate(df.columns):
|
||||
row = ci // C + 1
|
||||
col_num = ci % C + 1
|
||||
fig.add_trace(
|
||||
go.Scatter(
|
||||
x=df.index,
|
||||
y=df[col],
|
||||
name=col,
|
||||
mode="lines+markers",
|
||||
connectgaps=True,
|
||||
marker=dict(size=10, color=colors[ci]) if colors else dict(size=10),
|
||||
hovertext=hover_texts,
|
||||
hovertemplate="%{hovertext}<br><br><span style='color: black'>%{x} Value:</span> <span style='color: blue'>%{y}</span><extra></extra>",
|
||||
),
|
||||
row=row,
|
||||
col=col_num,
|
||||
)
|
||||
fig.update_layout(showlegend=False, height=height)
|
||||
|
||||
if state.alpha158_metrics is not None:
|
||||
for i in range(1, R + 1): # 行
|
||||
for j in range(1, C + 1): # 列
|
||||
fig.update_xaxes(
|
||||
tickvals=[df.index[0]] + list(df.index[1:]),
|
||||
ticktext=[f'<span style="color:blue; font-weight:bold">{df.index[0]}</span>'] + list(df.index[1:]),
|
||||
row=i,
|
||||
col=j,
|
||||
)
|
||||
st.plotly_chart(fig)
|
||||
|
||||
|
||||
def summary_window():
|
||||
if state.log_type in ["Qlib Model", "Data Mining", "Qlib Factor"]:
|
||||
st.header("Summary📊", divider="rainbow", anchor="_summary")
|
||||
with st.container():
|
||||
# TODO: not fixed height
|
||||
with st.container():
|
||||
bc, cc = st.columns([2, 2], vertical_alignment="center")
|
||||
with bc:
|
||||
st.subheader("Metrics📈", anchor="_metrics")
|
||||
with cc:
|
||||
show_true_only = st.toggle("successful hypotheses", value=False)
|
||||
|
||||
# hypotheses_c, chart_c = st.columns([2, 3])
|
||||
chart_c = st.container()
|
||||
hypotheses_c = st.container()
|
||||
|
||||
with hypotheses_c:
|
||||
st.subheader("Hypotheses🏅", anchor="_hypotheses")
|
||||
display_hypotheses(state.hypotheses, state.h_decisions, show_true_only)
|
||||
|
||||
with chart_c:
|
||||
if state.log_type == "Qlib Factor" and state.alpha158_metrics is not None:
|
||||
df = pd.DataFrame([state.alpha158_metrics] + state.metric_series)
|
||||
else:
|
||||
df = pd.DataFrame(state.metric_series)
|
||||
if show_true_only and len(state.hypotheses) >= len(state.metric_series):
|
||||
if state.alpha158_metrics is not None:
|
||||
selected = ["alpha158"] + [i for i in df.index if state.h_decisions[int(i[6:])]]
|
||||
else:
|
||||
selected = [i for i in df.index if state.h_decisions[int(i[6:])]]
|
||||
df = df.loc[selected]
|
||||
if df.shape[0] == 1:
|
||||
st.table(df.iloc[0])
|
||||
elif df.shape[0] > 1:
|
||||
if df.shape[1] == 1:
|
||||
# suhan's scenario
|
||||
fig = px.line(df, x=df.index, y=df.columns, markers=True)
|
||||
fig.update_layout(xaxis_title="Loop Round", yaxis_title=None)
|
||||
st.plotly_chart(fig)
|
||||
else:
|
||||
metrics_window(df, 1, 4, height=300, colors=["red", "blue", "orange", "green"])
|
||||
|
||||
elif state.log_type == "Model from Paper" and len(state.msgs[state.lround]["d.evolving code"]) > 0:
|
||||
with st.container(border=True):
|
||||
st.subheader("Summary📊", divider="rainbow", anchor="_summary")
|
||||
|
||||
# pass
|
||||
ws: list[FactorFBWorkspace | ModelFBWorkspace] = state.msgs[state.lround]["d.evolving code"][-1].content
|
||||
# All Tasks
|
||||
|
||||
tab_names = [
|
||||
w.target_task.factor_name if isinstance(w.target_task, FactorTask) else w.target_task.name for w in ws
|
||||
]
|
||||
for j in range(len(ws)):
|
||||
if state.msgs[state.lround]["d.evolving feedback"][-1].content[j].final_decision:
|
||||
tab_names[j] += "✔️"
|
||||
else:
|
||||
tab_names[j] += "❌"
|
||||
|
||||
wtabs = st.tabs(tab_names)
|
||||
for j, w in enumerate(ws):
|
||||
with wtabs[j]:
|
||||
# Evolving Code
|
||||
for k, v in w.code_dict.items():
|
||||
with st.expander(f":green[`{k}`]", expanded=False):
|
||||
st.code(v, language="python")
|
||||
|
||||
# Evolving Feedback
|
||||
evolving_feedback_window(state.msgs[state.lround]["d.evolving feedback"][-1].content[j])
|
||||
|
||||
|
||||
def tabs_hint():
|
||||
st.markdown(
|
||||
"<p style='font-size: small; color: #888888;'>You can navigate through the tabs using ⬅️ ➡️ or by holding Shift and scrolling with the mouse wheel🖱️.</p>",
|
||||
unsafe_allow_html=True,
|
||||
)
|
||||
|
||||
|
||||
# TODO: when tab names are too long, some tabs are not shown
|
||||
def tasks_window(tasks: list[FactorTask | ModelTask]):
|
||||
if isinstance(tasks[0], FactorTask):
|
||||
st.markdown("**Factor Tasks🚩**")
|
||||
tnames = [f.factor_name for f in tasks]
|
||||
if sum(len(tn) for tn in tnames) > 100:
|
||||
tabs_hint()
|
||||
tabs = st.tabs(tnames)
|
||||
for i, ft in enumerate(tasks):
|
||||
with tabs[i]:
|
||||
# st.markdown(f"**Factor Name**: {ft.factor_name}")
|
||||
st.markdown(f"**Description**: {ft.factor_description}")
|
||||
st.latex("Formulation")
|
||||
st.latex(f"{ft.factor_formulation}")
|
||||
|
||||
mks = "| Variable | Description |\n| --- | --- |\n"
|
||||
for v, d in ft.variables.items():
|
||||
mks += f"| ${v}$ | {d} |\n"
|
||||
st.markdown(mks)
|
||||
|
||||
elif isinstance(tasks[0], ModelTask):
|
||||
st.markdown("**Model Tasks🚩**")
|
||||
tnames = [m.name for m in tasks]
|
||||
if sum(len(tn) for tn in tnames) > 100:
|
||||
tabs_hint()
|
||||
tabs = st.tabs(tnames)
|
||||
for i, mt in enumerate(tasks):
|
||||
with tabs[i]:
|
||||
# st.markdown(f"**Model Name**: {mt.name}")
|
||||
st.markdown(f"**Model Type**: {mt.model_type}")
|
||||
st.markdown(f"**Description**: {mt.description}")
|
||||
st.latex("Formulation")
|
||||
st.latex(f"{mt.formulation}")
|
||||
|
||||
mks = "| Variable | Description |\n| --- | --- |\n"
|
||||
for v, d in mt.variables.items():
|
||||
mks += f"| ${v}$ | {d} |\n"
|
||||
st.markdown(mks)
|
||||
|
||||
|
||||
# Config Sidebar
|
||||
with st.sidebar:
|
||||
st.text_input("log path", key="log_path", on_change=refresh)
|
||||
st.selectbox("trace type", ["qlib_model", "qlib_factor", "model_extraction_and_implementation"], key="log_type")
|
||||
st.markdown(
|
||||
"""
|
||||
# RD-Agent🤖
|
||||
## [Scenario Description](#_scenario)
|
||||
## [Summary](#_summary)
|
||||
- [**Hypotheses**](#_hypotheses)
|
||||
- [**Metrics**](#_metrics)
|
||||
## [RD-Loops](#_rdloops)
|
||||
- [**Research**](#_research)
|
||||
- [**Development**](#_development)
|
||||
- [**Feedback**](#_feedback)
|
||||
"""
|
||||
)
|
||||
|
||||
st.multiselect("excluded log tags", ["llm_messages"], ["llm_messages"], key="excluded_tags")
|
||||
st.multiselect("excluded log types", ["str", "dict", "list"], ["str"], key="excluded_types")
|
||||
st.selectbox(
|
||||
":green[**Scenario**]", ["Qlib Model", "Data Mining", "Qlib Factor", "Model from Paper"], key="log_type"
|
||||
)
|
||||
|
||||
if st.button("refresh"):
|
||||
with st.popover(":orange[**Config⚙️**]"):
|
||||
with st.container(border=True):
|
||||
st.markdown(":blue[**log path**]")
|
||||
if main_log_path:
|
||||
if st.toggle("Manual Input"):
|
||||
st.text_input("log path", key="log_path", on_change=refresh)
|
||||
else:
|
||||
folders = [
|
||||
folder.relative_to(main_log_path) for folder in main_log_path.iterdir() if folder.is_dir()
|
||||
]
|
||||
st.selectbox(f"Select from `{main_log_path}`", folders, key="log_path", on_change=refresh)
|
||||
else:
|
||||
st.text_input("log path", key="log_path", on_change=refresh)
|
||||
|
||||
with st.container(border=True):
|
||||
st.markdown(":blue[**excluded configs**]")
|
||||
st.multiselect("excluded log tags", ["llm_messages"], ["llm_messages"], key="excluded_tags")
|
||||
st.multiselect("excluded log types", ["str", "dict", "list"], ["str"], key="excluded_types")
|
||||
|
||||
if st.button("All Loops"):
|
||||
if not state.fs:
|
||||
refresh()
|
||||
get_msgs_until(lambda m: False)
|
||||
|
||||
if st.button("Next Loop"):
|
||||
if not state.fs:
|
||||
refresh()
|
||||
get_msgs_until(lambda m: "ef.feedback" in m.tag)
|
||||
|
||||
if st.button("One Evolving"):
|
||||
if not state.fs:
|
||||
refresh()
|
||||
get_msgs_until(lambda m: "d.evolving feedback" in m.tag)
|
||||
|
||||
if st.button("refresh logs", help="clear all log messages in cache"):
|
||||
refresh()
|
||||
debug = st.checkbox("debug", value=False)
|
||||
debug = st.toggle("debug", value=False)
|
||||
|
||||
if debug:
|
||||
if st.button("Single Step Run"):
|
||||
if not state.fs:
|
||||
refresh()
|
||||
get_msgs_until()
|
||||
|
||||
|
||||
@@ -174,148 +504,70 @@ if debug:
|
||||
if isinstance(state.last_msg.content, list):
|
||||
st.write(state.last_msg.content[0])
|
||||
elif not isinstance(state.last_msg.content, str):
|
||||
st.write(state.last_msg.content)
|
||||
st.write(state.last_msg.content.__dict__)
|
||||
|
||||
|
||||
# Main Window
|
||||
header_c1, header_c3 = st.columns([1, 6], vertical_alignment="center")
|
||||
with st.container():
|
||||
with header_c1:
|
||||
st.image("https://img-prod-cms-rt-microsoft-com.akamaized.net/cms/api/am/imageFileData/RE1Mu3b?ver=5c31")
|
||||
with header_c3:
|
||||
st.markdown(
|
||||
"""
|
||||
<h1>
|
||||
RD-Agent:<br>LLM-based autonomous evolving agents for industrial data-driven R&D
|
||||
</h1>
|
||||
""",
|
||||
unsafe_allow_html=True,
|
||||
)
|
||||
|
||||
# Project Info
|
||||
with st.container():
|
||||
image_c, toc_c = st.columns([3, 3], vertical_alignment="center")
|
||||
image_c, scen_c = st.columns([3, 3], vertical_alignment="center")
|
||||
with image_c:
|
||||
st.image("./docs/_static/scen.jpg")
|
||||
with toc_c:
|
||||
st.markdown(
|
||||
"""
|
||||
# RD-Agent🤖
|
||||
## [Scenario Description](#_scenario)
|
||||
## [Summary](#_summary)
|
||||
## [RD-Loops](#_rdloops)
|
||||
### [Research](#_research)
|
||||
### [Development](#_development)
|
||||
### [Feedback](#_feedback)
|
||||
"""
|
||||
)
|
||||
with st.container(border=True):
|
||||
st.header("Scenario Description📖", divider=True, anchor="_scenario")
|
||||
# TODO: other scenarios
|
||||
if state.log_type == "qlib_model":
|
||||
st.markdown(QlibModelScenario().rich_style_description)
|
||||
elif state.log_type == "model_extraction_and_implementation":
|
||||
st.markdown(
|
||||
"""
|
||||
# General Model Scenario
|
||||
|
||||
## Overview
|
||||
|
||||
This demo automates the extraction and iterative development of models from academic papers, ensuring functionality and correctness.
|
||||
|
||||
### Scenario: Auto-Developing Model Code from Academic Papers
|
||||
|
||||
#### Overview
|
||||
|
||||
This scenario automates the development of PyTorch models by reading academic papers or other sources. It supports various data types, including tabular, time-series, and graph data. The primary workflow involves two main components: the Reader and the Coder.
|
||||
|
||||
#### Workflow Components
|
||||
|
||||
1. **Reader**
|
||||
- Parses and extracts relevant model information from academic papers or sources, including architectures, parameters, and implementation details.
|
||||
- Uses Large Language Models to convert content into a structured format for the Coder.
|
||||
|
||||
2. **Evolving Coder**
|
||||
- Translates structured information from the Reader into executable PyTorch code.
|
||||
- Utilizes an evolving coding mechanism to ensure correct tensor shapes, verified with sample input tensors.
|
||||
- Iteratively refines the code to align with source material specifications.
|
||||
|
||||
#### Supported Data Types
|
||||
|
||||
- **Tabular Data:** Structured data with rows and columns, such as spreadsheets or databases.
|
||||
- **Time-Series Data:** Sequential data points indexed in time order, useful for forecasting and temporal pattern recognition.
|
||||
- **Graph Data:** Data structured as nodes and edges, suitable for network analysis and relational tasks.
|
||||
"""
|
||||
)
|
||||
st.image("./docs/_static/flow.png")
|
||||
with scen_c:
|
||||
st.header("Scenario Description📖", divider="violet", anchor="_scenario")
|
||||
# TODO: other scenarios
|
||||
if state.log_type == "Qlib Model":
|
||||
st.markdown(QlibModelScenario().rich_style_description)
|
||||
elif state.log_type == "Data Mining":
|
||||
st.markdown(DMModelScenario().rich_style_description)
|
||||
elif state.log_type == "Qlib Factor":
|
||||
st.markdown(QlibFactorScenario().rich_style_description)
|
||||
elif state.log_type == "Model from Paper":
|
||||
st.markdown(GeneralModelScenario().rich_style_description)
|
||||
|
||||
|
||||
# Summary Window
|
||||
@st.experimental_fragment()
|
||||
def summary_window():
|
||||
if state.log_type in ["qlib_model", "qlib_factor"]:
|
||||
with st.container():
|
||||
st.header("Summary📊", divider=True, anchor="_summary")
|
||||
hypotheses_c, chart_c = st.columns([2, 3])
|
||||
# TODO: not fixed height
|
||||
with hypotheses_c.container(height=600):
|
||||
st.markdown("**Hypotheses🏅**")
|
||||
h_str = "\n".join(
|
||||
f"{id}. :green[**{h.hypothesis}**]\n\t>:green-background[*{h.__dict__.get('concise_reason', '')}*]"
|
||||
if state.h_decisions[id]
|
||||
else f"{id}. {h.hypothesis}\n\t>*{h.__dict__.get('concise_reason', '')}*"
|
||||
for id, h in state.hypotheses.items()
|
||||
)
|
||||
st.markdown(h_str)
|
||||
with chart_c.container(height=600):
|
||||
mt_c, ms_c = st.columns(2, vertical_alignment="center")
|
||||
with mt_c:
|
||||
st.markdown("**Metrics📈**")
|
||||
with ms_c:
|
||||
show_true_only = st.checkbox("True Decisions Only", value=False)
|
||||
|
||||
labels = [f"Round {i}" for i in range(1, len(state.metric_series) + 1)]
|
||||
df = pd.DataFrame(state.metric_series, index=labels)
|
||||
if show_true_only and len(state.hypotheses) >= len(state.metric_series):
|
||||
df = df.iloc[[i for i in range(df.shape[0]) if state.h_decisions[i + 1]]]
|
||||
if df.shape[0] == 1:
|
||||
st.table(df.iloc[0])
|
||||
elif df.shape[0] > 1:
|
||||
# TODO: figure label
|
||||
# TODO: separate into different figures
|
||||
if df.shape[1] == 1:
|
||||
# suhan's scenario
|
||||
fig = px.line(df, x=df.index, y=df.columns, markers=True)
|
||||
fig.update_layout(legend_title_text="Metrics", xaxis_title="Loop Round", yaxis_title=None)
|
||||
else:
|
||||
# 2*2 figure
|
||||
fig = make_subplots(rows=2, cols=2, subplot_titles=df.columns)
|
||||
for ci, col in enumerate(df.columns):
|
||||
row = ci // 2 + 1
|
||||
col_num = ci % 2 + 1
|
||||
fig.add_trace(
|
||||
go.Scatter(x=df.index, y=df[col], mode="lines+markers", name=col), row=row, col=col_num
|
||||
)
|
||||
fig.update_layout(title_text="Metrics", showlegend=False)
|
||||
st.plotly_chart(fig)
|
||||
|
||||
|
||||
summary_window()
|
||||
|
||||
# R&D Loops Window
|
||||
st.header("R&D Loops♾️", divider=True, anchor="_rdloops")
|
||||
button_c1, button_c2, round_s_c = st.columns([2, 3, 18], vertical_alignment="center")
|
||||
with button_c1:
|
||||
if st.button("Run One Loop"):
|
||||
get_msgs_until(lambda m: "ef.feedback" in m.tag)
|
||||
with button_c2:
|
||||
if st.button("Run One Evolving Step"):
|
||||
get_msgs_until(lambda m: "d.evolving feedback" in m.tag)
|
||||
if state.log_type in ["Qlib Model", "Data Mining", "Qlib Factor"]:
|
||||
st.header("R&D Loops♾️", divider="rainbow", anchor="_rdloops")
|
||||
|
||||
if len(state.msgs) > 1:
|
||||
with round_s_c:
|
||||
round = st.select_slider("Select RDLoop Round", options=state.msgs.keys(), value=state.lround)
|
||||
if state.log_type in ["Qlib Model", "Data Mining", "Qlib Factor"]:
|
||||
if len(state.msgs) > 1:
|
||||
r_options = list(state.msgs.keys())
|
||||
if 0 in r_options:
|
||||
r_options.remove(0)
|
||||
round = st_btn_select(options=r_options, index=state.lround - 1)
|
||||
else:
|
||||
round = 1
|
||||
else:
|
||||
round = 1
|
||||
|
||||
rf_c, d_c = st.columns([2, 2])
|
||||
|
||||
# Research & Feedback Window
|
||||
with rf_c:
|
||||
if state.log_type in ["qlib_model", "qlib_factor"]:
|
||||
# Research Window
|
||||
with st.container(border=True):
|
||||
st.subheader("Research🔍", divider=True, anchor="_research")
|
||||
def research_window():
|
||||
with st.container(border=True):
|
||||
title = "Research🔍" if state.log_type in ["Qlib Model", "Data Mining", "Qlib Factor"] else "Research🔍 (reader)"
|
||||
st.subheader(title, divider="blue", anchor="_research")
|
||||
if state.log_type in ["Qlib Model", "Data Mining", "Qlib Factor"]:
|
||||
# pdf image
|
||||
if pim := state.msgs[round]["r.extract_factors_and_implement.load_pdf_screenshot"]:
|
||||
for i in range(min(2, len(pim))):
|
||||
st.image(pim[i].content)
|
||||
st.image(pim[i].content, use_column_width=True)
|
||||
|
||||
# Hypothesis
|
||||
if hg := state.msgs[round]["r.hypothesis generation"]:
|
||||
@@ -328,37 +580,27 @@ with rf_c:
|
||||
)
|
||||
|
||||
if eg := state.msgs[round]["r.experiment generation"]:
|
||||
if isinstance(eg[0].content[0], FactorTask):
|
||||
st.markdown("**Factor Tasks**")
|
||||
fts = eg[0].content
|
||||
tabs = st.tabs([f.factor_name for f in fts])
|
||||
for i, ft in enumerate(fts):
|
||||
with tabs[i]:
|
||||
# st.markdown(f"**Factor Name**: {ft.factor_name}")
|
||||
st.markdown(f"**Description**: {ft.factor_description}")
|
||||
st.latex(f"Formulation: {ft.factor_formulation}")
|
||||
tasks_window(eg[0].content)
|
||||
|
||||
variables_df = pd.DataFrame(ft.variables, index=["Description"]).T
|
||||
variables_df.index.name = "Variable"
|
||||
st.table(variables_df)
|
||||
elif isinstance(eg[0].content[0], ModelTask):
|
||||
st.markdown("**Model Tasks**")
|
||||
mts = eg[0].content
|
||||
tabs = st.tabs([m.name for m in mts])
|
||||
for i, mt in enumerate(mts):
|
||||
with tabs[i]:
|
||||
# st.markdown(f"**Model Name**: {mt.name}")
|
||||
st.markdown(f"**Model Type**: {mt.model_type}")
|
||||
st.markdown(f"**Description**: {mt.description}")
|
||||
st.latex(f"Formulation: {mt.formulation}")
|
||||
elif state.log_type == "Model from Paper":
|
||||
# pdf image
|
||||
c1, c2 = st.columns([2, 3])
|
||||
with c1:
|
||||
if pim := state.msgs[round]["r.pdf_image"]:
|
||||
for i in range(len(pim)):
|
||||
st.image(pim[i].content, use_column_width=True)
|
||||
|
||||
variables_df = pd.DataFrame(mt.variables, index=["Value"]).T
|
||||
variables_df.index.name = "Variable"
|
||||
st.table(variables_df)
|
||||
# loaded model exp
|
||||
with c2:
|
||||
if mem := state.msgs[round]["d.load_experiment"]:
|
||||
me: QlibModelExperiment = mem[0].content
|
||||
tasks_window(me.sub_tasks)
|
||||
|
||||
# Feedback Window
|
||||
|
||||
def feedback_window():
|
||||
if state.log_type in ["Qlib Model", "Data Mining", "Qlib Factor"]:
|
||||
with st.container(border=True):
|
||||
st.subheader("Feedback📝", divider=True, anchor="_feedback")
|
||||
st.subheader("Feedback📝", divider="orange", anchor="_feedback")
|
||||
if fbr := state.msgs[round]["ef.Quantitative Backtesting Chart"]:
|
||||
st.markdown("**Returns📈**")
|
||||
fig = report_figure(fbr[0].content)
|
||||
@@ -375,104 +617,82 @@ with rf_c:
|
||||
- **Reason**: {h.reason}"""
|
||||
)
|
||||
|
||||
elif state.log_type == "model_extraction_and_implementation":
|
||||
# Research Window
|
||||
with st.container(border=True):
|
||||
# pdf image
|
||||
st.subheader("Research🔍", divider=True, anchor="_research")
|
||||
if pim := state.msgs[round]["r.pdf_image"]:
|
||||
for i in range(len(pim)):
|
||||
st.image(pim[i].content)
|
||||
|
||||
# loaded model exp
|
||||
if mem := state.msgs[round]["d.load_experiment"]:
|
||||
me: QlibModelExperiment = mem[0].content
|
||||
mts: list[ModelTask] = me.sub_tasks
|
||||
tabs = st.tabs([m.name for m in mts])
|
||||
for i, mt in enumerate(mts):
|
||||
with tabs[i]:
|
||||
# st.markdown(f"**Model Name**: {mt.name}")
|
||||
st.markdown(f"**Model Type**: {mt.model_type}")
|
||||
st.markdown(f"**Description**: {mt.description}")
|
||||
st.latex(f"Formulation: {mt.formulation}")
|
||||
if state.log_type in ["Qlib Model", "Data Mining", "Qlib Factor"]:
|
||||
rf_c, d_c = st.columns([2, 2])
|
||||
elif state.log_type == "Model from Paper":
|
||||
rf_c = st.container()
|
||||
d_c = st.container()
|
||||
|
||||
variables_df = pd.DataFrame(mt.variables, index=["Value"]).T
|
||||
variables_df.index.name = "Variable"
|
||||
st.table(variables_df)
|
||||
|
||||
# Feedback Window
|
||||
with st.container(border=True):
|
||||
st.subheader("Feedback📝", divider=True, anchor="_feedback")
|
||||
if fbr := state.msgs[round]["d.developed_experiment"]:
|
||||
st.markdown("**Returns📈**")
|
||||
result_df = fbr[0].content.result
|
||||
if result_df:
|
||||
fig = report_figure(result_df)
|
||||
st.plotly_chart(fig)
|
||||
else:
|
||||
st.markdown("Returns is None")
|
||||
with rf_c:
|
||||
research_window()
|
||||
feedback_window()
|
||||
|
||||
|
||||
# Development Window (Evolving)
|
||||
with d_c.container(border=True):
|
||||
st.subheader("Development🛠️", divider=True, anchor="_development")
|
||||
title = (
|
||||
"Development🛠️"
|
||||
if state.log_type in ["Qlib Model", "Data Mining", "Qlib Factor"]
|
||||
else "Development🛠️ (evolving coder)"
|
||||
)
|
||||
st.subheader(title, divider="green", anchor="_development")
|
||||
|
||||
# Evolving Status
|
||||
if state.erounds[round] > 0:
|
||||
st.markdown("**☑️ Evolving Status**")
|
||||
es = state.e_decisions[round]
|
||||
e_status_mks = "".join(f"| {ei} " for ei in range(1, state.erounds[round] + 1)) + "|\n"
|
||||
e_status_mks += "|--" * state.erounds[round] + "|\n"
|
||||
for ei, estatus in es.items():
|
||||
if not estatus:
|
||||
estatus = (0, 0)
|
||||
e_status_mks += "| " + "✔️<br>" * estatus[0] + "❌<br>" * estatus[1] + " "
|
||||
e_status_mks += "|\n"
|
||||
st.markdown(e_status_mks, unsafe_allow_html=True)
|
||||
|
||||
# Evolving Tabs
|
||||
if state.erounds[round] > 0:
|
||||
etabs = st.tabs([str(i) for i in range(1, state.erounds[round] + 1)])
|
||||
if state.erounds[round] > 1:
|
||||
st.markdown("**🔄️Evolving Rounds**")
|
||||
evolving_round = st_btn_select(
|
||||
options=range(1, state.erounds[round] + 1), index=state.erounds[round] - 1, key="show_eround"
|
||||
)
|
||||
else:
|
||||
evolving_round = 1
|
||||
|
||||
for i in range(0, state.erounds[round]):
|
||||
with etabs[i]:
|
||||
ws: list[FactorFBWorkspace | ModelFBWorkspace] = state.msgs[round]["d.evolving code"][i].content
|
||||
ws = [w for w in ws if w]
|
||||
# All Tasks
|
||||
ws: list[FactorFBWorkspace | ModelFBWorkspace] = state.msgs[round]["d.evolving code"][
|
||||
evolving_round - 1
|
||||
].content
|
||||
# All Tasks
|
||||
|
||||
tab_names = [
|
||||
w.target_task.factor_name if isinstance(w.target_task, FactorTask) else w.target_task.name for w in ws
|
||||
]
|
||||
wtabs = st.tabs(tab_names)
|
||||
for j, w in enumerate(ws):
|
||||
with wtabs[j]:
|
||||
# Evolving Code
|
||||
for k, v in w.code_dict.items():
|
||||
with st.expander(f":green[`{k}`]", expanded=True):
|
||||
st.code(v, language="python")
|
||||
tab_names = [
|
||||
w.target_task.factor_name if isinstance(w.target_task, FactorTask) else w.target_task.name for w in ws
|
||||
]
|
||||
if len(state.msgs[round]["d.evolving feedback"]) >= evolving_round:
|
||||
for j in range(len(ws)):
|
||||
if state.msgs[round]["d.evolving feedback"][evolving_round - 1].content[j].final_decision:
|
||||
tab_names[j] += "✔️"
|
||||
else:
|
||||
tab_names[j] += "❌"
|
||||
if sum(len(tn) for tn in tab_names) > 100:
|
||||
tabs_hint()
|
||||
wtabs = st.tabs(tab_names)
|
||||
for j, w in enumerate(ws):
|
||||
with wtabs[j]:
|
||||
# Evolving Code
|
||||
for k, v in w.code_dict.items():
|
||||
with st.expander(f":green[`{k}`]", expanded=True):
|
||||
st.code(v, language="python")
|
||||
|
||||
# Evolving Feedback
|
||||
if len(state.msgs[round]["d.evolving feedback"]) > i:
|
||||
wsf: list[FactorSingleFeedback | ModelCoderFeedback] = state.msgs[round]["d.evolving feedback"][
|
||||
i
|
||||
].content[j]
|
||||
if isinstance(wsf, FactorSingleFeedback):
|
||||
st.markdown(
|
||||
f"""#### :blue[Factor Execution Feedback]
|
||||
{wsf.execution_feedback}
|
||||
#### :blue[Factor Code Feedback]
|
||||
{wsf.code_feedback}
|
||||
#### :blue[Factor Value Feedback]
|
||||
{wsf.factor_value_feedback}
|
||||
#### :blue[Factor Final Feedback]
|
||||
{wsf.final_feedback}
|
||||
#### :blue[Factor Final Decision]
|
||||
This implementation is {'SUCCESS' if wsf.final_decision else 'FAIL'}.
|
||||
"""
|
||||
)
|
||||
elif isinstance(wsf, ModelCoderFeedback):
|
||||
st.markdown(
|
||||
f"""#### :blue[Model Execution Feedback]
|
||||
{wsf.execution_feedback}
|
||||
#### :blue[Model Shape Feedback]
|
||||
{wsf.shape_feedback}
|
||||
#### :blue[Model Value Feedback]
|
||||
{wsf.value_feedback}
|
||||
#### :blue[Model Code Feedback]
|
||||
{wsf.code_feedback}
|
||||
#### :blue[Model Final Feedback]
|
||||
{wsf.final_feedback}
|
||||
#### :blue[Model Final Decision]
|
||||
This implementation is {'SUCCESS' if wsf.final_decision else 'FAIL'}.
|
||||
"""
|
||||
)
|
||||
# Evolving Feedback
|
||||
if len(state.msgs[round]["d.evolving feedback"]) >= evolving_round:
|
||||
evolving_feedback_window(state.msgs[round]["d.evolving feedback"][evolving_round - 1].content[j])
|
||||
|
||||
# TODO: evolving tabs -> slider
|
||||
# TODO: multi tasks SUCCESS/FAIL
|
||||
# TODO: evolving progress bar, diff colors
|
||||
|
||||
with st.container(border=True):
|
||||
st.subheader("Disclaimer", divider="gray")
|
||||
st.markdown(
|
||||
"This content is AI-generated and may not be fully accurate or up-to-date; please verify with a professional for critical matters."
|
||||
)
|
||||
|
||||
+4
-1
@@ -61,4 +61,7 @@ python-dotenv
|
||||
# infrastructure related.
|
||||
docker
|
||||
|
||||
|
||||
# demo related
|
||||
streamlit
|
||||
plotly
|
||||
st-btn-select
|
||||
|
||||
Reference in New Issue
Block a user