如何在R中实现决策树节点合并(适配Yan 2004及决策流方法)
Solution for Adaptive Tree with Node Merging
Approach 1: Post-Hoc Merging with partykit
This method builds a full decision tree first, then traverses the tree to merge terminal sibling nodes where the target variable difference falls below your threshold.
Step-by-Step Code
set.seed(1) n_subjects <- 100 n_items <- 4 responses <- matrix(rep(c(0,1), times=(n_subjects/2)*n_items), ncol=n_items) responses <- as.data.frame(apply(responses, 2, function(x) sample(x))) weights <- c(20,20,20,10) responses$outcome <- rowSums(responses[,1:n_items] * weights) library(partykit) # Build initial full tree ct <- ctree(outcome~., data=responses) # Function to traverse and merge sibling terminal nodes traverse_merge <- function(node, threshold=1) { if (is.terminal(node)) return(node) left_child <- node$children[[1]] right_child <- node$children[[2]] # Check if both children are terminal nodes if (is.terminal(left_child) && is.terminal(right_child)) { left_mean <- left_child$info$prediction right_mean <- right_child$info$prediction # Merge if outcome difference is below threshold if (abs(left_mean - right_mean) <= threshold) { # Convert parent node to terminal with combined mean node$children <- NULL combined_n <- left_child$info$n + right_child$info$n node$info$prediction <- (left_mean * left_child$info$n + right_mean * right_child$info$n) / combined_n node$info$n <- combined_n return(node) } } # Recursively process child nodes node$children[[1]] <- traverse_merge(left_child, threshold) node$children[[2]] <- traverse_merge(right_child, threshold) return(node) } # Apply merging with threshold=1 merged_ct <- traverse_merge(ct, threshold=1) # Plot the merged tree plot(merged_ct) # Predict with merged tree predictions <- predict(merged_ct, newdata=responses)
Approach 2: Custom Recursive Tree with Merge-Before-Split Logic
This method follows your exact workflow: split nodes using CART logic first, check if children need merging, and if so, discard the split and try the next best split before proceeding recursively.
Step-by-Step Code
set.seed(1) n_subjects <- 100 n_items <- 4 responses <- matrix(rep(c(0,1), times=(n_subjects/2)*n_items), ncol=n_items) responses <- as.data.frame(apply(responses, 2, function(x) sample(x))) weights <- c(20,20,20,10) responses$outcome <- rowSums(responses[,1:n_items] * weights) # Function to find best CART split (minimizes RSS) find_best_split <- function(data, target_col="outcome") { predictors <- setdiff(colnames(data), target_col) best_split <- NULL min_rss <- Inf for (var in predictors) { # Handle binary predictors (0/1) split_val <- 0.5 left <- data[data[[var]] <= split_val, target_col] right <- data[data[[var]] > split_val, target_col] rss <- sum((left - mean(left))^2) + sum((right - mean(right))^2) if (rss < min_rss) { min_rss <- rss best_split <- list( var=var, split_val=split_val, left_mean=mean(left), right_mean=mean(right), rss=rss ) } } return(best_split) } # Function to find next best split (exclude a variable) find_next_best_split <- function(data, target_col="outcome", exclude_var=NULL) { predictors <- setdiff(colnames(data), c(target_col, exclude_var)) if (length(predictors) == 0) return(NULL) best_split <- NULL min_rss <- Inf for (var in predictors) { split_val <- 0.5 left <- data[data[[var]] <= split_val, target_col] right <- data[data[[var]] > split_val, target_col] rss <- sum((left - mean(left))^2) + sum((right - mean(right))^2) if (rss < min_rss) { min_rss <- rss best_split <- list( var=var, split_val=split_val, left_mean=mean(left), right_mean=mean(right), rss=rss ) } } return(best_split) } # Recursive tree building with merge logic build_merged_tree <- function(data, target_col="outcome", threshold=1, min_samples=5) { # Stop conditions: too few samples or pure node if (nrow(data) <= min_samples || length(unique(data[[target_col]])) == 1) { return(list( type="leaf", mean_outcome=mean(data[[target_col]]), n=nrow(data) )) } # Find best split best_split <- find_best_split(data, target_col) if (is.null(best_split)) { return(list( type="leaf", mean_outcome=mean(data[[target_col]]), n=nrow(data) )) } # Check if merge is needed if (abs(best_split$left_mean - best_split$right_mean) <= threshold) { # Try next best split next_split <- find_next_best_split(data, target_col, exclude_var=best_split$var) if (is.null(next_split)) { return(list( type="leaf", mean_outcome=mean(data[[target_col]]), n=nrow(data) )) } else { # Split data with next best split left_data <- data[data[[next_split$var]] <= next_split$split_val, ] right_data <- data[data[[next_split$var]] > next_split$split_val, ] # Recurse on children return(list( type="node", var=next_split$var, split_val=next_split$split_val, left=build_merged_tree(left_data, target_col, threshold, min_samples), right=build_merged_tree(right_data, target_col, threshold, min_samples), n=nrow(data) )) } } else { # Proceed with best split left_data <- data[data[[best_split$var]] <= best_split$split_val, ] right_data <- data[data[[best_split$var]] > best_split$split_val, ] return(list( type="node", var=best_split$var, split_val=best_split$split_val, left=build_merged_tree(left_data, target_col, threshold, min_samples), right=build_merged_tree(right_data, target_col, threshold, min_samples), n=nrow(data) )) } } # Build merged tree merged_tree <- build_merged_tree(responses, threshold=1) # Prediction function for custom tree predict_merged_tree <- function(tree, new_data) { predictions <- numeric(nrow(new_data)) for (i in 1:nrow(new_data)) { current_node <- tree while (current_node$type == "node") { var_val <- new_data[i, current_node$var] current_node <- if (var_val <= current_node$split_val) current_node$left else current_node$right } predictions[i] <- current_node$mean_outcome } return(predictions) } # Generate predictions custom_predictions <- predict_merged_tree(merged_tree, responses)
Handling Tree Stumps for Prediction
If you choose to use repeated tree stumps (single splits), you can chain them recursively:
- Fit a stump on the current node's dataset
- Check if the resulting children have outcome differences above your threshold
- If yes, keep the split and recurse on each child
- If no, discard the stump, fit the next best stump on the same dataset, and repeat
- Stop when no valid splits are left or you hit minimum sample size
内容的提问来源于stack exchange,提问作者user2173836
相关产品推荐
相关产品推荐

