Bayesian analysis with Stan

Bayesian fits use the same Stan likelihood as MLE but sample from the posterior with rstan::sampling.

Before you model

Brief exploratory context for the larynx data (see Getting started for full EDA):

library(spsurv)
library(generics)
library(KMsurv)
library(survival)
library(ggplot2)
data(larynx)
larynx$stage <- factor(larynx$stage)
km_stage <- survfit(Surv(time, delta) ~ stage, data = larynx)
km_long <- data.frame(
  time = km_stage$time,
  surv = km_stage$surv,
  stage = rep(levels(larynx$stage), km_stage$strata)
)
ggplot(km_long, aes(x = time, y = surv, color = stage)) +
  geom_step(linewidth = 0.6) +
  labs(x = "Time (years)", y = "Survival probability", color = "Stage") +
  theme_bw() +
  theme(legend.position = "bottom")
Kaplan-Meier by stage (unadjusted).
Kaplan-Meier by stage (unadjusted).

What to decide

Fit and summarize

fit <- bpph(
  Surv(time, delta) ~ age + stage,
  degree = 5,
  data = larynx,
  approach = "bayes",
  iter = 400,
  warmup = 200,
  chains = 1,
  cores = 1,
  init = 0,
  priors = list(
    beta = c("normal(0,4)"),
    gamma = c("lognormal(0,4)")
  )
)
#> 
#> SAMPLING FOR MODEL 'spbp' NOW (CHAIN 1).
#> Chain 1: 
#> Chain 1: Gradient evaluation took 6.5e-05 seconds
#> Chain 1: 1000 transitions using 10 leapfrog steps per transition would take 0.65 seconds.
#> Chain 1: Adjust your expectations accordingly!
#> Chain 1: 
#> Chain 1: 
#> Chain 1: Iteration:   1 / 400 [  0%]  (Warmup)
#> Chain 1: Iteration:  40 / 400 [ 10%]  (Warmup)
#> Chain 1: Iteration:  80 / 400 [ 20%]  (Warmup)
#> Chain 1: Iteration: 120 / 400 [ 30%]  (Warmup)
#> Chain 1: Iteration: 160 / 400 [ 40%]  (Warmup)
#> Chain 1: Iteration: 200 / 400 [ 50%]  (Warmup)
#> Chain 1: Iteration: 201 / 400 [ 50%]  (Sampling)
#> Chain 1: Iteration: 240 / 400 [ 60%]  (Sampling)
#> Chain 1: Iteration: 280 / 400 [ 70%]  (Sampling)
#> Chain 1: Iteration: 320 / 400 [ 80%]  (Sampling)
#> Chain 1: Iteration: 360 / 400 [ 90%]  (Sampling)
#> Chain 1: Iteration: 400 / 400 [100%]  (Sampling)
#> Chain 1: 
#> Chain 1:  Elapsed Time: 0.081 seconds (Warm-up)
#> Chain 1:                0.057 seconds (Sampling)
#> Chain 1:                0.138 seconds (Total)
#> Chain 1:
summary(fit)
#> Call:
#> bpph(formula = Surv(time, delta) ~ age + stage, degree = 5, data = larynx, 
#>     approach = "bayes", iter = 400, warmup = 200, chains = 1, 
#>     cores = 1, init = 0, priors = list(beta = c("normal(0,4)"), 
#>         gamma = c("lognormal(0,4)")), model = "ph")
#> 
#> Bayesian Bernstein PH model: 
#> Regression coefficients:
#>        Estimate    2.5%   97.5% Std. Error
#> age      0.0213 -0.0034  0.0453        0.0
#> stage2   0.2161 -0.6641  0.9496        0.4
#> stage3   0.6826 -0.0351  1.2380        0.4
#> stage4   1.7913  0.9005  2.4630        0.4
#> 
#> Exponentiated coefficients:
#>        Estimate 2.5% 97.5%
#> age        1.02 1.00   1.0
#> stage2     1.36 0.48   2.5
#> stage3     2.11 0.97   3.4
#> stage4     6.56 2.35  11.6
#> 
#> --- 
#> DIC = 293   WAIC = -148
fit_mle <- bpph(
  Surv(time, delta) ~ age + stage,
  degree = 5,
  data = larynx,
  approach = "mle",
  init = 0
)

Visual check — posterior forest plot

td_b <- tidy(fit, conf.int = TRUE, exponentiate = TRUE)
hpd <- credint(fit, prob = 0.95, type = "HPD")
td_b$hpd.low <- exp(hpd[, 1])
td_b$hpd.high <- exp(hpd[, 2])
td_b$term <- factor(td_b$term, levels = rev(td_b$term))

ggplot(td_b, aes(x = estimate, y = term)) +
  geom_vline(xintercept = 1, linetype = "dashed", color = "grey50") +
  geom_pointrange(aes(xmin = hpd.low, xmax = hpd.high), linewidth = 0.4, color = "steelblue") +
  labs(x = "Hazard ratio (posterior median)", y = NULL, title = "95% HPD intervals") +
  theme_bw()
Posterior medians with 95% HPD intervals (exponentiated).
Posterior medians with 95% HPD intervals (exponentiated).

Visual check — MLE vs Bayes

td_m <- tidy(fit_mle, conf.int = TRUE, exponentiate = TRUE)
cmp <- merge(
  td_m[, c("term", "estimate", "conf.low", "conf.high")],
  td_b[, c("term", "estimate", "hpd.low", "hpd.high")],
  by = "term",
  suffixes = c("_mle", "_bayes")
)
cmp_long <- rbind(
  data.frame(term = cmp$term, method = "MLE", low = cmp$conf.low, high = cmp$conf.high,
             est = cmp$estimate_mle),
  data.frame(term = cmp$term, method = "Bayes", low = cmp$hpd.low, high = cmp$hpd.high,
             est = cmp$estimate_bayes)
)
cmp_long$term <- factor(cmp_long$term, levels = rev(unique(cmp$term)))

ggplot(cmp_long, aes(x = est, y = term, color = method)) +
  geom_vline(xintercept = 1, linetype = "dashed", color = "grey50") +
  geom_pointrange(aes(xmin = low, xmax = high), position = position_dodge(width = 0.5), linewidth = 0.4) +
  labs(x = "Hazard ratio", y = NULL, color = NULL) +
  theme_bw()
MLE 95% CI vs Bayesian 95% HPD (exponentiated).
MLE 95% CI vs Bayesian 95% HPD (exponentiated).

Why this is insightful

Posterior intervals show whether data overwhelm priors. Agreement between MLE and Bayes supports data-driven conclusions; wide HPDs flag weak information.

What to decide

Credible intervals and criteria

credint(fit, prob = 0.95, type = "HPD")
#>               lower      upper
#> age    -0.003425295 0.04534822
#> stage2 -0.664082553 0.94959664
#> stage3 -0.035083416 1.23795738
#> stage4  0.900520035 2.46303145
#> attr(,"Probability")
#> [1] 0.95
sm <- summary(fit)
c(DIC = sm$dic, WAIC = sm$waic, LPML = sm$lpml)
#>       DIC      WAIC      LPML 
#>  292.7730 -147.7388 -147.7872

Convergence (optional deep dive)

Install bayesplot for trace and pairs plots:

library(bayesplot)
mcmc_pairs(fit$posterior$beta, off_diag_fun = "hex")

Always inspect divergences, split R-hat, and ESS before interpreting results.

tidybayes posterior draws

Install posterior and tidybayes for long-format draws and interval plots. After library(spsurv), spread_draws() and tidy_draws() dispatch on Bayes fits when tidybayes is loaded.

library(tidybayes)
dr <- as_draws_df.spbp(fit)
spread_draws(fit, `beta[age]`)
#> # A tibble: 200 x 4
#>    .chain .iteration .draw `beta[age]`
#>     <int>      <int> <int>       <dbl>
#>  1      1          1     1     0.0383 
#>  2      1          2     2     0.0196 
#>  3      1          3     3     0.0436 
#>  4      1          4     4     0.0226 
#>  5      1          5     5     0.0367 
#>  6      1          6     6     0.0318 
#>  7      1          7     7     0.0201 
#>  8      1          8     8    -0.00518
#>  9      1          9     9     0.0284 
#> 10      1         10    10     0.0192 
#> # i 190 more rows

Draw-level survival curves:

head(spread_surv_draws.spbp(fit, times = c(1, 2, 3), newdata = larynx[1, ]))
#>     stage time age diagyr delta .chain .iteration .draw time      surv id
#> 1       1  0.6  77     76     1      1          1     1    1 0.9288226  1
#> 1.1     1  0.6  77     76     1      1          1     1    2 0.8611402  1
#> 1.2     1  0.6  77     76     1      1          1     1    3 0.7826486  1
#> 1.3     1  0.6  77     76     1      1          2     2    1 0.9351974  1
#> 1.4     1  0.6  77     76     1      1          2     2    2 0.8845403  1
#> 1.5     1  0.6  77     76     1      1          2     2    3 0.8343337  1

See vignette("tidymodels", package = "spsurv") for workflows and parsnip engines.

Survival prediction

Posterior mean survival is a Bernstein polynomial in time, so the curve is smooth (geom_line). Kaplan–Meier in the EDA plot above remains a step function.

pr <- predict(fit, times = seq(0, max(larynx$time), length.out = 121))
head(pr)
#>   id       time      surv     lower     upper     cumhaz     std.err
#> 1  1 0.00000000 1.0000000 1.0000000 1.0000000 0.00000000 0.000000000
#> 2  1 0.08916667 0.9897779 0.9839384 0.9955287 0.01028049 0.003362817
#> 3  1 0.17833333 0.9797725 0.9681722 0.9903378 0.02045654 0.006454702
#> 4  1 0.26750000 0.9699671 0.9527796 0.9851377 0.03053899 0.009299796
#> 5  1 0.35666667 0.9603462 0.9385156 0.9800237 0.04053843 0.011920232
#> 6  1 0.44583333 0.9508951 0.9249226 0.9747961 0.05046516 0.014336244
ggplot(pr, aes(x = time, y = surv)) +
  geom_ribbon(aes(ymin = lower, ymax = upper), alpha = 0.2, colour = NA) +
  geom_line(linewidth = 0.6) +
  labs(x = "Time (years)", y = "Survival probability") +
  theme_bw()
Smooth posterior survival at the mean covariate profile with 95% HPD ribbon.
Smooth posterior survival at the mean covariate profile with 95% HPD ribbon.

See the Survival prediction and ggplot vignette for covariate-specific ribbons.

mirror server hosted at Truenetwork, Russian Federation.