"""
Beta-binomial hierarchical model -- the CONJUGATE binomial partial-pooling / shrinkage model, from
scratch. r_i ~ Binomial(n_i, p_i), p_i ~ Beta(alpha, beta). Used for the Efron-Morris baseball example.

This is the binomial analogue of the conjugate gamma-Poisson (Pumps): the per-unit success probability has
a CONJUGATE Beta posterior, so the p_i are exact draws; only the hyperparameters need Metropolis. Contrast
the LOGIT-NORMAL binomial GLMM (seeds_gibbs): there the unit effect is a non-conjugate normal on the logit.
Marginally, r_i ~ Beta-Binomial -- the binomial cousin of the NegBin.

Parameterisation
----------------
  mu = alpha/(alpha+beta)  = the population batting average (shrinkage target)
  kappa = alpha + beta     = concentration ("prior sample size"; larger = stronger shrinkage)
Sample (logit mu, log kappa) by RW-Metropolis on the COLLAPSED beta-binomial marginal
  log P(r_i | n_i, alpha, beta) = betaln(alpha+r_i, beta+n_i-r_i) - betaln(alpha, beta) + const,
then draw the shrunk probabilities p_i | r_i ~ Beta(alpha+r_i, beta+n_i-r_i) (exact, conjugate).

Hyperprior: the standard Gelman BDA prior p(alpha,beta) ∝ (alpha+beta)^(-5/2) (uniform on mu, and on
(alpha+beta)^(-1/2)) -- proper posterior, tames the weakly-identified kappa tail. In (u=logit mu,
lk=log kappa) coordinates this is  log p ∝ -0.5*lk + log mu + log(1-mu).
"""
import numpy as np
from scipy.special import betaln, expit


def betabin_gibbs(r, n, R=20000, burn=5000, seed=0, step=0.3):
    rng = np.random.default_rng(seed); r = np.asarray(r, float); n = np.asarray(n, float); I = len(r)
    p = r / n; mbar = p.mean(); vbar = max(p.var(), 1e-4)
    kappa0 = max(mbar * (1 - mbar) / vbar - 1, 1.0)
    u = np.log(mbar / (1 - mbar)); lk = np.log(kappa0)            # logit(mu), log(kappa)

    def marg(u, lk):
        mu = expit(u); kap = np.exp(lk); a = mu * kap; b = (1 - mu) * kap
        loglik = np.sum(betaln(a + r, b + n - r) - betaln(a, b))
        logprior = -0.5 * lk + np.log(mu) + np.log(1 - mu)       # BDA (alpha+beta)^(-5/2), uniform mu
        return loglik + logprior

    cur = marg(u, lk); keep = R - burn
    MU = np.zeros(keep); KA = np.zeros(keep); P = np.zeros((keep, I)); acc = 0
    for it in range(R):
        up = u + step * rng.standard_normal(); lkp = lk + step * rng.standard_normal()
        pr = marg(up, lkp)
        if np.log(rng.random()) < pr - cur:
            u, lk, cur = up, lkp, pr; acc += 1
        mu = expit(u); kap = np.exp(lk); a = mu * kap; b = (1 - mu) * kap
        if it >= burn:
            k = it - burn; MU[k] = mu; KA[k] = kap
            P[k] = rng.beta(a + r, b + n - r)                    # exact conjugate shrunk probabilities
    return dict(mu=MU, kappa=KA, alpha=MU * KA, beta=(1 - MU) * KA, p=P, accept=acc / R)


def summary(draws, names=None):
    d = draws.reshape(draws.shape[0], -1); q = d.shape[1]
    names = names or [f'p{j}' for j in range(q)]
    return [(nm, d[:, j].mean(), d[:, j].std(), *np.percentile(d[:, j], [2.5, 97.5]))
            for j, nm in enumerate(names)]
