garvitsachdeva commited on
Commit
44f11d2
Β·
1 Parent(s): 0c3df82

trained evironment UI fixes

Browse files
demo/assets/reward_curve.json CHANGED
@@ -1 +1 @@
1
- {"episodes": [0, 11, 22, 33, 44, 55, 66, 77, 88, 99, 110, 121, 132, 143, 154, 165, 176, 187, 198, 209, 220, 231, 242, 253, 264, 275, 286, 297, 308, 319, 330, 341, 352, 363, 374, 385, 396, 407, 418, 429, 440, 451, 462, 473, 484, 495, 506, 517, 528, 539, 550, 561, 572, 583, 594, 605, 616, 627, 638, 649, 660, 671, 682, 693, 704, 715, 726, 737, 748, 759, 770, 781, 792, 803, 814, 825, 836, 847, 858, 869, 880, 891, 902, 913, 924, 935, 946, 957, 968, 979, 990, 1001, 1012, 1023, 1034, 1045, 1056, 1067, 1078, 1089, 1100, 1111, 1122, 1133, 1144, 1155, 1166, 1177, 1188, 1199, 1210, 1221, 1232, 1243, 1254, 1265, 1276, 1287, 1298, 1309, 1320, 1331, 1342, 1353, 1364, 1375, 1386, 1397, 1408, 1419, 1430, 1441, 1452, 1463, 1474, 1485, 1496, 1507, 1518, 1529, 1540, 1551, 1562, 1573, 1584, 1595, 1606, 1617, 1628, 1639, 1650, 1661, 1672, 1683, 1694, 1705, 1716, 1727, 1738, 1749, 1760, 1771, 1782, 1793, 1804, 1815, 1826, 1837, 1848, 1859, 1870, 1881, 1892, 1903, 1914, 1925, 1936, 1947, 1958, 1969, 1980, 1991, 2002, 2013, 2024, 2035, 2046, 2057, 2068, 2079, 2090, 2101, 2112, 2123, 2134, 2145, 2156, 2167, 2178, 2189, 2200], "mean_rewards": [-2.6738038063049316, -1.705311691761017, -2.2153279781341553, -1.8650923013687133, -1.9583142399787903, -1.8090984106063843, -2.5727408647537233, -2.006777358055115, -2.0646845579147337, -1.1843333005905152, -0.8511799693107605, -1.2869279697537421, -2.5326566219329836, -0.7975572127848863, -2.2941975355148316, -1.4255218148231505, -1.9773519873619079, -1.829572582244873, -2.2942489624023437, -1.592001461982727, -1.8560773760080338, -2.144868350028992, -1.937927508354187, -1.1779373006895184, -1.5583532094955443, -1.5792918443679809, -1.2494795009493829, -2.25146803855896, -1.8984802484512329, -1.3299775309860706, -0.8860581159591675, -0.6782042820006609, -1.4215008795261384, -0.8339593816548586, -2.1198282480239867, -1.8454582929611205, -1.2758302211761474, -1.1315348207950593, -1.375254637002945, -2.120091676712036, -1.5853264234960078, -1.157479214668274, -1.266526734828949, -0.948374779522419, -2.1824836492538453, -1.2791759371757507, -0.9780700504779816, -1.8573646306991578, -1.4734271883964538, -0.45685309171676636, -1.6383790135383607, -1.0759720027446746, -1.504695177078247, -1.726955735683441, -1.088908851146698, -0.9255473613739014, -1.5862729153595865, -1.8054921865463256, -1.5902058459818362, -0.7862645149230957, -1.2847756624221802, -0.4538323223590851, -0.24534327983856202, -0.7213976144790649, -0.7808282060548664, -1.2140628814697265, -0.24957830905914308, -0.7205866644158959, -1.0317823708057403, -0.36452836729586124, -0.9707806944847107, -0.14061078652739525, -1.054512779880315, -0.4149759531021118, -1.2930978775024413, -0.8258169777691364, -1.356018888950348, -0.8899088740348816, -1.6979908108711244, -0.6806863307952881, -0.9120665602385998, 0.395650053024292, -1.86594614982605, -0.873254942893982, -1.5391783475875855, -0.7206376433372498, -0.5297608852386475, 0.46408586725592615, 0.21402924209833146, -0.24489773511886598, 0.08052548803389073, 0.6628240764141082, -1.275925225019455, -0.3005677070468664, -0.4723848819732666, -0.29810856431722643, -0.4034378886222839, -0.8178201481699944, 0.46010567545890807, -0.9913323003798723, 0.2993836283683777, 0.08219350576400757, -0.34826181530952455, -0.879417422413826, 0.40615544966422024, 0.9001223504543304, 0.5579557850956917, -0.18564149364829063, 0.05578359365463257, 0.38205742835998535, -1.4494811177253724, 0.04445687234401703, -0.3005406914278865, -0.7186087477952242, 0.023816481232643127, -0.3200356105342507, 0.1748729705810547, 0.49465489387512207, 0.09322566390037537, -0.20863972902297973, -0.013048544526100159, -0.2582117199897766, 0.30120803266763685, 0.13326873779296874, -1.7269521832466126, 0.22264335341751576, 0.2890779085457325, 0.25854286178946495, 0.028514337539672852, 0.15758876800537108, 0.9122146368026733, 0.025657114386558533, 0.8382625341415405, 0.8449460297822953, 0.7839016802608967, 0.33553348779678344, 0.6816077768802643, -0.13622485473752022, 0.8707041293382645, 1.0687336444854736, -0.34334572553634646, -0.43794297277927396, 0.515097776055336, -0.8650284081697464, -0.20771026611328125, 0.13080331087112426, 0.647852110862732, -0.26858361195772884, 0.09040446281433105, 0.5966767907142639, 0.7839245915412902, 0.9312916576862336, -0.8558926701545715, 0.8143998086452484, 1.2133472323417664, -0.05484856106340885, 0.693803608417511, 0.9091606378555298, 0.4998580813407898, 0.7885102093219757, 0.31582592204213145, 0.8510897813364864, 0.11140216141939163, 0.9307787224650383, 0.7449860155582428, 0.8639730155467987, 0.9730179116129876, -0.652894401550293, 0.30474201031029224, 0.7902945404872298, 0.7700751990079879, 0.5174719452857971, 0.9151068434119225, 0.84403036236763, 0.8516681623645127, 0.13887905478477477, 0.9150871947407723, -0.6614223957061768, 0.9483977686613798, 1.0316770553588868, 1.0025377452373505, 1.1537045121192933, 0.2673381119966507, 0.9019387006759644, 0.6476128563284874, 0.672609269618988, 0.9197988472878933, 0.9209991149604321, 1.0379021286964416, 0.8294112265110016, 0.9367486596107483, 0.5053324922919273, 0.5285568356513977, 0.5070471465587616, 0.6434216737747193, 0.3712703872472048, -0.25931897163391116, 0.49494273737072947, 0.8008696258068084, 0.8263677477836608, -0.2617871671915054]}
 
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
- os.environ.setdefault("HF_HUB_OFFLINE", "1")
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
- model = RecurrentPPO.load(
91
- hf_hub_download(hf_repo, "spindleflow_model.zip"), device="cpu"
92
- )
 
 
 
 
93
  obs_mean = obs_var = None
94
  clip_obs = 10.0
95
  try:
96
- stats_path = hf_hub_download(hf_repo, "vec_normalize.pkl")
 
 
 
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
- if _old_hf is not None:
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
- src = hf_hub_download(HF_MODEL_REPO, "reward_curve.json")
 
 
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.