"""Hierarchical Bayesian partial pooling for player skill estimation. Implements two modes: - PyMC: Full MCMC-based hierarchical model with role-level priors - Scipy fallback: James-Stein-style shrinkage with empirical Bayes estimates """ import logging from typing import Optional, Tuple import numpy as np import pandas as pd from .base_model import BaseModel logger = logging.getLogger(__name__) VALID_ROLES = {"P", "D", "C", "A"} def _has_pymc() -> bool: try: import pymc as pm # noqa: F401 return True except ImportError: return False class BayesianPlayerModel(BaseModel): """Hierarchical Bayesian model with partial pooling by player role. Two implementation modes: - PyMC (if installed): Full MCMC hierarchical model - Scipy fallback: Empirical Bayes with James-Stein shrinkage """ def __init__( self, model_dir: str = "models_trained", use_pymc: Optional[bool] = None, samples: int = 2000, tune: int = 1000, chains: int = 2, random_seed: int = 42, ): super().__init__(model_dir) self.use_pymc = use_pymc if use_pymc is not None else _has_pymc() self.samples = samples self.tune = tune self.chains = chains self.random_seed = random_seed self.trace = None self.player_indices = {} self.role_encoder = {} self.role_reverse = {} self.player_means = {} self.player_vars = {} self.role_means = {} self.role_vars = {} self.fitted = False self._pymc_mode = False def _validate_roles(self, X: pd.DataFrame): if "role" not in X.columns: raise ValueError("X must contain a 'role' column with values: P, D, C, A") unknown = set(X["role"].unique()) - VALID_ROLES if unknown: raise ValueError(f"Unknown role values: {unknown}. Allowed: {VALID_ROLES}") def _prepare_data(self, X: pd.DataFrame, y: pd.Series): self._validate_roles(X) roles = X["role"].values unique_roles = sorted(VALID_ROLES) self.role_encoder = {r: i for i, r in enumerate(unique_roles)} self.role_reverse = {i: r for r, i in self.role_encoder.items()} role_idx = np.array([self.role_encoder[r] for r in roles]) n_players = len(y) n_roles = len(unique_roles) player_map = {} player_id = np.zeros(n_players, dtype=int) for i in range(n_players): key = (roles[i], i) if key not in player_map: player_map[key] = len(player_map) player_id[i] = player_map[key] n_unique = len(player_map) return role_idx, player_id, n_players, n_roles, n_unique def _fit_pymc(self, X: pd.DataFrame, y: pd.Series): import pymc as pm role_idx, player_id, n_players, n_roles, n_unique = self._prepare_data(X, y) with pm.Model() as model: mu_role = pm.Normal("mu_role", mu=6.0, sigma=2.0, shape=n_roles) sigma_role = pm.HalfNormal("sigma_role", sigma=1.0, shape=n_roles) sigma_obs = pm.HalfNormal("sigma_obs", sigma=1.0) player_skill = pm.Normal( "player_skill", mu=mu_role[role_idx], sigma=sigma_role[role_idx], shape=n_players, ) pm.Normal( "observed_fv", mu=player_skill, sigma=sigma_obs, observed=y.values, ) self.trace = pm.sample( draws=self.samples, tune=self.tune, chains=self.chains, random_seed=self.random_seed, progressbar=False, ) logger.info( f"PyMC model fitted: {n_players} players, {n_roles} roles, " f"{len(self.trace.posterior.draw) * len(self.trace.posterior.chain)} posterior samples" ) self._pymc_mode = True def _fit_scipy(self, X: pd.DataFrame, y: pd.Series): role_idx, player_id, n_players, n_roles, n_unique = self._prepare_data(X, y) roles = X["role"].values y_vals = y.values.astype(np.float64) self.role_means = {} self.role_vars = {} for r_idx, r_name in self.role_reverse.items(): mask = role_idx == r_idx if mask.sum() > 0: self.role_means[r_name] = float(np.mean(y_vals[mask])) role_var = float(np.var(y_vals[mask], ddof=1)) if mask.sum() > 1 else 0.0 self.role_vars[r_name] = role_var else: self.role_means[r_name] = 6.0 self.role_vars[r_name] = 2.0 player_data = {} for i in range(n_players): role = roles[i] val = y_vals[i] if role not in player_data: player_data[role] = {} player_data[role][i] = val self.player_means = {} self.player_vars = {} for role, players_by_idx in player_data.items(): vals = list(players_by_idx.values()) role_mean = self.role_means[role] role_var = max(self.role_vars[role], 1e-8) for idx in players_by_idx: self.player_means[idx] = vals[0] self.player_vars[idx] = role_var logger.info( f"Scipy fallback fitted: {n_players} players, {n_roles} roles" ) self._pymc_mode = False def fit(self, X: pd.DataFrame, y: pd.Series, **kwargs): if len(X) == 0: raise ValueError("X cannot be empty") if len(X) != len(y): raise ValueError(f"X and y lengths must match: {len(X)} vs {len(y)}") if self.use_pymc and _has_pymc(): self._fit_pymc(X, y) else: if self.use_pymc and not _has_pymc(): logger.warning("PyMC requested but not installed. Falling back to scipy.") self.use_pymc = False self._fit_scipy(X, y) self.fitted = True return self def predict(self, X: pd.DataFrame) -> np.ndarray: if not self.fitted: raise RuntimeError("Model not fitted. Call fit() first.") mean, _ = self.predict_with_uncertainty(X) return mean def predict_with_uncertainty(self, X: pd.DataFrame) -> Tuple[np.ndarray, np.ndarray]: if not self.fitted: raise RuntimeError("Model not fitted. Call fit() first.") self._validate_roles(X) if self._pymc_mode: return self._predict_with_uncertainty_pymc(X) means = np.zeros(len(X)) stds = np.zeros(len(X)) roles = X["role"].values for i, role in enumerate(roles): role_mean = self.role_means.get(role, 6.0) role_var = self.role_vars.get(role, 2.0) player_mean = self.player_means.get(i, role_mean) player_var = self.player_vars.get(i, role_var) shrinkage = role_var / max(role_var + player_var, 1e-8) means[i] = role_mean + (1.0 - shrinkage) * (player_mean - role_mean) stds[i] = np.sqrt(role_var * (1.0 - shrinkage)) return means, stds def _predict_with_uncertainty_pymc(self, X: pd.DataFrame) -> Tuple[np.ndarray, np.ndarray]: import pymc as pm import arviz as az n_players = len(X) roles = X["role"].values with pm.Model() as pred_model: n_roles = len(self.role_encoder) mu_role = pm.Normal("mu_role", mu=6.0, sigma=2.0, shape=n_roles) sigma_role = pm.HalfNormal("sigma_role", sigma=1.0, shape=n_roles) sigma_obs = pm.HalfNormal("sigma_obs", sigma=1.0) player_skill = pm.Normal( "player_skill", mu=mu_role[[self.role_encoder.get(r, 0) for r in roles]], sigma=sigma_role[[self.role_encoder.get(r, 0) for r in roles]], shape=n_players, ) pm.Normal("observed_fv", mu=player_skill, sigma=sigma_obs, shape=n_players) ppc = pm.sample_posterior_predictive( self.trace, var_names=["observed_fv"], random_seed=self.random_seed, progressbar=False, ) observed_samples = ppc.posterior_predictive["observed_fv"].values draws_per_chain = observed_samples.shape[0] n_chains = observed_samples.shape[1] observed_flat = observed_samples.reshape(draws_per_chain * n_chains, n_players) means = observed_flat.mean(axis=0) stds = observed_flat.std(axis=0) return means, stds def posterior_predictive(self, X: pd.DataFrame, n_samples: int = 2000) -> np.ndarray: if not self.fitted: raise RuntimeError("Model not fitted. Call fit() first.") self._validate_roles(X) if self._pymc_mode: return self._posterior_predictive_pymc(X, n_samples) n_players = len(X) means, stds = self.predict_with_uncertainty(X) rng = np.random.RandomState(self.random_seed) draws = rng.normal( loc=means[np.newaxis, :], scale=stds[np.newaxis, :] + 1e-6, size=(n_samples, n_players), ) return np.clip(draws, -10, 20) def _posterior_predictive_pymc(self, X: pd.DataFrame, n_samples: int) -> np.ndarray: import pymc as pm n_players = len(X) roles = X["role"].values with pm.Model() as pred_model: n_roles = len(self.role_encoder) mu_role = pm.Normal("mu_role", mu=6.0, sigma=2.0, shape=n_roles) sigma_role = pm.HalfNormal("sigma_role", sigma=1.0, shape=n_roles) sigma_obs = pm.HalfNormal("sigma_obs", sigma=1.0) player_skill = pm.Normal( "player_skill", mu=mu_role[[self.role_encoder.get(r, 0) for r in roles]], sigma=sigma_role[[self.role_encoder.get(r, 0) for r in roles]], shape=n_players, ) pm.Normal("observed_fv", mu=player_skill, sigma=sigma_obs, shape=n_players) ppc = pm.sample_posterior_predictive( self.trace, var_names=["observed_fv"], random_seed=self.random_seed, progressbar=False, ) observed_samples = ppc.posterior_predictive["observed_fv"].values draws_per_chain = observed_samples.shape[0] n_chains = observed_samples.shape[1] observed_flat = observed_samples.reshape(draws_per_chain * n_chains, n_players) total = observed_flat.shape[0] if total > n_samples: rng = np.random.RandomState(self.random_seed) idx = rng.choice(total, size=n_samples, replace=False) return observed_flat[idx] return observed_flat def get_player_reliability(self, X: pd.DataFrame) -> np.ndarray: if not self.fitted: raise RuntimeError("Model not fitted. Call fit() first.") self._validate_roles(X) roles = X["role"].values scores = np.zeros(len(X)) if self._pymc_mode: means, stds = self.predict_with_uncertainty(X) for i, role in enumerate(roles): role_var = stds[i] ** 2 total_var = role_var + 1.0 scores[i] = np.clip(1.0 - (role_var / max(total_var, 1e-8)), 0.0, 1.0) return np.clip(scores, 0.0, 1.0) for i, role in enumerate(roles): role_var = self.role_vars.get(role, 2.0) player_var = self.player_vars.get(i, role_var) total_var = role_var + player_var scores[i] = np.clip(role_var / max(total_var, 1e-8), 0.0, 1.0) return np.clip(scores, 0.0, 1.0) def get_rookie_estimates(self, X: pd.DataFrame, min_observations: int = 5) -> pd.DataFrame: if not self.fitted: raise RuntimeError("Model not fitted. Call fit() first.") self._validate_roles(X) reliability = self.get_player_reliability(X) rookie_mask = reliability < (1.0 / max(min_observations, 1)) means, stds = self.predict_with_uncertainty(X) roles = X["role"].values results = [] for i in range(len(X)): if not rookie_mask[i]: continue role = roles[i] role_mean = self.role_means.get(role, 6.0) results.append({ "index": i, "role": role, "player_estimate": float(means[i]), "role_mean": float(role_mean), "naive_player_mean": self.player_means.get(i, role_mean), "shrunken_estimate": float(means[i]), "shrunken_std": float(stds[i]), "reliability": float(reliability[i]), "shrinkage_factor": float( (self.player_means.get(i, role_mean) - means[i]) / max(abs(self.player_means.get(i, role_mean) - role_mean), 1e-8) ) if abs(self.player_means.get(i, role_mean) - role_mean) > 1e-8 else 1.0, }) if not results: logger.info("No rookie players found (all have sufficient observations)") return pd.DataFrame(columns=[ "index", "role", "player_estimate", "role_mean", "naive_player_mean", "shrunken_estimate", "shrunken_std", "reliability", "shrinkage_factor", ]) df = pd.DataFrame(results) df = df.sort_values("reliability") logger.info(f"Found {len(results)} rookie players (heavy shrinkage toward role mean)") return df