"""
Module: MCMC et PyMC
Categorie: Bayesian Statistics
Difficulte: Avance

Genere 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

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

# Metropolis-Hastings code a la main
# Type: Code executable
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.
""")


# Diagnostics du MH manuel: burn-in, acceptation, autocorrelation
# Type: Code executable
print("=" * 70)
print("       DIAGNOSTIQUER UNE CHAINE MCMC (VERSION MANUELLE)")
print("=" * 70)

print("""
Un MCMC peut echouer SILENCIEUSEMENT: la chaine explore mal et
l'histogramme est faux, sans message d'erreur. Trois diagnostics
essentiels, illustres en comparant trois tailles de pas.
""")

# Donnees et log-posterior (identiques a la cellule precedente)
sub_a = df[df["variant"] == "A"]
n_obs = len(sub_a)
k_obs = int(sub_a["converted"].sum())
a0, b0 = 2, 2

def log_posterior(p):
    if p <= 0 or p >= 1:
        return -np.inf
    return ((a0 - 1) * np.log(p) + (b0 - 1) * np.log(1 - p)
            + k_obs * np.log(p) + (n_obs - k_obs) * np.log(1 - p))

def metropolis(pas, n_iterations=6000, seed=42):
    rng = np.random.default_rng(seed)
    chaine = np.zeros(n_iterations)
    p_courant, logpost_courant, acceptes = 0.5, log_posterior(0.5), 0
    for t in range(n_iterations):
        p_candidat = p_courant + rng.normal(0, pas)
        logpost_candidat = log_posterior(p_candidat)
        if np.log(rng.random()) < logpost_candidat - logpost_courant:
            p_courant, logpost_courant = p_candidat, logpost_candidat
            acceptes += 1
        chaine[t] = p_courant
    return chaine, acceptes / n_iterations

def autocorrelation(x, lag_max=60):
    x = x - x.mean()
    acf = np.correlate(x, x, mode="full")[len(x) - 1:]
    return acf[:lag_max] / acf[0]

configurations = [
    (0.002, "Pas MINUSCULE"),
    (0.03, "Pas BIEN REGLE"),
    (0.8, "Pas ENORME"),
]

fig, axes = plt.subplots(2, 3, figsize=(13, 7))
print(f"\n{'Configuration':>16} {'Acceptation':>12} {'Diagnostic':>34}")
print("-" * 66)

for j, (pas, nom) in enumerate(configurations):
    chaine, taux = metropolis(pas)
    acf = autocorrelation(chaine[1000:])

    if taux > 0.7:
        verdict = "trop timide: exploration escargot"
    elif taux < 0.10:
        verdict = "trop audacieux: rejets en serie"
    else:
        verdict = "bon compromis (~20-50% vise)"
    print(f"{nom:>16} {taux:>11.1%} {verdict:>34}")

    axes[0, j].plot(chaine[:1500], color="#9B7AC4", lw=0.6)
    axes[0, j].set_title(f"{nom} (acc. {taux:.0%})", fontsize=10)
    axes[0, j].set_ylim(0, 0.6)
    axes[1, j].bar(range(len(acf)), acf, color="#C09CF0", width=1.0)
    axes[1, j].set_title("Autocorrelation", fontsize=10)
    axes[1, j].set_ylim(-0.1, 1.05)

axes[0, 0].set_ylabel("Trace de p")
axes[1, 0].set_ylabel("ACF")
plt.tight_layout()
plt.show()

print("\n" + "-" * 40)
print("LES TROIS DIAGNOSTICS:")
print("-" * 40)
print("""
1. BURN-IN: le debut de chaine (depuis 0.5) ne represente pas le
   posterior, on le jette (ici: 1000 premiers pas).
2. TAUX D'ACCEPTATION: ~20-50% pour un MH bien regle. Trop haut =
   pas trop petits (la chaine rampe). Trop bas = pas trop grands
   (la chaine cale sur place).
3. AUTOCORRELATION: mesure combien les echantillons successifs se
   ressemblent. Une ACF qui decroit vite = beaucoup d'information
   par tirage. Une ACF qui traine = il faut des chaines plus longues.

Regler le pas a la main est penible... et on n'a qu'UN parametre.
Imaginez 50. C'est exactement le probleme que PyMC va resoudre.
""")


# Le meme modele en PyMC
# Type: Code executable
import logging
logging.getLogger("pymc").setLevel(logging.ERROR)  # sortie compacte

print("=" * 70)
print("       NOTRE MODELE BETA-BINOMIALE, VERSION PYMC")
print("=" * 70)

print("""
Le modele qui demandait 30 lignes de numpy tient en 3 lignes declaratives.
On garde la variante A pour pouvoir TOUT verifier contre l'exact.
""")

# Donnees de la variante A
sub_a = df[df["variant"] == "A"]
n_obs = len(sub_a)
k_obs = int(sub_a["converted"].sum())
print(f"Donnees: {k_obs} conversions / {n_obs} visiteurs\n")

# Declaration + echantillonnage (reglages sandbox, cf. section precedente)
with pm.Model():
    p = pm.Beta("p", alpha=2, beta=2)
    pm.Binomial("y", n=n_obs, p=p, observed=k_obs)
    idata = pm.sample(draws=500, tune=500, chains=2, cores=1,
                      progressbar=False, random_seed=42)

# Resume ArviZ
resume = az.summary(idata, kind="stats")
display(resume, title="az.summary: statistiques du posterior")

# Verification contre le posterior exact (conjugaison du module 17)
posterior_exact = stats.beta(2 + k_obs, 2 + n_obs - k_obs)
echantillons = idata.posterior["p"].values.ravel()

print(f"{'':>24} {'PyMC/NUTS':>10} {'Exact':>10}")
print("-" * 46)
print(f"{'Moyenne':>24} {echantillons.mean():>10.4f} {posterior_exact.mean():>10.4f}")
print(f"{'Ecart-type':>24} {echantillons.std():>10.4f} {posterior_exact.std():>10.4f}")

# Histogramme vs exact
p_grid = np.linspace(0.03, 0.16, 300)
fig, ax = plt.subplots(figsize=(9, 5))
ax.hist(echantillons, bins=40, density=True, color="#C09CF0",
        edgecolor="white", alpha=0.85, label="Echantillons NUTS (PyMC)")
ax.plot(p_grid, posterior_exact.pdf(p_grid), color="#27ae60", lw=2.5,
        label="Posterior exact Beta")
ax.set_xlabel("Taux de conversion p (variante A)")
ax.set_ylabel("Densite")
ax.set_title("PyMC retrouve le posterior exact (1000 tirages)")
ax.legend()
plt.tight_layout()
plt.show()

print("""
Meme resultat que notre MH artisanal et que la formule exacte, mais:
• zero reglage manuel (pas de taille de pas a choisir),
• une syntaxe qui LIT comme le modele mathematique,
• et ca marcherait pareil avec 50 parametres sans conjugaison.
""")


# Diagnostics ArviZ: r_hat, ESS, trace
# Type: Code executable
import logging
logging.getLogger("pymc").setLevel(logging.ERROR)

print("=" * 70)
print("       LIRE LES DIAGNOSTICS COMME UN PRO (ARVIZ)")
print("=" * 70)

print("""
Les diagnostics manuels de la cellule MH (burn-in, acceptation, ACF)
ont des equivalents industrialises dans ArviZ. Les deux a connaitre:

• r_hat  : compare les chaines entre elles. Si elles racontent la meme
           histoire, r_hat ~ 1.00. Au-dela de 1.01: suspect. 1.05: poubelle.
• ess    : Effective Sample Size, le nombre d'echantillons INDEPENDANTS
           equivalents (l'autocorrelation en deduit). 1000 tirages tres
           correles peuvent ne valoir que 50 tirages utiles.
""")

# Modele identique a la cellule precedente
sub_a = df[df["variant"] == "A"]
n_obs = len(sub_a)
k_obs = int(sub_a["converted"].sum())

with pm.Model():
    p = pm.Beta("p", alpha=2, beta=2)
    pm.Binomial("y", n=n_obs, p=p, observed=k_obs)
    idata = pm.sample(draws=500, tune=500, chains=2, cores=1,
                      progressbar=False, random_seed=42)

# Diagnostics chiffres
resume = az.summary(idata)  # inclut r_hat et ess
display(resume, title="Diagnostics complets (r_hat, ess_bulk, ess_tail)")

r_hat = float(resume["r_hat"].iloc[0])
ess_bulk = float(resume["ess_bulk"].iloc[0])
print(f"r_hat    = {r_hat:.3f}  (cible: <= 1.01)  "
      f"{'OK' if r_hat <= 1.01 else 'PROBLEME'}")
print(f"ess_bulk = {ess_bulk:.0f}  sur {2 * 500} tirages  "
      f"{'OK' if ess_bulk > 400 else 'FAIBLE'}")

# Trace plot: LE reflexe visuel apres chaque echantillonnage
az.plot_trace(idata, figsize=(11, 3.5))
plt.suptitle("az.plot_trace: densite par chaine (gauche), trace (droite)", y=1.04)
plt.tight_layout()
plt.show()

# Posterior plot avec intervalle de credibilite
az.plot_posterior(idata, hdi_prob=0.95, figsize=(7, 4), color="#9B7AC4")
plt.title("az.plot_posterior: moyenne et HDI 95%")
plt.tight_layout()
plt.show()

print("""
COMMENT LIRE LE TRACE PLOT:
• A droite: les 2 chaines doivent ressembler a du "bruit stable"
  (une chenille bien grasse), sans derive ni plateau.
• A gauche: les densites des 2 chaines doivent se superposer.
• Le HDI (Highest Density Interval) est le cousin de l'intervalle
  de credibilite du module 17, calcule sur les echantillons.

Ici tout est vert: NUTS sur un modele conjugue, c'est l'autoroute.
""")


# Le test A/B complet en PyMC
# Type: Code executable
import logging
logging.getLogger("pymc").setLevel(logging.ERROR)

print("=" * 70)
print("       TEST A/B BAYESIEN DE BOUT EN BOUT")
print("=" * 70)

print("""
Le graal du module 17, version PyMC: modeliser les DEUX variantes
ensemble et faire porter l'inference directement sur la difference
delta = pB - pA grace a pm.Deterministic.
""")

# Donnees des deux variantes
resume_donnees = df.groupby("variant")["converted"].agg(["count", "sum"])
n_a, k_a = int(resume_donnees.loc["A", "count"]), int(resume_donnees.loc["A", "sum"])
n_b, k_b = int(resume_donnees.loc["B", "count"]), int(resume_donnees.loc["B", "sum"])
print(f"A: {k_a}/{n_a} conversions | B: {k_b}/{n_b} conversions\n")

with pm.Model():
    # Un prior par variante
    p_a = pm.Beta("p_A", alpha=2, beta=2)
    p_b = pm.Beta("p_B", alpha=2, beta=2)
    # Les quantites qui nous interessent VRAIMENT, suivies par PyMC
    delta = pm.Deterministic("delta", p_b - p_a)
    uplift_relatif = pm.Deterministic("uplift_relatif", (p_b - p_a) / p_a)
    # Vraisemblances
    pm.Binomial("obs_A", n=n_a, p=p_a, observed=k_a)
    pm.Binomial("obs_B", n=n_b, p=p_b, observed=k_b)
    idata = pm.sample(draws=500, tune=500, chains=2, cores=1,
                      progressbar=False, random_seed=42)

display(az.summary(idata, kind="stats"), title="Posteriors: p_A, p_B, delta, uplift")

# Les reponses business, par simple comptage des echantillons
delta_samples = idata.posterior["delta"].values.ravel()
uplift_samples = idata.posterior["uplift_relatif"].values.ravel()
prob_b_meilleure = float(np.mean(delta_samples > 0))
prob_gain_1pt = float(np.mean(delta_samples > 0.01))

print(f"P(pB > pA)                      = {prob_b_meilleure:.1%}")
print(f"P(gain > 1 point de conversion) = {prob_gain_1pt:.1%}")
print(f"Uplift relatif moyen            = {float(np.mean(uplift_samples)):+.1%}")

# Forest plot des deux taux + histogramme de delta
fig, axes = plt.subplots(1, 2, figsize=(12, 4.5))
az.plot_forest(idata, var_names=["p_A", "p_B"], combined=True,
               hdi_prob=0.95, ax=axes[0], colors="#9B7AC4")
axes[0].set_title("Taux de conversion (HDI 95%)")
axes[1].hist(delta_samples, bins=60, color="#C09CF0", edgecolor="white", alpha=0.9)
axes[1].axvline(0, color="#e74c3c", lw=2, linestyle="--", label="delta = 0")
axes[1].axvline(float(np.mean(delta_samples)), color="#27ae60", lw=2,
                label=f"delta moyen {float(np.mean(delta_samples)):+.4f}")
axes[1].set_title(f"Posterior de delta: P(delta > 0) = {prob_b_meilleure:.1%}")
axes[1].set_xlabel("delta = pB - pA")
axes[1].legend()
plt.tight_layout()
plt.show()

print("""
LECTURE DECISIONNELLE:
• La reponse n'est pas "significatif oui/non" mais une probabilite,
  directement combinable avec les couts metier: deployer B si
  P(delta > 0) depasse votre seuil de confort (90% ? 95% ?).
• pm.Deterministic a fait porter TOUTE l'inference (moyenne, HDI,
  probabilites) sur delta et l'uplift sans aucun calcul manuel.
• Comparez avec l'exercice du module 17: memes conclusions, mais ici
  la recette s'etend telle quelle a 10 variantes ou a un modele
  hierarchique par pays. C'est ca, la programmation probabiliste.
""")


# Exercice: la sensibilite au prior, enfin mesurable
# Type: Exercice
import logging
logging.getLogger("pymc").setLevel(logging.ERROR)

# Exercice: mesurer l'impact du prior sur la conclusion du test A/B
#
# Le module 17 promettait qu'on pouvait "auditer la sensibilite au prior".
# C'est l'heure. Comparez deux analyses de la variante B:
#   1. Prior NEUTRE     : Beta(1, 1)    (uniforme, "je ne sais rien")
#   2. Prior PESSIMISTE : Beta(20, 180) (conviction forte: taux ~10%,
#                         poids equivalent a 200 observations virtuelles)

# Donnees de la variante B
sub_b = df[df["variant"] == "B"]
n_b = len(sub_b)
k_b = int(sub_b["converted"].sum())
print(f"Donnees B: {k_b} conversions / {n_b} visiteurs "
      f"(taux observe {k_b / n_b:.2%})")

# TODO: pour chacun des deux priors, construisez le modele PyMC
#       (pm.Beta + pm.Binomial) et echantillonnez avec les reglages
#       du sandbox: draws=500, tune=500, chains=2, cores=1,
#       progressbar=False, random_seed=42
# TODO: comparez les moyennes des deux posteriors de p
# TODO: tracez les deux histogrammes sur la meme figure
# TODO: question: avec 240 visiteurs, le prior pessimiste (qui "pese"
#       200 visiteurs virtuels) change-t-il la conclusion ? Que se
#       passerait-il avec 10 fois plus de donnees ?

