Bayesian fits use the same Stan likelihood as MLE
but sample from the posterior with rstan::sampling.
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")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 = -148td_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()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()Posterior intervals show whether data overwhelm priors. Agreement between MLE and Bayes supports data-driven conclusions; wide HPDs flag weak information.
adapt_delta, more iter).Install bayesplot for trace and pairs plots:
Always inspect divergences, split R-hat, and ESS before interpreting results.
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 rowsDraw-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 1See vignette("tidymodels", package = "spsurv") for
workflows and parsnip engines.
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()See the Survival prediction and ggplot vignette for covariate-specific ribbons.