## ----include = FALSE----------------------------------------------------------
knitr::opts_chunk$set(
  collapse = TRUE,
  comment = "#>",
  fig.width = 7,
  fig.height = 4.5,
  dpi = 96,
  message = FALSE,
  warning = FALSE
)

## ----setup--------------------------------------------------------------------
library(spsurv)
library(generics)
library(KMsurv)
library(survival)
library(ggplot2)
data(larynx)
larynx$stage <- factor(larynx$stage)

## ----eda-km, fig.cap = "Kaplan-Meier by stage (unadjusted)."------------------
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-bayes, cache = TRUE--------------------------------------------------
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)")
  )
)
summary(fit)

## ----fit-mle------------------------------------------------------------------
fit_mle <- bpph(
  Surv(time, delta) ~ age + stage,
  degree = 5,
  data = larynx,
  approach = "mle",
  init = 0
)

## ----forest-bayes, fig.cap = "Posterior medians with 95% HPD intervals (exponentiated)."----
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()

## ----mle-bayes-compare, fig.cap = "MLE 95% CI vs Bayesian 95% HPD (exponentiated)."----
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()

## ----credint------------------------------------------------------------------
credint(fit, prob = 0.95, type = "HPD")

## ----criteria-----------------------------------------------------------------
sm <- summary(fit)
c(DIC = sm$dic, WAIC = sm$waic, LPML = sm$lpml)

## ----stanfit, eval = FALSE----------------------------------------------------
# library(bayesplot)
# mcmc_pairs(fit$posterior$beta, off_diag_fun = "hex")

## ----tidybayes, eval = requireNamespace("tidybayes", quietly = TRUE) && requireNamespace("posterior", quietly = TRUE)----
library(tidybayes)
dr <- as_draws_df.spbp(fit)
spread_draws(fit, `beta[age]`)

## ----spread-surv, eval = requireNamespace("tidybayes", quietly = TRUE)--------
head(spread_surv_draws.spbp(fit, times = c(1, 2, 3), newdata = larynx[1, ]))

## ----predict-bayes, fig.cap = "Smooth posterior survival at the mean covariate profile with 95% HPD ribbon."----
pr <- predict(fit, times = seq(0, max(larynx$time), length.out = 121))
head(pr)
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()

