KidIQ: Bayesian R-squared

Source: KidIQ/kidiq_R2.Rmd

This page keeps the rstanarm idea directly: fit Gaussian regressions with lapylace, then compute the Gelman-style Bayesian \(R^2\) from posterior expected values and residual scale.

Setup

from pathlib import Path
import numpy as np
import pandas as pd
import matplotlib.pyplot as plt
import lapylace as lp

root = Path("../../ROS-Examples")

def coef_medians(fit):
    return pd.Series(
        [np.median(fit.alpha_draws()), *np.median(fit.beta_draws(), axis=0)],
        index=["Intercept", *fit.columns],
    )

def fit_glm(formula, data, family=None, seed=1, prior_scale=2.5, intercept_scale=5, aux_scale=10, **kwargs):
    return lp.stan_glm(
        formula,
        data=data,
        family=family or lp.gaussian(),
        prior=lp.normal(0, prior_scale),
        prior_intercept=lp.normal(0, intercept_scale),
        prior_aux=lp.exponential(aux_scale),
        chains=2,
        parallel_chains=2,
        iter_warmup=300,
        iter_sampling=500,
        seed=seed,
        refresh=100,
        **kwargs,
    )
kidiq = pd.read_csv(root / "KidIQ/data/kidiq.csv")
kidiq.head()
kid_score mom_hs mom_iq mom_work mom_age
0 65 1 121.117529 4 27
1 98 1 89.361882 4 25
2 85 1 115.443165 4 27
3 83 1 99.449639 3 25
4 115 1 92.745710 4 27
def bayes_r2(fit, data):
    mu = fit.posterior_epred(data)
    var_mu = np.var(mu, axis=1)
    sigma2 = fit.stan_variables()["sigma"] ** 2
    return var_mu / (var_mu + sigma2)

fits = {
    "mom_hs": fit_glm("kid_score ~ mom_hs", kidiq, seed=17201, prior_scale=10, intercept_scale=30, aux_scale=30),
    "mom_hs + mom_iq": fit_glm("kid_score ~ mom_hs + mom_iq", kidiq, seed=17202, prior_scale=10, intercept_scale=30, aux_scale=30),
}
rows = []
for label, fit in fits.items():
    r2 = bayes_r2(fit, kidiq)
    rows.append({"model": label, "r2_median": np.median(r2), "r2_10%": np.quantile(r2, .1), "r2_90%": np.quantile(r2, .9)})
pd.DataFrame(rows)
                                                                                                                                                                
                                                                                                                                                                
model r2_median r2_10% r2_90%
0 mom_hs 0.053506 0.029807 0.084202
1 mom_hs + mom_iq 0.214905 0.177137 0.257144