import numpy as np
import pandas as pd
from scipy.stats import gaussian_kde
from tqdm import tqdm

# Configuration
data_path = r'[your drive partition]\[df_6]'
output_path = r'[your drive partition]\manager_map_hdi_summary.csv'

print("Loading posterior draws...")
df = pd.read_csv(data_path)

# Convert to wins per 162 games, centered on average manager
df['prob_raw'] = 1 / (1 + np.exp(-df['u']))
mean_prob = df['prob_raw'].mean()  # Empirical mean across all draws
df['w162'] = (df['prob_raw'] - mean_prob) * 162

print(f"Loaded {len(df)} posterior draws for {df['mgrID'].nunique()} managers")
print()

# ========== FUNCTIONS ==========

def calculate_hdi(draws, credible_mass=0.95):
    """
    Calculate the Highest Density Interval (HDI) - the shortest interval
    containing the specified credible mass.
    
    Parameters:
    -----------
    draws : array
        Posterior draws
    credible_mass : float
        Probability mass to include (default 0.95)
    
    Returns:
    --------
    tuple : (lower, upper) bounds of HDI
    """
    sorted_draws = np.sort(draws)
    n = len(sorted_draws)
    interval_size = int(np.ceil(credible_mass * n))
    n_intervals = n - interval_size + 1
    
    # Calculate width of each possible interval
    interval_widths = sorted_draws[interval_size-1:] - sorted_draws[:n_intervals]
    
    # Find the narrowest interval
    min_idx = np.argmin(interval_widths)
    hdi_lower = sorted_draws[min_idx]
    hdi_upper = sorted_draws[min_idx + interval_size - 1]
    
    return hdi_lower, hdi_upper

def calculate_map(draws):
    """
    Calculate the Maximum A Posteriori (MAP) estimate using KDE.
    
    Parameters:
    -----------
    draws : array
        Posterior draws
    
    Returns:
    --------
    float : MAP estimate
    """
    kde = gaussian_kde(draws)
    x_grid = np.linspace(draws.min() - 1, draws.max() + 1, 1000)
    density_vals = kde(x_grid)
    map_estimate = x_grid[np.argmax(density_vals)]
    return map_estimate

# ========== CALCULATE STATISTICS FOR ALL MANAGERS ==========

print("Calculating MAP and HDI for all managers...")
managers = sorted(df['mgrID'].unique())

results = []
for mgr in tqdm(managers, desc="Processing managers"):
    mgr_draws = df[df['mgrID'] == mgr]['w162'].values
    
    # Calculate MAP on w162 scale
    map_w162 = calculate_map(mgr_draws)
    
    # Calculate 95% HDI on w162 scale
    hdi_lower, hdi_upper = calculate_hdi(mgr_draws, credible_mass=0.95)
    
    # Also calculate other useful statistics
    mean_w162 = mgr_draws.mean()
    median_w162 = np.median(mgr_draws)
    sd_w162 = mgr_draws.std()
    
    # Equal-tailed 95% credible interval (for comparison with HDI)
    ci_lower = np.percentile(mgr_draws, 2.5)
    ci_upper = np.percentile(mgr_draws, 97.5)
    
    results.append({
        'mgrID': mgr,
        'MAP_w162': map_w162,
        'HDI_lower_w162': hdi_lower,
        'HDI_upper_w162': hdi_upper,
        'HDI_width': hdi_upper - hdi_lower,
        'mean_w162': mean_w162,
        'median_w162': median_w162,
        'sd_w162': sd_w162,
        'CI_lower_w162': ci_lower,
        'CI_upper_w162': ci_upper,
        'CI_width': ci_upper - ci_lower,
        'n_draws': len(mgr_draws)
    })

# Create DataFrame
results_df = pd.DataFrame(results)

# Sort by MAP (descending)
results_df = results_df.sort_values('MAP_w162', ascending=False)

# ========== SAVE RESULTS ==========
results_df.to_csv(output_path, index=False)
print(f"\n✓ Results saved to: {output_path}")
print(f"  Total managers: {len(results_df)}")
print()

# ========== DISPLAY SUMMARY ==========
print("Summary Statistics:")
print(f"  Mean MAP: {results_df['MAP_w162'].mean():.3f}")
print(f"  Median MAP: {results_df['MAP_w162'].median():.3f}")
print(f"  Mean HDI width: {results_df['HDI_width'].mean():.3f}")
print(f"  Mean CI width: {results_df['CI_width'].mean():.3f}")
print()

print("Top 10 Managers by MAP:")
print(results_df[['mgrID', 'MAP_w162', 'HDI_lower_w162', 'HDI_upper_w162', 'mean_w162']].head(10).to_string(index=False))
print()

print("Bottom 10 Managers by MAP:")
print(results_df[['mgrID', 'MAP_w162', 'HDI_lower_w162', 'HDI_upper_w162', 'mean_w162']].tail(10).to_string(index=False))