Spec-Zone.ru › scikit-learn

Примечание

Перейти к концу, чтобы загрузить полный пример кода. или запустить этот пример в вашем браузере через JupyterLite или Binder

Визуализация поведения кросс-валидации в scikit-learn

Выбор правильного объекта кросс-валидации является важной частью правильного обучения модели. Существует множество способов разделить данные на обучающую и тестовую выборки, чтобы избежать переобучения модели, стандартизировать количество групп в тестовых выборках и т.д.

В этом примере визуализируется поведение нескольких распространённых объектов scikit-learn для сравнения.

# Authors: The scikit-learn developers
# SPDX-License-Identifier: BSD-3-Clause

import matplotlib.pyplot as plt
import numpy as np
from matplotlib.patches import Patch

from sklearn.model_selection import (
    GroupKFold,
    GroupShuffleSplit,
    KFold,
    ShuffleSplit,
    StratifiedGroupKFold,
    StratifiedKFold,
    StratifiedShuffleSplit,
    TimeSeriesSplit,
)

rng = np.random.RandomState(1338)
cmap_data = plt.cm.Paired
cmap_cv = plt.cm.coolwarm
n_splits = 4

Визуализируем наши данные

Сначала необходимо понять структуру наших данных. Она содержит 100 случайно сгенерированных точек данных, 3 класса, неравномерно распределённых по точкам данных, и 10 «групп», равномерно распределённых по точкам данных.

Как мы увидим, некоторые объекты кросс-валидации выполняют определённые действия с метками данных, другие ведут себя иначе с группированными данными, а другие не используют эту информацию.

Для начала визуализируем наши данные.

# Generate the class/group data
n_points = 100
X = rng.randn(100, 10)

percentiles_classes = [0.1, 0.3, 0.6]
y = np.hstack([[ii] * int(100 * perc) for ii, perc in enumerate(percentiles_classes)])

# Generate uneven groups
group_prior = rng.dirichlet([2] * 10)
groups = np.repeat(np.arange(10), rng.multinomial(100, group_prior))


def visualize_groups(classes, groups, name):
    # Visualize dataset groups
    fig, ax = plt.subplots()
    ax.scatter(
        range(len(groups)),
        [0.5] * len(groups),
        c=groups,
        marker="_",
        lw=50,
        cmap=cmap_data,
    )
    ax.scatter(
        range(len(groups)),
        [3.5] * len(groups),
        c=classes,
        marker="_",
        lw=50,
        cmap=cmap_data,
    )
    ax.set(
        ylim=[-1, 5],
        yticks=[0.5, 3.5],
        yticklabels=["Data\ngroup", "Data\nclass"],
        xlabel="Sample index",
    )


visualize_groups(y, groups, "no groups")
plot cv indices

Определение функции для визуализации поведения кросс-валидации

Определим функцию, которая позволяет визуализировать поведение каждого объекта кросс-валидации. Мы выполним 4 разбиения данных. При каждом разбиении мы визуализируем индексы, выбранные для обучающей выборки (синим цветом) и тестовой выборки (красным цветом).

def plot_cv_indices(cv, X, y, group, ax, n_splits, lw=10):
    """Create a sample plot for indices of a cross-validation object."""
    use_groups = "Group" in type(cv).__name__
    groups = group if use_groups else None
    # Generate the training/testing visualizations for each CV split
    for ii, (tr, tt) in enumerate(cv.split(X=X, y=y, groups=groups)):
        # Fill in indices with the training/test groups
        indices = np.array([np.nan] * len(X))
        indices[tt] = 1
        indices[tr] = 0

        # Visualize the results
        ax.scatter(
            range(len(indices)),
            [ii + 0.5] * len(indices),
            c=indices,
            marker="_",
            lw=lw,
            cmap=cmap_cv,
            vmin=-0.2,
            vmax=1.2,
        )

    # Plot the data classes and groups at the end
    ax.scatter(
        range(len(X)), [ii + 1.5] * len(X), c=y, marker="_", lw=lw, cmap=cmap_data
    )

    ax.scatter(
        range(len(X)), [ii + 2.5] * len(X), c=group, marker="_", lw=lw, cmap=cmap_data
    )

    # Formatting
    yticklabels = list(range(n_splits)) + ["class", "group"]
    ax.set(
        yticks=np.arange(n_splits + 2) + 0.5,
        yticklabels=yticklabels,
        xlabel="Sample index",
        ylabel="CV iteration",
        ylim=[n_splits + 2.2, -0.2],
        xlim=[0, 100],
    )
    ax.set_title("{}".format(type(cv).__name__), fontsize=15)
    return ax

Давайте посмотрим, как это выглядит для объекта кросс-валидации KFold:

fig, ax = plt.subplots()
cv = KFold(n_splits)
plot_cv_indices(cv, X, y, groups, ax, n_splits)
KFold
<Axes: title={'center': 'KFold'}, xlabel='Sample index', ylabel='CV iteration'>

Как вы можете видеть, по умолчанию итератор кросс-валидации KFold не учитывает ни класс точек данных, ни группы. Мы можем изменить это, используя:

  • StratifiedKFold для сохранения процента образцов для каждого класса.
  • GroupKFold для обеспечения того, что одна и та же группа не будет появляться в двух разных фолдах.
  • StratifiedGroupKFold для сохранения ограничения GroupKFold при попытке вернуть стратифицированные фолды.
cvs = [StratifiedKFold, GroupKFold, StratifiedGroupKFold]

for cv in cvs:
    fig, ax = plt.subplots(figsize=(6, 3))
    plot_cv_indices(cv(n_splits), X, y, groups, ax, n_splits)
    ax.legend(
        [Patch(color=cmap_cv(0.8)), Patch(color=cmap_cv(0.02))],
        ["Testing set", "Training set"],
        loc=(1.02, 0.8),
    )
    # Make the legend fit
    plt.tight_layout()
    fig.subplots_adjust(right=0.7)
  • StratifiedKFold
  • GroupKFold
  • StratifiedGroupKFold

Далее мы визуализируем это поведение для ряда итераторов CV.

Визуализация индексов кросс-валидации для множества объектов CV

Давайте визуально сравним поведение кросс-валидации для многих объектов кросс-валидации scikit-learn. Ниже мы пройдёмся по нескольким распространённым объектам кросс-валидации, визуализируя поведение каждого из них.

Обратите внимание, как некоторые используют информацию о группе/классе, а другие нет.

cvs = [
    KFold,
    GroupKFold,
    ShuffleSplit,
    StratifiedKFold,
    StratifiedGroupKFold,
    GroupShuffleSplit,
    StratifiedShuffleSplit,
    TimeSeriesSplit,
]


for cv in cvs:
    this_cv = cv(n_splits=n_splits)
    fig, ax = plt.subplots(figsize=(6, 3))
    plot_cv_indices(this_cv, X, y, groups, ax, n_splits)

    ax.legend(
        [Patch(color=cmap_cv(0.8)), Patch(color=cmap_cv(0.02))],
        ["Testing set", "Training set"],
        loc=(1.02, 0.8),
    )
    # Make the legend fit
    plt.tight_layout()
    fig.subplots_adjust(right=0.7)
plt.show()
  • KFold
  • GroupKFold
  • ShuffleSplit
  • StratifiedKFold
  • StratifiedGroupKFold
  • GroupShuffleSplit
  • StratifiedShuffleSplit
  • TimeSeriesSplit

Общее время выполнения скрипта: (0 минут 1,434 секунды)

Launch binder
Launch JupyterLite

Download Jupyter notebook: plot_cv_indices.ipynb

Download Python source code: plot_cv_indices.py

Download zipped: plot_cv_indices.zip

Связанные примеры

Вложенная и невложенная кросс-валидация

Пошаговое исключение признаков с кросс-валидацией

Кривая ROC (Receiver Operating Characteristic) с кросс-валидацией

Основные моменты выпуска scikit-learn 1.4

© 2007–2025 The scikit-learn developers
Licensed under the 3-clause BSD License.
https://scikit-learn.org/1.6/auto_examples/model_selection/plot_cv_indices.html

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API