R/minNodeSizePruningCompRisks.R

Defines functions minNodePruningCompRisks

Documented in minNodePruningCompRisks

#' Minimal Node Size Pruning For Competing Risks
#'
#' Computes optimal minimal node size of a discrete survival tree 
#' from a given vector of possible node sizes by cross-validation. 
#' Laplace-smoothing can be applied to the estimated hazards.
#' 
#' @param formulaVariable Model formula for tree fitting (class "formula") of the form "~ x1 + x2 + ..." without response.
#' @param data Discrete survival data in short format for which a survival tree is
#' to be fitted (class "data.frame").
#' @param treetype Type of tree to be fitted. Possible values are "rpart" or "ranger" (class "character"). The default
#' is to fit an rpart tree; when "ranger" is chosen, a ranger forest with a single tree is fitted.
#' @param splitruleranger String specifying the splitting rule of the ranger tree (class "character"). 
#' Possible values are either "gini" or "extratrees". Default is "gini".
#' @param sizes Vector of different node sizes to try (class "integer"). 
#' Values need to be non-negative.
#' @param indexList List of data partitioning indices for cross-validation (class "list").
#' Each element represents the test indices of one fold (class "integer").
#' @param timeColumn Character giving the column name of the observed times in
#' the "data"-argument (class "character").
#' @param eventColumns Character vector giving the column names of the event
#' indicators (excluding censoring column) in the "data"-argument (class "character").
#' @param alpha Parameter for laplace-smoothing. A value of 0 corresponds to 
#' no laplace-smoothing (class "numeric").
#' @param logOut Logical value (class "logical"). If True, computation progress will be written to
#' console.
#' @param eventColumnsAsFactor Should the argument eventColumns be intepreted
#' as column name of a factor variable (class "logical")? Default is FALSE.
#' @param ... Additional arguments to the estimation function. It is either "rpart" 
#' or "ranger" (see argument \emph{treetype}). 
#' @details Computes the out-of-sample log likelihood for all data partitionings
#' for each node size in \emph{sizes} and returns the node size for which the log 
#' likelihood was minimal. Also returns an rpart tree with the optimal minimal 
#' node size using the entire data set.
#' @note Note that depending on argument \emph{treetype} some arguments are fixed
#' and can not be changed:
#' \itemize{
#'   \item \emph{treetype}="rpart": formula, data, method, minbucket
#'   \item \emph{treetype}="ranger": formula, data, num.trees, mtry, 
#'   classification, splitrule, replace, sample.fraction, min.node.size
#' }
#' @return A list containing the two items
#' \itemize{
#'   \item OptimNodeSize - Node size with lowest out-of-sample log-likelihood
#'   \item OptimTree - A tree object with type corresponding to \emph{treetype} argument with the optimal minimal node size
#' }
#' @examples
#' # Example unemployment data
#' library(Ecdat)
#' library(caret)
#' data(UnempDur)
#' 
#' # Select training and testing subsample
#' subUnempDur <- UnempDur[which(UnempDur$spell < 10),]
#' subUnempDur <- subUnempDur[1:250,]
#' 
#' # Creating status variable for data partitioning
#' subUnempDur$status <- ifelse(subUnempDur$censor1, 1, 
#' ifelse(subUnempDur$censor2, 2, ifelse(
#' subUnempDur$censor3, 3, ifelse(subUnempDur$censor4, 4, 0))))
#' 
#' # Create cross validation sets
#' # Stratified by events and time distribution
#' set.seed(1972)
#' indexList <- createFolds(factor(paste(subUnempDur$status, 
#' subUnempDur$spell, sep="_")), k = 5)
#' 
#' # Perform minimal node size pruning
#' formula1 <- ~ timeInt + age + logwage
#' sizes <- 1:10
#' timeColumn <- "spell"
#' eventColumns <- c("censor1", "censor2", "censor3","censor4")
#' optiTree <- minNodePruningCompRisks(formula1, subUnempDur, treetype = "rpart", sizes = sizes, 
#' indexList = indexList, timeColumn = timeColumn, eventColumns = eventColumns, alpha = 1, 
#' logOut = TRUE)
#' plot(optiTree)
#' 
#' @export minNodePruningCompRisks
minNodePruningCompRisks <- function(formulaVariable, data, treetype = "rpart", splitruleranger = "gini", sizes, indexList, 
                                  timeColumn, eventColumns, alpha = 1, logOut = FALSE, 
                                  eventColumnsAsFactor=FALSE, ...)
{
  
  # Construct formula
  constructFormula <- formula(paste("responses ~", 
                                    paste(attr(terms(formulaVariable),"term.labels"), 
                                          collapse=" + "), sep = " "))
  
  #inputchecks
  if (!treetype %in% c("rpart", "ranger"))
  {
    stop("treetype must be either \"rpart\" or \"ranger\".")
  }
  mean_total_llh <- rep(NA, length(sizes))
  for (iNode in 1:length(sizes))
  {
    total_llh <- rep(NA, length(indexList))
    for (iTrainIndex in 1:length(indexList))
    {
      dataTrain <- data[-indexList[[iTrainIndex]], ]
      dataTest <- data[indexList[[iTrainIndex]], ]
      dataTrainLong <- dataLongCompRisks(dataTrain, timeColumn, eventColumns, 
                                         responseAsFactor = TRUE, 
                                         eventColumnsAsFactor=eventColumnsAsFactor)
      dataTrainLong$y <- as.numeric(factor(dataTrainLong$responses)) - 1
      dataTestLong <- dataLongCompRisks(dataTest, timeColumn, eventColumns, responseAsFactor = TRUE, 
                                        eventColumnsAsFactor=eventColumnsAsFactor)
      dataTestLong$y <- as.numeric(factor(dataTestLong$responses)) - 1
      if(treetype == "ranger")
      {
        tree <- ranger(constructFormula, dataTrainLong, num.trees = 1, mtry = length(attr(terms(constructFormula), "term.labels")),
                      classification = TRUE, splitrule = splitruleranger, replace = FALSE, 
                      sample.fraction = 1, min.node.size = sizes[iNode], ...)
        test_hazards <- survTreeLaplaceHazard(tree, dataTestLong, alpha, dataTrainLong)
      } else
      {
        tree <- rpart(constructFormula, dataTrainLong, method = "class", minbucket = sizes[iNode], ...)
        test_hazards <- survTreeLaplaceHazard(tree, dataTestLong, alpha)
      }
      lh <- test_hazards[dataTestLong$y*nrow(test_hazards)+c(1:nrow(test_hazards))]
      llh <- -1 * log(lh)
      # llh[which(is.infinite(llh))] = 10
      total_llh[iTrainIndex] = sum(llh)
    }
    mean_total_llh[iNode] <- mean(total_llh)
    if(logOut)
    {
      cat('\r', iNode/length(sizes)*100,"% finished")
      flush.console()
    }
  }
  selectSize1 <- max(sizes[mean_total_llh==min(mean_total_llh)])
  optimalNodeSize <- sizes[sizes==selectSize1]
  attr(optimalNodeSize, "llh") <- data.frame(sizes, mean_total_llh)
  dataLong = dataLongCompRisks(data, timeColumn, eventColumns, 
                               responseAsFactor = TRUE, 
                               eventColumnsAsFactor=eventColumnsAsFactor)
  if(treetype == "ranger")
  {
    optimalTree <- ranger(constructFormula, dataLong, num.trees = 1, mtry = length(attr(terms(constructFormula), "term.labels")),
                         classification = TRUE, splitrule = splitruleranger, replace = FALSE, 
                         sample.fraction = 1, min.node.size = optimalNodeSize)
  } else
  {
    optimalTree <- rpart(constructFormula, dataLong, method = "class", minbucket = optimalNodeSize)
  }
  optimalTree <- rpart(constructFormula, dataLong, method = "class", minbucket = optimalNodeSize)
  
  RES <- list("OptimNodeSize" = optimalNodeSize,
              "OptimTree" = optimalTree)
  class(RES) <- "discSurvMinNodeSizePrune"
  return(RES)
}

Try the discSurv package in your browser

Any scripts or data that you put into this service are public.

discSurv documentation built on April 29, 2026, 9:07 a.m.