Mathew K Analytics

Lesson 8 · Scikit-learn deep dive

Scikit-learn Tutorial #8: Splitting & Cross-Validation

Video eight of the eighteen-part series: honestly measuring how well a model generalizes. traintestsplit, KFold, StratifiedKFold, crossvalscore, and…

⬇ Download notebookOpen in Colab ↗

📓 Full notebook

Download .ipynb

Scikit-learn Deep-Dive, Video 8: Model Selection - Splitting and Cross-Validation#

  • Video eight of the eighteen-part series: honestly measuring how well a model generalizes.
  • train_test_split, KFold, StratifiedKFold, cross_val_score, and cross_validate.
  • Let's get into it.

Part 1: train_test_split Basics#

from sklearn.datasets import load_wine
from sklearn.model_selection import train_test_split
X, y = load_wine(return_X_y=True)
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42)
print(X_train.shape, X_test.shape)
(142, 13) (36, 13)

Part 2: stratify - Preserving Class Balance#

import numpy as np
X_tr_strat, X_te_strat, y_tr_strat, y_te_strat = train_test_split(
    X, y, test_size=0.2, random_state=42, stratify=y
)
print(np.bincount(y) / len(y))
print(np.bincount(y_te_strat) / len(y_te_strat))
[0.33146067 0.3988764  0.26966292]
[0.33333333 0.38888889 0.27777778]

Part 3: KFold - Basic K-Fold Splitting#

from sklearn.model_selection import KFold
kf = KFold(n_splits=5, shuffle=True, random_state=42)
for i, (train_idx, val_idx) in enumerate(kf.split(X)):
    print(i, len(train_idx), len(val_idx))
0 142 36
1 142 36
2 142 36
3 143 35
4 143 35

Part 4: StratifiedKFold - Preserving Class Balance Across Folds#

from sklearn.model_selection import StratifiedKFold
skf = StratifiedKFold(n_splits=5, shuffle=True, random_state=42)
for i, (train_idx, val_idx) in enumerate(skf.split(X, y)):
    fold_props = np.bincount(y[val_idx]) / len(val_idx)
    print(i, fold_props.round(2))
0 [0.33 0.39 0.28]
1 [0.33 0.39 0.28]
2 [0.33 0.39 0.28]
3 [0.34 0.4  0.26]
4 [0.31 0.43 0.26]

Part 5: cross_val_score with a cv Object#

from sklearn.model_selection import cross_val_score
from sklearn.linear_model import LogisticRegression
from sklearn.pipeline import make_pipeline
from sklearn.preprocessing import StandardScaler
model = make_pipeline(StandardScaler(), LogisticRegression(max_iter=1000))
scores = cross_val_score(model, X, y, cv=skf)
print(scores.round(3))
print(round(scores.mean(), 3), round(scores.std(), 3))
[0.972 0.972 0.972 1.    1.   ]
0.983 0.014

Part 6: cross_validate - Multiple Metrics and Timing#

from sklearn.model_selection import cross_validate
results = cross_validate(
    model, X, y, cv=skf,
    scoring=['accuracy', 'f1_macro'],
    return_train_score=True
)
print(sorted(results.keys()))
print(results['test_accuracy'].round(3))
print(results['test_f1_macro'].round(3))
['fit_time', 'score_time', 'test_accuracy', 'test_f1_macro', 'train_accuracy', 'train_f1_macro']
[0.972 0.972 0.972 1.    1.   ]
[0.972 0.972 0.971 1.    1.   ]

Part 7: shuffle and random_state - Why They Matter#

kf_unshuffled = KFold(n_splits=5, shuffle=False)
kf_shuffled_a = KFold(n_splits=5, shuffle=True, random_state=1)
kf_shuffled_b = KFold(n_splits=5, shuffle=True, random_state=1)
first_fold_a = next(kf_shuffled_a.split(X))[1]
first_fold_b = next(kf_shuffled_b.split(X))[1]
print(np.array_equal(first_fold_a, first_fold_b))
True

Part 8: ShuffleSplit - Repeated Random Subsampling#

from sklearn.model_selection import ShuffleSplit
ss = ShuffleSplit(n_splits=10, test_size=0.2, random_state=42)
ss_scores = cross_val_score(model, X, y, cv=ss)
print(ss_scores.round(3))
print(round(ss_scores.mean(), 3))
[1.    1.    1.    1.    0.972 1.    0.972 0.972 1.    1.   ]
0.992

Part 9: KFold vs StratifiedKFold on Imbalanced Data#

y_imbalanced = (y == 2).astype(int)
print(np.bincount(y_imbalanced))
kf_plain = KFold(n_splits=5, shuffle=True, random_state=42)
skf_strat = StratifiedKFold(n_splits=5, shuffle=True, random_state=42)
plain_props = [np.bincount(y_imbalanced[val])[1] for _, val in kf_plain.split(X, y_imbalanced)]
strat_props = [np.bincount(y_imbalanced[val])[1] for _, val in skf_strat.split(X, y_imbalanced)]
print('plain KFold rare-class counts:', plain_props)
print('StratifiedKFold rare-class counts:', strat_props)
[130  48]
plain KFold rare-class counts: [np.int64(8), np.int64(11), np.int64(9), np.int64(10), np.int64(10)]
StratifiedKFold rare-class counts: [np.int64(10), np.int64(10), np.int64(10), np.int64(9), np.int64(9)]

Part 10: A Real Pattern - a Reusable evaluate_model_cv Function#

def evaluate_model_cv(estimator, X, y, n_splits=5, metrics=('accuracy', 'f1_macro')):
    cv = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=42)
    results = cross_validate(estimator, X, y, cv=cv, scoring=list(metrics))
    summary = {}
    for metric in metrics:
        scores = results[f'test_{metric}']
        summary[metric] = (round(scores.mean(), 3), round(scores.std(), 3))
    return summary
print(evaluate_model_cv(model, X, y))
{'accuracy': (np.float64(0.983), np.float64(0.014)), 'f1_macro': (np.float64(0.983), np.float64(0.014))}

Wrap-Up: What You Learned#

  • train_test_split shuffles and splits data in one call; test_size and random_state control the split.
  • stratify=y keeps each class's proportion consistent across train and test subsets.
  • KFold divides data into k folds, using each once for validation while training on the rest.
  • StratifiedKFold extends KFold by keeping class balance consistent within every fold, the recommended default for classification.
  • Passing a configured splitter object as cv gives full control over the exact cross-validation strategy.
  • cross_validate returns a dictionary supporting multiple scoring metrics and fit/score timing at once.
  • shuffle=True plus random_state avoids splitting on ordered data while staying reproducible.
  • ShuffleSplit generates independent random train/test splits, useful for more repeats than KFold alone allows.
  • Plain KFold can under-represent a rare class in a fold; StratifiedKFold guards against that by construction.
  • That wraps up model selection. Next up: Hyperparameter Tuning with GridSearchCV and RandomizedSearchCV.

Found this useful?

All lessons, notebooks and datasets here are free. If they helped you, a coffee keeps new lessons coming.