"""
Module: MCMC et PyMC
Catégorie : Statistiques bayésiennes
Difficulté : Avancé

Généré depuis la plateforme ML Formation
"""

# Imports
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score, mean_squared_error, r2_score

from scipy import stats

# Pour le module MCMC uniquement (pip install pymc arviz)
try:
    import arviz as az
    import pymc as pm
except ImportError:
    pass

# Charger le dataset
df = pd.read_csv('ab_testing.csv')

# Metropolis-Hastings code à la main
# Type: Code exécutable
print("=" * 70)
print("       METROPOLIS-HASTINGS EN NUMPY PUR (~30 LIGNES)")
print("=" * 70)

print("""
On echantillonne le posterior du taux de conversion de la variante A
(prior Beta(2, 2), donnees k succes sur n). AVANTAGE PEDAGOGIQUE:
le module Inference Bayesienne nous donne la reponse EXACTE,
Beta(2 + k, 2 + n - k). On pourra donc verifier notre MCMC.
""")

# Donnees de la variante A
sub_a = df[df["variant"] == "A"]
n_obs = len(sub_a)
k_obs = int(sub_a["converted"].sum())
a0, b0 = 2, 2
print(f"Donnees: {k_obs} conversions / {n_obs} visiteurs, prior Beta({a0}, {b0})")

# Log-posterior non normalise: log-vraisemblance binomiale + log-prior beta
# (on travaille en log pour eviter les underflows numeriques)
def log_posterior(p):
    if p <= 0 or p >= 1:
        return -np.inf  # hors du domaine: probabilite nulle
    log_prior = (a0 - 1) * np.log(p) + (b0 - 1) * np.log(1 - p)
    log_vraisemblance = k_obs * np.log(p) + (n_obs - k_obs) * np.log(1 - p)
    return log_prior + log_vraisemblance

# L'algorithme de Metropolis-Hastings
rng = np.random.default_rng(42)
n_iterations = 10000
pas = 0.03           # ecart-type de la proposition gaussienne
chaine = np.zeros(n_iterations)
p_courant = 0.5      # point de depart volontairement mauvais
logpost_courant = log_posterior(p_courant)
n_acceptes = 0

for t in range(n_iterations):
    # 1. Proposer un candidat autour de la position courante
    p_candidat = p_courant + rng.normal(0, pas)
    logpost_candidat = log_posterior(p_candidat)
    # 2. Ratio d'acceptation (en log: une soustraction)
    log_r = logpost_candidat - logpost_courant
    # 3. Accepter ou rejeter
    if np.log(rng.random()) < log_r:
        p_courant = p_candidat
        logpost_courant = logpost_candidat
        n_acceptes += 1
    chaine[t] = p_courant

taux_acceptation = n_acceptes / n_iterations
print(f"Taux d'acceptation: {taux_acceptation:.1%}")

# Verification contre le posterior exact
posterior_exact = stats.beta(a0 + k_obs, b0 + n_obs - k_obs)
burn_in = 1000
echantillons = chaine[burn_in:]

print(f"\n{'':>24} {'MCMC':>10} {'Exact':>10}")
print("-" * 46)
print(f"{'Moyenne du posterior':>24} {echantillons.mean():>10.4f} {posterior_exact.mean():>10.4f}")
print(f"{'Ecart-type':>24} {echantillons.std():>10.4f} {posterior_exact.std():>10.4f}")
q_mcmc = np.percentile(echantillons, [2.5, 97.5])
q_exact = [posterior_exact.ppf(0.025), posterior_exact.ppf(0.975)]
print(f"{'Quantile 2.5%':>24} {q_mcmc[0]:>10.4f} {q_exact[0]:>10.4f}")
print(f"{'Quantile 97.5%':>24} {q_mcmc[1]:>10.4f} {q_exact[1]:>10.4f}")

# Trace de la chaine + histogramme vs posterior exact
fig, axes = plt.subplots(1, 2, figsize=(12, 4.5))
axes[0].plot(chaine[:2000], color="#9B7AC4", lw=0.7)
axes[0].axhline(posterior_exact.mean(), color="#27ae60", linestyle="--",
                label="Moyenne exacte")
axes[0].set_title("Trace de la chaine (2000 premiers pas)")
axes[0].set_xlabel("Iteration")
axes[0].set_ylabel("p")
axes[0].legend()

p_grid = np.linspace(0.03, 0.16, 300)
axes[1].hist(echantillons, bins=50, density=True, color="#C09CF0",
             edgecolor="white", alpha=0.85, label="Echantillons MCMC")
axes[1].plot(p_grid, posterior_exact.pdf(p_grid), color="#27ae60", lw=2.5,
             label="Posterior exact Beta")
axes[1].set_title("L'histogramme MCMC retrouve le posterior exact")
axes[1].set_xlabel("p")
axes[1].set_ylabel("Densite")
axes[1].legend()
plt.tight_layout()
plt.show()

print("""
La chaine part de 0.5 (tres improbable), degringole vers la zone du
posterior en quelques dizaines de pas, puis oscille dedans pour
toujours. L'histogramme des positions visitees = le posterior. Magie.
""")


# ----------------------------------------------------------------------
# La suite de ce module demande un compte.
# Les cellules de code restantes ne sont pas dans ce fichier.
# ----------------------------------------------------------------------
