Перейти к содержимому

Кросс-валидация

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

При обычном TrainTestSplit данные делятся один раз. Но результат может зависеть от случайности: в тестовую часть могут попасть более простые или более сложные объекты. Тогда качество модели будет выглядеть лучше или хуже, чем есть на самом деле.

Кросс-валидация решает эту проблему: данные делятся на несколько частей, и каждая часть по очереди используется для проверки, а остальные — для обучения модели.

Пусть данные делятся на 5 частей.

На первом шаге модель обучается на частях 2-5, а проверяется на части 1.

На втором шаге модель обучается на частях 1, 3, 4, 5, а проверяется на части 2.

И так далее, пока каждая часть не побывает тестовой.

Схема KFold: каждая часть по очереди используется для проверки

В результате получается несколько оценок качества. После этого обычно смотрят среднее значение.

Шаг 1: Accuracy = 0.83
Шаг 2: Accuracy = 0.79
Шаг 3: Accuracy = 0.86
Шаг 4: Accuracy = 0.81
Шаг 5: Accuracy = 0.84
Среднее качество: примерно 0.83

Так оценка меньше зависит от одного случайного разбиения.

KFold — обычное разбиение данных на k частей.

Оно подходит для регрессии и для простых случаев классификации, когда классы распределены достаточно равномерно.

var folds := Validation.KFold(X.RowCount, k := 5, seed := 42);

Обычно вручную работать с folds не нужно. Чаще используется CrossValidate, который сам выполняет обучение и проверку на каждом разбиении.

StratifiedKFold — вариант для классификации.

Он старается сохранить соотношение классов в каждой части данных. Например, если во всём датасете примерно поровну бульдогов и спаниелей, то в каждой части тоже будет примерно поровну объектов этих классов.

Это особенно важно, если один класс встречается заметно реже другого.

var folds := Validation.StratifiedKFold(y, k := 5, seed := 42);

Для классификации в большинстве случаев лучше начинать именно со стратифицированной кросс-валидации.

CrossValidate запускает весь цикл кросс-валидации:

  1. делит данные на несколько частей;
  2. несколько раз обучает модель;
  3. проверяет её на разных тестовых частях;
  4. возвращает среднее значение выбранной метрики.

Для обычного KFold используется Validation.CrossValidate.

Для классификации со сохранением пропорций классов StratifiedKFold используется Validation.StratifiedCrossValidate.

В примере модель несколько раз обучается на данных о породах собак. В качестве метрики используется Accuracy.

uses MLABC;
begin
var df := DataFrame.FromCsvText('''
Вес,ВысотаВХолке,Порода
20,33,бульдог
22,34,бульдог
24,35,бульдог
25,36,бульдог
19,35,бульдог
21,35,бульдог
26,34,бульдог
20,36,бульдог
14,39,спаниель
15,40,спаниель
16,41,спаниель
18,42,спаниель
17,43,спаниель
15,38,спаниель
19,40,спаниель
16,40,спаниель
''');
var X := df.ToMatrix(['Вес', 'ВысотаВХолке']);
var target := df.EncodeTarget('Порода');
var y := target.Labels;
var model := new KNNClassifier(5);
var acc := Validation.StratifiedCrossValidate(
model,
X,
y,
k := 4,
Metrics.Accuracy,
seed := 42
);
Println('Средняя Accuracy:', acc:0:3);
end.

Здесь используется StratifiedCrossValidate, потому что задача классификации: важно, чтобы в каждой части данных сохранились оба класса в соответствующих пропорциях.

Для регрессии можно использовать обычный CrossValidate.

uses MLABC;
begin
var df := DataFrame.FromCsvText('''
Площадь,Цена
35,4.8
42,5.7
50,7.9
58,7.9
65,8.2
72,10.4
80,10.5
90,12.5
''');
var X := df.ToMatrix(['Площадь']);
var y := df.ToVector('Цена');
var model := new LinearRegression;
var r2 := Validation.CrossValidate(
model,
X,
y,
k := 4,
Metrics.R2,
seed := 42
);
Println('Средний R2:', r2:0:3);
end.

Здесь модель проверяется на разных частях данных, а итогом является среднее значение R2.

ЗадачаЧто использовать
РегрессияValidation.CrossValidate
Классификация с равномерными классамиValidation.CrossValidate или Validation.StratifiedCrossValidate
Классификация с несбалансированными классамиValidation.StratifiedCrossValidate
Маленький датасет для классификациичаще Validation.StratifiedCrossValidate

Кросс-валидация даёт более устойчивую оценку качества, чем одно разбиение.

Если среднее качество высокое, модель в среднем работает хорошо.

Если качество сильно меняется от шага к шагу, модель нестабильна: её результат слишком зависит от того, какие объекты использовались для обучения, а какие — для проверки.