Train-test split
In the world of machine learning, data is the fuel that drives every model. However, how you use this fuel correctly determines whether your model becomes a smart engine that can accurately predict the future, or a parrot that can only memorize by rote.
Today, we will take a deep dive into a crucial and fundamental concept in machine learning:the split between training and test sets. This is the first step in building any reliable model and the key to evaluating a model's true capability.
In simple terms, splitting training and test sets is like studying and taking exams during school days:
- Training setis the student's textbooks and practice problems; the model uses it to learn the rules and patterns in the data.
- Test setis the final exam; the model uses it to check whether it has truly mastered the knowledge, rather than merely memorizing the answers to the practice problems (training set).
Why is it necessary to split training and test sets?
Imagine a student who only reviewed the mock questions given by the teacher, and the exam questions were exactly the same mock questions. He got a perfect score. Does this prove that he truly understood the subject? Obviously not. He may have just memorized the answers.
In machine learning, if we train a model onall the dataand then use thesame datato evaluate its performance, we make the same mistake. The model will appear exceptionally good because it has "seen" and "remembered" every detail of the data, including the noise and random coincidences. This phenomenon is calledoverfitting。
An overfitted model is like a student who can only recite example questions; once it encounters new, unseen questions (new data), it performs very poorly. Its "generalization ability" is very weak.
Therefore, we must divide the data into two parts:
- Training set: used toTeachtrain the model, letting it learn.
- Test set: used toexamineevaluate the model and assess how it handlesnew data it has never seen.
The test set must be completely isolated from the training set and mustnotbe seen by the model during the entire training process. Only in this way can the evaluation results on the test set objectively reflect the true generalization ability of the model.
How to split: common methods and strategies
Splitting data sounds simple, but there is a lot to know. Different splitting strategies are suitable for different scenarios.
1. Simple random split
This is the most basic and common method. Randomly shuffle the entire dataset, then split it into two parts according to a certain ratio.
Example
from sklearn.model_selection import train_test_split
# Assume X is the feature data and y is the label data
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)
print(f"Training set sample count: {len(X_train)}")
print(f"Test set sample count: {len(X_test)}")
Code explanation:
train_test_split: This is the core function in scikit-learn for splitting data.X, y: The input feature data and the corresponding labels.test_size=0.2: Specifies that the test set size ratio is 20% (i.e., training set is 80%). You can also usetrain_size=0.8to specify it.random_state=42: Set a random seed. This ensures that the split result is exactly the same every time the code runs, which is crucial for the reproducibility of experiments. You can set it to any integer.
2. Stratified sampling split
In classification problems, if the class distribution of the dataset is imbalanced (e.g., 90% class A, 10% class B), simple random splitting may cause the class proportions in the training and test sets to differ greatly, affecting the fairness of evaluation.
Stratified samplingensures that the proportions of each class in the training and test sets remain consistent with the original dataset.
Example
from sklearn.model_selection import train_test_split
# Assume y is the classification labels
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, stratify=y, random_state=42)
# Check the class distribution after splitting
from collections import Counter
print("Original data class distribution:", Counter(y))
print("Training set class distribution:", Counter(y_train))
print("Test set class distribution:", Counter(y_test))
Code explanation:
stratify=y: This is the key parameter. It tells the function to stratify according toythe class distribution of the labels when performing stratified sampling.
3. Time-series data split
For time-series data (e.g., stock prices, daily temperatures), there are temporal dependencies between data points. We cannot shuffle randomly, because future data cannot be used to predict the past.
The usual practice is to split in chronological order:Use the first 80% of the data by time as the training set, and the last 20% as the test set.。
Example
split_index = int(len(X) * 0.8) # Compute the index at the 80% position
X_train, X_test = X[:split_index], X[split_index:]
y_train, y_test = y[:split_index], y[split_index:]
print(f"Training set time range: first {split_index} samples")
print(f"Test set time range: last {len(X) - split_index} samples")
How to choose the split ratio?
This is a common question, but there is no fixed answer. Common ratios include:
| Ratio (training set:test set) | Applicable scenario | Advantages | Disadvantages |
|---|---|---|---|
| 70:30 | Classic choice for small to medium-sized datasets (thousands to tens of thousands of samples) | Balances training data size and evaluation reliability | For extremely small datasets, a 30% test set may have too few samples, making evaluation unstable |
| 80:20 | Currently a more popular default choice, especially suitable for deep learning | Provides more data for the model to learn from | The test set is relatively small, so the variance of evaluation may be slightly larger |
| 90:10 or 95:5 | When the amount of data is very limited | Maximizes the use of limited data for training | The test set is too small, so the evaluation results may be unreliable and have low confidence |
Core principles:
- Ensure the training set is large enough: the model needs enough data to learn effective patterns.
- Ensure the test set is large enough: the test set needs to provide a statistically reliable performance evaluation. Usually, the test set should contain at least a few hundred samples for stable evaluation results.
- The larger the data volume,the proportion allocated to the test set can be relativelysmaller, because even a very small proportion may represent a large number of samples.
Advanced concepts: validation set and cross-validation
In real projects, we not only need to evaluate the final model, but also need to tune the model'shyperparameters(such as learning rate, tree depth, etc.) during training. If we directly use the test set to tune parameters, the test set becomes "contaminated" again, losing its impartiality as the "final examiner."
To this end, we introduce thevalidation set。
Three-way data split: training, validation, and test sets
- Training set: used for learning model parameters.
- Validation set: used to tune hyperparameters, select models, or perform early stopping during training. It is equivalent to a "mock exam."
- Test set: used for the final, one-time performance evaluation after both the model and hyperparameters are determined. It is the "final exam."
Example
X_temp, X_test, y_temp, y_test = train_test_split(X, y, test_size=0.15, random_state=42) # First set aside 15% as the final test set
X_train, X_val, y_train, y_val = train_test_split(X_temp, y_temp, test_size=0.176, random_state=42) # Then from the remaining 85%, set aside about 15% as the validation set
# Calculate the ratio: 0.85 * 0.176 ≈ 0.15, final ratio is about 70:15:15
print(f"Training set: {len(X_train)}, Validation set: {len(X_val)}, Test set: {len(X_test)}")
K-fold cross-validation
When the amount of data is not large, setting aside a separate validation set further reduces the training data.K-fold cross-validationis a more powerful solution.
Its process is as follows, and it can effectively use limited data:
Example
A[Original Dataset] --> B[Randomly shuffle and evenly divide into K parts]
B --> C{Repeat for K rounds}
C --> D[Round i: Use the i-th part as the validation set]
D --> E[Merge the remaining K-1 parts as the training set]
E --> F[Train the model on the training set for this round]
F --> G[Evaluate score Si on the validation set for this round]
G --> C
C -- After K rounds are completed --> H[Calculate the average of the K scores as the final evaluation]
Example
from sklearn.model_selection import cross_val_score
from sklearn.linear_model import LogisticRegression
model = LogisticRegression()
scores = cross_val_score(model, X, y, cv=5) # cv=5 means 5-fold cross-validation
print(f"Scores for each fold: {scores}")
print(f"Average score: {scores.mean():.4f} (+/- {scores.std()*2:.4f})") # Output the mean and standard deviation
Advantages of Cross-Validation:
- Fully utilizes all data for training and validation.
- The evaluation results are more stable and reliable (because they are the average of multiple evaluations).
- It is the gold standard for model selection and hyperparameter tuning on small to medium-sized datasets.
Hands-on practice: experience data splitting yourself
Now, let's practice with a simple dataset.
Example
import numpy as np
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split
# 2. Load the Iris dataset
iris = load_iris()
X, y = iris.data, iris.target
print(f"Dataset shape: features {X.shape}, labels {y.shape}")
# 3. Simple random split (80% training, 20% testing)
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)
print(f"Random split -> Training set: {X_train.shape}, Test set: {X_test.shape}")
# 4. Stratified random split
X_train_s, X_test_s, y_train_s, y_test_s = train_test_split(X, y, test_size=0.2, stratify=y, random_state=42)
print(f"Stratified split -> Training set: {X_train_s.shape}, Test set: {X_test_s.shape}")
# 5. Check the effect of stratification
print("\n"Original data class distribution:", np.bincount(y))
print("Test set distribution after random split:", np.bincount(y_test)) # May be imbalanced
print("Test set distribution after stratified split:", np.bincount(y_test_s)) # Should be proportional to the original distribution
Your Task:
- Run the code above and observe the output results.
- Try modifying
test_sizeto 0.3, and observe the change in training and test set sizes. - Try modifying
random_stateto another number (e.g., 7), run again, and observe whether the split results change. - (Challenge) Do not set
random_statethe parameter, run the code multiple times, and observe whether the split results are the same each time.
Summary and key points
- Core purpose: Splitting the training and test sets is toevaluate the model's generalization ability, prevent overfitting, and ensure the model can handle new data.
- Golden Rule:The test set must be kept completely confidential throughout the entire training process, and used only for the final evaluation.
- Splitting Methods:
- Random splitting: The most commonly used.
- Stratified splitting: Suitable for imbalanced data in classification problems.
- Sequential splitting: Suitable for time series data.
- Split Ratio: There is no absolute standard; a trade-off is needed between "sufficient training" and "reliable evaluation". 80:20 or 70:30 are common starting points.
- Advanced Tools:
- Validation set: Used for model tuning, protecting the purity of the test set.
- K-fold cross-validation: A powerful tool for evaluating and tuning on small to medium-sized datasets, producing more robust results.