CoolFace
Apppublic

Wen1201/BayesianPyMc

sourceHugging Faceupdated 8mo agoView on Hugging Face
0likes
bayesian_core.py313 linesDownload Raw Back to root
1import pandas as pd2import numpy as np3import pymc as pm4import arviz as az5import threading6from datetime import datetime7import warnings8warnings.filterwarnings('ignore')9 10class BayesianHierarchicalAnalyzer:11    """12    貝氏階層模型分析器13    用於分析寶可夢速度對勝率的影響(跨屬性)14    """15    16    # 類別級的鎖,用於執行緒安全17    _lock = threading.Lock()18    19    # 儲存各 session 的分析結果20    _session_results = {}21    22    def __init__(self, session_id):23        """24        初始化分析器25        26        Args:27            session_id: 唯一的 session 識別碼28        """29        self.session_id = session_id30        self.df = None31        self.model = None32        self.trace = None33    34    def load_data(self, csv_path_or_df):35        """36        載入資料37        38        Args:39            csv_path_or_df: CSV 檔案路徑或 DataFrame40            41        Expected columns:42            - Trial_Type: 屬性名稱 (e.g., Water, Fire, Grass)43            - rc: 控制組(速度慢)的勝場數44            - nc: 控制組的總場數45            - rt: 實驗組(速度快)的勝場數46            - nt: 實驗組的總場數47        """48        if isinstance(csv_path_or_df, str):49            self.df = pd.read_csv(csv_path_or_df)50        else:51            self.df = csv_path_or_df.copy()52        53        # 驗證必要欄位54        required_cols = ['Trial_Type', 'rc', 'nc', 'rt', 'nt']55        missing_cols = [col for col in required_cols if col not in self.df.columns]56        57        if missing_cols:58            raise ValueError(f"資料缺少必要欄位: {missing_cols}")59        60        return True61    62    def validate_data(self):63        """驗證資料有效性"""64        if self.df is None:65            raise ValueError("請先載入資料")66        67        # 檢查數值欄位68        for col in ['rc', 'nc', 'rt', 'nt']:69            if not pd.api.types.is_numeric_dtype(self.df[col]):70                raise ValueError(f"欄位 {col} 必須是數值類型")71        72        # 檢查邏輯約束73        if (self.df['rc'] > self.df['nc']).any():74            raise ValueError("rc (勝場數) 不能大於 nc (總場數)")75        76        if (self.df['rt'] > self.df['nt']).any():77            raise ValueError("rt (勝場數) 不能大於 nt (總場數)")78        79        return True80    81    def run_analysis(self, n_samples=2000, n_tune=1000, n_chains=2, target_accept=0.95):82        """83        執行貝氏階層模型分析84        85        Args:86            n_samples: MCMC 抽樣數87            n_tune: 調整期樣本數88            n_chains: 鏈數89            target_accept: 目標接受率90            91        Returns:92            dict: 包含所有分析結果的字典93        """94        with self._lock:95            try:96                self.validate_data()97                98                # 準備資料99                trial_labels = self.df['Trial_Type'].values100                num_trials = len(self.df)101                102                # 建立模型103                with pm.Model() as self.model:104                    # --- 先驗分佈 (Priors) ---105                    d = pm.Normal('d', mu=0, sigma=10)106                    tau = pm.Gamma('tau', alpha=0.001, beta=0.001)107                    sigma = pm.Deterministic('sigma', 1 / pm.math.sqrt(tau))108                    109                    # --- 各屬性特定效應 (Trial-specific effects) ---110                    mu = pm.Normal('mu', mu=0, sigma=10, shape=num_trials)111                    delta = pm.Normal('delta', mu=d, sigma=1 / pm.math.sqrt(tau), shape=num_trials)112                    113                    # --- 轉換與似然函數 (Logit Link & Likelihood) ---114                    pc = pm.Deterministic('pc', pm.math.invlogit(mu))115                    pt = pm.Deterministic('pt', pm.math.invlogit(mu + delta))116                    117                    rc_obs = pm.Binomial('rc_obs', n=self.df['nc'].values, p=pc, observed=self.df['rc'].values)118                    rt_obs = pm.Binomial('rt_obs', n=self.df['nt'].values, p=pt, observed=self.df['rt'].values)119                    120                    # --- 其他統計量 ---121                    delta_new = pm.Normal('delta_new', mu=d, sigma=1 / pm.math.sqrt(tau))122                    or_speed = pm.Deterministic('or_speed', pm.math.exp(d))123                    124                    # 執行 MCMC 抽樣125                    self.trace = pm.sample(126                        draws=n_samples,127                        tune=n_tune,128                        chains=n_chains,129                        target_accept=target_accept,130                        return_inferencedata=True,131                        progressbar=False, # 在 Streamlit 中關閉進度條132                        discard_tuned_samples=False  # 👈 加這行!保留 tune 樣本133                    )134                135                # 生成摘要統計136                summary = az.summary(self.trace, var_names=['d', 'sigma', 'or_speed'], hdi_prob=0.95)137                138                # 計算各屬性的 delta 統計量139                delta_posterior = self.trace.posterior['delta'].values.reshape(-1, num_trials)140                delta_mean = delta_posterior.mean(axis=0)141                delta_std = delta_posterior.std(axis=0)142                delta_hdi = az.hdi(self.trace, var_names=['delta'], hdi_prob=0.95)['delta'].values143                144                # 判斷顯著性(HDI 不包含 0)145                delta_significant = (delta_hdi[:, 0] > 0) | (delta_hdi[:, 1] < 0)146                147                # 計算控制組和實驗組的勝率148                pc_posterior = self.trace.posterior['pc'].values.reshape(-1, num_trials)149                pt_posterior = self.trace.posterior['pt'].values.reshape(-1, num_trials)150                151                pc_mean = pc_posterior.mean(axis=0)152                pt_mean = pt_posterior.mean(axis=0)153                154                # 整理結果155                results = {156                    'timestamp': datetime.now().isoformat(),157                    'n_trials': num_trials,158                    'trial_labels': trial_labels.tolist(),159                    160                    # 整體效應161                    'overall': {162                        'd_mean': float(summary.loc['d', 'mean']),163                        'd_sd': float(summary.loc['d', 'sd']),164                        'd_hdi_low': float(summary.loc['d', 'hdi_2.5%']),165                        'd_hdi_high': float(summary.loc['d', 'hdi_97.5%']),166                        167                        'sigma_mean': float(summary.loc['sigma', 'mean']),168                        'sigma_sd': float(summary.loc['sigma', 'sd']),169                        'sigma_hdi_low': float(summary.loc['sigma', 'hdi_2.5%']),170                        'sigma_hdi_high': float(summary.loc['sigma', 'hdi_97.5%']),171                        172                        'or_mean': float(summary.loc['or_speed', 'mean']),173                        'or_sd': float(summary.loc['or_speed', 'sd']),174                        'or_hdi_low': float(summary.loc['or_speed', 'hdi_2.5%']),175                        'or_hdi_high': float(summary.loc['or_speed', 'hdi_97.5%']),176                    },177                    178                    # 各屬性的效應179                    'by_trial': {180                        'delta_mean': delta_mean.tolist(),181                        'delta_std': delta_std.tolist(),182                        'delta_hdi_low': delta_hdi[:, 0].tolist(),183                        'delta_hdi_high': delta_hdi[:, 1].tolist(),184                        'delta_significant': delta_significant.tolist(),185                        'pc_mean': pc_mean.tolist(),186                        'pt_mean': pt_mean.tolist(),187                    },188                    189                    # 原始資料190                    'data': self.df.to_dict('records'),191                    192                    # 模型參數193                    'model_params': {194                        'n_samples': n_samples,195                        'n_tune': n_tune,196                        'n_chains': n_chains,197                        'target_accept': target_accept198                    },199                    200                    # 收斂診斷201                    'diagnostics': self._compute_diagnostics(summary),202                    203                    # 解釋204                    'interpretation': self._interpret_results(205                        summary.loc['or_speed', 'mean'],206                        summary.loc['or_speed', 'hdi_2.5%'],207                        summary.loc['or_speed', 'hdi_97.5%'],208                        summary.loc['sigma', 'mean']209                    )210                }211                212                # 儲存到 session results213                self._session_results[self.session_id] = results214                215                return results216                217            except Exception as e:218                raise Exception(f"分析失敗: {str(e)}")219    220    def _compute_diagnostics(self, summary):221        """計算收斂診斷指標"""222        try:223            # R-hat (應該接近 1.0)224            rhat_d = float(summary.loc['d', 'r_hat']) if 'r_hat' in summary.columns else None225            rhat_sigma = float(summary.loc['sigma', 'r_hat']) if 'r_hat' in summary.columns else None226            227            # ESS (有效樣本數)228            ess_d = float(summary.loc['d', 'ess_bulk']) if 'ess_bulk' in summary.columns else None229            ess_sigma = float(summary.loc['sigma', 'ess_bulk']) if 'ess_bulk' in summary.columns else None230            231            return {232                'rhat_d': rhat_d,233                'rhat_sigma': rhat_sigma,234                'ess_d': ess_d,235                'ess_sigma': ess_sigma,236                'converged': (rhat_d is None or rhat_d < 1.1) and (rhat_sigma is None or rhat_sigma < 1.1)237            }238        except:239            return {240                'converged': None,241                'rhat_d': None,242                'rhat_sigma': None,243                'ess_d': None,244                'ess_sigma': None245            }246    247    def _interpret_results(self, or_mean, or_low, or_high, sigma_mean):248        """解釋分析結果"""249        # 整體效應顯著性       250        if or_low > 1:251            overall_effect = "火系寶可夢相對於水系顯著更容易獲勝"252            overall_significance = "顯著正效應"253        elif or_high < 1:254            overall_effect = "水系寶可夢相對於火系顯著更容易獲勝"255            overall_significance = "顯著負效應"256        else:257            overall_effect = "火系與水系勝率無顯著差異"258            overall_significance = "不顯著"        259 260        # 效果大小           261        if or_mean > 2:262            effect_size = "大效果 (OR > 2) - 火系有明顯優勢"263        elif or_mean > 1.5:264            effect_size = "中等效果 (OR > 1.5) - 火系有一定優勢"265        elif or_mean > 1:266            effect_size = "小效果 (OR > 1) - 火系略有優勢"267        elif or_mean == 1:268            effect_size = "無差異 (OR = 1) - 火系與水系勢均力敵"269        elif or_mean > 0.67:270            effect_size = "小效果 (OR < 1) - 水系略有優勢"271        elif or_mean > 0.5:272            effect_size = "中等效果 (OR < 0.67) - 水系有一定優勢"273        else:274            effect_size = "大效果 (OR < 0.5) - 水系有明顯優勢"          275                    276        277        # 異質性評估278        if sigma_mean > 0.5:279            heterogeneity = "高異質性 - 不同配對的勝率差異很大"280        elif sigma_mean > 0.3:281            heterogeneity = "中等異質性 - 不同配對的勝率有一定差異"282        else:283            heterogeneity = "低異質性 - 不同配對的勝率相對一致"        284              285        return {286            'overall_effect': overall_effect,287            'overall_significance': overall_significance,288            'effect_size': effect_size,289            'heterogeneity': heterogeneity290        }291    292    def get_model_graph(self):293        """生成模型 DAG 圖(返回 graphviz 物件)"""294        if self.model is None:295            raise ValueError("請先執行分析")296        297        try:298            gv = pm.model_to_graphviz(self.model)299            return gv300        except Exception as e:301            raise Exception(f"無法生成 DAG 圖: {str(e)}")302    303    @classmethod304    def get_session_results(cls, session_id):305        """獲取特定 session 的結果"""306        return cls._session_results.get(session_id)307    308    @classmethod309    def clear_session_results(cls, session_id):310        """清除特定 session 的結果"""311        if session_id in cls._session_results:312            del cls._session_results[session_id]313