Wen1201/BayesianPyMc1
0
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