## ----include = FALSE----------------------------------------------------------
knitr::opts_chunk$set(
  collapse = TRUE,
  comment = "#>",
  fig.width = 7,
  fig.height = 4.5
)

## -----------------------------------------------------------------------------
library(exnexSurv)
library(survival)

simulate_surv_data <- function(
  theta,
  sigma2,
  beta = NULL,
  n_per_group = 50,
  censor_min = 4,
  censor_max = 12,
  seed = NULL
) {
  if (!is.null(seed)) {
    set.seed(seed)
  }

  groups <- rep(seq_along(theta), each = n_per_group)
  age_std <- rnorm(length(groups), mean = 0, sd = 1)
  mean_log_time <- rep(theta, each = n_per_group)

  if (!is.null(beta)) {
    mean_log_time <- mean_log_time + beta * age_std
  }

  log_time <- rnorm(length(groups), mean = mean_log_time, sd = sqrt(sigma2))
  true_time <- exp(log_time)
  censor_time <- runif(length(groups), min = censor_min, max = censor_max)

  data.frame(
    time = pmin(true_time, censor_time),
    event = as.integer(true_time <= censor_time),
    group = factor(groups),
    age_std = age_std
  )
}

sim_data <- simulate_surv_data(
  theta = c(1.1, 1.6, 2.0),
  sigma2 = 0.25,
  beta = -0.30,
  n_per_group = 50,
  seed = 6421
)

head(sim_data)
mean(sim_data$event)

## -----------------------------------------------------------------------------
fit_formula <- exnex_surv(
  Surv(time, event) ~ group + age_std,
  data = sim_data,
  iter = 1200,
  warmup = 400,
  chains = 1,
  seed = 6421
)

print(fit_formula, show_trace = FALSE)
summary(fit_formula)

## -----------------------------------------------------------------------------
fit_parallel <- exnex_surv(
  Surv(time, event) ~ group + age_std,
  data = sim_data,
  iter = 1200,
  warmup = 400,
  chains = 2,
  parallel_chains = 2,
  seed = 6421
)

print(fit_parallel, show_trace = FALSE)
plot(
  fit_parallel,
  parameters = c("theta_1", "theta_2", "theta_3", "beta_1", "sigma2"),
  ask = FALSE
)

## -----------------------------------------------------------------------------
str(fit_formula$data)
plot(
  fit_formula,
  parameters = c("theta_1", "theta_2", "sigma2"),
  ask = FALSE
)

## -----------------------------------------------------------------------------
fit_xy <- exnex_surv(
  x = sim_data[c("group", "age_std")],
  y = Surv(sim_data$time, sim_data$event),
  iter = 1200,
  warmup = 400,
  chains = 1,
  seed = 6421
)

summary(fit_xy)
all.equal(fit_formula$draws, fit_xy$draws)

## -----------------------------------------------------------------------------
fit_no_cov <- exnex_surv(
  Surv(time, event) ~ group,
  data = sim_data,
  iter = 1200,
  warmup = 400,
  chains = 1,
  seed = 6421
)

summary(fit_no_cov)
plot(fit_no_cov, ask = FALSE)

## -----------------------------------------------------------------------------
summary(fit_formula)

## -----------------------------------------------------------------------------
fit_formula$resolved_priors

