Stratified sampling using createDataPartition drops small classes out of test

Viewed 139

I'm trying to do stratified sampling, and I realized that when I have classes with very few cases, I can end up with a test data set that has not a single case of these minority classes.

Here is some example code

library(caret)

# data set for debugging in RStudio
data("imports85")
input<-imports85
    
# settings
set.seed(1)
dependent <- make.names("make")
training.share <- 0.75
impute <- "no"
type <- "classification"

# save original column names for later and make R-friendly column names
original.names <- names(input)
names(input) <- make.names(original.names)
    
# create train and test data sets
input.labelled <- input[complete.cases(input[,dependent]),] #split off rows w/o dependent
if (impute=="no") { 
    input.clean <- input.labelled[complete.cases(input.labelled),] #drop cases w/ missing variables
} else if (impute=="yes") {
    input.clean <- rfImpute(input.labelled[,dependent] ~ .,input.labelled)[,-1] #or impute missing variables and remove added duplicate of dependent column
}

train.index <- createDataPartition(input.clean[,dependent], p=training.share, list=FALSE) #create row index for train data set using stratified sampling but very small classes might all go into train?!
rf.train <- input.clean[train.index,] #create train data set
rf.test <- input.clean[-train.index,] #create test data set from left-overs
if (type=="classification") { #balance train data set for classification (can be skipped if upsampling takes place as part of tuning settings cntrl)
    rf.train <- upSample(x=rf.train[, names(rf.train) != dependent], y=rf.train[, names(rf.train) == dependent], yname=dependent)
}

# define variables Y and dependent x
Y.train <- rf.train[, names(rf.train) == dependent]
x.train <- rf.train[, names(rf.train) != dependent]
Y.test <- rf.test[, names(rf.test) == dependent]
x.test <- rf.test[, names(rf.test) != dependent]

# train single RF model
rf <- randomForest(x.train, y=Y.train, xtest=x.test, ytest=Y.test, type=type, keep.forest=TRUE)

You will get a warning from createDataPartition, and you will see that e.g., "make"==chevrolet has 3 cases in rf.train and none in rf.test, which can cause issues downstream with the randomForest.

Any smart way how to avoid that w/o leaking data from the train into the test?

1 Answers

A lot of this is the same, but not all.

The same:

This is because of your dependent variable. You chose make. Did you inspect this field? You have training and testing; where do you put an outcome with only one observation, like make = "mercury"? How can you train with that? How could you test for it if you didn't train for it?

input %>% 
  group_by(make) %>% 
  summarise(count = n()) %>% 
  arrange(count) %>% 
  print(n = 22)

# # A tibble: 22 × 2
#    make        count
#    <fct>       <int>
#  1 mercury         1
#  2 renault         2
#  3 alfa-romero     3
#  4 chevrolet       3
#  5 jaguar          3
#  6 isuzu           4
#  7 porsche         5
#  8 saab            6
#  9 audi            7
# 10 plymouth        7
# 11 bmw             8
# 12 mercedes-benz   8
# 13 dodge           9
# 14 peugot         11
# 15 volvo          11
# 16 subaru         12
# 17 volkswagen     12
# 18 honda          13
# 19 mitsubishi     13
# 20 mazda          17
# 21 nissan         18
# 22 toyota         32

When you executed the function createDataPartition(), you also had warnings. I think the randomForest package requires a minimum of five per group. You can filter for the groups you'll include and use that data for testing and training.

Before the comment labeled settings, you can add the following to subset the groups and validate the results.

filtGrps <- input %>% 
  group_by(make) %>% 
  summarise(count = n()) %>% 
  filter(count >=5) %>% 
  select(make) %>% 
  unlist()

# filter for groups with sufficient observations for package
input <- input %>% 
  filter(make %in% filtGrps) %>% 
  droplevels() # then drop the empty levels

# check to see if it filtered as expected
input %>% 
  group_by(make) %>% 
  summarise(count = n()) %>% 
  arrange(-count) %>% 
  print(n = 16)

This only uses 5, which isn't ideal. (More would be better.)

Change here

In the caret model, you used imputation. You didn't do that for this model. You dropped another 34 observations when you created input.clean. At that point...

# you removed another 34 rows- need to check the classes, again
# you imputed for caret/train
input.clean %>% 
  group_by(make) %>% 
  summarise(count = n()) %>% 
  arrange(-count) %>% 
  print(n = 16)
# # A tibble: 16 × 2
#    make          count
#    <fct>         <int>
#  1 toyota           31
#  2 nissan           18
#  3 honda            13
#  4 subaru           12
#  5 mazda            11
#  6 volvo            11
#  7 mitsubishi       10
#  8 dodge             8
#  9 volkswagen        8
# 10 peugot            7
# 11 plymouth          6
# 12 saab              6
# 13 mercedes-benz     5
# 14 audi              4
# 15 bmw               4
# 16 porsche           1 

You need to drop three more classes now.

# there is an exclamation point to negate this
input.clean <- input.clean %>% 
  filter(!make %in% c("audi", "bmw", "porsche")) %>% 
  droplevels()

# validate changes
input.clean %>% 
  group_by(make) %>% 
  summarise(count = n()) %>% 
  arrange(-count) %>% 
  print(n = 16)
# 13 classes now

From here on, your code is good to go.

rf
# 
# Call:
#  randomForest(x = x.train, y = Y.train, xtest = x.test, ytest = Y.test,      keep.forest = TRUE, type = type) 
#                Type of random forest: classification
#                      Number of trees: 500
# No. of variables tried at each split: 5
# 
#         OOB estimate of  error rate: 1.92%
# Confusion matrix:
#               dodge honda mazda mercedes-benz mitsubishi nissan peugot
# dodge            24     0     0             0          0      0      0
# honda             0    22     0             0          2      0      0
# mazda             0     0    24             0          0      0      0
# mercedes-benz     0     0     0            24          0      0      0
# mitsubishi        0     0     0             0         23      0      0
# nissan            0     0     0             0          0     23      0
# peugot            0     0     0             0          0      0     24
# plymouth          0     0     0             0          0      0      0
# saab              0     0     0             0          0      0      0
# subaru            0     0     0             0          0      0      0
# toyota            0     0     0             0          0      1      0
# volkswagen        0     0     0             0          0      0      0
# volvo             0     0     0             0          0      0      0
#               plymouth saab subaru toyota volkswagen volvo class.error
# dodge                0    0      0      0          0     0  0.00000000
# honda                0    0      0      0          0     0  0.08333333
# mazda                0    0      0      0          0     0  0.00000000
# mercedes-benz        0    0      0      0          0     0  0.00000000
# mitsubishi           1    0      0      0          0     0  0.04166667
# nissan               0    0      0      1          0     0  0.04166667
# peugot               0    0      0      0          0     0  0.00000000
# plymouth            24    0      0      0          0     0  0.00000000
# saab                 0   24      0      0          0     0  0.00000000
# subaru               0    0     24      0          0     0  0.00000000
# toyota               0    0      0     22          0     1  0.08333333
# volkswagen           0    0      0      0         24     0  0.00000000
# volvo                0    0      0      0          0    24  0.00000000
#                 Test set error rate: 3.23%

Just a tip- if you have these calls in the same script file, use unique object names between the models, that way, you always know what data is in which object. It can be a hidden error that causes all sorts of issues.

Related