CoolFace
Apppublic

Wen1201/BayesianPyMc1

sourceHugging Faceupdated 8mo agoView on Hugging Face
0likes
bayesian_utils.py441 linesDownload Raw Back to root
1import plotly.graph_objects as go2import plotly.express as px3import pandas as pd4import numpy as np5import matplotlib.pyplot as plt6import matplotlib7matplotlib.use('Agg')  # 使用非互動式後端8import arviz as az9import io10import base6411from PIL import Image12 13def plot_trace(trace, var_names=['d', 'sigma']):14    """15    繪製 Trace Plot(MCMC 收斂診斷)16    包含完整的 warmup + posterior17    18    Args:19        trace: ArviZ InferenceData 物件20        var_names: 要繪製的變數名稱21        22    Returns:23        PIL Image24    """25    fig, axes = plt.subplots(len(var_names), 2, figsize=(14, 4 * len(var_names)))26    if len(var_names) == 1:27        axes = axes.reshape(1, -1)28    29    # 檢查是否有 warmup_posterior30    has_warmup = hasattr(trace, 'warmup_posterior') and trace.warmup_posterior is not None31    32    for idx, var_name in enumerate(var_names):33        # 左圖: KDE 密度圖(只用 posterior, 不用 warmup)34        post_data = trace.posterior[var_name].values35        for chain_idx in range(post_data.shape[0]):36            from scipy import stats37            data = post_data[chain_idx].flatten()38            density = stats.gaussian_kde(data)39            xs = np.linspace(data.min(), data.max(), 200)40            axes[idx, 0].plot(xs, density(xs), alpha=0.8, label=f'Chain {chain_idx+1}')41        axes[idx, 0].set_xlabel(var_name, fontsize=12)42        axes[idx, 0].set_ylabel('Density', fontsize=12)43        axes[idx, 0].set_title(f'{var_name}', fontsize=13, fontweight='bold')44        if idx == 0:45            axes[idx, 0].legend()46        47        # 右圖: Trace 圖(完整 warmup + posterior)48        if has_warmup:49            # 有 warmup: 合併繪製50            warmup_data = trace.warmup_posterior[var_name].values51            post_data = trace.posterior[var_name].values52            53            n_warmup = warmup_data.shape[1]54            n_post = post_data.shape[1]55            56            # 定義顏色,讓每條鏈用固定顏色57            colors = plt.cm.tab10.colors  # 使用 matplotlib 的顏色循環58            59            for chain_idx in range(warmup_data.shape[0]):60                chain_color = colors[chain_idx % len(colors)]  # 每條鏈一個固定顏色61                62                # 繪 warmup 部分63                x_warmup = np.arange(n_warmup)64                axes[idx, 1].plot(x_warmup, warmup_data[chain_idx].flatten(), 65                                color=chain_color,  # 👈 指定顏色66                                alpha=0.7, linewidth=0.5,67                                label=f'Chain {chain_idx+1}' if idx == 0 else '')68                69                # 繪 posterior 部分 (用同樣的顏色!)70                x_post = np.arange(n_warmup, n_warmup + n_post)71                axes[idx, 1].plot(x_post, post_data[chain_idx].flatten(), 72                                color=chain_color,  # 👈 同一個顏色73                                alpha=0.7, linewidth=0.5)74            75            # 加 Tune 結束的紅線76            axes[idx, 1].axvline(x=n_warmup, color='red', linestyle='--', 77                               linewidth=2, alpha=0.7, 78                               label='Tune結束' if idx == 0 else '')79                               80                               81                               82                               83        else:84            # 沒有 warmup: 只用 posterior85            post_data = trace.posterior[var_name].values86            for chain_idx in range(post_data.shape[0]):87                axes[idx, 1].plot(post_data[chain_idx].flatten(), 88                                alpha=0.7, linewidth=0.5,89                                label=f'Chain {chain_idx+1}' if idx == 0 else '')90        91        axes[idx, 1].set_xlabel('Iteration', fontsize=12)92        axes[idx, 1].set_ylabel(var_name, fontsize=12)93        axes[idx, 1].set_title(f'{var_name} trace', fontsize=13, fontweight='bold')94        if idx == 0:95            axes[idx, 1].legend(loc='upper right', fontsize=9)96        axes[idx, 1].grid(alpha=0.3)97    98    plt.tight_layout()99    100    # 轉換為圖片101    buf = io.BytesIO()102    plt.savefig(buf, format='png', dpi=300, bbox_inches='tight')103    buf.seek(0)104    img = Image.open(buf)105    plt.close()106    107    return img108 109 110# ============================================111# 替換說明:112# 在 bayesian_utils.py 中,把第 13-51 行的整個 plot_trace 函數113# 替換成上面這個版本114# ============================================115 116def plot_posterior(trace, var_names=['d', 'sigma', 'or_speed'], hdi_prob=0.95):117    """118    繪製後驗分佈圖119    120    Args:121        trace: ArviZ InferenceData 物件122        var_names: 要繪製的變數名稱123        hdi_prob: HDI 機率124        125    Returns:126        PIL Image127    """128    fig = az.plot_posterior(trace, var_names=var_names, hdi_prob=hdi_prob, figsize=(14, 5))129    plt.tight_layout()130    131    # 轉換為圖片132    buf = io.BytesIO()133    plt.savefig(buf, format='png', dpi=300, bbox_inches='tight')134    buf.seek(0)135    img = Image.open(buf)136    plt.close()137    138    return img139 140def plot_forest(trace, trial_labels, title='Effect of Speed on Win Rate by Type'):141    """142    繪製 Forest Plot(各屬性效應)143    144    Args:145        trace: ArviZ InferenceData 物件146        trial_labels: 屬性標籤列表147        title: 圖表標題148        149    Returns:150        PIL Image151    """152    num_trials = len(trial_labels)153    154    # 計算統計量155    delta_posterior = trace.posterior['delta'].values.reshape(-1, num_trials)156    delta_mean = delta_posterior.mean(axis=0)157    delta_hdi = az.hdi(trace, var_names=['delta'], hdi_prob=0.95)['delta'].values158    159    # 建立圖表160    fig, ax = plt.subplots(figsize=(12, max(10, num_trials * 0.4)))161    y_pos = np.arange(num_trials)162    163    # 繪製信賴區間(橫線)164    ax.hlines(y_pos, delta_hdi[:, 0], delta_hdi[:, 1], color='steelblue', linewidth=3, label='95% HDI')165    166    # 繪製平均值(點)167    ax.scatter(delta_mean, y_pos, color='darkblue', s=120, zorder=3, 168               edgecolors='white', linewidth=1.5, label='Mean')169    170    # 標註顯著的點171    for i, (mean, hdi) in enumerate(zip(delta_mean, delta_hdi)):172        if hdi[0] > 0:  # 顯著正效應173            ax.text(mean, i, ' ★', fontsize=15, ha='left', va='center', color='gold')174        elif hdi[1] < 0:  # 顯著負效應175            ax.text(mean, i, ' ☆', fontsize=15, ha='left', va='center', color='red')176    177    # 設定軸178    ax.set_yticks(y_pos)179    ax.set_yticklabels(trial_labels, fontsize=11)180    ax.invert_yaxis()181    ax.axvline(0, color='red', linestyle='--', linewidth=2, label='No Effect (δ=0)')182    ax.set_xlabel('Delta (Log Odds Ratio)', fontsize=13)183    ax.set_title(title, fontsize=15, fontweight='bold', pad=20)184    ax.legend(loc='lower right')185    ax.grid(axis='x', alpha=0.3)186    187    plt.tight_layout()188    189    # 轉換為圖片190    buf = io.BytesIO()191    plt.savefig(buf, format='png', dpi=300, bbox_inches='tight')192    buf.seek(0)193    img = Image.open(buf)194    plt.close()195    196    return img197 198def plot_model_dag(analyzer):199    """200    繪製模型 DAG 圖201    202    Args:203        analyzer: BayesianHierarchicalAnalyzer 物件204        205    Returns:206        PIL Image 或 None207    """208    try:209        gv = analyzer.get_model_graph()210        211        # 轉換為 PNG212        png_bytes = gv.pipe(format='png')213        214        # 轉換為 PIL Image215        img = Image.open(io.BytesIO(png_bytes))216        217        return img218    except Exception as e:219        print(f"無法生成 DAG 圖: {e}")220        return None221 222def create_summary_table(results):223    """224    創建結果摘要表格225    226    Args:227        results: 分析結果字典228        229    Returns:230        pandas DataFrame231    """232    overall = results['overall']233    234    summary_data = {235        '參數': ['d (整體效應)', 'sigma (配對間變異)', 'or_speed (勝算比)'],236        '平均值': [237            f"{overall['d_mean']:.4f}",238            f"{overall['sigma_mean']:.4f}",239            f"{overall['or_mean']:.4f}"240        ],241        '標準差': [242            f"{overall['d_sd']:.4f}",243            f"{overall['sigma_sd']:.4f}",244            f"{overall['or_sd']:.4f}"245        ],246        '95% HDI 下界': [247            f"{overall['d_hdi_low']:.4f}",248            f"{overall['sigma_hdi_low']:.4f}",249            f"{overall['or_hdi_low']:.4f}"250        ],251        '95% HDI 上界': [252            f"{overall['d_hdi_high']:.4f}",253            f"{overall['sigma_hdi_high']:.4f}",254            f"{overall['or_hdi_high']:.4f}"255        ]256    }257    258    return pd.DataFrame(summary_data)259 260 261def create_trial_results_table(results):262    """263    創建各配對結果表格 (使用動態欄位名稱)264    265    Args:266        results: 分析結果字典267        268    Returns:269        pandas DataFrame270    """271    trial_labels = results['trial_labels']272    by_trial = results['by_trial']273    data = results['data']274    col_names = results['column_names']275    276    # 動態獲取勝率欄位的鍵名277    control_key = f"p_{col_names['control_prefix']}_mean"278    treatment_key = f"p_{col_names['treatment_prefix']}_mean"279    280    trial_data = {281        '配對': trial_labels,282        'Delta (平均)': [f"{x:.4f}" for x in by_trial['delta_mean']],283        'Delta (標準差)': [f"{x:.4f}" for x in by_trial['delta_std']],284        '95% HDI 下界': [f"{x:.4f}" for x in by_trial['delta_hdi_low']],285        '95% HDI 上界': [f"{x:.4f}" for x in by_trial['delta_hdi_high']],286        '顯著性': ['★ 顯著' if sig else '不顯著' for sig in by_trial['delta_significant']],287        f"{col_names['control_prefix']}勝率": [f"{x:.2%}" for x in by_trial[control_key]],288        f"{col_names['treatment_prefix']}勝率": [f"{x:.2%}" for x in by_trial[treatment_key]],289        f"{col_names['control_prefix']} (勝/總)": [f"{d[col_names['control_win']]}/{d[col_names['control_total']]}" for d in data],290        f"{col_names['treatment_prefix']} (勝/總)": [f"{d[col_names['treatment_win']]}/{d[col_names['treatment_total']]}" for d in data]291    }292    293    return pd.DataFrame(trial_data)294 295 296def export_results_to_text(results):297    """298    匯出結果為純文字格式299    300    Args:301        results: 分析結果字典302        303    Returns:304        str: 格式化的文字報告305    """306    overall = results['overall']307    interp = results['interpretation']308    diag = results['diagnostics']309    col_names = results['column_names']310    311    report = f"""312==============================================313貝氏階層模型分析報告314==============================================315 316分析時間: {results['timestamp']}317配對數量: {results['n_trials']}318 319----------------------------------------------3201. 整體效應摘要321----------------------------------------------322d (整體效應 - Log OR):323  - 平均值: {overall['d_mean']:.4f}324  - 標準差: {overall['d_sd']:.4f}325  - 95% HDI: [{overall['d_hdi_low']:.4f}, {overall['d_hdi_high']:.4f}]326 327sigma (配對間變異):328  - 平均值: {overall['sigma_mean']:.4f}329  - 標準差: {overall['sigma_sd']:.4f}330  - 95% HDI: [{overall['sigma_hdi_low']:.4f}, {overall['sigma_hdi_high']:.4f}]331 332or_speed (勝算比):333  - 平均值: {overall['or_mean']:.4f}334  - 標準差: {overall['or_sd']:.4f}335  - 95% HDI: [{overall['or_hdi_low']:.4f}, {overall['or_hdi_high']:.4f}]336 337----------------------------------------------3382. 模型收斂診斷339----------------------------------------------340R-hat (d): {f"{diag['rhat_d']:.4f}" if diag['rhat_d'] is not None else 'N/A'}341R-hat (sigma): {f"{diag['rhat_sigma']:.4f}" if diag['rhat_sigma'] is not None else 'N/A'}342ESS (d): {int(diag['ess_d']) if diag['ess_d'] is not None else 'N/A'}343ESS (sigma): {int(diag['ess_sigma']) if diag['ess_sigma'] is not None else 'N/A'}344收斂狀態: {'✓ 已收斂' if diag['converged'] else '✗ 未收斂'}345 346----------------------------------------------3473. 結果解釋348----------------------------------------------349整體效應: {interp['overall_effect']}350顯著性: {interp['overall_significance']}351效果大小: {interp['effect_size']}352異質性: {interp['heterogeneity']}353 354----------------------------------------------3554. 各配對詳細結果356----------------------------------------------357"""358    359    # 添加各配對的詳細資訊360    trial_labels = results['trial_labels']361    by_trial = results['by_trial']362    363    # 動態獲取鍵名364    control_key = f"p_{col_names['control_prefix']}_mean"365    treatment_key = f"p_{col_names['treatment_prefix']}_mean"366    control_label = col_names['control_prefix'].capitalize()367    treatment_label = col_names['treatment_prefix'].capitalize()368    369    for i, label in enumerate(trial_labels):370        sig_marker = "★" if by_trial['delta_significant'][i] else " "371        report += f"""372{sig_marker} {label}:373  Delta (平均): {by_trial['delta_mean'][i]:.4f}374  95% HDI: [{by_trial['delta_hdi_low'][i]:.4f}, {by_trial['delta_hdi_high'][i]:.4f}]375  {control_label}勝率: {by_trial[control_key][i]:.2%}376  {treatment_label}勝率: {by_trial[treatment_key][i]:.2%}377  勝率差異: {(by_trial[treatment_key][i] - by_trial[control_key][i]):.2%}378"""379    380    report += """381==============================================382"""383    384    return report385    386 387 388def plot_odds_ratio_comparison(results):389    """390    繪製各屬性的勝算比比較圖(Plotly 版本)391    392    Args:393        results: 分析結果字典394        395    Returns:396        plotly figure397    """398    trial_labels = results['trial_labels']399    delta_mean = results['by_trial']['delta_mean']400    401    # 轉換為勝算比402    or_values = [np.exp(d) for d in delta_mean]403    404    # 排序405    sorted_indices = np.argsort(or_values)[::-1]406    sorted_labels = [trial_labels[i] for i in sorted_indices]407    sorted_or = [or_values[i] for i in sorted_indices]408    sorted_sig = [results['by_trial']['delta_significant'][i] for i in sorted_indices]409    410    # 顏色標記411    colors = ['#2ecc71' if sig else '#95a5a6' for sig in sorted_sig]412    413    fig = go.Figure()414    415    fig.add_trace(go.Bar(416        x=sorted_or,417        y=sorted_labels,418        orientation='h',419        marker=dict(420            color=colors,421            line=dict(color='white', width=1)422        ),423        text=[f'{or_val:.2f}' for or_val in sorted_or],424        textposition='outside',425        hovertemplate='%{y}<br>OR: %{x:.3f}<extra></extra>'426    ))427    428    # 參考線 (OR = 1)429    fig.add_vline(x=1, line_dash="dash", line_color="red", line_width=2)430    431    fig.update_layout(432        title='各屬性速度效應(勝算比)',433        xaxis_title='Odds Ratio',434        yaxis_title='',435        width=800,436        height=max(400, len(trial_labels) * 25),437        template='plotly_white',438        showlegend=False439    )440    441    return fig