Bayes - included with batteries

Bayesian framework comes packed with amazing features that frequentist had to come up with and deal with for each scenario.

1. Natural regularization

In this chapter, I assume you’re familiar with Ridge regression (L2) and Lasso regression (L1). Both approaches are great for dense and spars data structures, respectively. The elastic net makes a compromise between them. However, how do you choose the penalty parameters value? Most people use tuning that needs cross-validation. This makes it less appealing for small datasets that actually benefits the most from regularization!

In comes Bayesian regression. Here, all parameters are regularized by the prior. You set the regularization based on meaningful scientific knowledge and instantly gain regularization without any tuning.

2. Marginalization

#| label: marginal-data

library(rethinking)
library(tidybayes)
library(dplyr)
library(ggplot2)

set.seed(42)

N <- 100

x <- rnorm(N)
z <- rnorm(N)

y <- 1 + 1.5 * x + rnorm(N, 0, 1)

dat <- list(
  N = N,
  x = x,
  z = z,
  y = y
)

The data-generating model is

\(y = 1 + 1.5x + ϵ\)

There is deliberately no \(z\) effect.

Here is the simple model:

#| label: simple-model

set.seed(173329571)

m1 <- ulam(
  alist(
    y ~ dnorm(mu, sigma),

    mu <- alpha + beta_x * x,

    alpha ~ dnorm(0, 1),
    beta_x ~ dnorm(0, 1),
    sigma ~ dexp(1)
  ),
  data = dat,
  chains = 4,
  cores = 4
)

And the unnecessarily complicated model:

#| label: complex-model

set.seed(83150518)

m2 <- ulam(
  alist(
    y ~ dnorm(mu, sigma),

    mu <- alpha + beta_x * x + beta_z * z,

    alpha ~ dnorm(0, 1),
    beta_x ~ dnorm(0, 1),
    beta_z ~ dnorm(0, 1),
    sigma ~ dexp(1)
  ),
  data = dat,
  chains = 4,
  cores = 4
)
#| label: see-estimates

precis(m1)
precis(m2)
#| label: plot1

library(tidybayes)
library(tidybayes.rethinking)

post_m2 <- m2 |>
  spread_draws(beta_x, beta_z)

ggplot(post_m2, aes(x = beta_z)) +
  stat_halfeye(fill = "steelblue", alpha = .7) +
  geom_vline(xintercept = 0, linetype = 2) +
  labs(
    x = expression(beta[z]),
    y = NULL,
    title = "Marginal posterior for the unnecessary parameter"
  ) +
  theme_classic()
#| label: plo2

prior <- tibble(
  beta_z = seq(-3, 3, length.out = 1000)
) |>
  mutate(
    density = dnorm(beta_z, 0, 1)
  )

posterior <- m2 |>
  spread_draws(beta_z)

ggplot() +

  geom_line(
    data = prior,
    aes(beta_z, density),
    linewidth = 1,
    color = "grey40" #
  ) +

  stat_halfeye(
    data = posterior,
    aes(x = beta_z, y = after_stat(pdf)),
    fill = "steelblue",
    alpha = .6
  ) +

  geom_vline(
    xintercept = 0,
    linetype = 2
  ) +

  labs(
    x = expression(beta[z]),
    y = "density",
    title = "The data eliminate most of the plausible parameter space",
    subtitle = "Grey = prior; blue = marginal posterior"
  ) +

  theme_classic() 

the penalty isn’t an arbitrary punishment for having parameters; it emerges because the model has to average over all parameters.

No reliance of quadratic approximation

#| label: quap1

library(tidyr)
# Sample sizes and corresponding counts
dat <- tibble(
  n = c(9, 18, 36),
  dead = n * 6 / 9,
  alive = n - dead
)

dat <- tibble(
  n = c(9, 18, 36),
  dead = n * 6 / 9,
  alive = n * 3 / 9
) %>%
  mutate(
    p = dead / n,
    sd = sqrt(p * (1 - p) / n)
  )


# Create x values and calculate exact + quadratic approximation
plot_dat <- dat |> 
  crossing(x = seq(0, 1, length.out = 500)) |>
  mutate(
    exact = dbeta(x, dead + 1, alive + 1),
    exact = dbeta(x, dead + 1, alive + 1),
    #approx = dnorm(
    #  x,
    #  mean = (dead + 1) / (n + 2),
    #  sd = sqrt(
    #    ((dead + 1) * (alive + 1)) /
    #      ((n + 2)^2 * (n + 3))
    #  )
    #)
    approx = dnorm(x, mean = p, sd = sd)
  ) |>
  pivot_longer(
    cols = c(exact, approx),
    names_to = "distribution",
    values_to = "density"
  )

# Plot
plot_dat |> 
  ggplot(aes(x, density, linetype = distribution)) +
  geom_line(linewidth = 0.8) +
  facet_wrap(~ n, nrow = 1) +
  scale_linetype_manual(
    values = c(exact = "solid", approx = "dashed"),
    labels = c(
      exact = "Beta posterior",
      approx = "Normal approximation"
    )
  ) +
  labs(
    x = "Probability",
    y = "Density",
    linetype = NULL
  ) +
  theme_minimal()

No assumption of multivariate normal in GLM