import matplotlib.pyplot as plt
import numpy as np
import arviz as az
plt.style.use("arviz-darkgrid")

def gpd_log_survival(x, sigma, xi):
    return -(1.0 / xi) * np.log1p(xi * x / sigma)

def gpd_log_pdf(x, sigma, xi):
    return -np.log(sigma) - (1.0 / xi + 1.0) * np.log1p(xi * x / sigma)

def extgpd_log_pdf(x, sigma, xi, kappa):
    log_s = gpd_log_survival(x, sigma, xi)
    log_g = np.log1p(-np.exp(log_s))
    return np.log(kappa) + (kappa - 1.0) * log_g + gpd_log_pdf(x, sigma, xi)

def boundary_log_pdf(x, xi, lam):
    alpha = 1.0 / xi
    return (
        np.log(lam)
        - (alpha + 1.0) * np.log(xi)
        - (alpha + 1.0) * np.log(x)
        - lam * (xi * x) ** (-alpha)
    )

x = np.linspace(1e-3, 8.0, 1200)
xi = 1.5
lam = 2.0
alpha = 1.0 / xi
sigmas = [0.7, 0.2, 0.05, 0.01]

plt.plot(x, np.exp(boundary_log_pdf(x, xi, lam)), "k--", lw=2.5, label="limit")
colors = plt.cm.plasma(np.linspace(0.08, 0.82, len(sigmas)))
for sigma, color in zip(sigmas, colors):
    kappa = lam / sigma**alpha
    plt.plot(
        x,
        np.exp(extgpd_log_pdf(x, sigma, xi, kappa)),
        color=color,
        lw=1.9,
        label=rf"$\sigma={sigma:g}$, $\kappa={kappa:.1f}$",
    )
plt.axvline(0.0, color="0.25", lw=1.2, ls=":")
plt.xlabel("x", fontsize=12)
plt.ylabel("density", fontsize=12)
plt.title(r"ExtGPD non-identifiability between $\sigma$ and $\kappa$")
plt.xlim(0.0, 8.0)
plt.ylim(bottom=0.0)
plt.legend(
    fontsize=8,
    frameon=False,
    title=rf"$\mu=0$, $\xi={xi}$, $\kappa\sigma^{{1/\xi}}={lam:g}$",
    title_fontsize=9,
)
plt.show()