Scikit-learn中Stratify与StratifiedKFold的区别及分层采样疑问
Hey there! Let's tackle your two questions about Scikit-learn's stratification tools clearly and practically.
stratify parameter and StratifiedKFold in Scikit-learn These two tools both handle stratified sampling (keeping class ratios consistent across splits), but they serve different use cases:
- The
stratifyparameter is a convenience argument for functions liketrain_test_splitorcross_val_score. You just pass your target variableyto it, and the function automatically handles the stratified split under the hood. It's great for quick, straightforward splits where you don't need fine-grained control over the cross-validation process. StratifiedKFoldis a full cross-validation iterator class (fromsklearn.model_selection). It gives you full control over how splits are generated: you can set the number of folds (n_splits), choose to shuffle data before splitting, set a random state for reproducibility, etc. It's ideal when you need to customize cross-validation workflows—like pairing withGridSearchCVfor hyperparameter tuning, or building custom evaluation loops.
Here's a quick example of using StratifiedKFold:
from sklearn.model_selection import StratifiedKFold import numpy as np # Sample imbalanced data X = np.array([[1,2], [3,4], [5,6], [7,8], [9,10], [11,12]]) y = np.array([0, 0, 1, 1, 1, 1]) # Initialize the stratified k-fold iterator skf = StratifiedKFold(n_splits=2, shuffle=True, random_state=42) # Generate and use splits for train_idx, test_idx in skf.split(X, y): X_train, X_test = X[train_idx], X[test_idx] y_train, y_test = y[train_idx], y[test_idx] print(f"Train class counts: {np.bincount(y_train)}") print(f"Test class counts: {np.bincount(y_test)}\n")
stratify=y in train_test_split preserve class ratios for imbalanced data? Absolutely—it preserves the original class ratio, not equalizes the classes.
In your case (99% class 1, 1% class 0), using stratify=y will ensure both your training and test sets maintain that ~99:1 ratio. It won't force 50:50 splits (that's a different technique called oversampling/undersampling, which you'd handle separately).
To prove it, here's a code example matching your scenario:
from sklearn.model_selection import train_test_split import numpy as np # Create an extremely imbalanced dataset y = np.concatenate([np.zeros(10), np.ones(990)]) # 1% class 0, 99% class 1 X = np.random.rand(1000, 5) # Dummy feature data # Split with stratification X_train, X_test, y_train, y_test = train_test_split( X, y, test_size=0.2, stratify=y, random_state=42 ) # Calculate and print ratios original_ratio_0 = (10 / 1000) * 100 train_ratio_0 = (np.sum(y_train == 0) / len(y_train)) * 100 test_ratio_0 = (np.sum(y_test == 0) / len(y_test)) * 100 print(f"Original class 0 ratio: {original_ratio_0:.1f}%") print(f"Training set class 0 ratio: {train_ratio_0:.1f}%") print(f"Test set class 0 ratio: {test_ratio_0:.1f}%")
When you run this, you'll see the ratios across splits are nearly identical to the original dataset (minor differences might happen with tiny sample sizes, but the overall proportion stays consistent). This is critical for imbalanced data—without stratification, you might end up with a test set that has no class 0 samples at all, making your model evaluation meaningless.
内容的提问来源于stack exchange,提问作者user9238790

