Wen1201/BayesianPyMc
0
1import streamlit as st2import pandas as pd3import uuid4from datetime import datetime, timedelta5import atexit6import os7import sys8 9# 頁面配置10st.set_page_config(11 page_title="Bayesian Hierarchical Model - Pokémon Speed Analysis",12 page_icon="🎲",13 layout="wide",14 initial_sidebar_state="expanded"15)16 17# 自定義 CSS18st.markdown("""19<style>20 .streamlit-expanderHeader {21 background-color: #e8f1f8;22 border: 1px solid #b0cfe8;23 border-radius: 5px;24 font-weight: 600;25 color: #1b4f72;26 }27 .streamlit-expanderHeader:hover {28 background-color: #d0e7f8;29 }30 .stMetric {31 background-color: #f8fbff;32 padding: 10px;33 border-radius: 5px;34 border: 1px solid #d0e4f5;35 }36 .stButton > button {37 width: 100%;38 border-radius: 20px;39 font-weight: 600;40 transition: all 0.3s ease;41 }42 .stButton > button:hover {43 transform: translateY(-2px);44 box-shadow: 0 4px 8px rgba(0,0,0,0.2);45 }46 .success-box {47 background-color: #d4edda;48 border: 1px solid #c3e6cb;49 border-radius: 5px;50 padding: 10px;51 margin: 10px 0;52 }53 .warning-box {54 background-color: #fff3cd;55 border: 1px solid #ffeaa7;56 border-radius: 5px;57 padding: 10px;58 margin: 10px 0;59 }60</style>61""", unsafe_allow_html=True)62 63# 導入自定義模組64from bayesian_core import BayesianHierarchicalAnalyzer65# 注意:如果要啟用 DAG 動態生成功能,請將下行改為:66# from bayesian_llm_assistant_enhanced import BayesianLLMAssistant67from bayesian_llm_assistant import BayesianLLMAssistant68from bayesian_utils import (69 plot_trace,70 plot_posterior,71 plot_forest,72 plot_model_dag,73 create_summary_table,74 create_trial_results_table,75 export_results_to_text,76 plot_odds_ratio_comparison77)78 79# 清理函數80def cleanup_old_sessions():81 """清理超過 1 小時的 session"""82 current_time = datetime.now()83 for session_id in list(BayesianHierarchicalAnalyzer._session_results.keys()):84 result = BayesianHierarchicalAnalyzer._session_results.get(session_id)85 if result:86 result_time = datetime.fromisoformat(result['timestamp'])87 if current_time - result_time > timedelta(hours=1):88 BayesianHierarchicalAnalyzer.clear_session_results(session_id)89 90# 註冊清理函數91atexit.register(cleanup_old_sessions)92 93# 初始化 session state94if 'session_id' not in st.session_state:95 st.session_state.session_id = str(uuid.uuid4())96if 'analysis_results' not in st.session_state:97 st.session_state.analysis_results = None98if 'chat_history' not in st.session_state:99 st.session_state.chat_history = []100if 'analyzer' not in st.session_state:101 st.session_state.analyzer = None102if 'trace_img' not in st.session_state:103 st.session_state.trace_img = None104if 'posterior_img' not in st.session_state:105 st.session_state.posterior_img = None106if 'forest_img' not in st.session_state:107 st.session_state.forest_img = None108if 'dag_img' not in st.session_state:109 st.session_state.dag_img = None110 111# 標題112st.title("🎲 Bayesian Hierarchical Model Analysis")113st.markdown("### 火系 vs 水系寶可夢配對勝率的貝氏階層分析")114st.markdown("---")115 116# Sidebar117with st.sidebar:118 st.header("⚙️ 配置設定")119 120 # API 選擇121 api_choice = st.radio(122 "選擇 LLM API",123 options=["Google Gemini", "Anthropic Claude"],124 index=0,125 help="選擇要使用的 AI 助手"126 )127 128 # API Key 輸入129 if api_choice == "Google Gemini":130 api_key = st.text_input(131 "Google Gemini API Key",132 type="password",133 help="輸入您的 Google Gemini API Key"134 )135 else: # Claude136 api_key = st.text_input(137 "Anthropic Claude API Key",138 type="password",139 help="輸入您的 Anthropic API Key (https://console.anthropic.com)"140 )141 142 if api_key:143 st.session_state.api_key = api_key144 st.session_state.api_choice = api_choice # 新增:儲存 API 選擇145 st.success(f"✅ {api_choice} API Key 已載入")146 147 st.markdown("---")148 149 # MCMC 參數設定150 st.subheader("🔬 MCMC 參數")151 152 n_samples = st.number_input(153 "抽樣數 (Samples)",154 min_value=500,155 max_value=10000,156 value=2000,157 step=500,158 help="每條鏈的抽樣數量"159 )160 161 n_tune = st.number_input(162 "調整期 (Tune)",163 min_value=200,164 max_value=5000,165 value=1000,166 step=200,167 help="調整期的樣本數"168 )169 170 n_chains = st.selectbox(171 "鏈數 (Chains)",172 options=[1, 2, 4],173 index=1,174 help="平行運行的鏈數"175 )176 177 target_accept = st.slider(178 "目標接受率",179 min_value=0.80,180 max_value=0.99,181 value=0.95,182 step=0.01,183 help="NUTS 採樣器的目標接受率"184 )185 186 st.markdown("---")187 188 # 清理按鈕189 if st.button("🧹 清理過期資料"):190 cleanup_old_sessions()191 st.success("✅ 清理完成")192 st.rerun()193 194 st.markdown("---")195 196 # 資料來源選擇197 st.subheader("📊 資料來源")198 data_source = st.radio(199 "選擇資料來源:",200 ["使用預設資料集", "上傳您的資料"]201 )202 203 uploaded_file = None204 if data_source == "上傳您的資料":205 uploaded_file = st.file_uploader(206 "上傳 CSV 檔案",207 type=['csv'],208 help="上傳寶可夢速度對戰資料"209 )210 211 with st.expander("📖 資料格式說明"):212 st.markdown("""213 **必要欄位格式:**214 - `Trial_Type`: 配對名稱(例如:Pair_1, Pair_2)215 - `rt`: 火系(治療組)的勝場數216 - `nt`: 火系的總場數217 - `rc`: 水系(對照組)的勝場數218 - `nc`: 水系的總場數219 220 **範例:**221 ```222 Trial_Type,rt,nt,rc,nc223 Pair_1,122,133,22,145224 Pair_2,85,132,17,135225 Pair_3,52,129,41,134226 ```227 """) 228 229 st.markdown("---")230 231 232 # 關於系統233 with st.expander("ℹ️ 關於此系統"):234 st.markdown("""235 **貝氏階層模型分析系統**236 237 本系統使用貝氏階層模型來分析速度對寶可夢勝率的影響,238 並考慮不同屬性之間的異質性。239 240 **主要功能:**241 - 🎲 貝氏推論與後驗分佈242 - 📊 階層模型(借用資訊)243 - 📈 4 種視覺化圖表244 - 💬 AI 助手解釋245 - 🎮 屬性對抗策略建議246 247 **適用場景:**248 - 分析火系對水系的配對勝率249 - 理解不同配對間的異質性250 - 評估屬性優劣勢251 """)252 253# 主要內容區 - 雙 Tab254tab1, tab2 = st.tabs(["📊 貝氏分析", "💬 AI 助手"])255 256# Tab 1: 貝氏分析257with tab1:258 st.header("📊 貝氏階層模型分析")259 260 # 載入資料261 if data_source == "使用預設資料集":262 # 檢查預設資料是否存在263 default_data_path = "fire_water_converted.csv"264 if os.path.exists(default_data_path):265 df = pd.read_csv(default_data_path)266 st.success(f"✅ 已載入預設資料集({len(df)} 組配對)")267 else:268 st.warning("⚠️ 找不到預設資料集,請上傳您的資料")269 df = None270 else:271 if uploaded_file is not None:272 df = pd.read_csv(uploaded_file)273 st.success(f"✅ 已載入資料({len(df)} 組配對)")274 else:275 df = None276 st.info("📁 請在左側上傳 CSV 檔案")277 278 if df is not None:279 # 顯示資料預覽280 with st.expander("👀 資料預覽"):281 st.dataframe(df, use_container_width=True)282 283 st.markdown("---")284 285 # 分析按鈕286 col1, col2, col3 = st.columns([1, 2, 1])287 288 with col2:289 analyze_button = st.button(290 "🔬 開始貝氏分析",291 type="primary",292 use_container_width=True293 )294 295 # 執行分析296 if analyze_button:297 with st.spinner(f"正在執行貝氏分析... (抽樣 {n_samples} × {n_chains} 條鏈)"):298 try:299 # 初始化分析器300 if st.session_state.analyzer is None:301 st.session_state.analyzer = BayesianHierarchicalAnalyzer(st.session_state.session_id)302 303 # 載入資料304 st.session_state.analyzer.load_data(df)305 306 # 執行分析307 results = st.session_state.analyzer.run_analysis(308 n_samples=n_samples,309 n_tune=n_tune,310 n_chains=n_chains,311 target_accept=target_accept312 )313 314 st.session_state.analysis_results = results315 316 # 生成圖表317 with st.spinner("生成視覺化圖表..."):318 st.session_state.trace_img = plot_trace(st.session_state.analyzer.trace)319 st.session_state.posterior_img = plot_posterior(st.session_state.analyzer.trace)320 st.session_state.forest_img = plot_forest(321 st.session_state.analyzer.trace,322 results['trial_labels']323 )324 st.session_state.dag_img = plot_model_dag(st.session_state.analyzer)325 326 st.success("✅ 分析完成!")327 st.balloons()328 329 except Exception as e:330 st.error(f"❌ 分析失敗: {str(e)}")331 332 # 顯示結果333 if st.session_state.analysis_results is not None:334 results = st.session_state.analysis_results335 336 st.markdown("---")337 st.subheader("📊 分析結果")338 339 # 創建 4 個子頁面340 result_tabs = st.tabs([341 "📊 概覽",342 "📈 Trace & Posterior",343 "🌲 Forest Plot",344 "🔍 DAG 模型圖",345 "📋 詳細報告"346 ])347 348 # Tab: 概覽349 with result_tabs[0]:350 st.markdown("### 🎯 整體效應摘要")351 352 overall = results['overall']353 interp = results['interpretation']354 355 # 關鍵指標356 col1, col2, col3 = st.columns(3)357 358 with col1:359 st.metric(360 "d (整體效應)",361 f"{overall['d_mean']:.4f}",362 delta=f"HDI: [{overall['d_hdi_low']:.3f}, {overall['d_hdi_high']:.3f}]"363 )364 365 with col2:366 st.metric(367 "勝算比 (OR)",368 f"{overall['or_mean']:.3f}",369 delta=f"HDI: [{overall['or_hdi_low']:.3f}, {overall['or_hdi_high']:.3f}]"370 )371 372 with col3:373 st.metric(374 "sigma (異質性)",375 f"{overall['sigma_mean']:.4f}",376 delta=f"HDI: [{overall['sigma_hdi_low']:.3f}, {overall['sigma_hdi_high']:.3f}]"377 )378 379 st.markdown("---")380 381 # 結果解釋382 st.markdown("### 📖 結果解釋")383 384 st.info(f"""385 **整體效應**: {interp['overall_effect']}386 387 **顯著性**: {interp['overall_significance']}388 389 **效果大小**: {interp['effect_size']}390 391 **異質性**: {interp['heterogeneity']}392 """)393 394 st.markdown("---")395 396 # 收斂診斷397 st.markdown("### 🔍 模型收斂診斷")398 399 diag = results['diagnostics']400 401 col1, col2 = st.columns(2)402 403 with col1:404 st.markdown("**R-hat 診斷** (應 < 1.1):")405 if diag['rhat_d']:406 st.metric("R-hat (d)", f"{diag['rhat_d']:.4f}", 407 delta="✓ 良好" if diag['rhat_d'] < 1.1 else "✗ 需改善")408 if diag['rhat_sigma']:409 st.metric("R-hat (sigma)", f"{diag['rhat_sigma']:.4f}",410 delta="✓ 良好" if diag['rhat_sigma'] < 1.1 else "✗ 需改善")411 412 with col2:413 st.markdown("**有效樣本數 (ESS)**:")414 if diag['ess_d']:415 st.metric("ESS (d)", f"{int(diag['ess_d'])}")416 if diag['ess_sigma']:417 st.metric("ESS (sigma)", f"{int(diag['ess_sigma'])}")418 419 if diag['converged']:420 st.success("✅ 模型已收斂,結果可信")421 else:422 st.warning("⚠️ 模型可能未完全收斂,建議增加抽樣數或鏈數")423 424 st.markdown("---")425 426 # 摘要表格427 st.markdown("### 📊 統計摘要表")428 summary_df = create_summary_table(results)429 st.dataframe(summary_df, use_container_width=True)430 431 st.markdown("---")432 433 # 各屬性結果434 st.markdown("### 🎮 各屬性詳細結果")435 trial_df = create_trial_results_table(results)436 st.dataframe(trial_df, use_container_width=True)437 438 st.markdown("---")439 440 # 勝算比比較圖441 st.markdown("### 📊 各屬性速度效應比較")442 or_fig = plot_odds_ratio_comparison(results)443 st.plotly_chart(or_fig, use_container_width=True)444 445 # Tab: Trace & Posterior446 with result_tabs[1]:447 st.markdown("### 📈 Trace Plot(收斂診斷)")448 st.markdown("""449 **Trace Plot 用途**:450 - 檢查 MCMC 抽樣是否收斂451 - 左圖:抽樣軌跡(應該像「毛毛蟲」)452 - 右圖:後驗分佈密度453 """)454 455 if st.session_state.trace_img:456 st.image(st.session_state.trace_img, use_column_width=True)457 else:458 st.info("請先執行分析以生成 Trace Plot")459 460 st.markdown("---")461 462 st.markdown("### 📊 Posterior Plot(後驗分佈)")463 st.markdown("""464 **Posterior Plot 用途**:465 - 顯示參數的後驗分佈466 - 包含 95% HDI(最高密度區間)467 - 顯示平均值468 """)469 470 if st.session_state.posterior_img:471 st.image(st.session_state.posterior_img, use_column_width=True)472 else:473 st.info("請先執行分析以生成 Posterior Plot")474 475 # Tab: Forest Plot476 with result_tabs[2]:477 st.markdown("### 🌲 Forest Plot(各屬性效應)")478 st.markdown("""479 **Forest Plot 用途**:480 - 顯示每個屬性的速度效應(delta)481 - 點:平均效應482 - 線:95% HDI483 - ★ 標記:顯著正效應(HDI 不包含 0)484 - ☆ 標記:顯著負效應485 """)486 487 if st.session_state.forest_img:488 st.image(st.session_state.forest_img, use_column_width=True)489 else:490 st.info("請先執行分析以生成 Forest Plot")491 492 # Tab: DAG 模型圖493 with result_tabs[3]:494 st.markdown("### 🔍 模型結構圖 (DAG)")495 st.markdown("""496 **DAG(有向無環圖)用途**:497 - 視覺化模型的階層結構498 - 顯示變數之間的依賴關係499 - 圓形/橢圓:隨機變數500 - 矩形:觀測資料501 - 菱形:推導變數502 """)503 504 if st.session_state.dag_img:505 st.image(st.session_state.dag_img, use_column_width=True)506 else:507 st.warning("⚠️ 無法生成 DAG 圖(可能需要安裝 Graphviz)")508 st.markdown("""509 **安裝 Graphviz:**510 - Windows: `choco install graphviz`511 - Mac: `brew install graphviz`512 - Ubuntu: `sudo apt-get install graphviz`513 """)514 515 # Tab: 詳細報告516 with result_tabs[4]:517 st.markdown("### 📋 完整分析報告")518 519 # 生成文字報告520 text_report = export_results_to_text(results)521 522 st.text_area(523 "報告內容",524 text_report,525 height=500526 )527 528 # 下載按鈕529 st.download_button(530 label="📥 下載完整報告 (.txt)",531 data=text_report,532 file_name=f"bayesian_report_{results['timestamp'][:10]}.txt",533 mime="text/plain"534 )535 536# Tab 2: AI 助手537with tab2:538 st.header("💬 AI 分析助手")539 540 if not st.session_state.get('api_key'):541 st.warning("⚠️ 請在左側輸入您的 Google Gemini API Key 以使用 AI 助手")542 elif st.session_state.analysis_results is None:543 st.info("ℹ️ 請先在「貝氏分析」頁面執行分析")544 else:545 # 初始化 LLM 助手546 if 'llm_assistant' not in st.session_state:547 api_choice = st.session_state.get('api_choice', 'Google Gemini')548 st.session_state.llm_assistant = BayesianLLMAssistant(549 api_key=st.session_state.api_key,550 session_id=st.session_state.session_id,551 api_provider=api_choice # 新增:傳遞 API 選擇552 )553 554 # 聊天容器555 chat_container = st.container()556 557 with chat_container:558 for message in st.session_state.chat_history:559 with st.chat_message(message["role"]):560 st.markdown(message["content"])561 # 如果訊息包含 DAG 圖,顯示圖片562 if message.get("has_dag", False) and message.get("dag_image") is not None:563 st.image(message["dag_image"], caption="🎨 生成的 DAG 圖", use_column_width=True) 564 565 566 # 使用者輸入567 if prompt := st.chat_input("詢問關於分析結果的任何問題..."):568 # 添加使用者訊息569 st.session_state.chat_history.append({570 "role": "user",571 "content": prompt572 })573 574 with st.chat_message("user"):575 st.markdown(prompt)576 577 # AI 回應578 with st.chat_message("assistant"):579 with st.spinner("思考中..."):580 try:581 # 修改:接收回應和可能的 DAG 圖片582 response, dag_image = st.session_state.llm_assistant.get_response(583 user_message=prompt,584 analysis_results=st.session_state.analysis_results585 )586 st.markdown(response)587 588 # 如果有生成 DAG 圖,顯示它589 if dag_image is not None:590 st.image(dag_image, caption="🎨 AI 生成的 DAG 圖", use_column_width=True)591 st.success("✨ DAG 圖已生成!你可以繼續詢問圖表相關問題。")592 593 except Exception as e:594 error_msg = f"❌ 錯誤: {str(e)}\n\n請檢查 API key 或重新表達問題。"595 st.error(error_msg)596 response = error_msg597 dag_image = None598 599 # 添加助手回應(包含 DAG 標記)600 st.session_state.chat_history.append({601 "role": "assistant",602 "content": response,603 "has_dag": dag_image is not None,604 "dag_image": dag_image # 新增:保存圖片605 })606 607 st.markdown("---")608 609 # 快速問題按鈕610 st.subheader("💡 快速問題")611 612 # 添加使用提示613 st.info("💡 提示:你可以要求助手「畫一個 DAG 圖」來視覺化模型結構!")614 615 quick_questions = [616 "📊 給我這次分析的總結",617 "🎯 解釋 d 和勝算比",618 "🔍 解釋 sigma(異質性)",619 "❓ 什麼是階層模型?",620 "🎨 畫一個模型結構圖", # 新增 DAG 生成按鈕621 "🆚 貝氏 vs 頻率論",622 "⚔️ 對戰策略建議",623 "🎮 比較不同屬性"624 ]625 626 cols = st.columns(4)627 for idx, question in enumerate(quick_questions):628 col_idx = idx % 4629 if cols[col_idx].button(question, key=f"quick_{idx}"):630 # 根據問題選擇對應的方法631 if "總結" in question:632 response = st.session_state.llm_assistant.generate_summary(633 st.session_state.analysis_results634 )635 dag_image = None # 這些方法不返回圖片636 elif "d 和勝算比" in question:637 response = st.session_state.llm_assistant.explain_metric(638 'd',639 st.session_state.analysis_results640 )641 dag_image = None642 elif "sigma" in question or "異質性" in question:643 response = st.session_state.llm_assistant.explain_metric(644 'sigma',645 st.session_state.analysis_results646 )647 dag_image = None648 elif "階層模型" in question:649 response = st.session_state.llm_assistant.explain_hierarchical_model()650 dag_image = None651 elif "畫一個" in question or "結構圖" in question:652 # DAG 生成請求653 response, dag_image = st.session_state.llm_assistant.get_response(654 "請畫一個貝氏階層模型的 DAG 圖,並用繁體中文解釋每個節點的意義",655 st.session_state.analysis_results656 )657 elif "貝氏" in question and "頻率論" in question:658 response = st.session_state.llm_assistant.explain_bayesian_vs_frequentist()659 dag_image = None660 elif "策略" in question:661 response = st.session_state.llm_assistant.battle_strategy_advice(662 st.session_state.analysis_results663 )664 dag_image = None665 elif "比較" in question:666 response = st.session_state.llm_assistant.compare_types(667 st.session_state.analysis_results668 )669 dag_image = None670 else:671 response, dag_image = st.session_state.llm_assistant.get_response(672 question,673 st.session_state.analysis_results674 )675 676 # 添加到聊天歷史677 st.session_state.chat_history.append({678 "role": "user",679 "content": question680 })681 682 st.session_state.chat_history.append({683 "role": "assistant",684 "content": response,685 "has_dag": dag_image is not None if 'dag_image' in locals() else False,686 "dag_image": dag_image if 'dag_image' in locals() else None687 })688 689 st.rerun()690 691 # 重置對話按鈕692 st.markdown("---")693 if st.button("🔄 重置對話"):694 st.session_state.llm_assistant.reset_conversation()695 st.session_state.chat_history = []696 st.success("✅ 對話已重置")697 st.rerun()698 699# Footer700st.markdown("---")701st.markdown(702 f"""703 <div style='text-align: center'>704 <p>🎲 Bayesian Hierarchical Model Analysis for Pokémon Speed | Built with Streamlit & PyMC</p>705 <p>Session ID: {st.session_state.session_id[:8]} | Powered by Google Gemini 2.0 Flash</p>706 </div>707 """,708 unsafe_allow_html=True709)