Skip to main content
Create your own

Splitting Data for Machine Learning

Hello! Welcome to our next lesson in the "Core Machine Learning Concepts" module.

Introduction

In our last lesson, we explored powerful optimization algorithms like Adam, which are the engines that drive model training by updating parameters to minimize loss. We saw how to use optimizer.step() in a training loop to make our model learn from data.

But this raises a crucial question: how do we know if our model is actually learning to generalize, or just memorizing the data we're showing it? A model that perfectly memorizes its training data but fails on new, unseen data is useless in the real world. This is the problem of overfitting.

To build reliable models, we need a robust process for evaluating their performance. This brings us to today's learning outcome: Properly split data into training, validation, and test sets.

This lesson is fundamental to the entire practice of applied machine learning. We will cover:

  1. The purpose of each data split: Why we need three sets, not just two.
  2. The iterative development process: How these sets are used for training, tuning, and final evaluation.
  3. Practical implementation: How to perform these splits correctly using Python's scikit-learn library.
  4. Best practices: Key considerations like randomization, stratification, and avoiding the critical pitfall of "data leakage."

The Three-Set Split: Train, Validate, Test

At first glance, a simple train-test split seems logical. You train the model on a training set and evaluate its performance on a held-out test set. However, a problem arises when we start tuning our model.

Machine learning development is an iterative process. We experiment with different architectures (e.g., number of layers in a neural network), optimization settings (e.g., learning rate), or regularization techniques (e.g., dropout strength). These are called hyperparameters. If we use the test set to compare models with different hyperparameters, we are implicitly using the test set to make decisions. The model's configuration starts to become "fit" to the test set, and it is no longer a truly "unseen" dataset. Consequently, our final performance estimate on that test set will be overly optimistic and won't reflect real-world performance.

To solve this, we introduce a third set: the validation set. This creates a clear separation of concerns.

To understand the distinct roles of these three sets, let's watch a short, intuitive video.

Validation data: How it works and why you need it - Machine Learning Basics Explained

The video 'Validation data: How it works and why you need it' from Galaxy Inferno Codes uses a great analogy to explain the necessity and purpose of each data split.

Please watch the entire video (about 5.5 minutes). Pay close attention to the student-teacher analogy and how it maps to the roles of the training, validation, and test sets in machine learning.

As the video explained, the roles are as follows:

  • Training Set: The data the model learns from. The optimization algorithm (like Adam) uses this set to adjust the model's internal parameters (e.g., weights and biases) to minimize the loss function. This is the study material for the model.
  • Validation Set: An intermediate, held-out set used to evaluate the model during development. Its purpose is to guide hyperparameter tuning and model selection. For example, you would use the validation set to decide:
    • Is a learning rate of 0.001 better than 0.01?
    • Does a 3-layer network perform better than a 5-layer one?
    • Is my neural network better than a Gradient Boosting model for this task?
      This is like a practice exam for the model.
  • Test Set: A final, completely untouched set of data. It is used only once at the very end of the development process to provide an unbiased estimate of the final, chosen model's performance on truly unseen data. This is the final competition or real-world deployment scenario.

This workflow is beautifully captured in the following diagram:

Training, Validation, and Test Data Workflow
This flowchart illustrates the iterative cycle of model development. The model is trained on the training set, and its performance is evaluated on the validation set. This loop continues as you tweak the model (i.e., tune hyperparameters). The test set is only used once the best model has been selected.

Given your computer science background, you might appreciate thinking of this as a nested optimization loop.

What is the difference between test set and validation set?

A user named Penghe Geng on Stack Exchange provides a fantastic analogy that frames this process in terms of nested loops, which should be very intuitive for a programmer.

Please read the short answer by user 'Penghe Geng'. Focus on the nested while loop analogy and the distinction between who (or what) performs the inner and outer loops.

To summarize that excellent analogy:

  • Inner Loop (Machine): The model automatically tunes its parameters (weights) on the training set.
  • Outer Loop (Human/Developer): You manually tune the hyperparameters (architecture, learning rate) based on performance on the validation set.
  • Final Check (Unbiased Assessor): Once both loops are finished and you have your final model, you assess its true performance on the test set.
Test your understanding!

You have trained a neural network and it gets 99% accuracy on the training set but only 75% on the validation set. This suggests overfitting. You decide to add more dropout (a regularization hyperparameter) and retrain.

  1. Which dataset do you use to evaluate whether the new dropout value improved the model?
  2. After trying several dropout values, you find one that gives you 88% accuracy on the validation set. You are happy with this and want to report the final, expected performance of your model to stakeholders. Which dataset's accuracy do you report?
Show answer
  1. You use the validation set to evaluate the change in the dropout hyperparameter. The goal is to find a dropout value that improves performance on this "unseen" proxy data, not just on the training data.
  2. You report the model's accuracy on the test set. The 88% on the validation set helped you pick the best model, but it is a slightly biased estimate because you used it to make a decision. The test set provides the final, unbiased measure of how your chosen model is expected to perform in the real world.

Practical Implementation with Scikit-Learn

Now, let's move from theory to practice. In Python, the standard tool for splitting data is the train_test_split function from the scikit-learn library. Since this function only creates two splits at a time, we'll use a two-step process to get our three sets.

Sklearn - Split Data into 3 Sets (train, validation and test) in Python

The video 'Sklearn - Split Data into 3 Sets' by Koolac provides a very clear, code-first demonstration of this two-step process.

Please watch from the beginning up to 02:07. This will show you the conceptual two-step process and then the direct implementation in Python.

Let's consolidate that process. Here is how you would typically implement an 80-10-10 split, where 80% of the data is for training, 10% for validation, and 10% for testing.

Step 1: Split off the Test Set
First, we separate our final test set from the rest of the data. Let's reserve 10% of the data for testing.

import numpy as np
from sklearn.model_selection import train_test_split

# Assuming X is your feature matrix and y is your target vector
# First, split into a training+validation set (90%) and a test set (10%)
X_train_val, X_test, y_train_val, y_test = train_test_split(
    X, y, test_size=0.1, random_state=42
)

Step 2: Split the Remainder into Training and Validation Sets
Now, we split the X_train_val set. We want our validation set to be 10% of the original dataset. Since X_train_val is currently 90% of the original, we need to take 0.1 / 0.9 (approximately 11.1%) of it to get our validation set.

# Split the 90% into training (80% of original) and validation (10% of original)
# test_size = 0.1 / 0.9 = 0.1111...
X_train, X_val, y_train, y_val = train_test_split(
    X_train_val, y_train_val, test_size=0.1111, random_state=42
)

print(f"Original dataset shape: {X.shape}")
print(f"Training set shape: {X_train.shape}")   # Should be ~80%
print(f"Validation set shape: {X_val.shape}") # Should be ~10%
print(f"Test set shape: {X_test.shape}")     # Should be 10%

This two-step process ensures our test set remains isolated from the very beginning.

Best Practices and Common Pitfalls

Properly splitting data involves more than just calling a function. Here are some critical best practices.

Train-Test-Validation Split in 2025

The article 'Train-Test-Validation Split' from Analytics Vidhya provides a good overview of some important considerations and common mistakes.

Please read the sections titled 'Randomization in Data Splitting', 'Best Practices in Data Splitting', and 'Common Mistakes to Avoid'. These sections are short but highlight crucial concepts.

Let's expand on those key points:

  1. Randomization and Reproducibility (random_state)
    By default, train_test_split shuffles the data randomly before splitting. This is vital to ensure that your splits are representative and not biased by any initial ordering of the data. However, for development and debugging, you need your experiments to be reproducible. Setting the random_state parameter to an integer (the number itself doesn't matter, 42 is just a convention) ensures that the same random split is generated every time you run the code. This is a cornerstone of good software engineering in ML.

  2. Stratification (stratify)
    Imagine a classification task where 90% of your samples are Class A and 10% are Class B. A random split might, by chance, put very few or even zero samples of Class B into your validation or test set. This would make evaluation meaningless.
    To prevent this, you can use the stratify parameter. By setting stratify=y, you instruct train_test_split to preserve the same percentage of samples for each class in the splits as in the original dataset. This is crucial for classification tasks, especially with imbalanced classes.

    # Example of stratified splitting
    X_train_val, X_test, y_train_val, y_test = train_test_split(
        X, y, test_size=0.1, random_state=42, stratify=y
    )
    
    # You must also stratify the second split, using the corresponding y
    X_train, X_val, y_train, y_val = train_test_split(
        X_train_val, y_train_val, test_size=0.1111, random_state=42, stratify=y_train_val
    )
    
  3. Data Leakage
    This is the most critical and subtle pitfall. Data leakage occurs when information from outside the training set is used to create the model. We already discussed the primary form of leakage: using the test set for hyperparameter tuning. Another common form occurs during preprocessing.

    For example, if you need to scale your data (e.g., using min-max scaling or standardization), you must compute the scaling parameters (like mean and standard deviation) only from the training data. You then use these same parameters to transform the validation and test sets. If you compute the mean from the entire dataset before splitting, information about the distribution of the validation and test sets has "leaked" into your training process, leading to an over-optimistic performance estimate.

Conclusion

We've established a foundational practice for any supervised learning project. Getting your data splitting strategy right is non-negotiable for building models that you can trust.

Key Takeaways:

  • Three distinct sets are essential: The training set fits parameters, the validation set tunes hyperparameters, and the test set provides a final, unbiased performance estimate.
  • The process is iterative: You cycle between training on the training set and evaluating on the validation set until you have a model you are satisfied with.
  • The test set is sacred: It must be held out and used only once at the very end. Any use of the test set during development invalidates it as an unbiased evaluator.
  • Implementation is a two-step process: Use train_test_split twice to create the three sets.
  • Best practices are critical: Always use random_state for reproducibility and stratify for classification tasks to ensure your splits are valid and reliable. Be vigilant against all forms of data leakage.

Preview of the next lesson:
The train-validation-test split is a great strategy, but it has a potential downside. If your dataset is small, holding out a significant chunk for validation and testing can leave you with too little data for training. Furthermore, what if you get "unlucky" and your single validation set isn't representative of the overall data?

In our next lesson, we will learn how to "Implement K-fold cross-validation to assess model generalization." This powerful technique provides a more robust performance estimate and uses your data more efficiently, which is especially valuable for smaller datasets.

Can't find a good explanation? Sign up and we'll make it for you

Sign up