Как с нуля построить регрессионное дерево
В предыдущей статье мы построили классификатор на основе дерева решений, который определяет, есть ли у пациента болезнь сердца. Деревья решений подходят и для задач регрессии: вместо категории они предсказывают число — например, зарплату, температуру или цену дома. В этой статье мы возьмём набор данных Hitters и построим дерево, которое предсказывает зарплату игрока по числу сыгранных сезонов и попаданий в прошлом сезоне.
Классификатор оценивал разбиения по уменьшению критерия Джини. В регрессии с квадратичной ошибкой используем уменьшение дисперсии, а в листе предсказываем среднее. Умножив уменьшение дисперсии на число строк родителя, получим уменьшение общей суммы квадратов ошибок.
| классификатор | регрессор | |
|---|---|---|
| целевой столбец | категория — болезнь или её отсутствие | число — зарплата в тысячах долларов |
| измерение группы | примесь Джини | дисперсия |
| оценка разбиения | прирост информации | уменьшение суммарной квадратичной ошибки |
| что хранит лист | подсчёт меток, попавших в него | среднее строк, попавших в него |
| что отвечает дерево | класс с вероятностями | одно число |
Критерий разбиения и значение, хранимое в каждом листе, — две стороны одного решения. Примесь узла есть ошибка обучения того ответа, который даст его лист. Джини измеряет цену предсказания долей классов для группы меток; дисперсия измеряет цену предсказания среднего для группы зарплат. Меняется предсказание листа — меняется вместе с ним и подходящая мера ошибки.
Генерация кандидатов, рекурсивное разбиение и правило остановки остаются прежними. Изменить нужно только вычисление примеси и предсказание листа.
Скачайте полный пример на Python и запустите python3 hitters_tree.py.
Лист, предсказывающий среднее, ограничивает экстраполяцию: любой прогноз остаётся в диапазоне обучающих целевых значений. Дальнейшие примеры показывают, почему увеличение глубины не позволяет продолжить растущий тренд за пределы наблюдаемого диапазона признака.
Тот же механизм
Сохраняем упрощённый алгоритм CART из классификатора. Каждый внутренний узел выбирает бинарный вопрос, оценивает две полученные группы и продолжает рекурсию. Лист завершает её, когда ни один кандидат не даёт улучшения сверх численной погрешности.
Кандидаты генерируются ровно как раньше — каждый признак в паре с каждым значением, которое он принимает в имеющихся строках, по одной паре на вопрос, причём числовой столбец спрашивает >=, а категориальный ==. Изменился только тип целевого столбца, а предикторы по-прежнему могут быть любого вида. Генератор ничего не знает и знать не хочет о содержимом целевого столбца: пять строк с двумя столбцами по четыре различных значения дают восемь кандидатов независимо от того, предсказываете вы болезнь или зарплату.
Изменить нужно способ оценки группы — то, насколько её строки отличаются друг от друга. В классификаторе функция gini измеряла, насколько перемешаны метки. Для числовой цели используется variance: она показывает, насколько зарплаты отклоняются от своего среднего.
Интересно, что формула прироста информации, которой мы оценивали вопросы, меняться вместе с ней не обязана. Вот так мы оценивали разбиение в классификаторе:
def info_gain(left, right, current_uncertainty):
p = float(len(left)) / (len(left) + len(right))
return current_uncertainty - p * gini(left) - (1 - p) * gini(right)А вот так — в регрессоре, та же функция с одним изменённым именем:
def info_gain(left, right, current_uncertainty):
p = float(len(left)) / (len(left) + len(right))
return current_uncertainty - p * variance(left) - (1 - p) * variance(right)Функция вычитает из дисперсии родителя дисперсии детей, взвешенные по числу строк. Результат — уменьшение средней квадратичной ошибки в этом узле. Зарплаты заданы в тысячах долларов, поэтому единицы дисперсии и SSE — квадраты тысяч долларов.
Пять игроков
Набор данных — пять игроков из исследования Hitters: два числовых предиктора и одна числовая цель.
| # | игрок | годы | попадания | зарплата |
|---|---|---|---|---|
| 1 | BillyJo Robidoux | 2 | 41 | 67.5 |
| 2 | Jack Howell | 2 | 41 | 95.0 |
| 3 | Alvin Davis | 3 | 130 | 480.0 |
| 4 | Mike Marshall | 6 | 77 | 670.0 |
| 5 | Lloyd Moseby | 7 | 149 | 787.5 |
years — число сезонов, отыгранных в высшей лиге, hits — число попаданий в прошлом сезоне, salary — зарплата за сезон 1987 года, в тысячах долларов. Это настоящие строки из настоящего исследования, сверенные с полным файлом.
Измерение разброса набора данных
Для оценки кандидатов нужна мера разброса зарплат. Начнём с суммы квадратов ошибок (SSE), затем разделим её на число строк и получим дисперсию.
Допустим, нужно определить, в какой из двух команд зарплаты различаются сильнее:
| команда | зарплаты |
|---|---|
| A | 400, 410, 420, 430, 440 |
| B | 67.5, 95.0, 480.0, 670.0, 787.5 |
Команда B взята из нашего настоящего набора данных — это зарплаты пяти игроков выше; команда A — гипотетический клуб, где все получают примерно поровну.
Один способ — нанести обе на график:
Одного взгляда достаточно, чтобы понять: вторая команда разбросана сильнее. Но поиску разбиения нужно число, которое можно вычислить. Первое, что приходит в голову, — среднее, и оно уже есть на рисунке в виде пунктирной линии, попадающей у обеих команд ровно в одно и то же место.
Это говорит нам, что среднее не может измерять разброс группы — оно сообщает 420 в обоих случаях. Вместо этого можно измерить, насколько далеко каждая зарплата стоит от среднего:
| команда | отклонения от 420 | сумма |
|---|---|---|
| A | −20, −10, 0, +10, +20 | 0 |
| B | −352.5, −325.0, +60.0, +250.0, +367.5 | 0 |
Использовать отклонения как есть мы не можем, потому что оба набора дают в сумме ноль — свойство среднего: оно является точкой равновесия набора, поэтому всё, что стоит выше, ровно компенсирует всё, что стоит ниже.
Возведение в квадрат не даёт отклонениям разных знаков взаимно уничтожаться и сильнее учитывает большие ошибки. Абсолютная ошибка — другая возможность, но ей соответствует медиана вместо среднего. Для поиска разбиений дифференцируемость не нужна.
Возведите отклонения в квадрат и сложите — получится SSE; поделите это на их количество — получится дисперсия:
| команда | квадраты отклонений | сумма (SSE) | количество | дисперсия |
|---|---|---|---|---|
| A | 400, 100, 0, 100, 400 | 1,000 | 5 | 200.00 |
| B | 124,256.25, 105,625, 3,600, 62,500, 135,056.25 | 431,037.5 | 5 | 86,207.50 |
SSE — сумма, дисперсия — SSE на строку. Дублирование всех строк удваивает SSE, не меняя дисперсию. Добавление произвольных новых строк может изменить обе величины. Здесь дисперсия команды B примерно в 431 раз больше: 86,207.50 против 200.00.
Записанная формулой, дисперсия выглядит так:
где — сколько строк содержит группа, — зарплата одной из них, а — их среднее, так что есть отклонение одной строки, тот самый столбец, который мы свели в таблицу выше.
Уберите , и то, что останется, , — это SSE.
Только что построенная формула известна как дисперсия генеральной совокупности, и, посмотрев её, вы найдёте рядом вторую версию, выборочную дисперсию, отличающуюся только знаменателем:
Здесь делим на , поскольку измеряем среднюю квадратичную ошибку строк данного узла. Поправка нужна для другой задачи — оценки дисперсии популяции по случайной выборке. В вычислении этой обучающей ошибки она не требуется.
Почему это правильная примесь
Разберёмся теперь, почему именно эта мера подходит для оценки узла. Когда дерево построено, каждый лист хранит среднее своей группы, и на этапе предсказания это среднее — зарплата, которую он предсказывает каждому новому игроку, попавшему в него: одно число на всех. Группа с широким разбросом делает это одно число сильно неверным для многого, что в ней лежит, и именно поэтому поиск разбиения хочет групп с как можно меньшим разбросом.
Допустим, дерево вообще не разбилось: один лист, содержащий целую команду и предсказывающий среднюю зарплату 420.0 каждому игроку в ней.
Для команды A это хороший ответ: никто там не зарабатывает дальше 20 от него, так что худшее, чем может ошибиться лист, — это 20. Для команды B ответ плохой: Робиду зарабатывает 67.5, а Мозби 787.5, и обоим говорят 420.0, промахиваясь на 352.5 и 367.5. Возведите эти промахи в квадрат, усредните — и вы вернётесь к двум числам из таблицы, 200.00 и 86,207.50: те же дисперсии, прочитанные теперь как ошибка, которую сделал бы каждый лист. Итак, разброс внутри узла — это ошибка, которую сделал бы ответ этого узла, и именно это число дерево сравнивает, решая, какой вопрос задать, предпочитая тот, чьи две группы оставляют её меньше всего.
Всё это уже есть в формуле . Здесь — те 420.0, которые лист предсказывает каждому, кто в него попал. Каждое — ошибка для одного игрока: для Робиду, для Мозби и не хуже для любого в команде A. Возведение этих ошибок в квадрат и усреднение даёт нам 86,207.50 и 200.00. Дисперсия, таким образом, есть средняя квадратичная ошибка обучения листа, предсказывающего среднее — ошибка, которую всё равно сделал бы лучший возможный постоянный лист.
Та же функция потерь, что и в линейной регрессии, на другом семействе функций
Квадратичную ошибку минимизирует и метод наименьших квадратов (OLS), и дерево не делает с ней ничего иного. Обе меряют одинаково, относительно истинного значения: минус то, что модель предсказывает для этой строки. Отличается то, что модели разрешено предсказывать: линейная регрессия даёт каждой строке своё число, считанное с прямой в точке этой строки, тогда как дерево даёт всем строкам в листе одно и то же число.
Аккуратнее всего эту связь видно так: лист — это регрессия с одним свободным членом. Подгоните OLS вообще без предикторов, и оценкой окажется — та самая константа, которую хранит лист, минимизирующая ту же сумму квадратов. Дерево — набор таких регрессий, по одной на регион, а поиск разбиений делает работу по выбору регионов. Поэтому же дисперсия группы здесь является её ошибкой обучения: дисперсия — это средний квадрат расстояния до среднего, а среднее — ровно то, что предсказывает лист.
Линейная регрессия оптимизирует коэффициенты, а это дерево перебирает конечный набор разбиений. Пороги и средние в листьях — обученные числовые значения, но алгоритм не обновляет их градиентным спуском. Для фиксированного разбиения оптимальная константа листа сразу вычисляется как среднее.
Классификация и регрессия используют одну и ту же структуру. Лист хранит лучшее постоянное предсказание, а примесь измеряет ошибку этого предсказания. Для классификации это доли классов и примесь Джини; для регрессии — среднее и дисперсия. Примесь вычисляется для каждого кандидата в каждом узле, поэтому её замена может сдвинуть каждое разбиение в дереве. Статистика листа при выборе разбиений не используется: она определяет, что предсказывает готовое дерево, а не какой оно формы.
Среднее или медиана: выбор критерия
Квадратичную ошибку мы выбрали несколькими абзацами выше, когда у отклонений нужно было убрать знаки и мы возвели их в квадрат, а не взяли модули. Это то, что библиотеки ставят по умолчанию: например, scikit-learn использует criterion="squared_error", если не сказать иного. Но это не единственный вариант, и выбор заходит дальше, чем кажется: функция потерь, которой мы оцениваем разбиения при построении, задаёт и то, что лист обязан хранить для предсказания, потому что и то и другое — ответы на один вопрос: какая единственная константа минимизирует эту потерю.
- Минимизируем квадратичную ошибку → лист хранит среднее.
- Минимизируем абсолютную ошибку → лист хранит медиану.
Возьмём лист со значениями [10, 12, 14, 16, 200], где 200 — выброс или ошибка ввода данных:
| предсказание листа | суммарная квадратичная ошибка | суммарная абсолютная ошибка |
|---|---|---|
| среднее = 50.40 | 27,995.2 | 299.2 |
| медиана = 14.00 | 34,620.0 | 194.0 |
Среднее минимизирует квадратичную ошибку, медиана — абсолютную. Увеличение наибольшего значения сверх 200 сдвинет среднее, но оставит медиану равной 14. Такая устойчивость к экстремальному значению делает листья с абсолютной ошибкой полезными при выбросах.
Итак, выбранная нами примесь решает, какую константу хранит лист, и поэтому две функции ниже написаны парой.
def mean(rows):
return sum(row[-1] for row in rows) / float(len(rows))
def variance(rows):
targets = [row[-1] for row in rows]
m = sum(targets) / len(targets)
return sum((t - m) ** 2 for t in targets) / len(targets)variance оценивает узел, а mean задаёт прогноз. При абсолютной ошибке соответствующая мера — среднее абсолютное отклонение от медианы, и лист хранит медиану.
Оценка разбиения
Итак, мы умеем измерять разброс одной группы и можем посмотреть, как использовать это измерение для оценки вопроса. Наша цель — выяснить, сколько квадратичной ошибки остаётся после того, как вопрос разделил строки, или, что то же самое, сколько её вопрос убрал. Сделать это можно двумя путями, работая в суммах или в средних:
| суммы | средние | |
|---|---|---|
| ошибка одной группы | SSE | дисперсия, SSE делённая на , — её же называют MSE |
| ошибка, оставшаяся после разбиения | две дисперсии, взвешенные размерами групп | |
| ошибка, убранная разбиением | то же вычитание, взвешенное |
Они упорядочивают вопросы-кандидаты одинаково, так что от выбора ничего не зависит. Арифметика в суммах проще, с неё и начнём; усреднённая форма — то, что вычисляет код, и мы вернёмся к ней, когда критерий будет на месте.
Вопрос делит узел на две группы, и каждая отвечает своим средним — с одной стороны, с другой. Значит, формулу SSE можно применить к каждой группе отдельно, относительно её собственного среднего. Сложив эти две SSE, мы получим остаточную сумму квадратов (RSS) — ошибку, которую разбиение оставляет после себя:
Чтобы выбрать лучший вопрос из кандидатов, мы вычисляем это для каждого и берём тот, у кого RSS наименьшая, — и этого уже достаточно, чтобы построить регрессионное дерево: оцените каждого кандидата, оставьте наименьшего, повторите рекурсивно на двух получившихся группах.
Глядя на эту формулу, можно задаться вопросом, пытается ли критерий сохранить дерево маленьким или сгруппировать строки в аккуратные кучки. Он не делает ни того, ни другого: он жадно ищет заданные признаками разделения, после которых цели в каждом потомке проще предсказать одним значением листа, — меньше разнородности классов в классификации, меньше квадратичного разброса вокруг среднего в регрессии. Он оценивает две группы в одном узле, берёт победителя и рекурсивно повторяет; во что обойдётся готовое дерево целиком, никогда не рассматривается.
Оценка по тому, что разбиение убирает
Есть другой способ использовать SSE для оценки вопроса. Примените ту же формулу к родительской группе — строкам до разбиения, индекс — относительно её собственного среднего , и вы получите — ошибку, которую эта группа делает как есть, отвечая каждой строке одним числом. Так что вместо вопроса, сколько ошибки разбиение оставляет, можно спросить, сколько оно убрало: родительская ошибка минус то, что два потомка всё ещё несут:
Это форма прироста информации, записанная в суммах, а не во взвешенных средних.
SSE родителя фиксирована при сравнении кандидатов в одном узле. Вычитание остаточной ошибки каждого кандидата из этой константы обращает порядок: наименьший остаток даёт наибольший выигрыш.
| кандидат | осталось ошибки | убрано ошибки |
|---|---|---|
| A | 60 | 100 − 60 = 40 |
| B | 25 | 100 − 25 = 75 |
Кандидат, который оставляет меньше всех, убрал больше всех. Минимизация RSS и максимизация прироста делают одну работу с двух концов.
В точной арифметике выигрыш по квадратичной ошибке неотрицателен: собственное среднее ребёнка не хуже среднего родителя для его строк. Нулевой непосредственный выигрыш не означает бесполезности всех дальнейших разбиений. Наша жадная реализация останавливается при пренебрежимо малом выигрыше; глубина, размер листьев и обрезка дополнительно контролируют сложность.
Тот же критерий в записи через дисперсию
Прирост записан в суммах, тогда как info_gain — функция из статьи о классификаторе, с variance там, где она звала gini, — записана в средних, взвешивая каждого потомка его долей строк:
def info_gain(left, right, current_uncertainty):
p = float(len(left)) / (len(left) + len(right))
return current_uncertainty - p * variance(left) - (1 - p) * variance(right)p — это доля строк родителя, оказавшихся в левой группе: три строки из пяти дают p = 0.6, оставляя 1 - p = 0.4 правой группе, — и она нужна потому, что мы работаем в средних. Дисперсия ничего не говорит о том, сколько строк её породило, так что потомок из одной строки и потомок из сотни считались бы наравне; а если мы хотим и дальше пользоваться дисперсиями, размеры приходится возвращать руками, чем p и 1 - p ровно и занимаются.
В записи через суммы никакие веса не нужны вовсе:
def sse(rows):
m = mean(rows)
return sum((row[-1] - m) ** 2 for row in rows)
def gain_sse(rows, left, right):
return sse(rows) - (sse(left) + sse(right))Обе формы оценивают одних и тех же кандидатов в одном порядке, а усреднённая — просто то, что у классификатора уже было, ведь Джини тоже среднее. Друг в друга они переводятся одним тождеством, поскольку дисперсия есть SSE на строку:
Подставьте это для всех трёх групп, и прирост станет
а деление на — снова константу в этом узле, снова безвредно — даёт форму, которой пользовалась статья о классификации, и ту, что вычисляет info_gain:
Минимизация SSE детей, максимизация уменьшения SSE и максимизация уменьшения дисперсии выбирают одно разбиение внутри фиксированного родителя. Значения связаны константой или множителем, равным числу строк родителя. Код использует уменьшение дисперсии, сохраняя структуру оценки классификатора.
Оценка кандидатов в корне
Прогоним теперь этот критерий на первом узле настоящего дерева — на корне, который держит всех пятерых игроков, до того как задан хоть один вопрос. Его среднее равно , а дисперсия — 86,207.50, так что по квадратичная ошибка, которую мы пытаемся уменьшить, есть . Каждый кандидат ниже оценивается относительно этого числа в усреднённой форме — это gain, который печатает наш код, минус две взвешенные по размеру дисперсии потомков, — а победителя мы потом прочитаем как RSS.
Вопросы-кандидаты генерируются ровно как раньше — каждый признак в паре с каждым значением, которое он принимает. Два столбца по четыре различных значения дают восемь кандидатов, двое из которых вообще не делят строки. Здесь они выложены от лучшего, хотя код их никогда не сортирует; он просто держит текущего победителя:
| кандидат | прирост | слева / справа |
|---|---|---|
Is years >= 3? | 76501.0417 | 3 / 2 |
Is hits >= 77? | 76501.0417 | 3 / 2 |
Is years >= 6? | 63551.0417 | 2 / 3 |
Is years >= 7? | 33764.0625 | 1 / 4 |
Is hits >= 149? | 33764.0625 | 1 / 4 |
Is hits >= 130? | 30459.3750 | 2 / 3 |
Is years >= 2? | не делит | 5 / 0 |
Is hits >= 41? | не делит | 5 / 0 |
Два пропущенных кандидата — наименьшее значение в каждом столбце: years пробегает 2, 2, 3, 6, 7, а hits — 41, 41, 77, 130, 149, — так что каждая строка отвечает на них «да». Всё уходит на сторону True и ничего на сторону False, отсюда 5 / 0: разбиения, которое можно оценить, нет, и они отбрасываются.
Столбец gain — то, что сообщает код, поскольку оценку делает функция info_gain. Прочитайте победителя как суммарную квадратичную ошибку — и критерий станет виднее. Is hits >= 77? отправляет Дэвиса, Маршалла и Мозби в одну сторону, а двух совпавших игроков — в другую:
против в корне. Один вопрос убирает 89% квадратичной ошибки набора данных, и ни один другой предложенный вопрос не оставляет меньше. (Столбец прироста — тот же факт в пересчёте на наблюдение: убрано , и .)
Два кандидата дают ровную ничью на вершине — Is years >= 3? и Is hits >= 77?, оба 76501.0416666667, потому что режут пятерых игроков на те же две группы. >= в find_best_split отдаёт победу тому столбцу, который просматривается последним, ровно как и в статье о классификации.
Готовое дерево и то, что хранят его листья
Рекурсия заканчивается, когда ни один кандидат не улучшает результат сверх численной погрешности. Для этих пяти строк получается:
Три листа держат ровно по одному игроку и воспроизводят их зарплату в точности. Четвёртый держит совпавшую пару и отвечает 81.25, среднее 67.5 и 95.0.
Прогоните пятерых игроков обратно вниз по готовому дереву — тот же путь, который прошёл бы новый игрок, каждый следует за вопросами до листа и берёт хранимое им число, — и вот что вернётся:
BillyJo Robidoux факт 67.5 прогноз 81.25
Jack Howell факт 95.0 прогноз 81.25
Alvin Davis факт 480.0 прогноз 480.00
Mike Marshall факт 670.0 прогноз 670.00
Lloyd Moseby факт 787.5 прогноз 787.50Три обучающие зарплаты воспроизводятся точно, а два игрока с одинаковыми признаками получают общее среднее. Это качество на обучении; результат для новых игроков мы ещё не измеряли.
Эти 81.25 наглядно показывают, почему в листе мы держим среднее, а не что-нибудь другое — скажем, меньшую из двух зарплат или большую. Лист обязан выдать одно число, назовём его , а минимизируемая нами потеря — квадратичная ошибка, так что вопрос в том, какое делает как можно меньше. Это гладкая функция от , поэтому минимум там, где её производная обращается в ноль:
Значит, среднее — не один из нескольких разумных вариантов и не соглашение: это единственная константа, которая этому удовлетворяет, и это решение той же минимизации, которую выполняет критерий разбиения. Примесь и значение листа происходят из одной функции потерь.
На этом листе дерево и перестаёт улучшаться. У Робиду и Хауэлла одинаковые years и hits, так что никакой вопрос никогда их не разделит: они делят лист на любой глубине, и каким бы числом этот лист ни отвечал, оно неверно хотя бы для одного из них. Ответ 81.25 оставляет квадратичной ошибки, и никакое дерево, читающее только эти два столбца, не опустит её ниже — это пол под ошибкой обучения, который не пробить ростом вглубь.
Среди 263 игроков с известной зарплатой девять пар имеют одинаковые (Years, Hits). Максимальная разница зарплат в паре — 310 000 долларов, что даёт минимальную среднюю абсолютную ошибку 155 000 долларов для общей предсказанной величины. Дополнительные признаки могут различить игроков. Это ограничение записанных признаков, а не доказательство принципиальной непредсказуемости зарплаты.
Внутри данных: лестница
Посмотрим теперь на форму, которую рисует готовое дерево, а затем на то, что происходит за краем обучающих данных. Это поможет увидеть, что регрессионное дерево может выразить, а что нет. Регрессионное дерево кусочно-постоянно: его вопросы делят пространство входов на регионы, и каждая точка региона получает одно и то же предсказание. Как функция признаков его выход — набор плоских площадок с вертикальными скачками на порогах, без каких-либо наклонов, на любой глубине. Ниже мы будем называть это свойство плоскостью.
Чтобы показать это попроще, ниже мы рисуем синтетический набор данных, а не Hitters: 40 строк с одним признаком , равномерно расположенным от 0 до 10, и целью, идущей по гладкой волне с небольшим шумом.
| x | 0 | 0.256 | 0.513 | 0.769 | … | 9.744 | 10 |
|---|---|---|---|---|---|---|---|
| y | 45.90 | 52.29 | 57.36 | 58.86 | … | 79.48 | 82.41 |
Виджет обучает регрессионное дерево на этих строках. Пунктирные линии отмечают пороги, а каждый интервал получает среднее своих обучающих целей. Увеличьте предел глубины, чтобы увидеть больше интервалов:
На глубине 1 площадок две и подгонка ужасна. К глубине 6 их 27, а квадратичная ошибка упала с 5,816 до 89 — это ошибка всей подгонки: каждая из 40 точек прошла вниз по дереву и была оценена относительно среднего, хранимого листом, в который она попала, ровно та сумма, которую мы называли RSS, теперь по 27 листьям, а не по двум. Дерево сходится к кривой, но никогда не изгибается — оно приближает гладкую функцию, рубя её на всё более узкие постоянные куски.
Дерево может приближать нелинейные зависимости и взаимодействия признаков без заранее заданной формы. Но конечное дерево с постоянными листьями не воспроизведёт точно непостоянную прямую на непрерывном интервале: для лучшего приближения нужно больше ступеней.
С двумя признаками лестница становится рельефом
Один признак даёт лестницу, потому что есть одна ось, вдоль которой раскладывать ступени, и предсказание на другой. Добавьте второй признак — и обе оси заняты, так что предсказанию приходится идти куда-то ещё: дерево режет плоскость на прямоугольники, и то, что было высотой ступени, становится высотой плоской крыши над каждым из них.
Две панели ниже — одна и та же модель. Слева разбиение, увиденное сверху, где предсказание проявляется в закраске и числе внутри каждого прямоугольника, — картинка, которую рисовала статья о классификации для решающих регионов. Справа те же коробки подняты до этого числа, так что высота несёт то, что несла ось на лестнице:
Каждая крыша плоская, а каждая стена вертикальная — вот как выглядит кусочная постоянность, когда её видно. Поднимайте глубину, и рельеф набирает блоки так же, как лестница набирала ступени: 2 региона, потом 4, 8, 16 — приближаясь к форме данных плоскими гранями, никогда наклонами.
За пределами данных: потолок
В одномерной лестнице любое значение выше наибольшего обученного порога попадает в один крайний лист. При нескольких признаках увеличение одного сверх всех его порогов больше не меняет решения по нему, но остальные признаки всё ещё могут направлять строки в разные листья.
Механизм за этим не уникален для регрессии. Регионы классификатора точно так же плоски, и он тоже отвечает на всё за пределами обучающих данных тем, что держит его крайний лист. Для регрессии ограничение особенно заметно, потому что у целевых значений есть порядок. Если зарплата продолжает расти за пределами обучающего диапазона, дерево не может за ней последовать; оно продолжает возвращать значение, сохранённое в крайнем листе. У меток нет аналогичного направления — нет категории выше «болезни», — поэтому в классификации то же поведение менее очевидно.
Это поведение прямо следует из того, как работает предсказание: строка, приходящая с , отвечает «да» на каждый пороговый вопрос по пути вниз, попадает в самый правый лист и получает среднее обучающих строк, которые попали туда же. Нет механизма, которым значение листа могло бы зависеть от того, насколько далеко за порогом находится строка.
Вот максимально острая демонстрация — идеально линейная зависимость вообще без шума, , отобранная на , приближённая деревом глубины 3 и обычной линейной регрессией:
import numpy as np
from sklearn.linear_model import LinearRegression
from sklearn.tree import DecisionTreeRegressor
X = np.linspace(0, 10, 60).reshape(-1, 1)
y = 2.5 * X.ravel() + 3
tree = DecisionTreeRegressor(max_depth=3, random_state=0).fit(X, y)
linear = LinearRegression().fit(X, y)Обе модели обучены на одних и тех же 60 строках, все они внутри . Спросим теперь каждую о значениях внутри этого диапазона, а затем далеко за ним:
| x | истинное значение | дерево | линейная регрессия |
|---|---|---|---|
| 2 | 8.00 | 7.45 | 8.00 |
| 5 | 15.50 | 13.81 | 15.50 |
| 8 | 23.00 | 23.34 | 23.00 |
| 12 | 33.00 | 26.52 | 33.00 |
| 20 | 53.00 | 26.52 | 53.00 |
| 50 | 128.00 | 26.52 | 128.00 |
| 1000 | 2503.00 | 26.52 | 2503.00 |
Прорисовано до , при том что обучающий диапазон кончается на 10:
Линейная модель восстанавливает правило в точности и верна при . Дерево возвращает 26.52 при , при и при — одно число на любой вход за пределами его опыта, на данных без шума и с зависимостью, которую прямая ловит двумя параметрами.
Внутри обучающего диапазона дерево тоже ошибается из-за приближения: при выдаёт 7.45 вместо 8.00. В отдельных точках прогноз может быть точным, но восемь постоянных ступеней не воспроизводят прямую повсюду.
Хуже плоскости — граница. Лист хранит среднее обучающих целей, которые в него попали, а среднее не может лежать вне усредняемых значений. Так что любое предсказание, которое регрессионное дерево вообще способно сделать, заперто в диапазоне обучающих целей. Оно не может предсказать рекордный максимум или рекордный минимум ни на какой глубине и ни на каких данных.
Наибольшее значение листа в линейном примере — 26.5169, ниже обучающего максимума 28.0. Дерево может достичь максимума, если лист содержит только это значение. Для сравнения обучим модель на всех 19 признаках Hitters: 184 обучающие строки из разбиения 70/30 с random_state=0:
глубина 2 наибольшее возможное предсказание 2127.3 достигается 1 игроком
глубина 3 наибольшее возможное предсказание 2127.3 достигается 1 игроком
глубина 5 наибольшее возможное предсказание 2127.3 достигается 1 игроком
полная глубина наибольшее возможное предсказание 2127.3 достигается 1 игрокомВ этом обучении дерево глубины 2 уже отделяет самого высокооплачиваемого игрока. Это зависит от признаков и остальных строк, а не только от экстремальности цели. Граница неизменна: среднее в листе не может превысить наибольшую цель внутри него.
Больше всего это важно, когда экстраполяция входит в задачу:
- Тренды. При фиксированных остальных признаках дерево перестаёт менять прогноз, когда растущий временной признак проходит все свои пороги. Может помочь отдельная модель тренда.
- Случайные леса. Среднее прогнозов деревьев со средними в листьях остаётся в диапазоне обучающих целей.
- Бустинг деревьев. Сумма может выйти за этот диапазон, но конечный ансамбль деревьев с постоянными листьями остаётся кусочно-постоянным. Вдоль фиксированного направления прогноз перестаёт меняться после пересечения последней границы разбиения.
Ограничение происходит из того, что сидит в листе, а не из разбиения. Некоторые варианты деревьев подгоняют в каждом листе линейную модель вместо константы. Это позволяет экстраполировать, хотя предсказания далеко за данными тогда сильно зависят от подогнанного наклона.
От одного дерева к ансамблю
Превращение классификатора в регрессор потребовало всего двух изменений: использовать дисперсию для оценки узлов и хранить среднее в каждом листе. Получившаяся модель гибка внутри обучающего диапазона, но её кусочно-постоянные предсказания не могут продолжить тренд за его пределы.
Регрессионные деревья часто используют в ансамблях. Случайный лес обучает деревья на бутстрэп-выборках, рассматривает случайные подмножества признаков при разбиениях и усредняет прогнозы. Это снижает чувствительность к одному конкретному дереву.
Подгоняйте их вместо этого по очереди к ошибкам друг друга — и получится градиентный бустинг, где минимизируемая квадратичная ошибка принадлежит всему ансамблю. Каждый раунд измеряет, что накопленный ансамбль всё ещё делает неверно, и подгоняет следующее дерево к этим остаткам, так что каждое дерево выполняет ровно тот поиск разбиений, что и в этой статье, только против цели, составленной из текущих ошибок, а не из сырых зарплат. Поэтому же эти деревья делают работу даже тогда, когда задача — классификация: подгоняют их к столбцу вещественных градиентов, а не к меткам.
Ни большая глубина, ни дополнительные раунды бустинга не гарантируют нулевой обучающей ошибки: здесь этому уже мешают одинаковые признаки с разными целями. Глубину, размер листа и число моделей выбирают по валидации, а не только по обучающему качеству.