## ----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(KMsurv)
library(survival)
library(ggplot2)
data(larynx)
larynx$stage <- factor(larynx$stage)

## ----eda-km, fig.cap = "Kaplan-Meier by stage with censoring marks (+)."------
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) +
  geom_point(
    data = subset(larynx, delta == 0),
    aes(x = time, y = 0, color = stage),
    shape = 3,
    size = 1.5,
    alpha = 0.5,
    inherit.aes = FALSE
  ) +
  labs(x = "Time (years)", y = "Survival probability", color = "Stage") +
  theme_bw() +
  theme(legend.position = "bottom")

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

## ----times--------------------------------------------------------------------
plot_times <- seq(0, max(larynx$time), length.out = 121)
newdata <- data.frame(age = 70, stage = levels(larynx$stage))
newdata

## ----predict------------------------------------------------------------------
pr <- predict(fit, newdata = newdata, times = plot_times)
head(pr)

## ----km-model-overlay, fig.cap = "Observed KM (steps) vs smooth Bernstein prediction at age 70 (dashed)."----
pr$stage <- newdata$stage[match(as.character(pr$id), as.character(seq_len(nrow(newdata))))]

ggplot() +
  geom_step(
    data = km_long,
    aes(x = time, y = surv, color = stage),
    linewidth = 0.5,
    alpha = 0.8
  ) +
  geom_line(
    data = pr,
    aes(x = time, y = surv, color = stage),
    linetype = "dashed",
    linewidth = 0.6
  ) +
  labs(x = "Time (years)", y = "Survival probability", color = "Stage") +
  theme_bw() +
  theme(legend.position = "bottom")

## ----ggplot, fig.cap = "Smooth stage-specific Bernstein survival with 95% intervals."----
ggplot(pr, aes(x = time, y = surv, color = stage, fill = stage)) +
  geom_ribbon(aes(ymin = lower, ymax = upper), alpha = 0.2, colour = NA) +
  geom_line(linewidth = 0.5) +
  labs(x = "Time (years)", y = "Survival probability", color = NULL, fill = NULL) +
  theme_bw() +
  theme(legend.position = "bottom")

## ----survfit-tidy-------------------------------------------------------------
sf <- survfit(fit, newdata = newdata, times = plot_times, tidy = TRUE)
head(sf)

## ----bayes-note, eval = FALSE-------------------------------------------------
# fit_bayes <- bpph(
#   Surv(time, delta) ~ age + stage,
#   degree = 5,
#   data = larynx,
#   approach = "bayes",
#   iter = 400,
#   chains = 1,
#   cores = 1
# )
# predict(fit_bayes, newdata = newdata, times = plot_times, interval.type = "hpd")

## ----interval-type------------------------------------------------------------
loglog_times <- plot_times[plot_times > 0]
head(predict(fit, newdata = newdata, times = loglog_times, type = "log-log"))

