skip to content

Model Selection & CV

Choosing a validation scheme — plain KFold, stratified, grouped, or time-series splits — and searching hyperparameters with grid or randomized search. Expect to explain why a naive random split is wrong for grouped or temporal data, and what nested CV protects against.

on this pageshow

questions

6

In scikit-learn, what does the stratify argument of train_test_split do?

level: juniorimportance: must knowfreq 80%

answer

  1. shuffle first, then slice
  2. class ratio preserved on both sides
  3. rare classes vanish without it
  4. needs shuffle=True, min two per class
  5. no notion of groups or time

basics

~20 s

stratify=y tells train_test_split to preserve each class's proportion in both halves instead of splitting purely at random. Without it a rare class can land unevenly in the test set, or be missing from it entirely.

solid answer

~50 s

`train_test_split` shuffles rows and slices them by default (`shuffle=True`, and `test_size` defaults to 0.25 when neither size is given). That is a uniform random draw, so class proportions in the two halves only match in expectation — with an imbalanced target the test set's positive rate can drift badly, and with a very rare class it may contain none at all. Passing `stratify=y` makes the split sample within each class so the class ratio is approximately preserved on both sides. You can stratify on any array, not just the label — a combined key of label and site, for example. Two constraints matter in practice: `stratify` requires `shuffle=True` (combining it with `shuffle=False` raises), and every class must have at least two members or the call raises. `stratify` says nothing about grouped or time-ordered rows; those need a different scheme.

code

python · 12 lines
python
from sklearn.datasets import make_classification
from sklearn.model_selection import train_test_split

X, y = make_classification(n_samples=1000, weights=[0.97, 0.03], random_state=0)

X_tr, X_te, y_tr, y_te = train_test_split(
    X, y, test_size=0.2, random_state=0, stratify=y
)
print(y.mean(), y_tr.mean(), y_te.mean())

# ordered hold-out: no shuffling, and therefore no stratify
X_tr2, X_te2 = train_test_split(X, test_size=0.2, shuffle=False)

go deeper

for a junior

Know the call signature and that stratify=y keeps class proportions in both halves. Say plainly that shuffling is on by default and that random_state only makes the split repeatable.

for a middle

Explain why an imbalanced target makes an unstratified split high-variance, and name the two hard constraints: stratify requires shuffle=True, and each class needs at least two members.

for a senior

Show judgment about when the row-level split is the wrong tool at all — repeated entities or time ordering — and describe the hold-out discipline of touching the final test set once.

for a principal

Own the split policy for a team: what the frozen evaluation set is, how it is versioned and stored as explicit indices rather than a seed, and how it is refreshed as the data distribution drifts.

## What train_test_split actually does `sklearn.model_selection.train_test_split(*arrays, test_size=None, train_size=None, random_state=None, shuffle=True, stratify=None)` takes one or more equal-length arrays (X and y, or X, y and a weights vector) and returns each of them cut into two pieces, in the order X_train, X_test, y_train, y_test. By default it shuffles the row order first and then slices; if neither `test_size` nor `train_size` is supplied, the test share is 0.25. Because the shuffle is uniform, the composition of each half is a random draw. For a balanced binary problem with tens of thousands of rows that is fine — the law of large numbers hides the variance. For a 2%-positive fraud dataset with 3,000 rows, the test set holds roughly 15 positives and its measured recall swings wildly from seed to seed. For a 30-class problem where some classes have four examples, a random split can put all four in training and leave the test set with a class it never sees, which then breaks per-class metrics. ## Stratification `stratify=y` switches the sampling: instead of one draw over all rows, the split is performed within each distinct value of the array you pass, so each stratum contributes the same proportion to train and test. The resulting class ratio matches the source ratio up to rounding. Internally this is the same machinery as `StratifiedShuffleSplit`. The array you pass does not have to be the target. Stratifying on a composite key — say `y.astype(str) + '_' + site` — balances both label and site at once, as long as every combination has at least two rows. ## The rules you will trip over - **`stratify` needs shuffling.** `shuffle=False` with a non-None `stratify` raises a `ValueError`. Turning shuffling off is how you take the tail of an ordered dataset; that is incompatible with stratified sampling by construction. - **Every class needs at least two members.** Otherwise you get "The least populated class in y has only 1 member, which is too few." Fold the singleton classes into an 'other' bucket, or drop them, before splitting. - **`random_state` controls reproducibility, not balance.** Fixing the seed makes the same split come back; it does not make that split representative. People routinely confuse the two. - **Stratifying a continuous target is wrong.** Each distinct float is its own class, so the call either raises or produces nonsense. Bin the target first if you want balance across ranges. ## What stratification does not fix Stratification is per-row. It has no notion of *which rows belong together* and no notion of *time*. If your data has repeated entities — several images per patient, several sessions per user, several rows per invoice — a stratified split will happily place one patient's images on both sides. The model then memorizes the patient and the score is optimistic. That is what group-aware splitters exist for; the row-level `stratify` argument cannot express it. If your data is a time series, any shuffled split lets the model train on the future and test on the past. The correct hold-out is the tail: `shuffle=False` with `test_size=0.2` keeps the original order and takes the last 20% as the test set, and that call must not carry `stratify`. ## Where it sits in a workflow The conventional pattern is one call to carve off a final test set that you touch once, and cross-validation over the remaining data for everything else — model choice, hyperparameters, threshold selection. A second `train_test_split` on the training portion gives an explicit validation set when you need a single fixed one (early stopping, for example) rather than repeated folds. A practical detail: run the split *before* fitting any transformer. Fitting a scaler or an encoder on the full array and splitting afterwards lets test-set statistics reach the training data, which inflates every number that follows.

  • What happens if you pass both stratify=y and shuffle=False?
    It raises a ValueError. Stratified sampling means drawing within each class, which is inherently a randomized operation; `shuffle=False` means "take the tail of the array in its existing order". The two requests contradict each other, so scikit-learn refuses rather than silently ignoring one. If you need an ordered hold-out, drop `stratify`.
  • You have 300 images from 40 patients. Is stratify=y enough to build an honest test set?
    No. `stratify` balances labels row by row but has no idea that several rows come from the same patient, so images of one patient will appear in both halves and the model can score well by recognizing the patient rather than the condition. You need a group-aware split that keeps each patient entirely on one side.
  • Does train_test_split guarantee the same split across scikit-learn versions if random_state is fixed?
    It guarantees reproducibility within a version and, in practice, across most versions, but scikit-learn does not promise byte-identical splits forever — changes to the underlying permutation logic can shift them. If exact row membership must be stable long-term, materialize the index lists once and store them rather than re-deriving them from a seed.

saying these in an interview costs you the question

  • Thinks random_state makes a split representative rather than reproducible
  • Claims stratify keeps related rows together
  • Stratifies a continuous regression target
  • Assumes the default test_size is 0.2
  • Splits after fitting the scaler on all rows

context

open as a page

What splitter does scikit-learn use when you pass cv=5 to cross_val_score?

level: middleimportance: must knowfreq 68%

basics

~20 s

An integer cv is expanded for you: StratifiedKFold when the estimator is a classifier and the target is binary or multiclass, plain KFold otherwise. Neither shuffles — folds are contiguous blocks of the data in its current row order.

open as a page

When should you use RandomizedSearchCV instead of GridSearchCV in scikit-learn?

level: middleimportance: must knowfreq 74%

basics

~20 s

GridSearchCV fits every combination in param_grid, so its cost multiplies with each added parameter. RandomizedSearchCV draws a fixed n_iter samples from param_distributions, letting you cap the budget and sample continuous ranges instead of a hand-picked ladder.

open as a page

How does scikit-learn's TimeSeriesSplit differ from KFold, and what is gap for?

level: middleimportance: should knowfreq 55%

basics

~20 s

TimeSeriesSplit never puts later rows in a training fold: each split trains on a prefix of the rows and tests on the block immediately after, with the training window growing each split. gap drops a fixed number of rows between the train end and the test start.

open as a page

In scikit-learn, which CV splitter keeps all of one patient's rows in a single fold?

level: seniorimportance: should knowfreq 48%

basics

~20 s

GroupKFold, given a groups array of patient IDs, guarantees no group's rows appear in both the training and test side of a split. The groups array is passed at fit time — GroupKFold(n_splits=5).split(X, y, groups) or search.fit(X, y, groups=ids).

open as a page

In scikit-learn, how do you run nested cross-validation around a GridSearchCV?

level: seniorimportance: should knowfreq 40%

basics

~20 s

Pass the unfitted search as the estimator to an outer cross-validation: cross_val_score(GridSearchCV(est, grid, cv=inner), X, y, cv=outer). Each outer fold tunes on its own training portion and is scored on data no tuning decision ever saw.

open as a page