You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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:

  1. Fit a stump on the current node's dataset
  2. Check if the resulting children have outcome differences above your threshold
  3. If yes, keep the split and recurse on each child
  4. If no, discard the stump, fit the next best stump on the same dataset, and repeat
  5. Stop when no valid splits are left or you hit minimum sample size

内容的提问来源于stack exchange,提问作者user2173836

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.15 09:55:19