## ----include = FALSE----------------------------------------------------------
knitr::opts_chunk$set(
  collapse = TRUE,
  comment = "#>",
  fig.width = 7,
  fig.height = 4.5,
  warning = FALSE,
  message = FALSE
)

## -----------------------------------------------------------------------------
library(exnexSurv)
library(survival)

simulate_surv_data <- function(
  theta,
  sigma2,
  beta = NULL,
  n_per_group = 40,
  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_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.0, 1.6, 2.1),
  sigma2 = 0.25,
  beta = -0.30,
  seed = 9201
)

## -----------------------------------------------------------------------------
fit_parallel <- exnex_surv(
  Surv(time, event) ~ group + age_std,
  data = sim_data,
  iter = 60,
  warmup = 20,
  chains = 2,
  parallel_chains = 2,
  seed = 9201
)

print(fit_parallel, show_trace = FALSE)

## -----------------------------------------------------------------------------
fit_four_chains <- exnex_surv(
  Surv(time, event) ~ group + age_std,
  data = sim_data,
  iter = 60,
  warmup = 20,
  chains = 4,
  parallel_chains = 2,
  seed = 9201
)

print(fit_four_chains, show_trace = FALSE)

## -----------------------------------------------------------------------------
plot(
  fit_parallel,
  parameters = c("theta_1", "theta_2", "theta_3", "beta_1", "sigma2"),
  ask = FALSE
)

## -----------------------------------------------------------------------------
fit_sequential <- exnex_surv(
  Surv(time, event) ~ group + age_std,
  data = sim_data,
  iter = 60,
  warmup = 20,
  chains = 2,
  parallel_chains = 1,
  seed = 9201
)

identical(fit_parallel$draws, fit_sequential$draws)

