Scikit-Learn
data stratification
machine learning
training and testing data
Python

How to stratify the training and testing data in Scikit-Learn?

ML System Design practice on Codemia

Design recommenders, ranking systems and training pipelines the way ML interviews actually ask for them, with worked solutions.

Practice ML system design

Stratifying your dataset during the training and testing split is essential when working with classification problems where the distribution of classes can significantly impact the model’s performance. Stratification ensures that the training and testing sets maintain the same proportion of classes as in the original dataset, thereby providing a more consistent and reliable evaluation of the model's performance.

In this article, we'll delve into how to leverage Scikit-Learn's capabilities to efficiently stratify your training and testing datasets. We'll walk through a technical explanation and illustrate it with practical examples.

Understanding Stratification

In classification, stratification involves partitioning your data so that each subset reflects the statistical properties of the original data. Essentially, if your dataset has 70% of class A and 30% of class B, then the stratified splits should also follow this distribution. This aids in preventing random sampling errors that may occur in imbalanced datasets.

Stratified Data Splitting in Scikit-Learn

Scikit-Learn offers utilities to perform stratified splitting through the `train_test_split` function and the `StratifiedKFold` class. These tools are invaluable for ensuring your datasets are representative of the entire population.

Using `train_test_split`

The `train_test_split` function can create stratified splits by using the `stratify` parameter. Here's how it's typically used:

  • stratify: If not `None`, data is split in a stratified fashion using this as the class labels.
  • test_size: Represents the proportion of the dataset to include in the test split (e.g., 0.3 for 30%).
  • random_state: Controls the shuffling applied to the data before the split. Pass an integer for reproducible results.
  • n_splits: Number of folds. Must be at least 2.
  • shuffle: Whether to shuffle each class’s samples before splitting into batches.
  • random_state: Controls the randomization of the splits.
  • Representation: Ensures that all classes are represented proportionally in both training and testing sets.
  • Stability: Reduces variability in performance evaluation by maintaining the class distribution.
  • Balanced Learning: Mitigates the impact of class imbalance during model training.

Related reading
Free course
Beginner
7 lessons
2 hours
Tackling System Design Interview Problems

A short course that equips you with the skills to approach system design interviews methodically.

Start the free course
Track what you have practised

A free account saves your progress, solutions and study plan across every problem on Codemia.

ML System Design practice on Codemia

Design recommenders, ranking systems and training pipelines the way ML interviews actually ask for them, with worked solutions.

Practice ML system design

All Rights Reserved.