使用sklearn的StratifiedShuffleSplit出现test_size参数重复赋值错误
TypeError with StratifiedShuffleSplit in scikit-learn Hey there! Let's break down what's going wrong and fix this quickly.
The Root Cause
When you switched from the deprecated sklearn.cross_validation to sklearn.model_selection, the API for StratifiedShuffleSplit changed drastically.
In the old cross_validation version, you could pass labels directly as the first constructor argument and iterate over the instance. But the new model_selection implementation has key differences:
- The constructor no longer accepts labels as a positional argument — the first positional parameter is
n_splits. - Parameters like
test_sizeandrandom_stateare keyword-only (marked by the*in the method signature), meaning they can't be passed as positional values. - You need to call the
.split()method to feed in your features and labels when generating indices.
Your current code passes labels as the first positional argument (which gets incorrectly mapped to n_splits), then 10 as the second positional argument (which tries to set test_size). When you explicitly add test_size=0.2 afterward, Python throws an error because you're assigning two values to the same parameter.
Fixed Code
Here's the corrected version that aligns with the new API:
from sklearn.model_selection import StratifiedShuffleSplit # Initialize the splitter with correct parameter formatting sss = StratifiedShuffleSplit(n_splits=10, test_size=0.2, random_state=23) # Use .split() to pass your training data and labels for train_index, test_index in sss.split(train, labels): X_train, X_test = train.values[train_index], train.values[test_index] y_train, y_test = labels[train_index], labels[test_index]
Key Changes:
- We explicitly name the
n_splitsparameter instead of passing it after labels. - We call
sss.split(train, labels)to generate train/test indices, rather than iterating directly over thesssinstance. - All optional parameters (
test_size,random_state) are passed as keyword arguments, which matches the new API's requirements.
With these fixes, your code should run without the TypeError and correctly perform stratified shuffle splits on your data.
内容的提问来源于stack exchange,提问作者Naidu Venkat

