Spaces:
Runtime error
Runtime error
Commit Β·
44f11d2
1
Parent(s): 0c3df82
trained evironment UI fixes
Browse files- demo/assets/reward_curve.json +1 -1
- demo/streamlit_app.py +255 -21
demo/assets/reward_curve.json
CHANGED
|
@@ -1 +1 @@
|
|
| 1 |
-
{"episodes": [0,
|
|
|
|
| 1 |
+
{"episodes": [0, 67, 134, 201, 268, 335, 402, 469, 536, 603, 670, 737, 804, 871, 938, 1005, 1072, 1139, 1206, 1273, 1340, 1407, 1474, 1541, 1608, 1675, 1742, 1809, 1876, 1943, 2010, 2077, 2144, 2211, 2278, 2345, 2412, 2479, 2546, 2613, 2680, 2747, 2814, 2881, 2948, 3015, 3082, 3149, 3216, 3283, 3350, 3417, 3484, 3551, 3618, 3685, 3752, 3819, 3886, 3953, 4020, 4087, 4154, 4221, 4288, 4355, 4422, 4489, 4556, 4623, 4690, 4757, 4824, 4891, 4958, 5025, 5092, 5159, 5226, 5293, 5360, 5427, 5494, 5561, 5628, 5695, 5762, 5829, 5896, 5963, 6030, 6097, 6164, 6231, 6298, 6365, 6432, 6499, 6566, 6633, 6700, 6767, 6834, 6901, 6968, 7035, 7102, 7169, 7236, 7303, 7370, 7437, 7504, 7571, 7638, 7705, 7772, 7839, 7906, 7973, 8040, 8107, 8174, 8241, 8308, 8375, 8442, 8509, 8576, 8643, 8710, 8777, 8844, 8911, 8978, 9045, 9112, 9179, 9246, 9313, 9380, 9447, 9514, 9581, 9648, 9715, 9782, 9849, 9916, 9983, 10050, 10117, 10184, 10251, 10318, 10385, 10452, 10519, 10586, 10653, 10720, 10787, 10854, 10921, 10988, 11055, 11122, 11189, 11256, 11323, 11390, 11457, 11524, 11591, 11658, 11725, 11792, 11859, 11926, 11993, 12060, 12127, 12194, 12261, 12328, 12395, 12462, 12529, 12596, 12663, 12730, 12797, 12864, 12931, 12998, 13065, 13132, 13199, 13266, 13333, 13400, 13467], "mean_rewards": [7.869385480880737, -0.4680378864902784, -0.48984615507501145, -0.5551410619040379, -0.4882915579094153, -0.4724334484614831, -0.4706858817982504, -0.49536512678172046, -0.488181437265121, -0.49218673040220723, -0.4815726703055059, -0.4820088524676782, -0.43784926100558375, -0.38558788626885354, -0.3533458443251239, -0.3040934353599124, -0.26831979214033574, -0.19140849363449494, -0.14203320211143491, -0.03154456203850192, 0.036356829438342037, 0.10923927782606911, 0.18537376434205202, 0.23818426830707837, 0.2971860210863843, 0.351879875832877, 0.41131269029737827, 0.49316205722733814, 0.54418244560461, 0.5529096998612051, 0.5803396979899199, 0.647491346514733, 0.6767948179212062, 0.7278867457176142, 0.7801977058292352, 0.8324114113203805, 0.8846536156650867, 0.9135522807000074, 0.9906498014824725, 1.0747783365352257, 1.118064828834938, 1.1254296454959651, 1.1729068560170468, 1.223296715608557, 1.250830467714188, 1.2658592661535932, 1.3026064717935804, 1.3210587293828564, 1.3291228487317284, 1.3327255037611838, 1.3617869722608757, 1.4155327298567029, 1.43804301345443, 1.4533835403824795, 1.4490587244572413, 1.4650314738802017, 1.4970641783053438, 1.5264675998406834, 1.5249301680553555, 1.5104940346199014, 1.5259575240262075, 1.5392910738888663, 1.5554828937112133, 1.5631932486042976, 1.577678209906162, 1.578992944977183, 1.6025471957099515, 1.5807723147838604, 1.5952618680505644, 1.6292361341937072, 1.6326987860754143, 1.6245606879480392, 1.6551997169478658, 1.6609219776152042, 1.6956165906986145, 1.7281981333328353, 1.7382102399385817, 1.7609851391969418, 1.7657828451398392, 1.7754888206320893, 1.7645542004733432, 1.7736322304512135, 1.7687026667687216, 1.7580264367949372, 1.7511381573787308, 1.738839277673473, 1.722614758535346, 1.736745614141461, 1.7516905580270015, 1.7799245542025521, 1.7700678259391585, 1.7663207697615921, 1.7562163755505256, 1.7587610971747853, 1.7777937432416644, 1.7704596316888854, 1.8123356010955924, 1.8279217910185117, 1.8466502024516873, 1.8298634645036476, 1.864942445150443, 1.8738156217091924, 1.866000112849446, 1.8781902273140842, 1.8416169193525835, 1.8387443213458785, 1.8485462165543736, 1.8278536795119362, 1.8172942939169872, 1.8037021071630313, 1.7578254819032049, 1.789927663865958, 1.8083126330250383, 1.8363163345167517, 1.8767271982310745, 1.8823451488214322, 1.8882187248528846, 1.879882797724712, 1.8682739835334294, 1.8768755783047664, 1.8853264416580362, 1.8954450143514918, 1.9106359317722665, 1.9006636210578318, 1.8993877949828675, 1.9237823593464614, 1.9344475316007566, 1.9503340135305085, 1.9860419661475228, 1.9742594653096717, 2.001614794778903, 1.9832609987142256, 1.9856477235292567, 1.9926783719954344, 1.9750121485619383, 1.9528438657518075, 1.9334705030714798, 1.9176346147756205, 1.9197328341457625, 1.929320812971638, 1.935611099781306, 1.9301533763026026, 1.9524929260035577, 1.9774023529251636, 1.9670411776322891, 1.9650732083037827, 1.9756527485480824, 1.9960569345849135, 1.9994612800762286, 2.0023671949464004, 1.968471810134468, 1.9351308788716286, 1.9281577662479867, 1.9180797577689113, 1.941868303399712, 1.9551983473480623, 1.9582474042363658, 1.9446796976185323, 1.9599162018053136, 1.9677243051538726, 1.9905415299682818, 2.0126337819571773, 2.0055850803634727, 1.9728688634663787, 1.942376174849831, 1.961291219857542, 1.9382520422395841, 1.9331532626002175, 1.9315321168749136, 1.9510658689675868, 1.93959757961361, 1.9583952448716349, 1.9957628746242215, 2.0171623558352496, 2.045212839087822, 2.027753655145466, 2.023227911412426, 2.0044681133162108, 1.9950296741726627, 1.9736256399910699, 1.9867524237810816, 1.9778640382239483, 1.9617218394162603, 1.9373436287414156, 1.9262697115193839, 1.9367491794142053, 1.9615602505914702, 1.9952850935889883, 1.9757408594996002, 1.9619378888624635, 1.9616936761638262, 1.9569105818823662, 1.966847003359166, 1.995661675253198, 2.011519161810991, 2.030814020499653, 2.015335946653363, 2.010846167103975, 2.036694835887966, 2.063149935127053, 2.0468815370260613, 2.0691288812374875]}
|
demo/streamlit_app.py
CHANGED
|
@@ -13,8 +13,7 @@ from dotenv import load_dotenv
|
|
| 13 |
|
| 14 |
load_dotenv() # load OPENAI_API_KEY (and any other vars) from .env
|
| 15 |
|
| 16 |
-
|
| 17 |
-
os.environ.setdefault("TRANSFORMERS_OFFLINE", "1")
|
| 18 |
|
| 19 |
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
| 20 |
sys.path.insert(0, str(Path(__file__).resolve().parent))
|
|
@@ -81,19 +80,24 @@ def _load_trained_model(hf_repo: str):
|
|
| 81 |
Temporarily lifts the HF_HUB_OFFLINE flag set at module level.
|
| 82 |
"""
|
| 83 |
import pickle
|
| 84 |
-
_old_hf = os.environ.pop("HF_HUB_OFFLINE", None)
|
| 85 |
-
_old_tf = os.environ.pop("TRANSFORMERS_OFFLINE", None)
|
| 86 |
try:
|
| 87 |
from huggingface_hub import hf_hub_download
|
| 88 |
from sb3_contrib import RecurrentPPO
|
| 89 |
|
| 90 |
-
|
| 91 |
-
|
| 92 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 93 |
obs_mean = obs_var = None
|
| 94 |
clip_obs = 10.0
|
| 95 |
try:
|
| 96 |
-
|
|
|
|
|
|
|
|
|
|
| 97 |
with open(stats_path, "rb") as f:
|
| 98 |
vn = pickle.load(f)
|
| 99 |
obs_mean = vn.obs_rms.mean.copy()
|
|
@@ -105,10 +109,7 @@ def _load_trained_model(hf_repo: str):
|
|
| 105 |
except Exception as exc:
|
| 106 |
return None, None, None, 10.0, str(exc)
|
| 107 |
finally:
|
| 108 |
-
|
| 109 |
-
os.environ["HF_HUB_OFFLINE"] = _old_hf
|
| 110 |
-
if _old_tf is not None:
|
| 111 |
-
os.environ["TRANSFORMERS_OFFLINE"] = _old_tf
|
| 112 |
|
| 113 |
|
| 114 |
def _predict(model, obs: np.ndarray, lstm_states, episode_starts,
|
|
@@ -1293,23 +1294,18 @@ def tab_training():
|
|
| 1293 |
|
| 1294 |
c_fetch, _ = st.columns([2, 5])
|
| 1295 |
if c_fetch.button("π₯ Fetch latest curve from HF Hub", key="fetch_curve"):
|
| 1296 |
-
_old_hf = os.environ.pop("HF_HUB_OFFLINE", None)
|
| 1297 |
-
_old_tf = os.environ.pop("TRANSFORMERS_OFFLINE", None)
|
| 1298 |
try:
|
| 1299 |
import shutil
|
| 1300 |
from huggingface_hub import hf_hub_download
|
| 1301 |
-
|
|
|
|
|
|
|
| 1302 |
ASSETS.mkdir(parents=True, exist_ok=True)
|
| 1303 |
shutil.copy(src, ASSETS / "reward_curve.json")
|
| 1304 |
st.success("reward_curve.json updated β chart will refresh.")
|
| 1305 |
st.cache_data.clear()
|
| 1306 |
except Exception as exc:
|
| 1307 |
st.error(f"Download failed: {exc}")
|
| 1308 |
-
finally:
|
| 1309 |
-
if _old_hf is not None:
|
| 1310 |
-
os.environ["HF_HUB_OFFLINE"] = _old_hf
|
| 1311 |
-
if _old_tf is not None:
|
| 1312 |
-
os.environ["TRANSFORMERS_OFFLINE"] = _old_tf
|
| 1313 |
|
| 1314 |
st.plotly_chart(fig_training_curve(), use_container_width=True)
|
| 1315 |
|
|
@@ -1536,6 +1532,242 @@ def tab_architecture():
|
|
| 1536 |
)""", language="python")
|
| 1537 |
|
| 1538 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1539 |
# βββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 1540 |
# Entry point
|
| 1541 |
# βββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
|
@@ -1545,13 +1777,14 @@ def main():
|
|
| 1545 |
S = _S()
|
| 1546 |
render_live_stats(S)
|
| 1547 |
|
| 1548 |
-
t1, t2, t3, t4, t5, t6 = st.tabs([
|
| 1549 |
"β‘ Live Demo",
|
| 1550 |
"π€ Specialists",
|
| 1551 |
"π Training",
|
| 1552 |
"π Quality Demo",
|
| 1553 |
"π§ͺ Reward Lab",
|
| 1554 |
"π Architecture",
|
|
|
|
| 1555 |
])
|
| 1556 |
with t1: tab_live_demo()
|
| 1557 |
with t2: tab_specialists()
|
|
@@ -1559,6 +1792,7 @@ def main():
|
|
| 1559 |
with t4: tab_quality()
|
| 1560 |
with t5: tab_reward_lab()
|
| 1561 |
with t6: tab_architecture()
|
|
|
|
| 1562 |
|
| 1563 |
|
| 1564 |
# Guard allows safe imports for testing without triggering the UI.
|
|
|
|
| 13 |
|
| 14 |
load_dotenv() # load OPENAI_API_KEY (and any other vars) from .env
|
| 15 |
|
| 16 |
+
# HF_HUB_OFFLINE intentionally NOT set β manual HF Hub downloads must work
|
|
|
|
| 17 |
|
| 18 |
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
| 19 |
sys.path.insert(0, str(Path(__file__).resolve().parent))
|
|
|
|
| 80 |
Temporarily lifts the HF_HUB_OFFLINE flag set at module level.
|
| 81 |
"""
|
| 82 |
import pickle
|
|
|
|
|
|
|
| 83 |
try:
|
| 84 |
from huggingface_hub import hf_hub_download
|
| 85 |
from sb3_contrib import RecurrentPPO
|
| 86 |
|
| 87 |
+
_tok = os.getenv("HF_TOKEN") or None
|
| 88 |
+
# Try final model first, fall back to latest periodic checkpoint
|
| 89 |
+
try:
|
| 90 |
+
_model_path = hf_hub_download(hf_repo, "spindleflow_model.zip", token=_tok)
|
| 91 |
+
except Exception:
|
| 92 |
+
_model_path = hf_hub_download(hf_repo, "spindleflow_model_latest.zip", token=_tok)
|
| 93 |
+
model = RecurrentPPO.load(_model_path, device="cpu")
|
| 94 |
obs_mean = obs_var = None
|
| 95 |
clip_obs = 10.0
|
| 96 |
try:
|
| 97 |
+
try:
|
| 98 |
+
stats_path = hf_hub_download(hf_repo, "vec_normalize.pkl", token=_tok)
|
| 99 |
+
except Exception:
|
| 100 |
+
stats_path = hf_hub_download(hf_repo, "vec_normalize_latest.pkl", token=_tok)
|
| 101 |
with open(stats_path, "rb") as f:
|
| 102 |
vn = pickle.load(f)
|
| 103 |
obs_mean = vn.obs_rms.mean.copy()
|
|
|
|
| 109 |
except Exception as exc:
|
| 110 |
return None, None, None, 10.0, str(exc)
|
| 111 |
finally:
|
| 112 |
+
pass
|
|
|
|
|
|
|
|
|
|
| 113 |
|
| 114 |
|
| 115 |
def _predict(model, obs: np.ndarray, lstm_states, episode_starts,
|
|
|
|
| 1294 |
|
| 1295 |
c_fetch, _ = st.columns([2, 5])
|
| 1296 |
if c_fetch.button("π₯ Fetch latest curve from HF Hub", key="fetch_curve"):
|
|
|
|
|
|
|
| 1297 |
try:
|
| 1298 |
import shutil
|
| 1299 |
from huggingface_hub import hf_hub_download
|
| 1300 |
+
_tok = os.getenv("HF_TOKEN") or None
|
| 1301 |
+
src = hf_hub_download(HF_MODEL_REPO, "reward_curve.json",
|
| 1302 |
+
token=_tok, force_download=True)
|
| 1303 |
ASSETS.mkdir(parents=True, exist_ok=True)
|
| 1304 |
shutil.copy(src, ASSETS / "reward_curve.json")
|
| 1305 |
st.success("reward_curve.json updated β chart will refresh.")
|
| 1306 |
st.cache_data.clear()
|
| 1307 |
except Exception as exc:
|
| 1308 |
st.error(f"Download failed: {exc}")
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1309 |
|
| 1310 |
st.plotly_chart(fig_training_curve(), use_container_width=True)
|
| 1311 |
|
|
|
|
| 1532 |
)""", language="python")
|
| 1533 |
|
| 1534 |
|
| 1535 |
+
# βββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 1536 |
+
# Tab 7 β Output (Trained Policy)
|
| 1537 |
+
# βββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 1538 |
+
def tab_output():
|
| 1539 |
+
"""Run the trained LSTM PPO policy on a custom task and show every specialist's output."""
|
| 1540 |
+
st.markdown(
|
| 1541 |
+
'<div style="font-size:12px;color:#64748b;margin-bottom:16px;">'
|
| 1542 |
+
'Enter any software engineering task. The trained LSTM PPO policy decides which '
|
| 1543 |
+
'specialists to delegate to β each specialist\'s individual output and the collective '
|
| 1544 |
+
'synthesis are shown below.</div>',
|
| 1545 |
+
unsafe_allow_html=True,
|
| 1546 |
+
)
|
| 1547 |
+
|
| 1548 |
+
col_input, col_ctrl = st.columns([3, 1], gap="large")
|
| 1549 |
+
with col_input:
|
| 1550 |
+
sec("Task")
|
| 1551 |
+
task_input = st.text_area(
|
| 1552 |
+
"Task description",
|
| 1553 |
+
height=110,
|
| 1554 |
+
key="output_task_input",
|
| 1555 |
+
placeholder=(
|
| 1556 |
+
"Build a real-time collaborative code review tool with inline comments, "
|
| 1557 |
+
"role-based access control, GitHub webhook integration, and CI/CD pipeline "
|
| 1558 |
+
"status display. Include authentication with OAuth2."
|
| 1559 |
+
),
|
| 1560 |
+
)
|
| 1561 |
+
with col_ctrl:
|
| 1562 |
+
sec("Config")
|
| 1563 |
+
out_phase = st.selectbox("Curriculum phase", [1, 2, 3], index=1, key="output_phase")
|
| 1564 |
+
st.markdown('<div style="height:8px"></div>', unsafe_allow_html=True)
|
| 1565 |
+
run_btn = st.button(
|
| 1566 |
+
"π Run Trained Policy",
|
| 1567 |
+
type="primary",
|
| 1568 |
+
use_container_width=True,
|
| 1569 |
+
key="output_run_btn",
|
| 1570 |
+
)
|
| 1571 |
+
|
| 1572 |
+
if run_btn:
|
| 1573 |
+
_task = (task_input or "").strip()
|
| 1574 |
+
if not _task:
|
| 1575 |
+
st.warning("Please enter a task description.")
|
| 1576 |
+
return
|
| 1577 |
+
|
| 1578 |
+
with st.spinner("Loading trained model from HF Hubβ¦"):
|
| 1579 |
+
model, obs_mean, obs_var, clip_obs, model_err = _load_trained_model(HF_MODEL_REPO)
|
| 1580 |
+
if model_err:
|
| 1581 |
+
st.error(f"Model load failed: {model_err}")
|
| 1582 |
+
return
|
| 1583 |
+
|
| 1584 |
+
st.success("Trained policy loaded β")
|
| 1585 |
+
|
| 1586 |
+
with st.spinner("Running episode with trained policyβ¦"):
|
| 1587 |
+
try:
|
| 1588 |
+
env = SpindleFlowEnv(
|
| 1589 |
+
config_path=CONFIG, catalog_path=CATALOG,
|
| 1590 |
+
use_real_spindleflow=False, phase=int(out_phase),
|
| 1591 |
+
)
|
| 1592 |
+
# Inject custom task so the env uses the user's input
|
| 1593 |
+
env.task_bank.sample = lambda: _task
|
| 1594 |
+
|
| 1595 |
+
obs, info = env.reset()
|
| 1596 |
+
task_used = info.get("task", _task)
|
| 1597 |
+
|
| 1598 |
+
lstm_states = None
|
| 1599 |
+
episode_starts = np.array([True])
|
| 1600 |
+
done = False
|
| 1601 |
+
rewards: list[float] = []
|
| 1602 |
+
|
| 1603 |
+
for _ in range(15):
|
| 1604 |
+
if done:
|
| 1605 |
+
break
|
| 1606 |
+
obs_arr = obs[np.newaxis, :].copy().astype(np.float32)
|
| 1607 |
+
if obs_mean is not None and obs_var is not None:
|
| 1608 |
+
obs_arr = np.clip(
|
| 1609 |
+
(obs_arr - obs_mean) / np.sqrt(obs_var + 1e-8),
|
| 1610 |
+
-clip_obs, clip_obs,
|
| 1611 |
+
)
|
| 1612 |
+
action_batch, lstm_states = model.predict(
|
| 1613 |
+
obs_arr,
|
| 1614 |
+
state=lstm_states,
|
| 1615 |
+
episode_start=episode_starts,
|
| 1616 |
+
deterministic=True,
|
| 1617 |
+
)
|
| 1618 |
+
action = action_batch[0]
|
| 1619 |
+
obs, r, term, trunc, _ = env.step(action)
|
| 1620 |
+
rewards.append(float(r))
|
| 1621 |
+
done = term or trunc
|
| 1622 |
+
episode_starts = np.array([done])
|
| 1623 |
+
|
| 1624 |
+
called = list(env.called_ids)
|
| 1625 |
+
edges = [(e.caller_id, e.callee_id)
|
| 1626 |
+
for e in env.delegation_graph.get_delegation_path()]
|
| 1627 |
+
spawned = list(getattr(env, "spawned_specialists", []))
|
| 1628 |
+
|
| 1629 |
+
st.session_state.output_results = {
|
| 1630 |
+
"task": task_used,
|
| 1631 |
+
"rewards": rewards,
|
| 1632 |
+
"called": called,
|
| 1633 |
+
"edges": edges,
|
| 1634 |
+
"specialist_results": [
|
| 1635 |
+
{
|
| 1636 |
+
"id": sr.specialist_id,
|
| 1637 |
+
"output": sr.output,
|
| 1638 |
+
"status": sr.status,
|
| 1639 |
+
"latency_ms": sr.latency_ms,
|
| 1640 |
+
}
|
| 1641 |
+
for sr in env.specialist_results
|
| 1642 |
+
],
|
| 1643 |
+
"spawned": spawned,
|
| 1644 |
+
}
|
| 1645 |
+
# Keep env alive for delegation-graph rendering
|
| 1646 |
+
st.session_state.output_env = env
|
| 1647 |
+
|
| 1648 |
+
except Exception as exc:
|
| 1649 |
+
import traceback
|
| 1650 |
+
st.error(f"Episode failed: {exc}")
|
| 1651 |
+
st.code(traceback.format_exc(), language=None)
|
| 1652 |
+
return
|
| 1653 |
+
|
| 1654 |
+
st.rerun()
|
| 1655 |
+
|
| 1656 |
+
# ββ Display results ββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 1657 |
+
results = st.session_state.get("output_results")
|
| 1658 |
+
env_obj = st.session_state.get("output_env")
|
| 1659 |
+
|
| 1660 |
+
if results is None:
|
| 1661 |
+
st.markdown(
|
| 1662 |
+
'<div style="color:#334155;font-size:12px;padding:40px;text-align:center;">'
|
| 1663 |
+
'Enter a task and click "Run Trained Policy" to see delegation and specialist outputs.'
|
| 1664 |
+
'</div>',
|
| 1665 |
+
unsafe_allow_html=True,
|
| 1666 |
+
)
|
| 1667 |
+
return
|
| 1668 |
+
|
| 1669 |
+
# Task banner
|
| 1670 |
+
st.markdown(
|
| 1671 |
+
f'<div style="background:rgba(0,212,255,0.04);'
|
| 1672 |
+
f'border:1px solid rgba(0,212,255,0.18);border-radius:10px;'
|
| 1673 |
+
f'padding:14px 18px;margin:10px 0 16px;">'
|
| 1674 |
+
f'<div style="font-size:9px;font-weight:700;color:#475569;'
|
| 1675 |
+
f'text-transform:uppercase;letter-spacing:1px;margin-bottom:5px;">Task</div>'
|
| 1676 |
+
f'<div style="font-size:13px;color:#e2e8f0;">{_html.escape(results["task"])}</div>'
|
| 1677 |
+
f'</div>',
|
| 1678 |
+
unsafe_allow_html=True,
|
| 1679 |
+
)
|
| 1680 |
+
|
| 1681 |
+
# Metrics strip
|
| 1682 |
+
total_r = sum(results["rewards"])
|
| 1683 |
+
mc1, mc2, mc3, mc4 = st.columns(4)
|
| 1684 |
+
mc1.metric("Total Reward", f"{total_r:+.3f}")
|
| 1685 |
+
mc2.metric("Steps", len(results["rewards"]))
|
| 1686 |
+
mc3.metric("Specialists Called", len(results["called"]))
|
| 1687 |
+
mc4.metric("Auto-Spawned", len(results["spawned"]))
|
| 1688 |
+
|
| 1689 |
+
# Delegation graph
|
| 1690 |
+
sec("Delegation Graph")
|
| 1691 |
+
if env_obj is not None:
|
| 1692 |
+
class _GraphProxy:
|
| 1693 |
+
registry = env_obj.registry
|
| 1694 |
+
spawned_specialists = results["spawned"]
|
| 1695 |
+
env = env_obj
|
| 1696 |
+
|
| 1697 |
+
st.plotly_chart(
|
| 1698 |
+
fig_delegation_graph(
|
| 1699 |
+
_GraphProxy(),
|
| 1700 |
+
results["called"],
|
| 1701 |
+
results["edges"],
|
| 1702 |
+
highlight_latest=False,
|
| 1703 |
+
spawned_ids=results["spawned"],
|
| 1704 |
+
),
|
| 1705 |
+
use_container_width=True,
|
| 1706 |
+
key="output_dag",
|
| 1707 |
+
)
|
| 1708 |
+
|
| 1709 |
+
# Auto-spawn alert
|
| 1710 |
+
if results["spawned"]:
|
| 1711 |
+
st.markdown(
|
| 1712 |
+
'<div style="background:rgba(251,191,36,0.06);'
|
| 1713 |
+
'border:1px solid rgba(251,191,36,0.22);border-radius:10px;'
|
| 1714 |
+
'padding:10px 16px;margin:8px 0;">'
|
| 1715 |
+
'<span style="font-size:10px;font-weight:700;color:#fbbf24;'
|
| 1716 |
+
'text-transform:uppercase;letter-spacing:1px;">β‘ Auto-Spawned: </span>'
|
| 1717 |
+
'<span style="font-size:12px;color:#e2e8f0;">'
|
| 1718 |
+
+ ", ".join(results["spawned"])
|
| 1719 |
+
+ '</span></div>',
|
| 1720 |
+
unsafe_allow_html=True,
|
| 1721 |
+
)
|
| 1722 |
+
|
| 1723 |
+
# Individual specialist outputs
|
| 1724 |
+
spec_results = results["specialist_results"]
|
| 1725 |
+
sec(f"Individual Specialist Outputs Β· {len(spec_results)} called")
|
| 1726 |
+
|
| 1727 |
+
if not spec_results:
|
| 1728 |
+
st.markdown(
|
| 1729 |
+
'<div style="color:#475569;font-size:12px;padding:16px;'
|
| 1730 |
+
'background:rgba(0,0,0,0.2);border-radius:8px;">'
|
| 1731 |
+
'The policy issued STOP without delegating to any specialists.</div>',
|
| 1732 |
+
unsafe_allow_html=True,
|
| 1733 |
+
)
|
| 1734 |
+
else:
|
| 1735 |
+
for sr in spec_results:
|
| 1736 |
+
sid = sr["id"]
|
| 1737 |
+
color = SPEC_COLORS.get(sid, "#7c3aed")
|
| 1738 |
+
ok_clr = "#10b981" if sr["status"] == "success" else "#ef4444"
|
| 1739 |
+
lat = sr.get("latency_ms", 0)
|
| 1740 |
+
label = (
|
| 1741 |
+
f"π€ {sid.replace('_', ' ').title()}"
|
| 1742 |
+
f" Β· {sr['status']} Β· {lat:.0f} ms"
|
| 1743 |
+
)
|
| 1744 |
+
with st.expander(label, expanded=True):
|
| 1745 |
+
st.markdown(
|
| 1746 |
+
f'<div style="border-left:3px solid {color};'
|
| 1747 |
+
f'padding:4px 0 4px 12px;margin-bottom:8px;">'
|
| 1748 |
+
f'<span style="font-size:10px;color:{color};font-weight:700;">{sid}</span>'
|
| 1749 |
+
f'<span style="font-size:10px;color:#475569;"> Β· status: </span>'
|
| 1750 |
+
f'<span style="font-size:10px;color:{ok_clr};">{sr["status"]}</span>'
|
| 1751 |
+
f'<span style="font-size:10px;color:#475569;"> Β· {lat:.0f} ms</span>'
|
| 1752 |
+
f'</div>',
|
| 1753 |
+
unsafe_allow_html=True,
|
| 1754 |
+
)
|
| 1755 |
+
st.code(sr["output"] or "(no output)", language=None)
|
| 1756 |
+
|
| 1757 |
+
# Synthesized / collective output
|
| 1758 |
+
sec("Synthesized Output Β· Collective Response")
|
| 1759 |
+
st.caption("All specialist outputs combined β this is what the orchestrator received.")
|
| 1760 |
+
if spec_results:
|
| 1761 |
+
parts = [
|
| 1762 |
+
f"{'β'*52}\n[{sr['id'].upper()}]\n{'β'*52}\n{sr['output'] or '(empty)'}"
|
| 1763 |
+
for sr in spec_results
|
| 1764 |
+
]
|
| 1765 |
+
synthesis = "\n\n".join(parts)
|
| 1766 |
+
else:
|
| 1767 |
+
synthesis = "(no specialists called β policy chose STOP on first step)"
|
| 1768 |
+
st.code(synthesis, language=None)
|
| 1769 |
+
|
| 1770 |
+
|
| 1771 |
# βββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
| 1772 |
# Entry point
|
| 1773 |
# βββββββββββββββββββββββββββββββββββββββββββββββββββββββββ
|
|
|
|
| 1777 |
S = _S()
|
| 1778 |
render_live_stats(S)
|
| 1779 |
|
| 1780 |
+
t1, t2, t3, t4, t5, t6, t7 = st.tabs([
|
| 1781 |
"β‘ Live Demo",
|
| 1782 |
"π€ Specialists",
|
| 1783 |
"π Training",
|
| 1784 |
"π Quality Demo",
|
| 1785 |
"π§ͺ Reward Lab",
|
| 1786 |
"π Architecture",
|
| 1787 |
+
"π― Output",
|
| 1788 |
])
|
| 1789 |
with t1: tab_live_demo()
|
| 1790 |
with t2: tab_specialists()
|
|
|
|
| 1792 |
with t4: tab_quality()
|
| 1793 |
with t5: tab_reward_lab()
|
| 1794 |
with t6: tab_architecture()
|
| 1795 |
+
with t7: tab_output()
|
| 1796 |
|
| 1797 |
|
| 1798 |
# Guard allows safe imports for testing without triggering the UI.
|