Team Ai
Apppublic

Wen1201/BayesianPyMc

sourceHugging Faceupdated 9mo agoView on Hugging Face
0likes
app_bayesian.py709 linesDownload Raw Back to root
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)