Кросс-валидация
Кросс-валидация нужна, чтобы оценить качество модели надёжнее, чем при одном разбиении на тренировочную и тестовую выборки.
При обычном TrainTestSplit данные делятся один раз. Но результат может зависеть от случайности: в тестовую часть могут попасть более простые или более сложные объекты. Тогда качество модели будет выглядеть лучше или хуже, чем есть на самом деле.
Кросс-валидация решает эту проблему: данные делятся на несколько частей, и каждая часть по очереди используется для проверки, а остальные — для обучения модели.
Основная идея
Заголовок раздела «Основная идея»Пусть данные делятся на 5 частей.
На первом шаге модель обучается на частях 2-5, а проверяется на части 1.
На втором шаге модель обучается на частях 1, 3, 4, 5, а проверяется на части 2.
И так далее, пока каждая часть не побывает тестовой.
В результате получается несколько оценок качества. После этого обычно смотрят среднее значение.
Шаг 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
Заголовок раздела «StratifiedKFold»StratifiedKFold — вариант для классификации.
Он старается сохранить соотношение классов в каждой части данных. Например, если во всём датасете примерно поровну бульдогов и спаниелей, то в каждой части тоже будет примерно поровну объектов этих классов.
Это особенно важно, если один класс встречается заметно реже другого.
var folds := Validation.StratifiedKFold(y, k := 5, seed := 42);Для классификации в большинстве случаев лучше начинать именно со стратифицированной кросс-валидации.
CrossValidate
Заголовок раздела «CrossValidate»CrossValidate запускает весь цикл кросс-валидации:
- делит данные на несколько частей;
- несколько раз обучает модель;
- проверяет её на разных тестовых частях;
- возвращает среднее значение выбранной метрики.
Для обычного 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 |
Как читать результат
Заголовок раздела «Как читать результат»Кросс-валидация даёт более устойчивую оценку качества, чем одно разбиение.
Если среднее качество высокое, модель в среднем работает хорошо.
Если качество сильно меняется от шага к шагу, модель нестабильна: её результат слишком зависит от того, какие объекты использовались для обучения, а какие — для проверки.