import numpy as np
import matplotlib.pyplot as plt
from scipy.optimize import minimize


# ============================================================
# Stable sigmoid
# ============================================================

def _sigmoid_stable(z):
    z = np.asarray(z, float)
    z = np.clip(z, -50.0, 50.0)
    return 1.0 / (1.0 + np.exp(-z))


# ============================================================
# 1) Model
# ============================================================

def model_p(x, b):
    """p(x|b) = sigmoid(b0 + b1*x)."""
    x = np.asarray(x, float).reshape(-1)
    b0, b1 = np.asarray(b, float).reshape(2)
    return _sigmoid_stable(b0 + b1 * x)


def design_matrix(x):
    """Design matrix X = [1, x]."""
    x = np.asarray(x, float).reshape(-1)
    return np.column_stack([np.ones_like(x), x])


# ============================================================
# 2) Likelihood
# ============================================================

def nll(x, y, b, l2=0.0):
    """
    Penalized negative log-likelihood:
        NLL(b) = -sum[y log p + (1-y) log(1-p)] + 0.5*l2*||b||^2
    """
    x = np.asarray(x, float).reshape(-1)
    y = np.asarray(y, float).reshape(-1)
    b = np.asarray(b, float).reshape(2)

    p = model_p(x, b)
    eps = 1e-12
    p = np.clip(p, eps, 1 - eps)

    base = -np.sum(y * np.log(p) + (1 - y) * np.log(1 - p))
    pen = 0.5 * l2 * float(np.dot(b, b))
    return base + pen


def llf(x, y, b):
    """
    Ordinary (unpenalized) log-likelihood at fitted parameters.
    """
    x = np.asarray(x, float).reshape(-1)
    y = np.asarray(y, float).reshape(-1)
    b = np.asarray(b, float).reshape(2)

    p = model_p(x, b)
    eps = 1e-12
    p = np.clip(p, eps, 1 - eps)
    return float(np.sum(y * np.log(p) + (1 - y) * np.log(1 - p)))


# ============================================================
# 3) Gradient / Hessian / Covariance
# ============================================================

def grad_nll(x, y, b, l2=0.0):
    """
    Gradient of penalized NLL:
        g(b) = X^T (p - y) + l2*b
    """
    X = design_matrix(x)
    y = np.asarray(y, float).reshape(-1)
    b = np.asarray(b, float).reshape(2)

    p = model_p(x, b)
    return X.T @ (p - y) + l2 * b


def hess_nll(x, b, l2=0.0):
    """
    Hessian of penalized NLL:
        H(b) = X^T W X + l2*I
        W = diag(p*(1-p))
    """
    X = design_matrix(x)
    b = np.asarray(b, float).reshape(2)

    p = model_p(x, b)
    w = p * (1 - p)
    return X.T @ (w[:, None] * X) + l2 * np.eye(2)


def covariance(x, b, l2=0.0):
    """
    Cov(b) ≈ H(b)^(-1), where H is the penalized Hessian if l2 > 0.
    Robust to near-singular Hessians.
    """
    H = hess_nll(x, b, l2=l2)
    try:
        return np.linalg.inv(H)
    except np.linalg.LinAlgError:
        return np.linalg.pinv(H)


def standard_errors(x, b, l2=0.0):
    """
    SE = sqrt(diag(Cov)).
    """
    C = covariance(x, b, l2=l2)
    return np.sqrt(np.maximum(np.diag(C), 0.0))


# Compatibility alias
def logit_poly_cov(x, b, l2=0.0):
    return covariance(x, b, l2=l2)


# ============================================================
# 4) Fit
# ============================================================
def fit_newton(x, y, b_start=None, max_iter=50, tol=1e-8, l2=0.0):
    """
    Newton updates for penalized NLL with backtracking line-search.

    Update:
        b_new = b - alpha * H^{-1} g
    alpha shrinks until NLL decreases.
    """
    x = np.asarray(x, float).reshape(-1)
    y = np.asarray(y, int).reshape(-1)

    if b_start is None:
        b = np.array([0.0, 0.0], float)
    else:
        b = np.asarray(b_start, float).reshape(2)

    f = nll(x, y, b, l2=l2)

    for _ in range(max_iter):
        g = grad_nll(x, y, b, l2=l2)
        H = hess_nll(x, b, l2=l2)

        try:
            step = np.linalg.solve(H, g)
        except np.linalg.LinAlgError:
            step = np.linalg.pinv(H) @ g

        alpha = 1.0
        while alpha > 1e-6:
            b_new = b - alpha * step
            f_new = nll(x, y, b_new, l2=l2)
            if np.isfinite(f_new) and f_new <= f:
                break
            alpha *= 0.5

        if alpha <= 1e-6:
            break

        if np.max(np.abs(b_new - b)) < tol:
            b = b_new
            break

        b, f = b_new, f_new

    return b


# ============================================================
# 14) Overlay plot (LOG left, RAW right)
# ============================================================
def plot_overlay_two_panels_final(
    r_log_full, r_log_trim, r_raw_full, r_raw_trim,
    dy_full=-0.010, dy_trim=0.010
):
    import numpy as np
    import matplotlib.pyplot as plt

    COL_NC   = "#4C78A8"
    COL_AE   = "#F58518"
    COL_FULL = "#1f77b4"
    COL_TRIM = "#ff7f0e"

    fig, axes = plt.subplots(1, 2, figsize=(9.5, 3.5), sharey=True)
    ax1, ax2 = axes

    def draw_panel(ax, r_full, r_trim, xlabel, panel_label,
                   show_legend=False):

        xF = np.asarray(r_full["x"], float)
        yF = np.asarray(r_full["y"], int)

        xT = np.asarray(r_trim["x"], float)
        yT = np.asarray(r_trim["y"], int)

        xx = np.linspace(
            min(xF.min(), xT.min()),
            max(xF.max(), xT.max()),
            500
        )

        # keep x-values unchanged
        xF_plot = xF
        xT_plot = xT

        # vertical offsets only
        yF_plot = yF + np.where(yF == 0, dy_full, -dy_full)
        yT_plot = yT + np.where(yT == 0, dy_trim, -dy_trim)

        # FULL = filled markers
        ax.scatter(
            xF_plot[yF == 0], yF_plot[yF == 0],
            s=16,
            color=COL_NC,
            alpha=0.70,
            edgecolors="none",
            label="data: FULL NC",
            zorder=3
        )

        ax.scatter(
            xF_plot[yF == 1], yF_plot[yF == 1],
            s=16,
            color=COL_AE,
            alpha=0.80,
            edgecolors="none",
            label="data: FULL AE",
            zorder=3
        )

        # TRIM = outlined markers
        ax.scatter(
            xT_plot[yT == 0], yT_plot[yT == 0],
            s=24,
            facecolors=COL_NC,
            edgecolors="black",
            linewidths=0.45,
            alpha=0.95,
            label="data: TRIM NC",
            zorder=4
        )

        ax.scatter(
            xT_plot[yT == 1], yT_plot[yT == 1],
            s=24,
            facecolors=COL_AE,
            edgecolors="black",
            linewidths=0.45,
            alpha=0.95,
            label="data: TRIM AE",
            zorder=4
        )

        # logistic fits
        ax.plot(
            xx,
            model_p(xx, r_full["b"]),
            lw=1.8,
            color=COL_FULL,
            label="fit FULL",
            zorder=2
        )

        ax.plot(
            xx,
            model_p(xx, r_trim["b"]),
            lw=1.8,
            color=COL_TRIM,
            label="fit TRIM",
            zorder=2
        )

        # legend only in panel B
        if show_legend:
            ax.legend(
                loc="lower right",
                fontsize=7,
                markerscale=0.9,
                frameon=True,
                framealpha=1.0,
                edgecolor="0.7",
                handlelength=1.8,
                borderpad=0.4,
                labelspacing=0.4,
                handletextpad=0.5
            )

        ax.set_xlabel(xlabel)
        ax.set_ylim(-0.08, 1.08)

        ax.text(
            0.05, 0.90,
            panel_label,
            transform=ax.transAxes,
            fontsize=11
        )

    draw_panel(
        ax1,
        r_log_full,
        r_log_trim,
        "log(X)",
        "A",
        show_legend=False
    )

    draw_panel(
        ax2,
        r_raw_full,
        r_raw_trim,
        "X",
        "B",
        show_legend=True
    )

    ax1.set_ylabel("P(AE | X = x)")

    for ax in axes:
        ax.grid(False)
        ax.tick_params(labelsize=8)

    plt.tight_layout()
    plt.show()

    return fig, axes

    return fig, axes
# ============================================================
# 5) Goodness of fit
# ============================================================

def goodness_of_fit(x, y, b, thresh=0.5, l2=0.0):
    """
    Returns:
      LLF, NLL, AIC, BIC, Accuracy, n, k

    Notes
    -----
    Fit may use l2 > 0, but GOF metrics below are computed from the
    ordinary (unpenalized) likelihood evaluated at the fitted parameters.
    The argument l2 is kept only for interface consistency.
    """
    x = np.asarray(x, float).reshape(-1)
    y = np.asarray(y, int).reshape(-1)
    b = np.asarray(b, float).reshape(2)

    p = model_p(x, b)
    eps = 1e-12
    p = np.clip(p, eps, 1 - eps)

    LLF = np.sum(y * np.log(p) + (1 - y) * np.log(1 - p))
    NLL = -LLF

    n = len(x)
    k = len(b)

    AIC = 2 * k - 2 * LLF
    BIC = k * np.log(n) - 2 * LLF

    yhat = (p >= thresh).astype(int)
    acc = np.mean(yhat == y)

    return {
        "LLF": float(LLF),
        "NLL": float(NLL),
        "AIC": float(AIC),
        "BIC": float(BIC),
        "A": float(acc),
        "n": int(n),
        "k": int(k),
    }


# ============================================================
# 6) x50 / Wald helpers / compact fit
# ============================================================
def x50(b):
    """
    Model-scale midpoint:
        x50 = -b0 / b1

    For LOG panels, this is on the log(x) scale.
    Raw-scale SUV50 is exp(x50).
    """
    b0, b1 = np.asarray(b, float).reshape(2)
    return np.nan if np.abs(b1) < 1e-12 else (-b0 / b1)


def check_x50_consistency(P):
    """
    Diagnostic check for x50 consistency.

    For LOG models:
        x50_model is on log(X) scale
        SUV50 is on raw X scale = exp(x50_model)

    For RAW models:
        x50_model = SUV50

    Correct result:
        P(x50_model) should be approximately 0.5
    """
    for key, pk in P.items():
        b = np.asarray(pk["b"], float).reshape(2)
        trans = pk.get("transform", "")

        x50_model = x50(b)
        suv50 = np.exp(x50_model) if trans == "log" else x50_model
        p_at_x50 = model_p(np.array([x50_model]), b)[0]

        print(
            key,
            "| transform =", trans,
            "| x50_model =", x50_model,
            "| SUV50 =", suv50,
            "| P(x50) =", p_at_x50
        )


def x50_wald_ci(b, cov, z=1.959963984540054):
    """
    Wald CI for x50 = -b0/b1 via delta method.
    Returned on MODEL scale.
    """
    b = np.asarray(b, float).reshape(2)
    cov = np.asarray(cov, float).reshape(2, 2)
    b0, b1 = b

    if np.abs(b1) < 1e-12:
        return np.nan, np.nan

    xhat = -b0 / b1
    grad = np.array([-1.0 / b1, b0 / (b1 ** 2)], float)
    var = float(grad.T @ cov @ grad)
    se = np.sqrt(max(var, 0.0))
    return float(xhat - z * se), float(xhat + z * se)


def wald_ci(b, cov, z=1.959963984540054):
    """
    Wald CI for parameters: b_i ± z*SE_i.
    """
    b = np.asarray(b, float).reshape(2)
    cov = np.asarray(cov, float).reshape(2, 2)
    se = np.sqrt(np.maximum(np.diag(cov), 0.0))
    return b - z * se, b + z * se


def fit_pack(x, y, name="", thresh=0.5, l2=0.0, z=1.959963984540054):
    """
    Fit + covariance + GOF + parameter Wald CI.
    x should already be on the MODEL scale.
    """
    x = np.asarray(x, float).reshape(-1)
    y = np.asarray(y, int).reshape(-1)

    b = fit_newton(x, y, l2=l2)
    cov = covariance(x, b, l2=l2)
    gof = goodness_of_fit(x, y, b, thresh=thresh, l2=l2)
    lcl, ucl = wald_ci(b, cov, z=z)

    return {
        "name": name,
        "x": x,
        "y": y,
        "b": b,
        "cov": cov,
        "gof": gof,
        "LCL": lcl,
        "UCL": ucl,
        "l2": float(l2),
    }


def trim_nc_by_value(x_raw, y, target=2.48, tol=0.05):
    """
    Remove ONE NC sample (y==0) with x_raw closest to target.
    """
    x_raw = np.asarray(x_raw, float).reshape(-1)
    y = np.asarray(y, int).reshape(-1)

    nc_idx = np.where(y == 0)[0]
    if len(nc_idx) == 0:
        raise ValueError("No NC samples found (y==0).")

    j = nc_idx[np.argmin(np.abs(x_raw[nc_idx] - target))]
    diff = float(np.abs(x_raw[j] - target))
    if diff > tol:
        print(f"[trim warning] closest NC to {target} is {x_raw[j]:.6f} (diff={diff:.6f}) > tol={tol}")

    mask = np.ones_like(y, dtype=bool)
    mask[j] = False
    print(f"[trim] removed index={j}, x_raw={x_raw[j]:.6f}, y={y[j]}")
    return x_raw[mask], y[mask]

def s50_from_b(b, transform="raw"):
    """
    Local slope s50 = dp/dx at x50 on RAW x scale.
    """
    b = np.asarray(b, float).reshape(2)
    b0, b1 = b

    if np.abs(b1) < 1e-12:
        return np.nan

    x50_model = x50(b)

    if transform == "raw":
        return float(b1 / 4.0)

    elif transform == "log":
        x50_raw = np.exp(x50_model)
        return float(b1 / (4.0 * x50_raw))

    else:
        raise ValueError("transform must be 'raw' or 'log'")


def s50_normal_ci_from_mvnorm(
    b, cov, transform="raw",
    M=200000, seed=123, alpha=0.05,
    enforce_positive_slope=True, slope_eps=1e-10
):
    """
    Normal-on-MLE CI for s50.
    """
    rng = np.random.default_rng(seed)

    b = np.asarray(b, float).reshape(2)
    cov = np.asarray(cov, float).reshape(2, 2)

    vals = []
    tries = 0
    max_tries = 20 * M

    while len(vals) < M and tries < max_tries:
        tries += 1

        bb = rng.multivariate_normal(mean=b, cov=cov)

        if not np.all(np.isfinite(bb)):
            continue

        if enforce_positive_slope and bb[1] <= slope_eps:
            continue

        val = s50_from_b(bb, transform=transform)

        if np.isfinite(val):
            vals.append(val)

    if len(vals) == 0:
        return np.nan, np.nan, np.nan, 0

    vals = np.asarray(vals, float)
    q = np.quantile(vals, [alpha / 2, 0.5, 1.0 - alpha / 2])

    return float(q[0]), float(q[1]), float(q[2]), int(len(vals))


def s50_wald_ci_numeric(
    b, cov, transform="raw",
    z=1.959963984540054,
    eps=1e-5
):
    """
    Delta-method CI for s50 using numerical derivatives.
    """
    b = np.asarray(b, float).reshape(2)
    cov = np.asarray(cov, float).reshape(2, 2)

    s_hat = s50_from_b(b, transform=transform)

    if not np.isfinite(s_hat):
        return np.nan, np.nan

    grad = np.zeros(2, float)

    for j in range(2):
        step = eps * max(1.0, abs(b[j]))

        bp = b.copy()
        bm = b.copy()

        bp[j] += step
        bm[j] -= step

        sp = s50_from_b(bp, transform=transform)
        sm = s50_from_b(bm, transform=transform)

        grad[j] = (sp - sm) / (2.0 * step)

    var = float(grad.T @ cov @ grad)
    se = np.sqrt(max(var, 0.0))

    return float(s_hat - z * se), float(s_hat + z * se)


# ============================================================
# 7) Alternative confidence-interval estimation
# ============================================================
# Final terminology:
#   Wald          : analytical approximation using the fitted covariance;
#   MC            : Monte Carlo propagation from the local Gaussian approximation;
#   Nonparametric : ordinary patient-level nonparametric bootstrap;
#   Stratified    : class-stratified nonparametric bootstrap, retained for comparison;
#   Parametric    : model-based Bernoulli bootstrap.
#
# The delta method is used internally for analytical propagation under
# the Wald approximation; it is not treated as a separate method.

from collections import OrderedDict


CI_METHODS = [
    "Wald",
    "MC",
    "Nonparametric",
    "Stratified",
    "Parametric",
]

MC_DRAWS_BANDS = 20_000
MC_DRAWS_TABLE = 200_000


def eta_se_grid(x_grid, cov):
    """Standard error of eta(x) = b0 + b1*x on a model-scale grid."""
    x_grid = np.asarray(x_grid, float).reshape(-1)
    cov = np.asarray(cov, float).reshape(2, 2)

    Xg = design_matrix(x_grid)
    var_eta = np.einsum("ij,jk,ik->i", Xg, cov, Xg)
    return np.sqrt(np.maximum(var_eta, 0.0))


def ci_band_wald(x_grid, b, cov, z=1.959963984540054):
    """
    Pointwise Wald confidence band for p(x).

    The fitted-parameter covariance is propagated to the probability scale
    using the first-order delta method.
    """
    x_grid = np.asarray(x_grid, float).reshape(-1)
    b = np.asarray(b, float).reshape(2)

    p = model_p(x_grid, b)
    se_eta = eta_se_grid(x_grid, cov)
    se_p = p * (1.0 - p) * se_eta

    lo = np.clip(p - z * se_p, 0.0, 1.0)
    hi = np.clip(p + z * se_p, 0.0, 1.0)
    return lo, p, hi


# Historical alias retained for notebook compatibility.
def ci_band_delta(x_grid, b, cov, z=1.959963984540054):
    return ci_band_wald(x_grid, b, cov, z=z)


def gaussian_parameter_draws(
    b,
    cov,
    M=MC_DRAWS_BANDS,
    seed=123,
    enforce_positive_slope=True,
    slope_eps=1e-10,
    x50_bounds=None,
):
    """
    Draw beta* ~ N(beta_hat, Cov_hat) for Monte Carlo propagation.

    Parameters
    ----------
    x50_bounds : tuple(float, float) or None
        Optional admissible interval for model-scale x50. When supplied,
        draws with x50 outside [lower, upper] are rejected.
    """
    rng = np.random.default_rng(seed)

    b = np.asarray(b, float).reshape(2)
    cov = np.asarray(cov, float).reshape(2, 2)

    draws = []
    attempts = 0
    rejected_nonfinite = 0
    rejected_slope = 0
    rejected_x50 = 0

    max_attempts = max(50 * int(M), 1000)

    while len(draws) < int(M) and attempts < max_attempts:
        attempts += 1

        try:
            bb = rng.multivariate_normal(mean=b, cov=cov)
        except Exception:
            break

        if not np.all(np.isfinite(bb)):
            rejected_nonfinite += 1
            continue

        if enforce_positive_slope and bb[1] <= slope_eps:
            rejected_slope += 1
            continue

        if x50_bounds is not None:
            x50_draw = x50(bb)

            if not np.isfinite(x50_draw):
                rejected_x50 += 1
                continue

            x50_lower, x50_upper = x50_bounds

            if not (x50_lower <= x50_draw <= x50_upper):
                rejected_x50 += 1
                continue

        draws.append(bb)

    arr = np.asarray(draws, float) if draws else np.empty((0, 2), float)

    diagnostics = {
        "attempted": int(attempts),
        "successful": int(len(arr)),
        "rejected": int(attempts - len(arr)),
        "rejected_nonfinite": int(rejected_nonfinite),
        "rejected_slope": int(rejected_slope),
        "rejected_x50": int(rejected_x50),
        "success_rate": (
            float(len(arr) / attempts)
            if attempts > 0 else np.nan
        ),
    }

    return arr, diagnostics

def bootstrap_band_from_params(x_grid, pars, alpha=0.05):
    """Convert parameter draws to pointwise confidence bands."""
    x_grid = np.asarray(x_grid, float).reshape(-1)
    pars = np.asarray(pars, float)

    if pars.ndim != 2 or pars.shape[0] == 0:
        nan = np.full_like(x_grid, np.nan, dtype=float)
        return nan, nan, nan

    curves = np.asarray([model_p(x_grid, bb) for bb in pars], float)
    q = np.quantile(curves, [alpha / 2, 0.5, 1.0 - alpha / 2], axis=0)
    return q[0], q[1], q[2]


def ci_band_normal_mle_sim(
    x_grid,
    b,
    cov,
    M=MC_DRAWS_BANDS,
    seed=123,
    alpha=0.05,
    enforce_positive_slope=True,
    enforce_x50_in_grid=False,
    slope_eps=1e-10,
):
    """
    Monte Carlo confidence band from the local Gaussian approximation.

    The historical function name and signature are retained. The old
    x50-in-grid filter is intentionally ignored.
    """
    draws, _ = gaussian_parameter_draws(
        b,
        cov,
        M=M,
        seed=seed,
        enforce_positive_slope=enforce_positive_slope,
        slope_eps=slope_eps,
    )
    return bootstrap_band_from_params(x_grid, draws, alpha=alpha)


def x50_normal_ci_from_mvnorm(
    b,
    cov,
    M=MC_DRAWS_TABLE,
    seed=123,
    alpha=0.05,
    enforce_positive_slope=True,
    slope_eps=1e-10,
):
    """Monte Carlo interval for model-scale x50 = -b0/b1."""
    draws, _ = gaussian_parameter_draws(
        b,
        cov,
        M=M,
        seed=seed,
        enforce_positive_slope=enforce_positive_slope,
        slope_eps=slope_eps,
    )

    if len(draws) == 0:
        return np.nan, np.nan, np.nan, 0

    vals = np.asarray([x50(bb) for bb in draws], float)
    vals = vals[np.isfinite(vals)]

    if len(vals) == 0:
        return np.nan, np.nan, np.nan, 0

    q = np.quantile(vals, [alpha / 2, 0.5, 1.0 - alpha / 2])
    return float(q[0]), float(q[1]), float(q[2]), int(len(vals))


def s50_mc_ci_from_mvnorm(
    b,
    cov,
    transform="raw",
    M=MC_DRAWS_TABLE,
    seed=123,
    alpha=0.05,
    enforce_positive_slope=True,
    slope_eps=1e-10,
):
    """Monte Carlo interval for raw-scale midpoint slope s50."""
    draws, _ = gaussian_parameter_draws(
        b,
        cov,
        M=M,
        seed=seed,
        enforce_positive_slope=enforce_positive_slope,
        slope_eps=slope_eps,
    )

    if len(draws) == 0:
        return np.nan, np.nan, np.nan, 0

    vals = np.asarray(
        [s50_from_b(bb, transform=transform) for bb in draws],
        float,
    )
    vals = vals[np.isfinite(vals)]

    if len(vals) == 0:
        return np.nan, np.nan, np.nan, 0

    q = np.quantile(vals, [alpha / 2, 0.5, 1.0 - alpha / 2])
    return float(q[0]), float(q[1]), float(q[2]), int(len(vals))


# Historical alias retained.
def s50_normal_ci_from_mvnorm(
    b,
    cov,
    transform="raw",
    M=MC_DRAWS_TABLE,
    seed=123,
    alpha=0.05,
    enforce_positive_slope=True,
    slope_eps=1e-10,
):
    return s50_mc_ci_from_mvnorm(
        b,
        cov,
        transform=transform,
        M=M,
        seed=seed,
        alpha=alpha,
        enforce_positive_slope=enforce_positive_slope,
        slope_eps=slope_eps,
    )


# ============================================================
# 8) Bootstrap parameter generators
# ============================================================

def bootstrap_params_nonparametric(
    x,
    y,
    B=2000,
    seed=123,
    l2=0.0,
    b_start=None,
):
    """Ordinary patient-level nonparametric bootstrap."""
    rng = np.random.default_rng(seed)

    x = np.asarray(x, float).reshape(-1)
    y = np.asarray(y, int).reshape(-1)
    n = len(y)

    out = []
    failed = 0

    for _ in range(int(B)):
        idx = rng.choice(n, size=n, replace=True)
        xb = x[idx]
        yb = y[idx]

        if np.unique(yb).size < 2:
            failed += 1
            continue

        try:
            bb = fit_newton(xb, yb, b_start=b_start, l2=l2)
            if np.all(np.isfinite(bb)):
                out.append(bb)
            else:
                failed += 1
        except Exception:
            failed += 1

    arr = np.asarray(out, float) if out else np.empty((0, 2), float)
    diagnostics = {
        "attempted": int(B),
        "successful": int(len(arr)),
        "failed": int(failed),
    }
    return arr, diagnostics



# Historical alias retained for older notebook cells.
def bootstrap_params_ordinary(
    x,
    y,
    B=2000,
    seed=123,
    l2=0.0,
    b_start=None,
):
    return bootstrap_params_nonparametric(
        x,
        y,
        B=B,
        seed=seed,
        l2=l2,
        b_start=b_start,
    )


def bootstrap_params_stratified(
    x,
    y,
    B=2000,
    seed=123,
    l2=0.0,
    b_start=None,
):
    """Class-stratified nonparametric bootstrap preserving class counts."""
    rng = np.random.default_rng(seed)

    x = np.asarray(x, float).reshape(-1)
    y = np.asarray(y, int).reshape(-1)

    x0 = x[y == 0]
    x1 = x[y == 1]
    n0 = len(x0)
    n1 = len(x1)

    if n0 == 0 or n1 == 0:
        return np.empty((0, 2), float), {
            "attempted": int(B),
            "successful": 0,
            "failed": int(B),
        }

    out = []
    failed = 0

    for _ in range(int(B)):
        xb0 = rng.choice(x0, size=n0, replace=True)
        xb1 = rng.choice(x1, size=n1, replace=True)

        xb = np.concatenate([xb0, xb1])
        yb = np.concatenate([
            np.zeros(n0, dtype=int),
            np.ones(n1, dtype=int),
        ])

        try:
            bb = fit_newton(xb, yb, b_start=b_start, l2=l2)
            if np.all(np.isfinite(bb)):
                out.append(bb)
            else:
                failed += 1
        except Exception:
            failed += 1

    arr = np.asarray(out, float) if out else np.empty((0, 2), float)
    diagnostics = {
        "attempted": int(B),
        "successful": int(len(arr)),
        "failed": int(failed),
        "success_rate": float(len(arr) / B) if B > 0 else np.nan,
    }
    return arr, diagnostics


def bootstrap_params_parametric(
    x,
    b,
    B=2000,
    seed=123,
    l2=0.0,
    min_ae=2,
):
    """
    Parametric bootstrap with y* ~ Bernoulli[p_hat(x)].

    ``min_ae`` is retained for compatibility.
    """
    rng = np.random.default_rng(seed)

    x = np.asarray(x, float).reshape(-1)
    b = np.asarray(b, float).reshape(2)

    p = model_p(x, b)
    n = len(x)

    out = []
    tries = 0
    max_tries = max(10 * int(B), 1000)

    while len(out) < int(B) and tries < max_tries:
        tries += 1
        yb = rng.binomial(1, p, size=n).astype(int)

        n1 = int(np.sum(yb))
        n0 = n - n1

        if n1 < int(min_ae) or n0 < 1:
            continue

        try:
            bb = fit_newton(x, yb, b_start=b, l2=l2)
            if np.all(np.isfinite(bb)):
                out.append(bb)
        except Exception:
            pass

    arr = np.asarray(out, float) if out else np.empty((0, 2), float)
    diagnostics = {
        "attempted": int(tries),
        "successful": int(len(arr)),
        "failed_or_rejected": int(tries - len(arr)),
        "success_rate": float(len(arr) / tries) if tries > 0 else np.nan,
    }
    return arr, diagnostics


# ============================================================
# 9) High-level wrapper for one panel
# ============================================================

def fit_ci_pack_rawgrid(
    x_raw,
    y,
    transform="raw",
    xmax_raw=None,
    grid_n=500,
    name="",
    l2=0.0,
    B=2000,
    seed=123,
    min_ae=2,
    z=1.959963984540054,
):
    """
    Fit one validated logistic model and construct five uncertainty summaries.
    """
    x_raw = np.asarray(x_raw, float).reshape(-1)
    y = np.asarray(y, int).reshape(-1)

    if transform not in ("raw", "log"):
        raise ValueError("transform must be 'raw' or 'log'")

    x_raw = np.clip(x_raw, 1e-12, None)
    x_model = x_raw if transform == "raw" else np.log(x_raw)

    # Single validated fitting path.
    validated = fit_pack(
        x_model,
        y,
        name=name,
        l2=l2,
        z=z,
    )
    b = validated["b"]
    cov = validated["cov"]
    gof = validated["gof"]

    xmin_raw = float(np.min(x_raw))
    xmax0 = float(np.max(x_raw))
    xmax_use = xmax0 if xmax_raw is None else max(float(xmax_raw), xmax0)

    x_grid_raw = np.linspace(xmin_raw, xmax_use, int(grid_n))
    x_grid_raw = np.clip(x_grid_raw, 1e-12, None)
    x_grid_model = x_grid_raw if transform == "raw" else np.log(x_grid_raw)

    mc_x50_bounds = (
        float(np.min(x_grid_model)),
        float(np.max(x_grid_model)),
    )

    mc_draws, diag_mc = gaussian_parameter_draws(
        b,
        cov,
        M=MC_DRAWS_BANDS,
        seed=seed + 10,
        enforce_positive_slope=True,
        x50_bounds=mc_x50_bounds,
    )

    pars_np, diag_np = bootstrap_params_nonparametric(
        x_model,
        y,
        B=B,
        seed=seed + 1,
        l2=l2,
        b_start=b,
    )

    pars_str, diag_str = bootstrap_params_stratified(
        x_model,
        y,
        B=B,
        seed=seed + 2,
        l2=l2,
        b_start=b,
    )

    pars_pm, diag_pm = bootstrap_params_parametric(
        x_model,
        b,
        B=B,
        seed=seed + 3,
        l2=l2,
        min_ae=min_ae,
    )

    return {
        "name": name,
        "transform": transform,
        "l2": float(l2),
        "x_raw": x_raw,
        "x_model": x_model,
        "y": y,
        "x_grid_raw": x_grid_raw,
        "x_grid_model": x_grid_model,
        "b": b,
        "cov": cov,
        "gof": gof,
        "LCL": validated["LCL"],
        "UCL": validated["UCL"],
        "bands": OrderedDict([
            ("Wald", ci_band_wald(x_grid_model, b, cov, z=z)),
            ("MC", bootstrap_band_from_params(x_grid_model, mc_draws)),
            ("Nonparametric", bootstrap_band_from_params(x_grid_model, pars_np)),
            ("Stratified", bootstrap_band_from_params(x_grid_model, pars_str)),
            ("Parametric", bootstrap_band_from_params(x_grid_model, pars_pm)),
        ]),
        "pars_mc": mc_draws,
        "pars_nonparametric": pars_np,
        "pars_stratified": pars_str,
        "pars_parametric": pars_pm,
        # Historical aliases
        "pars_normal": mc_draws,
        "pars_nonparam": pars_np,
        "pars_nonparam_ordinary": pars_np,
        "pars_nonparam_stratified": pars_str,
        "bootstrap_diagnostics": OrderedDict([
            ("MC", diag_mc),
            ("Nonparametric", diag_np),
            ("Stratified", diag_str),
            ("Parametric", diag_pm),
        ]),
    }


# ============================================================
# 10) Model-band table with LL / UL
# ============================================================

def model_ci_table_4methods(
    P,
    keys=("FULL-RAW", "TRIM-RAW", "FULL-LOG", "TRIM-LOG"),
):
    """Long pointwise model-band table."""
    import pandas as pd

    rows = []

    for key in keys:
        pk = P[key]
        xg_raw = np.asarray(pk["x_grid_raw"], float)
        xg_mod = np.asarray(pk["x_grid_model"], float)
        trans = pk.get("transform", "")

        for method in CI_METHODS:
            if method not in pk["bands"]:
                continue

            lo, md, hi = pk["bands"][method]
            lo = np.asarray(lo, float)
            md = np.asarray(md, float)
            hi = np.asarray(hi, float)

            for i in range(len(xg_raw)):
                rows.append({
                    "Panel": key,
                    "Method": method,
                    "transform": trans,
                    "x_grid_raw": float(xg_raw[i]),
                    "x_grid_model": float(xg_mod[i]),
                    "fit": float(md[i]),
                    "LL": float(lo[i]),
                    "UL": float(hi[i]),
                    "width": float(hi[i] - lo[i]),
                })

    return pd.DataFrame(rows)


# ============================================================
# 11) Parameter, x50, and s50 CI summary table
# ============================================================

def param_ci_table_4methods(
    P,
    keys=("FULL-RAW", "TRIM-RAW", "FULL-LOG", "TRIM-LOG"),
    z=1.959963984540054,
    alpha=0.05,
    include_point_est=True,
    M_normal=MC_DRAWS_TABLE,
    seed_normal=123,
):
    """
    Notebook-compatible CI summary using:
      Wald, MC, Nonparametric, Stratified, Parametric.
    """
    import pandas as pd

    def _quantile_ci(values):
        values = np.asarray(values, float)
        values = values[np.isfinite(values)]

        if len(values) == 0:
            return np.nan, np.nan, np.nan

        q = np.quantile(values, [alpha / 2, 0.5, 1.0 - alpha / 2])
        return float(q[0]), float(q[1]), float(q[2])

    def _to_raw_x50_scalar(value, trans):
        if not np.isfinite(value):
            return np.nan
        return float(np.exp(value)) if trans == "log" else float(value)

    def _width(lo, hi):
        if np.isfinite(lo) and np.isfinite(hi):
            return float(hi - lo)
        return np.nan

    rows = []

    for ik, key in enumerate(keys):
        pk = P[key]

        b = np.asarray(pk["b"], float).reshape(2)
        cov = np.asarray(pk["cov"], float).reshape(2, 2)
        trans = pk.get("transform", "")
        l2 = float(pk.get("l2", 0.0))

        b0_hat = float(b[0])
        b1_hat = float(b[1])
        x50_hat = float(x50(b))
        suv50_hat = _to_raw_x50_scalar(x50_hat, trans)
        s50_hat = float(s50_from_b(b, transform=trans))

        def add_row(method, b0_ci, b1_ci, x_ci, suv_ci, s_ci, n_used):
            row = {
                "Panel": key,
                "Method": method,

                "b0_hat": b0_hat,
                "b0_LCL": b0_ci[0],
                "b0_UCL": b0_ci[2],
                "b0_width": _width(b0_ci[0], b0_ci[2]),

                "b1_hat": b1_hat,
                "b1_LCL": b1_ci[0],
                "b1_UCL": b1_ci[2],
                "b1_width": _width(b1_ci[0], b1_ci[2]),

                "x50_hat": x50_hat,
                "x50_med": x_ci[1],
                "x50_LCL": x_ci[0],
                "x50_UCL": x_ci[2],
                "x50_width": _width(x_ci[0], x_ci[2]),

                "SUV50_hat": suv50_hat,
                "SUV50_med": suv_ci[1],
                "SUV50_LCL": suv_ci[0],
                "SUV50_UCL": suv_ci[2],
                "SUV50_width": _width(suv_ci[0], suv_ci[2]),

                "s50_hat": s50_hat,
                "s50_med": s_ci[1],
                "s50_LCL": s_ci[0],
                "s50_UCL": s_ci[2],
                "s50_width": _width(s_ci[0], s_ci[2]),

                "transform": trans,
                "l2": l2,
                "B_used": n_used,
            }

            if not include_point_est:
                for col in (
                    "b0_hat",
                    "b1_hat",
                    "x50_hat",
                    "x50_med",
                    "SUV50_hat",
                    "SUV50_med",
                    "s50_hat",
                    "s50_med",
                    "transform",
                    "l2",
                ):
                    row.pop(col, None)

            rows.append(row)

        # Wald
        lcl, ucl = wald_ci(b, cov, z=z)
        x_l, x_u = x50_wald_ci(b, cov, z=z)
        suv_l = _to_raw_x50_scalar(x_l, trans)
        suv_u = _to_raw_x50_scalar(x_u, trans)
        s_l, s_u = s50_wald_ci_numeric(b, cov, transform=trans, z=z)

        add_row(
            "Wald",
            (float(lcl[0]), b0_hat, float(ucl[0])),
            (float(lcl[1]), b1_hat, float(ucl[1])),
            (float(x_l), x50_hat, float(x_u)),
            (float(suv_l), suv50_hat, float(suv_u)),
            (float(s_l), s50_hat, float(s_u)),
            np.nan,
        )

        # Draw-based methods
        mc_draws = pk.get("pars_mc", pk.get("pars_normal"))

        if mc_draws is None or len(mc_draws) < int(M_normal):
            mc_draws, _ = gaussian_parameter_draws(
                b,
                cov,
                M=M_normal,
                seed=seed_normal + 1000 * ik,
                enforce_positive_slope=True,
            )

        draw_sets = {
            "MC": mc_draws,
            "Nonparametric": pk.get(
                "pars_nonparametric",
                pk.get("pars_nonparam", np.empty((0, 2))),
            ),
            "Stratified": pk.get(
                "pars_stratified",
                pk.get("pars_nonparam_stratified", np.empty((0, 2))),
            ),
            "Parametric": pk.get(
                "pars_parametric",
                np.empty((0, 2)),
            ),
        }

        for method in ("MC", "Nonparametric", "Stratified", "Parametric"):
            pars = np.asarray(draw_sets[method], float)

            if pars.ndim != 2 or len(pars) == 0:
                nan3 = (np.nan, np.nan, np.nan)
                add_row(method, nan3, nan3, nan3, nan3, nan3, 0)
                continue

            b0_ci = _quantile_ci(pars[:, 0])
            b1_ci = _quantile_ci(pars[:, 1])

            xvals = np.asarray([x50(bb) for bb in pars], float)
            suvvals = np.asarray(
                [_to_raw_x50_scalar(v, trans) for v in xvals],
                float,
            )
            svals = np.asarray(
                [s50_from_b(bb, transform=trans) for bb in pars],
                float,
            )

            add_row(
                method,
                b0_ci,
                b1_ci,
                _quantile_ci(xvals),
                _quantile_ci(suvvals),
                _quantile_ci(svals),
                int(len(pars)),
            )

    return pd.DataFrame(rows)


def combined_x50_model_bounds_table(
    P,
    keys=("FULL-RAW", "TRIM-RAW", "FULL-LOG", "TRIM-LOG"),
    z=1.959963984540054,
    alpha=0.05,
    M_normal=MC_DRAWS_TABLE,
    seed_normal=123,
):
    """Combine characteristic intervals with global curve-band summaries."""
    import pandas as pd

    param_df = param_ci_table_4methods(
        P,
        keys=keys,
        z=z,
        alpha=alpha,
        include_point_est=True,
        M_normal=M_normal,
        seed_normal=seed_normal,
    ).copy()

    model_df = model_ci_table_4methods(P, keys=keys).copy()

    global_df = (
        model_df
        .groupby(["Panel", "Method"], as_index=False)
        .agg(
            global_LL=("LL", "min"),
            global_UL=("UL", "max"),
            fit_min=("fit", "min"),
            fit_max=("fit", "max"),
            mean_width=("width", "mean"),
            max_width=("width", "max"),
        )
    )

    global_df["global_width"] = global_df["global_UL"] - global_df["global_LL"]

    out = pd.merge(
        param_df,
        global_df,
        on=["Panel", "Method"],
        how="left",
    )

    preferred = [
        "Panel", "Method", "transform",
        "x50_hat", "x50_LCL", "x50_UCL", "x50_width",
        "SUV50_hat", "SUV50_LCL", "SUV50_UCL", "SUV50_width",
        "s50_hat", "s50_LCL", "s50_UCL", "s50_width",
        "global_LL", "global_UL", "global_width",
        "fit_min", "fit_max", "mean_width", "max_width", "B_used",
    ]

    cols = [c for c in preferred if c in out.columns] + [
        c for c in out.columns if c not in preferred
    ]

    return out[cols]


def bootstrap_diagnostics_table(P):
    """Return success information for MC and bootstrap methods."""
    import pandas as pd

    rows = []
    for panel, pk in P.items():
        for method, diag in pk.get("bootstrap_diagnostics", {}).items():
            rows.append({"Panel": panel, "Method": method, **diag})

    return pd.DataFrame(rows)


# ============================================================
# 12) CI figure
# ============================================================
def plot_ci_four_panels(P):
    import numpy as np
    import matplotlib.pyplot as plt
    import matplotlib.lines as mlines

    plt.style.use("default")

    COL_NC = "#4c9ed9"
    COL_AE = "#f28e2b"
    COL_FIT = "#000000"

    # Methods displayed in the main figure
    FIGURE_METHODS = [
        "Wald",
        "MC",
        "Nonparametric",
        "Parametric",
    ]

    styles = OrderedDict([
        ("Wald",          ("#2ca02c", "-.", 0.12)),
        ("MC",            ("#d62728", "--", 0.14)),
        ("Nonparametric", ("#1f77b4", ":", 0.16)),
        ("Parametric",    ("#17becf", (0, (6, 2)), 0.14)),
    ])

    panel_order = [
        "FULL-LOG",
        "FULL-RAW",
        "TRIM-LOG",
        "TRIM-RAW",
    ]
    panel_letters = ["A", "B", "C", "D"]

    fig, axs = plt.subplots(
        2,
        2,
        figsize=(15, 10),
        dpi=180,
        sharex="col",
        sharey=True,
    )

    for ax, key, letter in zip(
        axs.flat,
        panel_order,
        panel_letters,
    ):
        pk = P[key]

        x_raw = np.asarray(pk["x_raw"], float)
        y = np.asarray(pk["y"], int)
        xg_raw = np.asarray(pk["x_grid_raw"], float)
        xg_model = np.asarray(pk["x_grid_model"], float)
        transform = pk["transform"]

        if transform == "raw":
            xs = x_raw
            xg = xg_raw
        else:
            xs = np.log(x_raw)
            xg = xg_model

        rng = np.random.default_rng(123 + ord(letter))
        jit = (rng.random(len(y)) - 0.5) * 0.04

        ax.scatter(
            xs[y == 0],
            (y + jit)[y == 0],
            s=22,
            alpha=0.45,
            color=COL_NC,
            edgecolors="none",
            zorder=5,
        )

        ax.scatter(
            xs[y == 1],
            (y + jit)[y == 1],
            s=24,
            alpha=0.85,
            color=COL_AE,
            edgecolors="none",
            zorder=5,
        )

        for method in FIGURE_METHODS:
            if method not in pk["bands"]:
                continue

            lo, _, hi = pk["bands"][method]
            color, linestyle, fill_alpha = styles[method]

            ax.fill_between(
                xg,
                lo,
                hi,
                color=color,
                alpha=fill_alpha,
                zorder=1,
            )

            ax.plot(
                xg,
                lo,
                color=color,
                linestyle=linestyle,
                lw=1.6,
                zorder=2,
            )

            ax.plot(
                xg,
                hi,
                color=color,
                linestyle=linestyle,
                lw=1.6,
                zorder=2,
            )

        fit_curve = model_p(
            xg_model,
            pk["b"],
        )

        ax.plot(
            xg,
            fit_curve,
            color=COL_FIT,
            lw=2.5,
            zorder=6,
        )

        ax.text(
            0.03,
            0.95,
            letter,
            transform=ax.transAxes,
            fontsize=15,
            ha="left",
            va="top",
        )

        ax.set_ylim(-0.05, 1.05)
        ax.grid(False)

        ax.tick_params(
            axis="both",
            which="major",
            labelsize=11,
            length=4,
            width=0.8,
            direction="out",
        )

    axs[0, 0].set_ylabel(
        r"$\mathrm{P(AE \mid X = x)}$",
        fontsize=13,
    )
    axs[1, 0].set_ylabel(
        r"$\mathrm{P(AE \mid X = x)}$",
        fontsize=13,
    )

    axs[1, 0].set_xlabel(
        r"$\log(\mathrm{X})$",
        fontsize=13,
    )
    axs[1, 1].set_xlabel(
        r"$\mathrm{X}$",
        fontsize=13,
    )

    for ax in axs[0, :]:
        ax.tick_params(
            axis="x",
            which="both",
            labelbottom=False,
        )

    for ax in axs[:, 1]:
        ax.tick_params(
            axis="y",
            which="both",
            labelleft=False,
        )

    labels = {
        "Wald": "CI: Wald 95%",
        "MC": "CI: MC propagation 95%",
        "Nonparametric": "CI: nonparametric bootstrap 95%",
        "Parametric": "CI: parametric bootstrap 95%",
    }

    handles = [
        mlines.Line2D(
            [],
            [],
            marker="o",
            color=COL_NC,
            linestyle="None",
            markersize=8,
            label="data: NC",
        ),
        mlines.Line2D(
            [],
            [],
            marker="o",
            color=COL_AE,
            linestyle="None",
            markersize=8,
            label="data: AE",
        ),
        mlines.Line2D(
            [],
            [],
            color=COL_FIT,
            lw=2.5,
            label="fit",
        ),
    ]

    for method in FIGURE_METHODS:
        color, linestyle, _ = styles[method]

        handles.append(
            mlines.Line2D(
                [],
                [],
                color=color,
                lw=2,
                linestyle=linestyle,
                label=labels[method],
            )
        )

    leg = axs[1, 1].legend(
        handles=handles,
        loc="lower right",
        bbox_to_anchor=(0.98, 0.04),
        fontsize=8.5,
        frameon=True,
    )

    leg.get_frame().set_facecolor("white")
    leg.get_frame().set_edgecolor("#bdbdbd")
    leg.get_frame().set_linewidth(0.8)

    fig.subplots_adjust(
        left=0.08,
        right=0.98,
        bottom=0.08,
        top=0.98,
        wspace=0.06,
        hspace=0.06,
    )

    return fig, axs
# ============================================================
# ELASTICITY ANALYSIS (x50 and s50)
# ============================================================

import numpy as np
import matplotlib.pyplot as plt


# ------------------------------------------------------------
# Core elasticity computation
# ------------------------------------------------------------
def elasticity_x50_s50(theta, mode="raw"):
    """
    Elasticity for x50 and s50 with respect to theta0 and theta1.

    mode
    ----
    'raw' : eta = theta0 + theta1*x
    'log' : eta = theta0 + theta1*log(x)

    Returns
    -------
    dict with:
        theta0, theta1,
        x50, s50,
        E_x50_theta0, E_x50_theta1,
        E_s50_theta0, E_s50_theta1
    """
    theta0, theta1 = map(float, np.asarray(theta).reshape(2))

    if np.abs(theta1) < 1e-12:
        return dict(
            theta0=theta0,
            theta1=theta1,
            x50=np.nan,
            s50=np.nan,
            E_x50_theta0=np.nan,
            E_x50_theta1=np.nan,
            E_s50_theta0=np.nan,
            E_s50_theta1=np.nan,
        )

    # =========================
    # RAW MODEL
    # =========================
    if mode == "raw":
        x50 = -theta0 / theta1
        s50 = theta1 / 4.0

        E_x50_theta0 = 1.0
        E_x50_theta1 = -1.0

        E_s50_theta0 = 0.0
        E_s50_theta1 = 1.0

    # =========================
    # LOG MODEL
    # =========================
    elif mode == "log":
        x50 = float(np.exp(-theta0 / theta1))
        s50 = theta1 / (4.0 * x50)

        E_x50_theta0 = -theta0 / theta1
        E_x50_theta1 = theta0 / theta1

        E_s50_theta0 = -E_x50_theta0
        E_s50_theta1 = 1.0 - E_x50_theta1

    else:
        raise ValueError("mode must be 'raw' or 'log'")

    return dict(
        theta0=theta0,
        theta1=theta1,
        x50=x50,
        s50=s50,
        E_x50_theta0=E_x50_theta0,
        E_x50_theta1=E_x50_theta1,
        E_s50_theta0=E_s50_theta0,
        E_s50_theta1=E_s50_theta1,
    )


# ------------------------------------------------------------
# Table for 4 panels
# ------------------------------------------------------------
def elasticity_table_4panels(P, keys=None, make_plots=True):
    import pandas as pd
    import numpy as np

    if keys is None:
        keys = list(P.keys())

    rows = []

    for key in keys:
        pk = P[key]

        if "b" not in pk:
            print(f"[skip] {key}: no fitted parameter key 'b'")
            continue

        theta = np.asarray(pk["b"], float).reshape(2)
        transform = pk.get("transform", "raw")

        res = elasticity_x50_s50(theta, mode=transform)

        rows.append({
            "Panel": key,
            "transform": transform,
            **res
        })

    df = pd.DataFrame(rows)

    if make_plots and len(df) > 0:
        plot_x50_values(df)
        plot_s50_values(df)
        plot_x50_theta1_elasticity(df)
        plot_s50_theta1_elasticity(df)

    return df


# ------------------------------------------------------------
# Plots
# ------------------------------------------------------------
def plot_x50_values(df):
    fig, ax = plt.subplots(figsize=(7, 4))
    ax.bar(df["Panel"], df["x50"])
    ax.set_ylabel("x50")
    ax.set_title("x50 across panels")
    plt.xticks(rotation=30)
    plt.tight_layout()
    plt.show()


def plot_s50_values(df):
    fig, ax = plt.subplots(figsize=(7, 4))
    ax.bar(df["Panel"], df["s50"])
    ax.set_ylabel("s50")
    ax.set_title("s50 across panels")
    plt.xticks(rotation=30)
    plt.tight_layout()
    plt.show()


def plot_x50_theta1_elasticity(df):
    fig, ax = plt.subplots(figsize=(7, 4))
    ax.bar(df["Panel"], df["E_x50_theta1"])
    ax.set_ylabel("Elasticity")
    ax.set_title("Elasticity of x50 w.r.t. theta1")
    plt.xticks(rotation=30)
    plt.tight_layout()
    plt.show()


def plot_s50_theta1_elasticity(df):
    fig, ax = plt.subplots(figsize=(7, 4))
    ax.bar(df["Panel"], df["E_s50_theta1"])
    ax.set_ylabel("Elasticity")
    ax.set_title("Elasticity of s50 w.r.t. theta1")
    plt.xticks(rotation=30)
    plt.tight_layout()
    plt.show()
# ------------------------------------------------------------
# logistic helpers for noise analysis
# ------------------------------------------------------------

# ============================================================
# NOISE ANALYSIS FOR LOGISTIC MODEL
# Correct x50 for RAW and LOG models
# ============================================================

import os
import numpy as np
import matplotlib.pyplot as plt
import matplotlib.lines as mlines


# ------------------------------------------------------------
# logistic fit / prediction / x50
# CONSISTENT WITH MAIN LOGISTIC ANALYSIS
# ------------------------------------------------------------

def fit_logistic_x(x_raw, y, transform="raw", l2=1e-8):
    x_raw = np.clip(np.asarray(x_raw, float).ravel(), 1e-12, None)
    y = np.asarray(y, int).ravel()

    if transform == "raw":
        x_model = x_raw
    elif transform == "log":
        x_model = np.log(x_raw)
    else:
        raise ValueError("transform must be 'raw' or 'log'")

    return fit_newton(x_model, y, l2=l2)


def predict_curve_x(b, x_grid_raw, transform="raw"):
    x_grid_raw = np.clip(np.asarray(x_grid_raw, float), 1e-12, None)

    if transform == "raw":
        x_model = x_grid_raw
    elif transform == "log":
        x_model = np.log(x_grid_raw)
    else:
        raise ValueError("transform must be 'raw' or 'log'")

    return model_p(x_model, b)


def x50_from_b(b, transform="raw"):
    b = np.asarray(b, float).reshape(2)
    x50_model = x50(b)

    if not np.isfinite(x50_model):
        return np.nan

    if transform == "raw":
        return float(x50_model)
    elif transform == "log":
        return float(np.exp(x50_model))
    else:
        raise ValueError("transform must be 'raw' or 'log'")


def check_noise_x50(pack):
    b = np.asarray(pack["b_clean"], float).reshape(2)
    transform = pack["transform"]
    x50_raw = pack["x50"]

    x50_model = x50_raw if transform == "raw" else np.log(x50_raw)
    p50 = model_p(np.array([x50_model]), b)[0]

    print(
        "transform =", transform,
        "| x50_raw =", x50_raw,
        "| P(x50) =", p50
    )


# ------------------------------------------------------------
# noise helpers
# ------------------------------------------------------------

def add_noise_mult(x, sigma, rng):
    x = np.asarray(x, float)
    return np.clip(x * np.exp(rng.normal(0, sigma, size=x.shape)), 1e-12, None)


def add_noise_add(x, sigma, rng):
    x = np.asarray(x, float)
    return np.clip(x + rng.normal(0, sigma, size=x.shape), 1e-12, None)


def band_quantiles(curves):
    C = np.vstack(curves)
    return np.quantile(C, [0.025, 0.5, 0.975], axis=0)


# ------------------------------------------------------------
# build noise bands
# ------------------------------------------------------------

def noise_logistic_bands(
    x_raw,
    y,
    transform="raw",
    sigma_mult=0.129,
    sigma_add=0.144,
    x_max=5,
    grid_n=1000,
    n_refit=200,
    n_tta=3000,
    seed=1234,
    l2=1e-8,
):
    rng = np.random.default_rng(seed)

    x_raw = np.clip(np.asarray(x_raw, float).ravel(), 1e-12, None)
    y = np.asarray(y).astype(int).ravel()

    xc = np.linspace(1e-12, x_max, grid_n)

    b_clean = fit_logistic_x(x_raw, y, transform=transform, l2=l2)
    clean = predict_curve_x(b_clean, xc, transform=transform)
    x50_val = x50_from_b(b_clean, transform=transform)

    curves = []
    for _ in range(n_refit):
        xn = add_noise_mult(x_raw, sigma_mult, rng)
        bn = fit_logistic_x(xn, y, transform=transform, l2=l2)
        curves.append(predict_curve_x(bn, xc, transform=transform))
    mult_refit = band_quantiles(curves)

    curves = []
    for _ in range(n_tta):
        xn = add_noise_mult(xc, sigma_mult, rng)
        curves.append(predict_curve_x(b_clean, xn, transform=transform))
    mult_tta = band_quantiles(curves)

    curves = []
    for _ in range(n_refit):
        xn = add_noise_add(x_raw, sigma_add, rng)
        bn = fit_logistic_x(xn, y, transform=transform, l2=l2)
        curves.append(predict_curve_x(bn, xc, transform=transform))
    add_refit = band_quantiles(curves)

    curves = []
    for _ in range(n_tta):
        xn = add_noise_add(xc, sigma_add, rng)
        curves.append(predict_curve_x(b_clean, xn, transform=transform))
    add_tta = band_quantiles(curves)

    return {
        "xc": xc,
        "clean": clean,
        "x50": x50_val,
        "b_clean": b_clean,
        "transform": transform,
        "l2": float(l2),
        "mult_refit": mult_refit,
        "mult_tta": mult_tta,
        "add_refit": add_refit,
        "add_tta": add_tta,
    }

    from scipy.ndimage import gaussian_filter1d

    lo = np.quantile(curves, 0.025, axis=0)
    md = np.quantile(curves, 0.500, axis=0)
    hi = np.quantile(curves, 0.975, axis=0)

# smooth boundaries
    lo = gaussian_filter1d(lo, sigma=8)
    md = gaussian_filter1d(md, sigma=8)
    hi = gaussian_filter1d(hi, sigma=8)

    return lo, md, hi

# ------------------------------------------------------------
# legend
# ------------------------------------------------------------

def noise_legend_handles():
    return [
        mlines.Line2D([], [], marker="o", color="#2b8cbe",
                      linestyle="None", markersize=7, label="NC data"),
        mlines.Line2D([], [], marker="o", color="#d7301f",
                      linestyle="None", markersize=7, label="AE data"),
        mlines.Line2D([], [], color="black", lw=2.2, label="Initial fit"),
        mlines.Line2D([], [], color="#1f78b4", lw=6, alpha=0.24,
                      label="refit band, multiplicative noise"),
        mlines.Line2D([], [], color="#1f78b4", lw=6, alpha=0.10,
                      label="fixed-model band, multiplicative noise"),
        mlines.Line2D([], [], color="#e66101", lw=6, alpha=0.24,
                      label="refit band, additive noise"),
        mlines.Line2D([], [], color="#e66101", lw=6, alpha=0.10,
                      label="fixed-model band, additive noise"),
        mlines.Line2D([], [], color="#666666", ls="--", lw=1.2,
                      label=r"$x_{50}$"),
    ]


# ------------------------------------------------------------
# plot one panel
# ------------------------------------------------------------

def plot_noise_panel(ax, pack, kind="mult", label="A", X=None, y=None):
    COL_MULT = "#1f78b4"
    COL_ADD  = "#e66101"
    COL_NC   = "#2b8cbe"
    COL_AE   = "#d7301f"

    xc = pack["xc"]
    clean = pack["clean"]
    x50_val = pack["x50"]

    if kind == "mult":
        refit = pack["mult_refit"]
        tta = pack["mult_tta"]
        color = COL_MULT
    elif kind == "add":
        refit = pack["add_refit"]
        tta = pack["add_tta"]
        color = COL_ADD
    else:
        raise ValueError("kind must be 'mult' or 'add'")

    lo_r, _, hi_r = refit
    lo_t, _, hi_t = tta

    ax.fill_between(xc, lo_t, hi_t, color=color, alpha=0.10, zorder=1)
    ax.fill_between(xc, lo_r, hi_r, color=color, alpha=0.24, zorder=2)

    ax.plot(xc, lo_r, color=color, lw=1.0, alpha=0.65, zorder=3)
    ax.plot(xc, hi_r, color=color, lw=1.0, alpha=0.65, zorder=3)

    ax.plot(xc, clean, color="black", lw=2.2, zorder=5)
    ax.axvline(x50_val, color="#666666", ls="--", lw=1.2, alpha=0.9, zorder=4)

    lo_r_x = np.interp(x50_val, xc, lo_r)
    hi_r_x = np.interp(x50_val, xc, hi_r)
    lo_t_x = np.interp(x50_val, xc, lo_t)
    hi_t_x = np.interp(x50_val, xc, hi_t)

    if X is not None and y is not None:
        X = np.asarray(X).ravel()
        y = np.asarray(y).astype(int).ravel()

        ax.scatter(
            X[y == 0], np.zeros(np.sum(y == 0)),
            color=COL_NC, s=24, alpha=0.75,
            edgecolors="none", zorder=7
        )
        ax.scatter(
            X[y == 1], np.ones(np.sum(y == 1)),
            color=COL_AE, s=24, alpha=0.75,
            edgecolors="none", zorder=7
        )

    ax.text(0.03, 0.97, label, transform=ax.transAxes,
            ha="left", va="top", fontsize=15)

    variant_txt = "FULL" if label in ["A", "B"] else "TRIM"
    d_ref = hi_r_x - lo_r_x
    d_tta = hi_t_x - lo_t_x

    info_txt = (
        f"{variant_txt}\n"
        f"$x_{{50}}$={x50_val:.2f}\n"
        f"$\\Delta r$={d_ref:.2f}  $\\Delta t$={d_tta:.2f}"
    )

    ax.text(
        0.02, 0.14,
        info_txt,
        transform=ax.transAxes,
        fontsize=10,
        color="#222",
        ha="left", va="bottom",
        bbox=dict(facecolor="white", edgecolor=color,
                  boxstyle="square,pad=0.25", alpha=0.9)
    )

    ax.set_xlim(0, xc.max())
    ax.set_ylim(-0.05, 1.05)
    ax.grid(alpha=0.25)
    ax.tick_params(axis="both", labelsize=10)


# ------------------------------------------------------------
# full noise figure
# ------------------------------------------------------------

def plot_noise_figure(
    pack_full,
    pack_trim,
    X_full,
    y_full,
    X_trim,
    y_trim,
    figsize=(12, 9),
    dpi=300,
):
    fig, axes = plt.subplots(
        2, 2,
        figsize=figsize,
        dpi=dpi,
        sharex=True,
        sharey=True
    )

    axes = axes.ravel()

    plot_noise_panel(axes[0], pack_full, kind="mult", label="A", X=X_full, y=y_full)
    plot_noise_panel(axes[1], pack_full, kind="add",  label="B", X=X_full, y=y_full)
    plot_noise_panel(axes[2], pack_trim, kind="mult", label="C", X=X_trim, y=y_trim)
    plot_noise_panel(axes[3], pack_trim, kind="add",  label="D", X=X_trim, y=y_trim)

    axes[0].set_ylabel(r"$P(\mathrm{AE}\mid X=x)$", fontsize=12)
    axes[2].set_ylabel(r"$P(\mathrm{AE}\mid X=x)$", fontsize=12)
    axes[2].set_xlabel(r"$X$", fontsize=12)
    axes[3].set_xlabel(r"$X$", fontsize=12)

    handles = noise_legend_handles()

    leg = axes[3].legend(
        handles=handles,
        loc="lower right",
        bbox_to_anchor=(0.97, 0.05),
        fontsize=9,
        frameon=True
    )

    frame = leg.get_frame()
    frame.set_facecolor("white")
    frame.set_edgecolor("#bdbdbd")
    frame.set_linewidth(0.8)

    fig.tight_layout()
    return fig, axes


# ------------------------------------------------------------
# wrapper
# ------------------------------------------------------------

def make_noise_figure(
    X_full,
    y_full,
    X_trim,
    y_trim,
    transform="raw",
    sigma_mult=0.129,
    sigma_add=0.144,
    x_max=5,
    grid_n=1000,
    n_refit=1100,
    n_tta=10000,
    seed=1234,
    l2=1e-8,
    save_path=None,
):
    pack_full = noise_logistic_bands(
        X_full, y_full,
        transform=transform,
        sigma_mult=sigma_mult,
        sigma_add=sigma_add,
        x_max=x_max,
        grid_n=grid_n,
        n_refit=n_refit,
        n_tta=n_tta,
        seed=seed,
        l2=l2,
    )

    pack_trim = noise_logistic_bands(
        X_trim, y_trim,
        transform=transform,
        sigma_mult=sigma_mult,
        sigma_add=sigma_add,
        x_max=x_max,
        grid_n=grid_n,
        n_refit=n_refit,
        n_tta=n_tta,
        seed=seed + 100,
        l2=l2,
    )

    print("FULL n:", len(X_full), "x50:", pack_full["x50"])
    print("TRIM n:", len(X_trim), "x50:", pack_trim["x50"])
    check_noise_x50(pack_full)
    check_noise_x50(pack_trim)

    fig, axes = plot_noise_figure(
        pack_full, pack_trim,
        X_full, y_full,
        X_trim, y_trim,
        figsize=(12, 9),
        dpi=300,
    )

    if save_path is not None:
        folder = os.path.dirname(save_path)
        if folder:
            os.makedirs(folder, exist_ok=True)

        fig.savefig(f"{save_path}.png", dpi=300, bbox_inches="tight")
        fig.savefig(f"{save_path}.pdf", bbox_inches="tight")

    return fig, axes, pack_full, pack_trim