Python3.7下Sklearn train_test_split的stratify参数异常问题
Hey folks, I recently hit a confusing issue where train_test_split with the stratify parameter worked exactly as expected in Python 3.5, but produced totally wrong results in Python 3.7—even though all my library versions were identical. Let me break down what happened, share the code to reproduce it, and the fix I found.
The Goal
I was trying to split a dataset with a 7:3 class distribution (70% class 1, 30% class 0) while preserving that ratio in both the training and test sets.
Reproducible Code
Here's the code that shows the discrepancy:
import numpy as np from sklearn.model_selection import train_test_split # Create sample data and target with 7:3 distribution data = np.random.rand(1000000).reshape(100000, 10) y_0 = [0]*30000 y_1 = [1]*70000 y_2 = y_0 + y_1 # Split with stratification to preserve class ratio x_train, x_test, y_train, y_test = train_test_split( data, y_2, test_size=0.2, random_state=0, stratify=y_2 ) # Print results print('Train set size : {}'.format(len(y_train))) print('Value 1 repartition in train set : {}'.format(sum(y_train)/len(y_train))) print('Test set size : {}'.format(len(y_test))) print('Value 1 repartition in test set : {}'.format(sum(y_test)/len(y_test)))
Python 3.7 Output (Unexpected)
Running this in Python 3.7 gave completely off results—way smaller training set and messed up class ratios:
Train set size : 24102 Value 1 repartition in train set : 0.5414903327524687 Test set size : 20000 Value 1 repartition in test set : 0.53775
Python 3.5 Output (Expected)
In Python 3.5, everything worked perfectly—training set is 80% of the data, and class ratios stay at 70%:
Train set size : 80000 Value 1 repartition in train set : 0.7 Test set size : 20000 Value 1 repartition in test set : 0.7
Identical Library Versions
To make this even more confusing, both environments had exactly the same package versions:
- Python 3.7 setup: Python 3.7.2, numpy1.16.1, pandas0.24.1, python-dateutil2.8.0, pytz2018.9, scikit-learn0.20.2, scipy1.2.1, six==1.12.0
- Python 3.5 setup: Python 3.5.1, numpy1.16.1, pandas0.24.1, python-dateutil2.8.0, pytz2018.9, scikit-learn0.20.2, scipy1.2.1, six==1.12.0
Root Cause & Fix
After troubleshooting, I found the problem comes down to how train_test_split handles list-type target variables in Python 3.7 with numpy 1.16.1. There's an implicit type conversion issue under the hood that causes incorrect sampling when using a list for stratify.
The simple fix is to convert your target list to a numpy array before passing it to the function:
y_2 = np.array(y_0 + y_1) # Convert list to numpy array
This ensures the stratification logic works consistently across both Python versions, as sklearn's internal handling is more reliable with numpy arrays than native lists in this specific version combination.
内容的提问来源于stack exchange,提问作者Nico

