Decision Tree Algorithm Overview
Decision trees are supervised learning models known for their interpretability and versatility. They are widely applied in finance, healthcare, education, and consumer technology. The algorithm produces a tree-like structure where internal nodes represent feature tests and leaf nodes represent class labels.
Algorithm Fundamentals
Decision trees operate by recursively partitioning data based on feature values. The key considerations include:
Feature Selection Metrics
- Information Gain: Measures reduction in entropy after splitting
- Gain Ratio: Adjusts information gain to prevent bias toward multi-valued features
- Gini Index: Measures impurity reduction using Gini impurity
Tree Pruning Techniques
Pruning prevents overfitting by removing unreliable branches. Common control parameters:
- Maximum tree depth (max_depth)
- Minimum samples for split (min_samples_split)
- Minimum samples per leaf (min_samples_leaf)
Practical Implementation
Using decision trees to classify Scooby-Doo monster authenticity:
Data Preparation
library(tidyverse)
library(tidymodels)
monster_data <- read_csv("https://raw.githubusercontent.com/rfordatascience/tidytuesday/master/data/2021/2021-07-13/scoobydoo.csv") %>%
filter(monster_amount > 0) %>%
mutate(
imdb_score = parse_number(imdb),
air_year = lubridate::year(date_aired),
monster_status = if_else(monster_real == "FALSE", "fake", "real") %>% factor()
) %>%
select(air_year, imdb_score, monster_status, title)
Model Configuraton
set.seed(123)
data_split <- initial_split(monster_data, strata = monster_status)
train_data <- training(data_split)
test_data <- testing(data_split)
tree_model <- decision_tree(
cost_complexity = tune(),
tree_depth = tune(),
min_n = tune()
) %>%
set_mode("classification") %>%
set_engine("rpart")
Hyperparameter Tuning
param_grid <- grid_regular(cost_complexity(), tree_depth(), min_n(), levels = 4)
set.seed(456)
tuned_tree <- tune_grid(
tree_model,
monster_status ~ air_year + imdb_score,
resamples = bootstraps(train_data, strata = monster_status),
grid = param_grid,
metrics = metric_set(accuracy, roc_auc, sensitivity, specificity)
)
Model Evaluation
best_params <- select_best(tuned_tree, metric = "roc_auc")
final_model <- finalize_model(tree_model, best_params) %>%
fit(monster_status ~ air_year + imdb_score, train_data)
test_results <- last_fit(final_model, monster_status ~ air_year + imdb_score, data_split)
collect_metrics(test_results)
Visualization
library(parttree)
train_data %>%
ggplot(aes(imdb_score, air_year)) +
geom_parttree(data = final_model, aes(fill = monster_status), alpha = 0.2) +
geom_point(aes(color = monster_status), alpha = 0.7, position = position_jitter(width = 0.05, height = 0.2))