--- title: "nbsurv Workflow" output: rmarkdown::html_vignette vignette: > %\VignetteIndexEntry{nbsurv Workflow} %\VignetteEngine{knitr::rmarkdown} %\VignetteEncoding{UTF-8} --- ```{r, include = FALSE} knitr::opts_chunk$set( collapse = TRUE, comment = "#>" ) ``` ## Overview `nbsurv` implements a conditional naive Bayes model for right-censored survival data. At each prediction horizon, the method treats survival past the horizon as a binary classification problem and combines: - a marginal Kaplan-Meier survival estimate, - inverse-probability of censoring weighting for the event class, - Gaussian likelihoods for continuous predictors, and - Laplace-smoothed categorical likelihoods. The package also includes utilities for horizon-specific evaluation, cross-validation, and hyper-parameter tuning. ## Fit a model ```{r} library(nbsurv) library(survival) lung <- stats::na.omit(lung) lung$status <- as.integer(lung$status == 2) fit <- nbsurv( Surv(time, status) ~ age + sex + ph.ecog, data = lung ) fit ``` ## Generate predictions ```{r} times <- c(100, 200, 400, 800) surv_pred <- predict( fit, newdata = lung[1:5, ], times = times ) event_pred <- predict( fit, newdata = lung[1:5, ], times = times, type = "event" ) surv_pred event_pred ``` The returned survival matrix is post-processed to be monotone in time. ## Evaluate the fitted model ```{r} metrics <- evaluate_nbsurv( fit, newdata = lung, times = times ) metrics ``` `evaluate_nbsurv()` reports horizon-specific IPCW Brier scores and concordance values. ## Cross-validation ```{r} cv_fit <- cv_nbsurv( Surv(time, status) ~ age + sex + ph.ecog, data = lung, folds = 3, times = times, seed = 1 ) cv_fit$summary ``` ## Tune hyper-parameters ```{r} grid <- data.frame( scale = c(TRUE, FALSE), laplace = c(1, 2), min_sd = c(0.05, 0.10) ) grid$time_grid <- I(list(NULL, NULL)) tuned <- tune_nbsurv( Surv(time, status) ~ age + sex + ph.ecog, data = lung, param_grid = grid, folds = 3, times = c(100, 200, 400), seed = 1 ) tuned$results tuned$best_params ``` ## Plot fitted curves ```{r} plot(fit, times = c(100, 200, 400), n_curves = 3) ```