Nothing
#' Model Comparison, Bootstrapped WAIC, and Cross-Validation for MultiFrailty Models
#'
#' Compares multiple fitted shared frailty models using AIC, BIC, AICc, HQIC,
#' optional bootstrapped WAIC, and k-fold cross-validation.
#'
#' @param ... Fitted \code{"multifrailty_fit"} objects passed individually or as a list.
#' @param criteria Character vector of selection criteria to include.
#' @param compute_waic Logical; if TRUE computes bootstrapped WAIC. Default is FALSE.
#' @param waic_B Number of bootstrap samples for WAIC. Default is 100.
#'
#' @return A data frame summarizing model comparisons ranked by AIC.
#'
#' @references
#' Pandey, A., Hanagal, D. D., & Tyagi, S. (2022). Shared Frailty Models Based on Cancer Data. International Journal of Statistics and Reliability Engineering, 9(3), 461-474.
#'
#' Pandey, A., & Tyagi, S. (2021). Comparison of Multiplicative Frailty Models Under Weibull Baseline Distribution. Lobachevskii Journal of Mathematics, 42(13), 3184-3195.
#'
#' @export
#' @examples
#' set.seed(123)
#' dat <- r_frailty(n = 60, baseline = "weibull", bpar = c(2, 1.5), frailty = "gamma", fpar = c(0.8))
#' fit1 <- fit_frailty(time = dat$time, status = dat$status, baseline = "weibull", frailty = "none")
#' fit2 <- fit_frailty(time = dat$time, status = dat$status, baseline = "weibull", frailty = "gamma")
#' cmp <- compare_models(fit1, fit2)
#' print(cmp)
compare_models <- function(..., criteria = c("AIC", "BIC", "AICc", "HQIC"), compute_waic = FALSE, waic_B = 100L) {
objs <- list(...)
if (length(objs) == 1 && is.list(objs[[1]]) && !inherits(objs[[1]], "multifrailty_fit")) {
objs <- objs[[1]]
}
n_models <- length(objs)
if (n_models == 0) stop("No models provided for comparison.")
mod_names <- names(objs)
if (is.null(mod_names)) mod_names <- paste0("Model_", 1:n_models)
res_df <- data.frame(
Model = mod_names,
Baseline = character(n_models),
Frailty = character(n_models),
logLik = numeric(n_models),
K = integer(n_models),
AIC = numeric(n_models),
BIC = numeric(n_models),
AICc = numeric(n_models),
HQIC = numeric(n_models),
FrailtyVar = numeric(n_models),
stringsAsFactors = FALSE
)
if (compute_waic) res_df$WAIC <- numeric(n_models)
for (i in 1:n_models) {
fit <- objs[[i]]
if (!inherits(fit, "multifrailty_fit")) stop(paste0("Item ", i, " is not a 'multifrailty_fit' object."))
k_total <- fit$n_par_base + fit$n_par_frailty + fit$n_cov
res_df$Baseline[i] <- fit$baseline
res_df$Frailty[i] <- fit$frailty
res_df$logLik[i] <- fit$logLik
res_df$K[i] <- k_total
res_df$AIC[i] <- fit$AIC
res_df$BIC[i] <- fit$BIC
res_df$AICc[i] <- fit$AICc
res_df$HQIC[i] <- fit$HQIC
res_df$FrailtyVar[i] <- fit$frailty_var
if (compute_waic) {
res_df$WAIC[i] <- bootstrap_waic(fit, B = waic_B)
}
}
res_df <- res_df[order(res_df$AIC), ]
rownames(res_df) <- NULL
res_df
}
#' Bootstrapped WAIC Computation
#' @param fit Fitted \code{multifrailty_fit} object.
#' @param B Number of bootstrap replications. Default is 200.
#' @return Numeric value of bootstrapped WAIC.
#' @export
bootstrap_waic <- function(fit, B = 200) {
n <- fit$n
x <- fit$x
time <- fit$time
status <- fit$status
b_type <- fit$baseline
f_type <- fit$frailty
# Compute pointwise log-likelihood matrix under bootstrap samples
log_lik_boot <- matrix(NA_real_, nrow = n, ncol = B)
for (b in 1:B) {
idx <- sample(1:n, replace = TRUE)
fit_b <- tryCatch({
fit_frailty(time = time[idx], status = status[idx], x = x[idx, , drop = FALSE],
baseline = b_type, frailty = f_type)
}, error = function(e) NULL)
if (!is.null(fit_b) && fit_b$converged) {
# Evaluate pointwise log-likelihood of original data under fit_b
p_est <- fit_b$raw_est
for (i in 1:n) {
log_lik_boot[i, b] <- loglik_frailty(p_est, time = time[i], status = status[i],
x = x[i, , drop = FALSE], baseline = b_type, frailty = f_type)
}
}
}
# p_waic = sum of sample variances across bootstrap draws
p_waic <- sum(apply(log_lik_boot, 1, stats::var, na.rm = TRUE), na.rm = TRUE)
lppd <- sum(log(rowMeans(exp(log_lik_boot), na.rm = TRUE)), na.rm = TRUE)
waic_val <- -2 * lppd + 2 * p_waic
waic_val
}
#' K-Fold Cross Validation for MultiFrailty Models
#' @param time Time vector.
#' @param status Status vector.
#' @param x Covariate matrix.
#' @param baseline Baseline distribution.
#' @param frailty Frailty distribution.
#' @param k Number of folds. Default is 5L.
#' @param time2 Optional upper interval bound vector.
#' @return Out-of-sample total log-likelihood score.
#' @export
cv_frailty <- function(time, status, x = matrix(nrow = length(time), ncol = 0),
baseline = "weibull", frailty = "gamma", k = 5L, time2 = NULL) {
n <- length(time)
folds <- sample(rep(1:k, length.out = n))
cv_loglik <- 0.0
for (j in 1:k) {
idx_train <- (folds != j)
idx_test <- (folds == j)
fit_k <- tryCatch({
fit_frailty(time = time[idx_train], status = status[idx_train],
x = x[idx_train, , drop = FALSE], baseline = baseline, frailty = frailty,
time2 = if (!is.null(time2)) time2[idx_train] else NULL)
}, error = function(e) NULL)
if (!is.null(fit_k) && fit_k$converged) {
ll_test <- loglik_frailty(fit_k$raw_est, time = time[idx_test], status = status[idx_test],
x = x[idx_test, , drop = FALSE], baseline = baseline, frailty = frailty,
time2 = if (!is.null(time2)) time2[idx_test] else NULL)
if (is.finite(ll_test) && ll_test > -1e10) {
cv_loglik <- cv_loglik + ll_test
}
}
}
cv_loglik
}
Any scripts or data that you put into this service are public.
Add the following code to your website.
For more information on customizing the embed code, read Embedding Snippets.