Как построить дерево решений с нуля
Нейросети прекрасны для неструктурированных данных — пикселей, звуковых волн, последовательностей символов — таких, где отдельное входное значение само по себе ничего не говорит. Дайте глубокой сети достаточно этого, и она выучит собственные признаки: пиксельные узоры, значения слов и звуковые текстуры, которые было бы очень трудно сконструировать вручную, — а затем проведёт между классами любую границу, какую потребуют данные, сколь угодно замысловатую.
Но огромная доля реального машинного обучения работает не на таких данных. Она работает на табличных данных — данных, живущих в таблице, в форме электронной таблицы или результата запроса к базе. Каждая строка представляет один образец или наблюдение — транзакцию, пациента, доставку, игрока; каждый столбец — признак с собственным именем и смыслом.
Хороший пример — кливлендское исследование сердечных заболеваний, где каждая строка — пациент, а последний столбец — то, что мы хотим предсказать:
| age | sex | chest_pain | cholesterol | max_heart_rate | vessels | disease |
|---|---|---|---|---|---|---|
| 63 | 1 | typical | 233 | 150 | 0 | No |
| 67 | 1 | asymptomatic | 286 | 108 | 3 | Yes |
| 37 | 1 | nonanginal | 250 | 187 | 0 | No |
Последний столбец — подкрашенный выше — содержит наблюдённое событие: у этого пациента сердечное заболевание обнаружилось, у того — нет. Такой столбец называется меткой; каждый столбец до него описывает случай, и только этот говорит, чем случай кончился. Выучить мы хотим связь между тем и другим — на случаях, где метка уже известна, чтобы применять её к случаям, где нет. Когда метка — категория, как здесь, это задача классификации.
Метод, который чаще всего побеждает на данных такой формы, — градиентный бустинг: обучать одну небольшую модель за другой, каждая из которых исправляет ошибки, оставленные предшественниками. Рецепту всё равно, что это за небольшая модель, но на практике сочетание градиентного бустинга и деревьев решений показывает лучшие результаты на табличных данных. Дерево решений — модель, которая предсказывает, задавая вопросы о столбцах и следуя за ответами вниз до вердикта; бустинг, как и случайный лес, — ансамбль: сотни деревьев, чьи ответы объединяются в один, что бьёт любое одиночное дерево.
Обычно с градиентным бустингом знакомятся под именем реализации, а не метода — XGBoost и его собратья LightGBM и CatBoost, — и даже в эпоху LLM они тихо крутят немалую долю продакшен-ML. Uber оценивает время прибытия распределённым XGBoost, Stripe ловит им мошеннические сети, а Dropbox гоняет XGBoost-ранжировщик внутри своего корпоративного поиска на LLM.
На табличных данных встречаются две задачи: классификация, где ответ — категория, и регрессия, где ответ — число: зарплата, цена. Дерево решений справляется с обеими. Эта статья строит классификатор; направьте тот же код на числовую цель, замените меру примеси дисперсией этой цели, а подсчёты меток в листе — их средним, и вы получите регрессионную модель, предсказывающую число. Больше в коде ничего не меняется, хотя в получившейся модели меняется многое — этому посвящена статья-компаньон.
Бустингу посвящена отдельная статья. Эта сосредоточена на том, что он повторяет сотни раз: как из таблицы строк строится одно дерево решений — откуда берутся его вопросы, как один из них выбирается среди прочих и когда разбиение останавливается. Мы пишем это на чистом Python, без NumPy и без scikit-learn, на пяти строках, достаточно маленьких, чтобы проверить каждое число вручную.
Два типа узлов и области, которые они вырезают
Мы будем работать с пятью пациентами из кливлендского исследования сердечных заболеваний. Полная таблица содержит 303 пациента и тринадцать признаков (предикторов), но мы возьмём лишь пять строк и два признака плюс метку:
| # | stress_test | vessels | disease |
|---|---|---|---|
| 1 | normal | 0 | No |
| 2 | fixed | 0 | Yes |
| 3 | reversable | 2 | Yes |
| 4 | reversable | 1 | Yes |
| 5 | fixed | 0 | No |
По результату стресс-теста и числу сосудов дерево должно предсказать, есть ли у пациента сердечное заболевание, — так что стоит понимать, что записывают эти два столбца.
stress_test (Thal в исходном файле) — таллиевый стресс-тест, который отображает кровоток в сердечной мышце в покое и под нагрузкой: normal означает, что кровоток выглядит нормально, fixed-дефект голодает в обоих состояниях (ткань уже мертва после ранее перенесённого инфаркта), а reversable-дефект голодает только под нагрузкой — суженный, но ещё живой сосуд.
vessels (Ca) — сколько крупных коронарных сосудов, от 0 до 3, оказались поражёнными по данным флюороскопии.
Мы возьмём эти пять строк и будем делить их на всё более мелкие группы, используя признаки и их значения, чтобы решать, в какую группу попадает каждая строка. Последовательность разбиений рисуется как ветвящаяся диаграмма — отсюда и название «дерево».
На этой диаграмме два типа узлов.
Ромбы — узлы решения: каждый содержит один вопрос, задающий логику ветвления, с двумя исходящими рёбрами, True и False, и каждая пришедшая строка отправляется по одному из них.
В зависимости от алгоритма у узла может быть больше двух ветвей: ID3 и C4.5 дали бы stress_test по ветви на значение и разветвили бы его сразу в три потомка. Мы реализуем CART, который задаёт только вопросы «да/нет», так что двух рёбер узлу здесь всегда достаточно. Это же используют и все мейнстримные реализации — бинарные разбиения, оцениваемые мерой примеси, — от DecisionTreeClassifier из sklearn до деревьев внутри случайных лесов и XGBoost.
Прямоугольники — листья: путь здесь заканчивается, больше ничего не спрашивается, и всё, что лист хранит, — это подсчёт меток обучающих пациентов, пришедших по тому же пути. Эти счётчики и есть его ответ любому новому пациенту, который сюда попадёт, — поэтому прямоугольники выше читаются как Yes: 2, Yes: 1, No: 1 и No: 1: двое пациентов с поражённым сосудом, оба больны; двое неразличимых, которые расходятся; один ясный случай без болезни. Поделите эти счётчики на их сумму — и получите вероятность: 100% Yes в первом листе, 100% No в последнем и 50/50 в среднем.
Эти два типа узлов и есть вся модель: узлы решения маршрутизируют строку, листья на неё отвечают, и каждый пациент, входящий сверху, оказывается ровно в одном листе. Дерево решений разбивает строки таблицы на группы — и эти группы можно нарисовать как области пространства всех возможных пар (стресс-тест, число сосудов). Нарисованные, они придают дереву его фирменный вид: каждая граница — разрез, параллельный оси, потому что каждый вопрос называет один столбец и одно значение. Ниже он стоит рядом с тремя другими способами разделить те же точки — прямой линией, гладкой кривой, которую изогнула бы через них нейросеть, и рваным контуром, который получает метод k ближайших соседей, когда каждую точку пространства опрашивают ближайшие к ней обучающие точки.
Так что же на самом деле нужно, чтобы построить такое? Нам понадобятся две вещи.
- Способ порождать вопросы-кандидаты из таблицы, ведь дерево должно их откуда-то брать.
- Способ оценивать этих кандидатов, чтобы можно было выбрать лучшего, — узел решения содержит один вопрос, и не более.
Кроме них нужно правило, когда останавливаться, поскольку обычно мы не хотим делить, пока у каждой строки не появится собственный лист. Предоставленный сам себе, алгоритм именно туда и приходит, ведь разбиение кончается только тогда, когда ни один вопрос больше не делит кучу. Такое дерево запомнило таблицу, а не научилось на ней: лист с одним обучающим пациентом может только повторить исход этого пациента, так что любой новый пациент, попавший туда, получит результат одного человека, а не закономерность, увиденную на многих.
От строк к вопросам
Каждый узел решения содержит вопрос, так что построить дерево — значит выбирать вопросы. Чтобы собрать список кандидатов, соедините каждый признак с каждым значением, которое этот признак принимает в данных, — каждая пара и есть один вопрос. Тип столбца выбирает сравнение. Категориальный столбец требует точного совпадения: Is stress_test == fixed? истинно для пациентов с фиксированным дефектом и больше ни для кого. Числовой столбец вместо этого задаёт порог: Is vessels >= 1? истинно для пациента с одним поражённым сосудом и для всех, кто выше, — именно это делает значение точкой отсечения на числовой прямой, а не именем для сопоставления.
Это соединение и есть весь генератор, так что стоит выписать его целиком. В нашем наборе два столбца по три различных значения в каждом, так что получается список из шести пар:
| тип | столбец | значение | какой вопрос получается |
|---|---|---|---|
| категориальный | stress_test | normal | Is stress_test == normal? |
| категориальный | stress_test | fixed | Is stress_test == fixed? |
| категориальный | stress_test | reversable | Is stress_test == reversable? |
| числовой | vessels | 0 | Is vessels >= 0? |
| числовой | vessels | 1 | Is vessels >= 1? |
| числовой | vessels | 2 | Is vessels >= 2? |
Один столбец, одно значение, одно сравнение — вот и всё, чем когда-либо является вопрос, здесь и на любом другом наборе данных.
CART никогда не соединяет два условия в один вопрос, например stress_test == normal AND vessels >= 1, никогда не взвешивает один столбец против другого и никогда не использует значение, отсутствующее в данных: никакого vessels >= 1.5 между двумя наблюдёнными счётчиками и никакого vessels >= 3, потому что 3 в нашем наборе из пяти строк не встречается.
Эти шесть — всё, что дерево может спросить, и не все они попадут в готовое дерево. Большинство кандидатов пробуются, оцениваются и отвергаются; здесь только два доживают до того, чтобы стать вопросами в дереве, которое мы построим, а остальные четыре оцениваются и отбрасываются. И список конечен — никогда не длиннее числа различных значений в таблице, — поэтому следующий шаг может просто перебрать их все.
В коде вопросы определяются одним маленьким классом. Question хранит индекс столбца и значение, а его метод match решает, какое из двух сравнений применить, глядя на тип найденного значения:
class Question:
def __init__(self, column, value):
self.column = column
self.value = value
def match(self, example):
val = example[self.column]
if is_numeric(val):
return val >= self.value # numeric: threshold
else:
return val == self.value # categorical: equality
def __repr__(self):
condition = ">=" if is_numeric(self.value) else "=="
return "Is %s %s %s?" % (header[self.column], condition, str(self.value))
def is_numeric(value):
return isinstance(value, int) or isinstance(value, float)Эту логику сопоставления можно реализовать по-разному. Наша — трёхстрочная ветка is_numeric внутри match, и именно она позволяет этому дереву работать с текстовым и числовым столбцом бок о бок вообще без предобработки. Библиотеки идут не все этим путём.
Дерево из scikit-learn требует числового входа, так что stress_test придётся сначала закодировать one-hot — по одному столбцу 0/1 на значение:
| stress_test | is_normal | is_fixed | is_reversable | |
|---|---|---|---|---|
| normal | → | 1 | 0 | 0 |
| fixed | → | 0 | 1 | 0 |
| reversable | → | 0 | 0 | 1 |
Тогда дерево спрашивает is_fixed >= 0.5 там, где наше спрашивает stress_test == fixed, — то же разбиение, размазанное по трём столбцам. Само 0.5 ничего не значит: столбец содержит только 0 и 1, так что любой разрез между ними отделяет те же строки, а sklearn ставит пороги в середину между двумя соседними значениями. Столбец с четырьмя категориями просто стал бы четырьмя такими столбцами 0/1, о каждом из которых спрашивают всё так же на 0.5 — кодировка растёт вширь, а каждый вопрос остаётся проверкой «да/нет» одного значения.
LightGBM, CatBoost и XGBoost делят по подмножествам: они проверяют группу категорий разом, что по-прежнему один вопрос об одном столбце, — разница в том, что проверяемое значение является множеством, а не одной категорией:
ours: Is stress_test == fixed?
theirs: Is stress_test in {normal, reversable}?Измеряем разнородность набора данных
Теперь мы знаем, как построить список вопросов-кандидатов, поэтому следующее, что нужно понять, — как их оценивать и решать, который станет вопросом узла. Для этого нужна мера того, насколько перемешана куча меток, — она называется примесью Джини, — и способ оценить каждый вопрос по тому, сколько перемешанности он снимает: насколько его две кучи менее перемешаны, чем та, из которой они вышли. Эта оценка называется приростом информации, и побеждает вопрос с наибольшим приростом.
Сначала посмотрим на примесь Джини — меру того, насколько перемешана коллекция: одно число, говорящее, всё ли в ней одного рода или это мешанина из многих.
Допустим, вам нужно сказать, какая из двух коллекций перемешана сильнее. Одного взгляда на картинку ниже достаточно, чтобы понять, что второй набор разнообразнее: четыре рода вместо двух и распределены ровнее. Глаз решает это мгновенно.
Теперь допустим, что состав ни одного из наборов нам неизвестен — ни подсчётов, ни списка родов, только возможность запустить руку и что-нибудь вынуть. Сможем ли мы выразить разнообразие числом?
Один из способов получить такое число — сэмплировать: взять наугад два элемента, отметить, одного они рода или разного, вернуть обратно и повторить. Доля пар, оказавшихся разными, и есть оценка перемешанности набора, и для неё не нужно знать о наборе ничего, кроме того, что покажут вытягивания.
Допустим, мы сделали это по десять раз на каждом наборе. Вот что получилось:
Четыре пары левого набора содержат два разных рода; у правого таких семь. Поделите на число вытягиваний — и получите оценку; шляпка над обозначает величину, оценённую по выборке, в отличие от вычисленной по всей генеральной совокупности:
И число ведёт себя так, как мы и хотели. Чем оно меньше, тем чаще два случайных выбора оказывались одного рода — тем однороднее набор. Чем больше, тем чаще они различались — тем разнообразнее набор. Оценка к тому же уточняется по мере продолжения вытягиваний: десяти пар уже хватает, чтобы разделить эти два набора, а сотня зафиксировала бы каждое число. Продолжайте тянуть — и оно устоится на одном точном значении.
Вообще говоря, сэмплировать не нужно вовсе: когда вы знаете состав набора, немного теории вероятностей даёт это точное значение напрямую. Вычислите шанс, что два выбора совпадут, и вычтите его из 1.
Возьмём левый набор. Семь из десяти его элементов — синие квадраты, так что один выбор оказывается квадратом с вероятностью 0.7, а вероятность вытянуть два квадрата подряд — . Круги дают . Это единственные два способа совпасть, так что совпадение происходит в случаев. Но нам нужно обратное — как часто два выбора оказываются разными, — а поскольку каждое вытягивание либо совпадает, либо нет, это единица минус шанс совпадения: 0.42.
Правый набор — тот же расчёт, только с четырьмя родами вместо двух:
| род | доля | оба выбора попадают сюда |
|---|---|---|
| квадрат | 0.4 | 0.16 |
| круг | 0.3 | 0.09 |
| треугольник | 0.2 | 0.04 |
| звезда | 0.1 | 0.01 |
| совпадение 0.30 |
Два выбора совпадают в 30% случаев, значит различаются в 0.70 случаев — совпадая с семью из десяти, которые дало сэмплирование, но без единого вытягивания.
Ту же логику можно показать геометрически. Разложите каждую упорядоченную пару выборов как клетку сетки — первый выбор по горизонтали, второй по вертикали. Десять элементов дают сто клеток, и эта сетка — все возможные исходы:
У этого точного значения есть имя. Вероятность того, что два случайно вытянутых из набора элемента окажутся разного рода, называется примесью Джини этого набора, и записывается она так:
где — доля набора, принадлежащая роду . Две половины — два способа сказать одно и то же: — вероятность совпадения вытягиваний (для каждого рода шанс, что оба попадут в него, сложенный по родам), а единица минус это — шанс, что они различаются.
Наши два набора, пропущенные через формулу, — это арифметика минутной давности в сжатом виде:
Эта статистика старше машинного обучения и встречается в других областях под другими именами — индекс Симпсона в экологии, индекс Херфиндаля — Хиршмана в экономике.
Джини на наших пяти строках
Теперь мы готовы вычислить примесь Джини для наших пяти строк. Сначала нужен подсчёт родов, которыми в нашем случае служат метки: где наборы выше содержали квадраты, круги, треугольники и звезду, куча строк содержит Yes и No. Так что сосчитайте их — сколько каких меток в данной куче, потому что каждая величина в этой статье вытекает из этого словаря.
def class_counts(rows):
"""Counts the number of each type of example in a dataset."""
counts = {} # label -> count
for row in rows:
label = row[-1] # the label is always the last column
if label not in counts:
counts[label] = 0
counts[label] += 1
return countsПервый прогон по всему набору, class_counts(training_data), даёт {'No': 2, 'Yes': 3} — наши пять пациентов, подсчитанные по родам.
Теперь, имея подсчёт под рукой, можно вычислить примесь Джини — четыре строки Python:
def gini(rows):
"""Calculate the Gini Impurity for a list of rows."""
counts = class_counts(rows)
impurity = 1
for lbl in counts:
prob_of_lbl = counts[lbl] / float(len(rows))
impurity -= prob_of_lbl**2
return impurityЦикл и есть формула, по слагаемому на метку: подайте ему кучу из одних только Yes — вернёт 0.0, подайте один Yes и один No — вернёт 0.5. Наш собственный обучающий набор, три Yes против двух No, стартует с:
gini(training_data) → 0.48Мы пройдёмся по каждому вопросу-кандидату и посмотрим, кто оставит после себя меньше всего перемешанности, так что 0.48 — число, которое нужно побить. Это ещё и высокая стартовая точка: при двух метках Джини достигает максимума 0.5 при их равном делении, так что три Yes против двух No оставляют нас на 0.48 — примерно настолько же перемешанно, насколько это возможно для пяти строк.
Прирост информации — оценка разбиения
Наша цель — оценивать вопросы, и мы уже умеем вычислять примесь кучи строк. Значит, можно сделать разбиение вопросом-кандидатом, измерить примесь каждой из двух получившихся куч и сравнить с тем, с чего начинали. Это и есть рецепт прироста информации — оценки, по которой мы судим о вопросе.
Записанный, он занимает одну строку:
И состоит из четырёх шагов:
- разбить кучу вопросом, получив две кучи: — строки, ответившие True, и — строки, ответившие False;
- прогнать
giniна каждой из них; - свести эти два числа в одно, взвесив по тому, сколько строк ушло на каждую сторону: даёт вес левой кучи, а — правой, каждый из них есть доля строк родителя, ушедшая в эту сторону;
- вычесть это из примеси родителя, .
Остаётся примесь, которую убрал вопрос: чем она выше, тем лучше вопрос.
Читайте это как сделку в валюте неопределённости. В словаре разнообразия из предыдущего раздела взвешенная сумма — это оставшееся разнообразие после разбиения, а прирост — убранное разнообразие, за которое заплатили вопросом. Прирост 0 означает, что две кучи перемешаны ровно так же, как та, из которой они вышли, то есть вопрос ничего не разделил; вопрос, у которого оба потомка выходят чистыми, убрал всю перемешанность, что была. Так что мы охотимся за вопросами с наибольшим приростом — теми, что покупают больше всего убранного разнообразия за один вопрос, которого они стоят.
Примесь — это замаскированная функция потерь
Интересный вопрос: что играет роль функции потерь в дереве решений — компонента, явного в нейросети и нигде не видного в коде, который мы написали. Ответ: это примесь — эта функция и есть обучающая потеря кучи при её лучшем константном ответе: дисперсия — это квадратичная ошибка предсказания среднего, энтропия — логарифмические потери предсказания долей классов, а Джини — квадратичная ошибка их предсказания, , та же величина, что измеряет оценка Брайера. Значит, прирост информации — это снижение потерь, и дерево обучается минимизацией потерь, как и всё остальное в машинном обучении, с двумя оговорками. Потери минимизируются перебором, а не дифференцированием, потому что нет непрерывных параметров, сквозь которые можно брать градиенты. И минимизируются они жадно, а не глобально — не из лени, а потому что построение оптимального дерева NP-полно, результат, восходящий к Хьяфилу и Ривесту в 1976 году; по одному разбиению за раз — цена вычислимости, а ничьи, с которыми мы вот-вот столкнёмся, — её видимый шрам.
Стоит подчеркнуть, зачем вообще нужно взвешивание из шага 3, потому что без него оценку легко обмануть. Два наших кандидата, Is stress_test == normal? и Is vessels >= 1?, каждый делит пять строк на одного идеально чистого потомка с Джини ровно 0 и одного всё ещё перемешанного. Отличаются они тем, сколько данных уносит этот чистый потомок: один отслаивает единственного пациента и оставляет позади четыре перемешанные строки, другой забирает двоих и оставляет три. Только взвешивание видит эту разницу. Оно заставляет чистого потомка считаться ровно на столько, сколько он весит, так что потомок из одной строки едва заметен, а счёт задаёт оставленный беспорядок.
Вот оба кандидата, расписанные полностью, каждый со своими двумя потомками, сведёнными двумя способами: посчитанными поровну и взвешенными по доле строк, которую держит каждый потомок:
Теперь то же самое для другого кандидата. Is vessels >= 1? тоже отрезает идеально чистого потомка, но в нём два пациента, а не один, а куча, которую он оставляет, — три строки, а не четыре, и она грязнее: 0.444 вместо 0.375:
Так что числа на двух рисунках показывают нечто более сильное, чем изменение масштаба.
Посчитанные поровну, Is stress_test == normal? даёт 0.293, а Is vessels >= 1? — 0.258, так что побеждает первый вопрос. Взвешенные, они выходят 0.180 и 0.213, и побеждает уже второй. Взвешивание не просто уменьшает оценки — оно переворачивает порядок, а поскольку это корень, два ответа дают деревья, различающиеся сверху донизу.
Вот как мы реализуем разбиение и его прирост информации, вместе со взвешиванием. partition выполняет шаг 1, раскладывая строки по двум кучам, которые делает вопрос, а info_gain выполняет шаги 2–4, оценивая то, что получилось, против того, что было:
def partition(rows, question):
"""Split rows into those matching the question, and those that don't."""
true_rows, false_rows = [], []
for row in rows:
if question.match(row):
true_rows.append(row)
else:
false_rows.append(row)
return true_rows, false_rows
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)Прогоните их на двух кандидатах выше — и они вернут 0.180 и 0.213, те же числа, что рисунки посчитали вручную, только теперь вычисленные, а не нарисованные.
Механизм — разбить, потом рекурсировать
Теперь у нас есть все части: способ порождать вопросы, способ измерять перемешанность кучи и способ оценивать, что вопрос с ней делает. Вот процедура, которая их соединяет. Дерево решений растится одним рецептом, применяемым к одной куче обучающих строк:
- Попробовать каждый вопрос, который допускают данные, — каждый признак, каждое значение, которое этот признак принимает.
- Оценить каждый вопрос по тому, насколько он расслаивает метки в куче, — это прирост информации, построенный на примеси Джини, ровно как мы только что вывели.
- Если ни один вопрос не помогает, остановиться: куча становится листом, а её подсчёты меток — предсказанием.
- Иначе разбить кучу лучшим вопросом на две кучи поменьше.
- Запустить ту же процедуру на каждой из двух куч.
У этой процедуры есть каноническое имя — рекурсивное бинарное разбиение: нисходящий жадный алгоритм построения деревьев решений последовательным делением набора данных на две группы. Он начинает со всех данных в корне, оценивает каждый признак и точку разбиения, чтобы минимизировать ошибку или максимизировать чистоту, и повторяет процесс на каждой новой подгруппе, пока не достигнут предел остановки.
Назвать алгоритм жадным значит сказать, что он решает каждое разбиение, глядя только на кучу перед собой. Он берёт вопрос с лучшей оценкой прямо здесь и больше к нему не возвращается: выбор не пересматривается, когда потомки оказываются плохими, и никогда не согласуется с разбиениями в других местах дерева. Локально лучший на каждом шаге, без гарантии, что готовое дерево — лучшее дерево, и следующий раздел показывает, как мало нужно, чтобы обнажить этот разрыв, когда два корневых вопроса набирают ровно одинаковый счёт, а выбор между ними меняет всё внизу.
Быть рекурсивным — это то, что вырезает прямоугольники с рисунка во вступлении: каждый вызов владеет одной областью пространства признаков (строками, пережившими вопросы выше) и либо подразделяет эту область, либо запечатывает её как лист. Прямоугольники — это кучи на дне рекурсии.
Выбор корневого разбиения — и ничья
Прежде чем строить всю рекурсивную конструкцию, быстро посмотрим на реализацию части, работающей в одном узле, — поиска лучшего разбивающего вопроса. В корне этот узел держит всех пятерых пациентов, а ищет функция find_best_split, которая пробует каждое значение каждого признака и оставляет лучшее.
Это два вложенных цикла — каждый столбец снаружи, каждое различное значение этого столбца внутри, — и каждая порождённая ими пара проходит четыре шага:
- построить
Questionиз столбца и значения; partitionстрок им, на две кучи, которые он делает;- оценить эти кучи через
info_gain; - сравнить оценку с лучшей на данный момент и оставить вопрос, если он побеждает.
Когда оба цикла заканчиваются, функция возвращает вопрос, всё ещё держащий лучшую оценку.
def find_best_split(rows):
best_gain = 0
best_question = None
current_uncertainty = gini(rows)
n_features = len(rows[0]) - 1
for col in range(n_features):
values = set([row[col] for row in rows])
for val in values:
question = Question(col, val)
true_rows, false_rows = partition(rows, question)
if len(true_rows) == 0 or len(false_rows) == 0:
continue # this split doesn't divide the data
gain = info_gain(true_rows, false_rows, current_uncertainty)
if gain >= best_gain:
best_gain, best_question = gain, question
return best_gain, best_questionОдин вызов в корне оценивает каждый вопрос, порождённый генератором, и возвращает вот это:
| вопрос | прирост |
|---|---|
Is stress_test == reversable? | 0.2133 |
Is vessels >= 1? | 0.2133 |
Is stress_test == normal? | 0.1800 |
Is vessels >= 2? | 0.0800 |
Is stress_test == fixed? | 0.0133 |
Is vessels >= 0? | пропущен — был бы 0 |
Шесть вопросов — два столбца по три различных значения, — но число получают только пять из них. Is vessels >= 0? истинно для каждого пациента, потому что 0 — наименьшее значение этого столбца, так что он отправляет все пять строк по ветви True и ни одной по False. Один потомок держит всю кучу, а другой — ничего, и это не разбиение, а копия: ничего не разделено, значит нечего оценивать. Защита len(true_rows) == 0 or len(false_rows) == 0 отбрасывает его до того, как info_gain его увидит. Его прирост всё равно был бы ровно 0 — один потомок ничего не весит, другой равен родителю, — но пропуск ещё и не даёт вернуть «неразбиение» как лучший вопрос, когда ничто другое не набирает больше нуля.
Теперь посмотрите на верх этой таблицы, потому что там самое интересное.
Два разных вопроса вернули одинаковую оценку 0.2133.
Ничьи и почти-ничьи обычны на реальных данных, и когда так случается, поздний кандидат перезаписывает раннего.
Такое поведение — деталь реализации, и в нашем алгоритме оно следует из двух вещей: столбцы просматриваются в порядке индексов, так что stress_test (столбец 0) добирается первым, а затем его тихо вытесняет столь же хороший vessels (столбец 1); и сравнение написано через >=, а не >, что и позволяет вытеснению случиться:
if gain >= best_gain:Ничьи и почти-ничьи возникают на реальных данных постоянно, и именно поэтому эта причуда — проблема. Ничто в данных не предпочло vessels перед stress_test — это сделал оператор сравнения, а поскольку это корень, всё ниже построено на этом выборе. Небольшого изменения строк достаточно, чтобы перевернуть почти-ничью и перестроить всё поддерево под ней. Именно это делает одиночное дерево моделью с высокой дисперсией: его форма зависит от конкретной выборки, на которой оно обучалось. Мы вернёмся к этому ближе к концу статьи, где это будет одной из двух неудач, объясняющих, почему одно дерево редко бывает моделью, которую вы выкатываете.
Рекурсия — строим дерево
Теперь, когда отдельный узел умеет находить свой вопрос, мы готовы построить всё дерево — рекурсию, которая гоняет этот поиск по куче за кучей и сохраняет находки. Для хранения нужен один класс на каждый тип узла из вступительной диаграммы: Leaf держит подсчёты меток тех строк, что до него дошли, а Decision_Node держит вопрос и две ветви. На языке учебника вопрос — это правило разбиения, один предикат по одному признаку, а узел решения — это правило, встроенное в блок-схему, где двум ветвям есть куда вести свои ответы «да» и «нет». Готовое дерево можно мыслить как серию правил разбиения.
Начиная с верхушки дерева и применяясь по пути вниз — find_best_split учит эти правила, а узлы — место, где живут выбранные.
class Leaf:
def __init__(self, rows):
self.predictions = class_counts(rows)
class Decision_Node:
def __init__(self, question, true_branch, false_branch):
self.question = question
self.true_branch = true_branch
self.false_branch = false_branch
def build_tree(rows):
gain, question = find_best_split(rows)
if gain == 0:
return Leaf(rows) # base case: no question helps anymore
true_rows, false_rows = partition(rows, question)
true_branch = build_tree(true_rows)
false_branch = build_tree(false_rows)
return Decision_Node(question, true_branch, false_branch)Когда мы запускаем это на пяти пациентах, получается вот такое дерево — нарисованное так, что кучи видны на каждой ветви:
Дерево вышло глубиной в два вопроса, с тремя листьями. Прочитаем его сверху вниз, начиная с корня: vessels >= 1, победитель ничьей. Каждый пациент с поражённым сосудом болен, и эта ветвь немедленно завершается чистым листом — оба пациента, одним вопросом.
На этом стоит задержаться: целая группа выпала из данных вообще без примеси — Джини 0, полученный одним вопросом. Уровнем ниже одинокий normal-пациент делает то же самое, — и заметьте, что у stress_test три значения, а дерево спрашивает лишь об одном из них. stress_test == fixed? отслаивает пациентов с фиксированным дефектом, а всё, что не fixed, едет по ветви False вместе, неразличённое. Здесь это оказывается единственный normal-пациент, потому что оба reversable-пациента ушли ещё в корне.
Из трёх листьев дерева два чистые; каждая обучающая строка, кроме сталкивающейся пары, попадает в группу с нулевой примесью, и рекурсия останавливается в каждой из них именно потому, что убирать больше нечего.
Остаётся третий лист, держащий одну метку Yes и одну No при одинаковом наборе значений признаков. Никакой вопрос не смог бы разделить этих двух пациентов — и никакая другая модель тоже, потому что то, что их различает, в данных попросту отсутствует. Это можно было бы разрешить, взяв больше предикторов, ведь одиннадцать выброшенных столбцов вполне могут содержать то, что этих двух пациентов разделяет.
Этот процесс оценки вопросов, разбиения кучи по победителю и рекурсии, по сути, и есть процедура обучения.
Архитектура нейросети проектируется заранее — число слоёв, их ширина, связи, — а градиентный спуск тысячи раз подталкивает значения внутри этой фиксированной рамки. У дерева нет фиксированной рамки и нечего подталкивать: обучение изобретает, о каком признаке спрашивает каждый узел, при каком пороге, в каком порядке и на какой глубине. Нейросеть обучает значения внутри фиксированной структуры; дерево обучает саму структуру, а его значения выпадают как сводки — подсчёты тех строк, что случайно пришли. Не было ни эпох, ни сходимости: каждый узел оценил своих кандидатов один раз, оставил лучшего и больше к выбору не возвращался, так что когда корневой вызов build_tree вернулся, обучение закончилось. Модель не становилась постепенно лучше; она постепенно строилась.
Классификация — считываем вероятность с листа
Предсказание — снова рекурсия, и короче обучающего кода:
def classify(row, node):
if isinstance(node, Leaf):
return node.predictions
if node.question.match(row):
return classify(row, node.true_branch)
else:
return classify(row, node.false_branch)
def print_leaf(counts):
total = sum(counts.values()) * 1.0
return {lbl: str(int(counts[lbl] / total * 100)) + "%" for lbl in counts}Каждый Decision_Node хранит один Question — индекс столбца плюс значение, — а match сравнивает запись строки в этом столбце с ним, возвращая обычный True или False. Этот булев результат — единственное, что нужно classify: True отправляет строку по true_branch, False — по false_branch, и рекурсия останавливается, как только попадает на Leaf.
Возьмём данные одного пациента, ['fixed', 0, 'Yes'], и посмотрим, как дерево предсказывает, есть ли у него сердечное заболевание:
- корень спрашивает
Is vessels >= 1?;matchчитает записьvesselsпациента — это0— и, поскольку значение числовое, вычисляет0 >= 1, получая False, так что строка идёт по ложной ветви; - этот узел спрашивает
Is stress_test == fixed?;matchчитает записьstress_testпациента —'fixed'— и, поскольку значение строковое, вычисляет'fixed' == 'fixed', получая True, так что строка идёт по истинной ветви; - эта ветвь —
Leaf, так чтоclassifyвозвращает хранящиеся там подсчёты: одинYesи одинNo.
Эти возвращённые classify подсчёты и есть предсказание в сыром виде. Их можно прочесть как одну метку, взяв ту, что чаще встречается в листе, — так делает predict в библиотеке вроде sklearn, и в чистых листьях это однозначно, где {'Yes': 2} означает Yes. А можно прочесть как вероятность, поделив каждый счётчик на сумму, — это predict_proba, и то, что здесь делает print_leaf.
Для этого конкретного пациента счётчики по одному, так что два прочтения дают «нет большинства» и «50/50» — один и тот же факт, дважды. И 50/50 — верный ответ: два обучающих пациента имеют ровно такие признаки и расходятся, так что модель, заявляющая уверенность, лгала бы, а в этой предметной области это не фигура речи. Это неустранимая ошибка, и подсчёты в листе сообщают о ней бесплатно, без всякой дополнительной машинерии для неопределённости.
Виджет ниже — чуть более крупная игрушка (два числовых признака, пороги вместо наших смешанных типов), но механизм идентичен, и он показывает два взгляда на дерево сразу. Левая панель — разбиение; правая — обход. Два ползунка — значения признаков и ; их перетаскивание составляет новую строку и двигает её по пространству признаков. Момент, когда точка пересекает пунктирную линию, — ровно тот момент, когда меняется путь по дереву, потому что область и лист — один и тот же объект в разных нарядах.
Заметьте также, что это дерево использует признак, обозначенный , дважды — один раз в корне и ещё раз двумя уровнями ниже при другом пороге. Это два разных вопроса по одному столбцу — тот же признак, другое значение, — потому что признак не расходуется от того, что по нему разбили: первый разрез разделяет, что может, а оставшиеся строки могут по-прежнему разделяться вдоль той же оси.
Именно поэтому число узлов решения и число признаков независимы. Признаки лишь поставляют меню; данные решают, какие вопросы будут заданы и как часто, — и это меню перестраивается в каждом узле, а не фиксируется раз на всё дерево. find_best_split считывает его со строк перед собой, values = set([row[col] for row in rows]), так что список вопросов-кандидатов сжимается вместе с кучами: наш корень может спросить Is vessels >= 2?, а узел под ним — уже нет, потому что ни у одной дошедшей туда строки нет двойки.
Теперь дайте ему пять пациентов, которых оно никогда не видело, — все они настоящие строки из того же файла:
patient 2 ['normal', 3] Actual: Yes. Predicted: {'Yes': '100%'}
patient 9 ['reversable', 1] Actual: Yes. Predicted: {'Yes': '100%'}
patient 266 ['fixed', 0] Actual: Yes. Predicted: {'Yes': '50%', 'No': '50%'}
patient 88 ['NA', 0] Actual: No. Predicted: {'No': '100%'}
patient 267 ['NA', 0] Actual: Yes. Predicted: {'No': '100%'}Первая тестовая строка — ['normal', 3, 'Yes'], а значение vessels, равное 3, в обучении никогда не встречалось: наши пять пациентов показывали только 0, 1 и 2. Она всё равно попадает в лист, потому что vessels >= 1 — это порог, а не таблица поиска: 3 его проходит и идёт тем же путём, что и 1. Пороги обобщаются за пределы своих обучающих значений бесплатно.
С пропущенными значениями история другая: модель просто обрабатывает их неверно. У пациентов 88 и 267 нет записанного стресс-теста — в файле стоит NA, — и, поскольку NA в обучении не встречалось, проверка stress_test == fixed? не проходит и строка съезжает по ветви False, из-за чего пациент 267 возвращается как {'No': '100%'}, хотя болезнь у него есть.
На практике с этим разбираются до обучения — выбрасывают неполные строки, заполняют пропуски или делают "missing" отдельной категорией.
Но чистка обучающих данных дыру не закрывает, потому что пациент, пришедший завтра, всё ещё может принести значение stress_test, которого дерево не видело. Сравнение ==, которое match использует на категориальном столбце, не проходит, строка съезжает по ветви False, и ответ возвращается выглядящим ровно так же, как обоснованный. Именно это молчание — настоящий дефект: не то, что дерево ошибается, а то, что ничто в {'No': '100%'} не отличает три уверенные обучающие строки от значения, которого модель никогда не встречала.
Настоящие библиотеки деревьев спроектированы обрабатывать пропущенные значения без всякой предобработки набора данных. XGBoost выучивает направление по умолчанию для каждого разбиения, отправляя строки, на которые не может ответить, в ту сторону, которая лучше показала себя на обучающих данных, а исходная формулировка CART хранит суррогатные разбиения — резервные вопросы, коррелирующие с основным, задаваемые любой строке, которая не может на него ответить.
Почему одно дерево — не конец истории
Модель, которую мы построили, — это настоящий CART примерно в 200 строках чистого Python, и, направив её на реальные данные, вы вырастите настоящее дерево. Однако у нашей реализации есть две проблемы, и весь остальной мир деревьев существует, чтобы с ними справляться.
Во-первых, дерево переобучится, если ничто не остановит его рост. Переобучение — это когда модель заучивает обучающие данные вместо того, чтобы на них учиться, и тем самым теряет способность обобщать на что-либо ещё. В нейросети это происходит через веса: при достаточной ёмкости и слишком слабой регуляризации градиентный спуск продолжает их подстраивать, пока сеть не воспроизведёт обучающий набор почти точно. В дереве это происходит через разбиение: без присмотра build_tree продолжает резать, пока почти каждой обучающей строке не достанется собственный чистый лист, потому что gain == 0 — единственное, что его останавливает. Ось ёмкости здесь не «сколько вы обучали», а «насколько глубоко вырастили», — поэтому всякий регуляризатор для деревьев структурный: ограничения глубины, минимальные размеры листа, обрезка.
Посмотрим, как эта проблема проявляется на реальном наборе данных, взяв sklearn как честного дублёра нашего кода: оставьте DecisionTreeClassifier(criterion="gini") на настройках по умолчанию — без ограничения глубины, без минимального размера листа, без обрезки — и у него тоже не будет правила остановки, кроме чистоты, так что дерево, которое он вырастит, — это дерево, которое вырастил бы build_tree, только посчитанное быстрее.
Теперь направьте его на набор данных о раке груди (398 обучающих строк, 171 тестовая, 30 признаков):
from sklearn.datasets import load_breast_cancer
from sklearn.model_selection import train_test_split
from sklearn.tree import DecisionTreeClassifier
data = load_breast_cancer()
X_tr, X_te, y_tr, y_te = train_test_split(
data.data, data.target, test_size=0.3, random_state=0
)
unpruned = DecisionTreeClassifier(criterion="gini", random_state=0).fit(X_tr, y_tr)
print(unpruned.get_depth(), unpruned.get_n_leaves())
print(unpruned.score(X_tr, y_tr), unpruned.score(X_te, y_te))
for depth in (1, 2, 3, 5, 7):
pruned = DecisionTreeClassifier(
criterion="gini", max_depth=depth, random_state=0
).fit(X_tr, y_tr)
print(depth, pruned.score(X_tr, y_tr), pruned.score(X_te, y_te))Неудержанное дерево — то, которое ничто не сдерживает — выходит таким:
UNPRUNED (what build_tree does)
depth 7
leaves 19
train acc 1.000
test acc 0.912Точность здесь — просто доля строк, чья предсказанная метка совпадает с записанной, измеренная на 398 строках, из которых дерево построено, и отдельно на 171, которых оно не видело. Точность на обучении 1.000 значит, что оно угадало все 398, и добилось этого разбиением до тех пор, пока отстающие — строки, отказавшиеся группироваться с чем бы то ни было, — не осели каждая в собственном листе: то же поведение gain == 0, за которым мы наблюдали на пяти пациентах, только с 398 строками вместо 5.
Теперь добавим единственную ручку, которой в нашей версии нет, — max_depth, жёсткий потолок на то, на сколько вопросов вглубь может уйти дерево, останавливающий разбиение независимо от того, остался ли прирост. Прогоняя то же дерево с рядом ограничений глубины и собирая результаты, получаем:
| max_depth | обучение | тест | разрыв |
|---|---|---|---|
| 1 | 0.930 | 0.895 | +0.035 |
| 2 | 0.960 | 0.947 | +0.012 |
| 3 | 0.967 | 0.947 | +0.020 |
| 5 | 0.987 | 0.936 | +0.052 |
| 7 (без обрезки) | 1.000 | 0.912 | +0.088 |
Читайте две последние строки друг против друга, потому что в этом весь урок. Переход с глубины 2 на глубину 7 улучшает точность на обучении с 0.960 до 1.000 и делает модель хуже — тестовая точность падает с 0.947 до 0.912. Необрезанное дерево не просто расточительно. Его бьёт дерево втрое меньшей глубины, которое ошибается на 4% обучающих данных. Эти лишние пять уровней глубины — заучивание 398 конкретных строк, и build_tree не может об этом узнать, потому что изнутри каждое из этих разбиений снижало примесь.
Вторая проблема в том, что дерево неустойчиво: измените данные немного — и оно может выйти другой формы. Что возвращает нас к >=, которым мы сравнивали прирост каждого кандидата с лучшим на данный момент, и к порядку, в котором цикл случайно посещает этих кандидатов. Жадная оценка на реальных данных постоянно порождает ничьи и почти-ничьи, а то, чья сторона победит, сводится к недокументированной детали реализации — вы видели это в корне: два вопроса с идентичными оценками и один символ, решающий между ними. Направьте тот же код на полное кливлендское исследование — 297 пациентов без пропусков, все тринадцать предикторов, 207 строк на обучение и 90 отложенных — и в десяти из 35 узлов решения готового дерева два вопроса-кандидата набирают ровно одинаковый счёт, разрезая пациентов по-разному. Ничто в данных их не разделяет, так что какой из них оставит алгоритм — произвольно, и всё, что возмущает оценки, переворачивает выбор и перестраивает всё под ним.
И цена ложится на реальных пациентов. Удалите одну обучающую строку и переобучите — и вернувшееся дерево отправит домой с другим диагнозом до 12 из 90 отложенных пациентов. Не меняйте вообще ничего — и предсказания всё равно сдвинутся: find_best_split посещает связанных ничьей кандидатов в том порядке, в каком итерируется set, а он различается от запуска к запуску, так что десять запусков на идентичных пациентах дали четыре разные модели, расходящиеся между собой до 4 диагнозов из 90. Два запуска на одних и тех же данных вручают вам две разные модели, обе корректные с точки зрения самого алгоритма.
Эти два поведения — две стороны компромисса смещения и дисперсии. В статистическом обучении смещение — это ошибка, вносимая приближением сложной реальности более простой моделью: модель, слишком жёсткая, чтобы представить закономерность, будет ошибаться, сколько данных ей ни дай. Дисперсия — это то, насколько сильно изменилась бы подогнанная модель, если оценить её по другому обучающему набору: переобучите на другой выборке пациентов, и метод с высокой дисперсией даст вам заметно другую модель, делающую заметно другие ошибки.
Как смещение и дисперсия складываются в ошибку предсказания
Ошибка предсказания — это то, что вы реально измеряете: расхождение между тем, что говорит модель, и тем, что случилось, — и, какова бы ни была модель, она приходит из трёх мест: из предположений, неверных при любом объёме данных; из чувствительности к тому, на каких строках вы случайно обучились; и из случайности, которую ничто не может предсказать. Когда цель — число, а ошибка измеряется квадратом разности — для любой модели, от линейной регрессии до дерева, — эти трое разделяются точно:
Смещение входит в квадрате, потому что это знаковая величина — насколько далеко среднее предсказание модели отстоит от истины, — которая иначе взаимно уничтожалась бы, а не накапливалась. Для ошибки классификации работают те же три источника, но складываются они не так аккуратно; интуиция переносится, арифметика — нет.
Первые два — те части, которыми владеете вы. Смещение — систематическая часть: модель промахивается мимо истинной связи каждый раз в одну и ту же сторону, и больше данных её не спасёт. Дисперсия — несистематическая часть: модель не ошибается в среднем, но любая её конкретная подгонка сбита, потому что слишком плотно последовала за конкретными строками, на которых обучалась. Только последнее слагаемое, шум в самих данных, недосягаемо.
Возьмите любое отдельное предсказание, которое модель делает неверно. Часть этой ошибки есть потому, что модель — неверной формы для задачи: дерево, ограниченное глубиной 1, не может выразить «сосуды и стресс-тест вместе», так что оно промахивается одинаково на любом наборе, который вы ему дадите. Часть есть потому, что именно это дерево выросло из именно этих строк, а другая выборка вырастила бы другое дерево, промахивающееся иначе. А часть была в данных ещё до того, как появилась хоть какая-то модель: два пациента, идентичные по каждому записанному признаку, один с болезнью, другой без. Над первыми двумя можно работать, меняя модель; третья задаёт потолок того, насколько хорошо может выступить любая модель.
Эти двое движутся в противоположных направлениях по мере роста гибкости метода, а для дерева гибкость — это глубина. Держите его мелким — и оно слишком простое, чтобы схватить закономерность (высокое смещение), но устойчивое: обучите на другой выборке и вернётся примерно то же дерево. Дайте ему расти — и оно сможет подогнать что угодно, включая шум (низкое смещение), ценой ровно той неустойчивости, что описана выше, то есть дисперсии. Точность на обучении вознаграждает только первое из двух, поскольку более глубокое дерево всегда лучше подгоняет собственные строки, а тестовая точность отвечает за оба — поэтому в таблице она достигает пика на глубине 2 и оттуда сползает. И никакая глубина не сведёт тестовую ошибку к нулю, потому что под обоими слагаемыми лежит неустранимая ошибка, встреченная нами у листа 50/50, — та часть исхода, которую признаки просто не определяют.
Есть устоявшиеся методы борьбы с этими проблемами: потолок глубины, минимальное число строк на лист, минимальный прирост, ради которого стоит разбивать, и обрезка ветвей постфактум. На практике их редко применяют к одинокому дереву; это ручки, которые крутят внутри ансамбля — модели, построенной из многих деревьев, чьи ответы объединяются в один, а именно этим и являются случайный лес и градиентный бустинг.
Посмотрим на единственное решение, применимое на уровне одного дерева, — обрезку, которая сводится к замене gain == 0 правилом остановки, знающим, когда пора закончить. Средства делятся на два именованных семейства. Предобрезка (ранняя остановка) отказывается расти изначально: max_depth, минимум строк на разбиение или на лист, пороги минимального прироста — ползунок глубины выше был предобрезкой в самом грубом виде. Она дёшева, но жадна во втором смысле: слабое разбиение может быть дверью к сильному под ним, а рано остановленное дерево этого никогда не узнает (известно как эффект горизонта). Постобрезка даёт дереву вырасти полностью, а затем срезает ветви, которые не окупаются на отложенных данных; каноническая версия CART — обрезка по стоимости и сложности: оцените дерево как его ошибку плюс цена за лист и срежьте всё, что не оправдывает содержания. Это и есть наконец явно выписанная функция потерь дерева — слагаемое подгонки плюс штраф за сложность, та же форма, которую регуляризация принимает где угодно ещё, — а sklearn выставляет эту цену как ccp_alpha.
Неустойчивость же обычно не лечится внутри одного дерева вообще. Вместо поиска более умного разрешения ничьих вы перестаёте полагаться на одно дерево. Вырастите много, каждое на слегка иной выборке строк и столбцов, чтобы они падали на разные стороны ничьих, а затем усредните их ответы: это случайный лес, и именно усреднение гасит дисперсию. Вырастите их вместо этого последовательно, каждое исправляя ошибки предыдущего, — и это градиентный бустинг, построенный из тех же деревьев, что мы только что написали, с потолком глубины и другой целью, в статье From one tree to XGBoost.
Где наша версия медленнее настоящей
В нашей реализации не хватает одной важной техники оптимизации, которая есть в каждой настоящей библиотеке. Всё остальное совпадает с тем, что делает продакшен-реализация — те же вопросы-кандидаты, та же примесь, тот же прирост, то же выбранное разбиение, — но find_best_split в том виде, как мы его написали, — это перебор в буквальном смысле.
Цикл проходит по каждому признаку и каждому его различному значению, так что число кандидатов равно признаки × значения.
Затем каждый кандидат стоит полного прохода по данным: partition обходит каждую строку, раскладывая её по двум кучам, а info_gain вызывает gini на каждой куче, которая считает её метки с нуля. Это . На пяти строках с тремя значениями на столбец — незаметно. На непрерывном признаке — скажем, холестерине — почти каждая строка несёт своё значение, так что число кандидатов растёт вместе с данными, а каждый кандидат по-прежнему стоит полного прохода: квадратично по числу строк и безнадёжно при сотне тысяч.
Возьмём пять строк столбца cholesterol — 210 (No), 233 (No), 250 (Yes), 286 (Yes), 300 (Yes). Генератор превращает их в пять вопросов-кандидатов, по одному на наблюдённое значение, из которых четыре действительно делят кучу:
| кандидат | ниже порога | на пороге или выше |
|---|---|---|
>= 210 | пусто | все пять |
>= 233 | 210 | 233, 250, 286, 300 |
>= 250 | 210, 233 | 250, 286, 300 |
>= 286 | 210, 233, 250 | 286, 300 |
>= 300 | 210, 233, 250, 286 | 300 |
Проследим двух из них, >= 233 и >= 250, через наш код.
Для >= 233 partition обходит все пять строк и кладёт 210 в список False, а остальные четыре — в список True. Затем info_gain вызывает gini на каждом, и gini обходит кучу из одной строки, считая метки, потом кучу из четырёх строк, считая метки. Пять посещений на разбиение, пять на подсчёт. Для >= 250 всё начинается заново с тех же пяти строк, и так далее по списку:
>= 233: partition 5 rows → gini({210}) + gini({233,250,286,300}) = 10 visits
>= 250: partition 5 rows → gini({210,233}) + gini({250,286,300}) = 10 visits
>= 286: partition 5 rows → gini({210,233,250}) + gini({286,300}) = 10 visits
>= 300: partition 5 rows → gini({210,233,250,286}) + gini({300}) = 10 visitsСорок посещений строк, и между строками ничего не переносится — хотя каждая пара куч отличается от пары выше ровно одной строкой.
Вот те же пять строк, отсортированные по холестерину и несущие метку disease, которая у каждого пациента оказалась, — так их держала бы настоящая реализация:
| cholesterol | disease |
|---|---|
| 210 | No |
| 233 | No |
| 250 | Yes |
| 286 | Yes |
| 300 | Yes |
Настоящие реализации получают оба ответа из первого же обхода. Отсортируйте строки по признаку, а затем, оценивая самого первого кандидата, пройдите по ним один раз, ведя нарастающий подсчёт увиденных меток: к концу этого единственного прохода ответ получен и для каждого последующего кандидата:
| cholesterol, disease | нарастающий подсчёт |
|---|---|
| 210, No | {No: 1} |
| 233, No | {No: 2} |
| 250, Yes | {No: 2, Yes: 1} |
| 286, Yes | {No: 2, Yes: 2} |
Каждая строка — это строка выше плюс метка только что пройденной строки: 210 — это No, так что подсчёт открывается на {No: 1}; 233 — ещё один No, доводящий до {No: 2}; 250 — Yes, добавляющий первый Yes; и так далее. Один инкремент на строку, пять строк, один проход.
Поскольку строки приходят по возрастанию, каждый из этих подсчётов — это ещё и группа, попадающая ниже конкретного порога, что и превращает его в ответ:
| подсчёт | обслуживает вопрос | группа ниже | группа на пороге или выше |
|---|---|---|---|
{No: 1} | >= 233 | {No: 1} | {No: 1, Yes: 3} |
{No: 2} | >= 250 | {No: 2} | {Yes: 3} |
{No: 2, Yes: 1} | >= 286 | {No: 2, Yes: 1} | {Yes: 2} |
{No: 2, Yes: 2} | >= 300 | {No: 2, Yes: 2} | {Yes: 1} |
Поскольку формула прироста взвешивает обоих потомков, нужны обе группы, — но по ходу прохода отслеживается только первая. Группу на пороге и выше считать не нужно никогда; её можно вывести. Собственные итоги узла были посчитаны при его создании — здесь {No: 2, Yes: 3}, — так что всё, что не ниже порога, находится выше. Ни один вопрос ничего не перечитывает: по две подстановки и вычитание на каждый, а Джини следует из четырёх целых чисел.
Одна такая таблица строится на признак, поскольку у каждого столбца свой порядок и свои пороги: отсортировать по холестерину и пройти его, потом отсортировать по возрасту и пройти его, и так далее, а лучшая строка среди всех таблиц становится вопросом узла. В этом вся разница. Наша версия платит полный проход на каждого кандидата; проход-развёртка платит один проход на признак, а затем считывает каждого кандидата с таблицы, построенной по пути. На пяти строках это незаметно; на сотне тысяч это разница между секундой и неделей.