"""
Module: Inference Bayesienne
Categorie: Bayesian Statistics
Difficulte: Intermediaire

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')

# Explorer les donnees
# Type: Code executable
print("=" * 70)
print("       EXPLORATION DES DONNEES DU TEST A/B")
print("=" * 70)

print("""
Une equipe web teste deux variantes d'une page (A = actuelle, B = nouvelle).
Chaque ligne est un visiteur: sa variante et s'il a converti (1) ou non (0).
Objectif du module: estimer les taux de conversion AVEC leur incertitude,
puis decider si B est vraiment meilleure que A.
""")

display(df.head(10), title="Apercu du dataset A/B")

print("\n" + "=" * 70)
print("1. COMPTAGES PAR VARIANTE")
print("=" * 70)

resume = df.groupby("variant")["converted"].agg(["count", "sum", "mean"])
resume.columns = ["visiteurs", "conversions", "taux_observe"]
display(resume.round(4), title="Resume par variante")

for variant in ["A", "B"]:
    sub = df[df["variant"] == variant]
    n = len(sub)
    k = int(sub["converted"].sum())
    print(f"  Variante {variant}: {k:3d} conversions sur {n} visiteurs "
          f"→ taux observe = {k / n:.2%}")

print("\n" + "=" * 70)
print("2. LA QUESTION QUE TOUT LE MONDE SE POSE")
print("=" * 70)
print("""
B affiche un meilleur taux que A. Mais avec quelques centaines de
visiteurs seulement, est-ce un vrai effet ou de la chance ?

L'approche bayesienne va repondre directement a la vraie question:
"Quelle est la probabilite que B soit meilleure que A ?"
""")

# Visualisation simple des comptages
fig, ax = plt.subplots(figsize=(8, 5))
resume["taux_observe"].plot(kind="bar", ax=ax, color=["#9B7AC4", "#F7E64D"],
                            edgecolor="#3A3A3A")
ax.set_title("Taux de conversion observes (estimations ponctuelles)")
ax.set_ylabel("Taux de conversion")
ax.set_xlabel("Variante")
ax.tick_params(axis="x", rotation=0)
for i, v in enumerate(resume["taux_observe"]):
    ax.text(i, v + 0.003, f"{v:.2%}", ha="center", fontweight="bold")
plt.tight_layout()
plt.show()

print("Ces barres cachent l'essentiel: l'INCERTITUDE autour de chaque taux.")
print("La suite du module va la rendre visible.")


# Le posterior beta-binomiale avec scipy.stats
# Type: Code executable
print("=" * 70)
print("       POSTERIOR BETA-BINOMIALE (CALCUL EXACT)")
print("=" * 70)

print("""
Recette de la conjugaison:
  Prior      p ~ Beta(a, b)
  Donnees    k succes sur n essais
  Posterior  p ~ Beta(a + k, b + n - k)

On applique a la variante A avec un prior peu informatif Beta(2, 2).
""")

# Donnees de la variante A
sub_a = df[df["variant"] == "A"]
n_a = len(sub_a)
k_a = int(sub_a["converted"].sum())

# Prior
a0, b0 = 2, 2
prior = stats.beta(a0, b0)

# Posterior par conjugaison
a_post = a0 + k_a
b_post = b0 + (n_a - k_a)
posterior = stats.beta(a_post, b_post)

print(f"Donnees variante A : {k_a} conversions / {n_a} visiteurs")
print(f"Prior              : Beta({a0}, {b0})          moyenne = {prior.mean():.3f}")
print(f"Posterior          : Beta({a_post}, {b_post})   moyenne = {posterior.mean():.4f}")
print(f"Frequence observee : {k_a / n_a:.4f}")

# Tracer prior et posterior
p_grid = np.linspace(0, 0.5, 500)
fig, ax = plt.subplots(figsize=(9, 5))
ax.plot(p_grid, prior.pdf(p_grid), color="#9B7AC4", lw=2,
        linestyle="--", label=f"Prior Beta({a0}, {b0})")
ax.plot(p_grid, posterior.pdf(p_grid), color="#27ae60", lw=2.5,
        label=f"Posterior Beta({a_post}, {b_post})")
ax.axvline(k_a / n_a, color="#e67e22", linestyle=":",
           label=f"Frequence observee ({k_a / n_a:.3f})")
ax.fill_between(p_grid, posterior.pdf(p_grid), alpha=0.15, color="#27ae60")
ax.set_xlabel("Taux de conversion p")
ax.set_ylabel("Densite de probabilite")
ax.set_title("Variante A: du prior au posterior")
ax.legend()
plt.tight_layout()
plt.show()

print("\n" + "-" * 40)
print("INTERPRETATION:")
print("-" * 40)
print(f"""
• Le prior (violet, plat) disait: "p est quelque part entre 0 et 1,
  plutot vers le milieu". Une croyance tres vague.
• Les {n_a} visiteurs ont apporte de l'information: le posterior (vert)
  est une distribution ETROITE centree pres de {posterior.mean():.3f}.
• Toute la suite (intervalles, decision A/B) se lit sur cette courbe.
""")


# Mise a jour sequentielle: apprendre au fil de l'eau
# Type: Code executable
print("=" * 70)
print("       MISE A JOUR SEQUENTIELLE DU POSTERIOR")
print("=" * 70)

print("""
Propriete magique de l'inference bayesienne: traiter les donnees
d'un coup ou par petits lots donne EXACTEMENT le meme posterior.
Le posterior d'une etape devient le prior de la suivante.

On simule l'arrivee des visiteurs de la variante B par lots de 60.
""")

sub_b = df[df["variant"] == "B"].reset_index(drop=True)
batch_size = 60

a_cur, b_cur = 2, 2  # prior initial Beta(2, 2)
p_grid = np.linspace(0, 0.4, 500)

fig, ax = plt.subplots(figsize=(10, 6))
ax.plot(p_grid, stats.beta(a_cur, b_cur).pdf(p_grid),
        color="#9B7AC4", lw=1.5, linestyle="--", label="Prior Beta(2, 2)")

colors = ["#E5D7F5", "#C09CF0", "#9B7AC4", "#6B4E9B"]
n_batches = int(np.ceil(len(sub_b) / batch_size))

print(f"{'Lot':>4} {'Visiteurs':>10} {'Conversions':>12} {'Posterior':>16} {'Moyenne':>9}")
print("-" * 56)

for i in range(n_batches):
    batch = sub_b.iloc[i * batch_size:(i + 1) * batch_size]
    k = int(batch["converted"].sum())
    n = len(batch)
    # Le posterior du lot precedent sert de prior a ce lot
    a_cur += k
    b_cur += n - k
    color = colors[min(i, len(colors) - 1)]
    ax.plot(p_grid, stats.beta(a_cur, b_cur).pdf(p_grid), color=color, lw=2,
            label=f"Apres lot {i + 1} ({(i + 1) * batch_size if (i + 1) * batch_size < len(sub_b) else len(sub_b)} visiteurs)")
    moyenne = a_cur / (a_cur + b_cur)
    print(f"{i + 1:>4} {n:>10} {k:>12} {'Beta(' + str(a_cur) + ', ' + str(b_cur) + ')':>16} {moyenne:>9.4f}")

ax.set_xlabel("Taux de conversion p (variante B)")
ax.set_ylabel("Densite")
ax.set_title("Le posterior s'affine a mesure que les visiteurs arrivent")
ax.legend()
plt.tight_layout()
plt.show()

print("\n" + "-" * 40)
print("CE QU'IL FAUT RETENIR:")
print("-" * 40)
print("""
• Chaque lot RETRECIT la distribution: plus de donnees = moins d'incertitude.
• L'ordre des lots ne change rien au resultat final: additionner les succes
  lot par lot ou d'un coup revient au meme.
• C'est un mode d'apprentissage naturellement INCREMENTAL, precieux pour
  les systemes en production (le modele n'oublie rien, il se met a jour).
""")


# Intervalle de credibilite vs intervalle de confiance
# Type: Code executable
print("=" * 70)
print("       CREDIBILITE (BAYESIEN) vs CONFIANCE (FREQUENTISTE)")
print("=" * 70)

print("""
Deux objets qui se ressemblent mais ne disent PAS la meme chose.
On les calcule tous les deux pour la variante B.
""")

sub_b = df[df["variant"] == "B"]
n = len(sub_b)
k = int(sub_b["converted"].sum())
p_hat = k / n

# 1) Intervalle de credibilite bayesien a 95% (quantiles du posterior)
posterior = stats.beta(2 + k, 2 + n - k)
cred_low, cred_high = posterior.ppf(0.025), posterior.ppf(0.975)

# 2) Intervalle de confiance frequentiste a 95% (approximation normale de Wald)
se = np.sqrt(p_hat * (1 - p_hat) / n)
conf_low, conf_high = p_hat - 1.96 * se, p_hat + 1.96 * se

print(f"Donnees variante B: {k} conversions / {n} visiteurs (taux {p_hat:.2%})\n")
print(f"Intervalle de CREDIBILITE 95% : [{cred_low:.4f}, {cred_high:.4f}]")
print(f"Intervalle de CONFIANCE   95% : [{conf_low:.4f}, {conf_high:.4f}]")

# Visualisation
p_grid = np.linspace(0.05, 0.25, 500)
fig, ax = plt.subplots(figsize=(10, 5.5))
ax.plot(p_grid, posterior.pdf(p_grid), color="#27ae60", lw=2.5, label="Posterior de p")
mask = (p_grid >= cred_low) & (p_grid <= cred_high)
ax.fill_between(p_grid[mask], posterior.pdf(p_grid)[mask], alpha=0.25,
                color="#27ae60", label="Credibilite 95% (aire verte)")
ax.hlines(-1.0, conf_low, conf_high, color="#e67e22", lw=4,
          label="Confiance 95% (segment orange)")
ax.axvline(p_hat, color="#3498db", linestyle=":", label=f"Taux observe {p_hat:.3f}")
ax.set_ylim(bottom=-2.5)
ax.set_xlabel("Taux de conversion p (variante B)")
ax.set_ylabel("Densite du posterior")
ax.set_title("Deux intervalles proches en valeurs, opposes en interpretation")
ax.legend()
plt.tight_layout()
plt.show()

print("\n" + "-" * 40)
print("INTERPRETATION (la partie importante!):")
print("-" * 40)
print(f"""
• CREDIBILITE: "il y a 95% de chances que p soit entre {cred_low:.3f} et
  {cred_high:.3f}". Enonce direct sur p, car p est une variable aleatoire.

• CONFIANCE: "si on repetait l'experience un grand nombre de fois, 95%
  des intervalles ainsi construits contiendraient le vrai p". L'enonce
  porte sur la PROCEDURE, pas sur cet intervalle-ci.

Les valeurs sont proches ici (beaucoup de donnees, prior faible), mais
seule la version bayesienne autorise la phrase que tout decideur veut
entendre. La section theorique suivante detaille ce point.
""")


# Exercice: quelle est la probabilite que B batte A ?
# Type: Exercice
# Exercice: repondre a LA question du test A/B
# "Quelle est la probabilite que la variante B soit meilleure que A ?"

# L'astuce: on sait echantillonner chaque posterior avec stats.beta.rvs.
# En tirant un grand nombre de couples (pA, pB), la proportion de tirages
# ou pB > pA estime P(pB > pA).

# Preparation des donnees
sub_a = df[df["variant"] == "A"]
sub_b = df[df["variant"] == "B"]
n_a, k_a = len(sub_a), int(sub_a["converted"].sum())
n_b, k_b = len(sub_b), int(sub_b["converted"].sum())

print(f"A: {k_a}/{n_a} conversions | B: {k_b}/{n_b} conversions")

# Les deux posteriors (prior Beta(2, 2) pour chacun)
post_a = stats.beta(2 + k_a, 2 + n_a - k_a)
post_b = stats.beta(2 + k_b, 2 + n_b - k_b)

rng = np.random.default_rng(42)

# TODO: tirez 100000 echantillons de chaque posterior
#       (indice: post_a.rvs(size=..., random_state=rng))
# TODO: calculez P(pB > pA) = proportion de tirages ou pB > pA
# TODO: calculez aussi l'"uplift" moyen espere: la moyenne de (pB - pA)
# TODO: tracez l'histogramme de la difference (pB - pA) et marquez le zero

