Как построить классификатор на дереве решений с нуля

Дерево решений делает прогноз, последовательно задавая вопросы о строке данных. Построим его на чистом Python для пяти пациентов и проверим каждое разбиение вручную.

Хороший пример — кливлендское исследование сердечных заболеваний, где каждая строка — пациент, а последний столбец — то, что мы хотим предсказать:

agesexchest_paincholesterolmax_heart_ratevesselsdisease
631typical2331500No
671asymptomatic2861083Yes
371nonanginal2501870No

Последний столбец — подкрашенный выше — содержит наблюдённое событие: у этого пациента сердечное заболевание обнаружилось, у того — нет. Такой столбец называется меткой; каждый столбец до него описывает случай, и только этот говорит, чем случай кончился. Выучить мы хотим связь между тем и другим — на случаях, где метка уже известна, чтобы применять её к случаям, где нет. Когда метка — категория, как здесь, это задача классификации.

Деревья решений также служат базовыми моделями в случайных лесах и градиентном бустинге. Разобравшись с выбором разбиений в одном дереве, проще понять эти ансамбли.

На табличных данных встречаются две задачи: классификация, где ответ — категория, и регрессия, где ответ — число: зарплата, цена. Дерево решений справляется с обеими. Эта статья строит классификатор; направьте тот же код на числовую цель, замените меру примеси дисперсией этой цели, а подсчёты меток в листе — их средним, и вы получите регрессионную модель, предсказывающую число. Больше в коде ничего не меняется, хотя в получившейся модели меняется многое — этому посвящена статья-компаньон.

Бустингу посвящена отдельная статья. Эта сосредоточена на том, что он повторяет сотни раз: как из таблицы строк строится одно дерево решений — откуда берутся его вопросы, как один из них выбирается среди прочих и когда разбиение останавливается. Мы пишем это на чистом Python, без NumPy и без scikit-learn, на пяти строках, достаточно маленьких, чтобы проверить каждое число вручную.

Скачайте полный пример на Python и запустите python3 heart_tree.py.

Два типа узлов и области, которые они вырезают

Мы будем работать с пятью пациентами из кливлендского исследования сердечных заболеваний. Полная таблица содержит 303 пациента и тринадцать признаков (предикторов), но мы возьмём лишь пять строк и два признака плюс метку:

#stress_testvesselsdisease
1normal0No
2fixed0Yes
3reversable2Yes
4reversable1Yes
5fixed0No
training_data = [
    ["normal", 0, "No"],
    ["fixed", 0, "Yes"],
    ["reversable", 2, "Yes"],
    ["reversable", 1, "Yes"],
    ["fixed", 0, "No"],
]

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

stress_test (Thal) — результат исследования с таллием: норма, фиксированный или обратимый дефект. В коде сохраняем исходное написание reversable. vessels (Ca) — число крупных сосудов, от 0 до 3, визуализированных при флюороскопии, а не число поражённых сосудов. Определения взяты из документации UCI.

Мы возьмём эти пять строк и будем делить их на всё более мелкие группы, используя признаки и их значения, чтобы решать, в какую группу попадает каждая строка. Последовательность разбиений рисуется как ветвящаяся диаграмма — отсюда и название «дерево».

flowchart TD n0{"Is vessels >= 1?"} n1["Yes: 2"] n2{"Is stress_test == fixed?"} n3["Yes: 1, No: 1"] n4["No: 1"] n0 -->|True| n1 n0 -->|False| n2 n2 -->|True| n3 n2 -->|False| n4 classDef pure fill:#dcf5e3,stroke:#3ba55c,color:#1c1c22; classDef mixed fill:#fdf0d0,stroke:#d9a514,color:#1c1c22; class n1,n4 pure; class n3 mixed;

На этой диаграмме два типа узлов.

Ромбы — узлы решений. Каждый хранит вопрос с ответом да/нет и направляет строку в ветвь True или False. Наше дерево по принципу CART всегда создаёт двух детей; другие алгоритмы могут использовать больше ветвей.

Прямоугольники — листья. В каждом хранятся количества меток попавших в него обучающих строк: Yes: 2, Yes: 1, No: 1 или No: 1. Деление на общее число даёт оценки частот классов: 100% Yes, 50/50 и 100% No. Эти крошечные выборки не определяют вероятность болезни у нового пациента.

Узлы направляют строки в листья, разбивая пространство входов. Для числовых признаков порог задаёт разрез вдоль одной оси. На рисунке отдельный синтетический набор иллюстрирует формы границ: прямая, гладкая кривая в стиле нейросети и разбиения дерева схематичны; панель k ближайших соседей вычислена по показанным точкам.

Four ways to draw a boundary class A class B
linear model
one straight cut — the tips are lost
neural network
stacked nonlinearities bend the boundary
k-nearest neighbours
the 3 nearest points vote — no model at all
decision tree
yes/no questions carve rectangles

Так что же на самом деле нужно, чтобы построить такое? Нам понадобятся две вещи.

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

Нужно и правило остановки. Здесь останавливаемся, если ни один кандидат не уменьшает неоднородность больше небольшой численной погрешности. Лист может остаться смешанным, даже если последовательность дальнейших разбиений разделила бы строки: пример — XOR. Ограничения глубины и размера листа дополнительно сдерживают переобучение.

От строк к вопросам

Создаём кандидатов, сочетая каждый признак с каждым его различным значением в текущих строках. Категориальное правило проверяет равенство, например stress_test == fixed. Числовое правило проверяет порог: vessels >= 1 истинно для значений 1, 2 и 3.

Это соединение и есть весь генератор, так что стоит выписать его целиком. В нашем наборе два столбца по три различных значения в каждом, так что получается список из шести пар:

типстолбецзначениекакой вопрос получается
категориальныйstress_testnormalIs stress_test == normal?
категориальныйstress_testfixedIs stress_test == fixed?
категориальныйstress_testreversableIs stress_test == reversable?
числовойvessels0Is vessels >= 0?
числовойvessels1Is vessels >= 1?
числовойvessels2Is vessels >= 2?

Наша реализация проверяет один признак за раз и берёт числовые пороги из наблюдаемых значений. Это упрощение, а не требование CART: scikit-learn использует середины между соседними различными значениями. Обучающие разбиения совпадают, но новое значение между наблюдениями может пойти по другой ветви. CART также допускает разбиения по подмножествам категорий; здесь проверяем одну категорию против остальных.

Эти шесть — всё, что дерево может спросить, и не все они попадут в готовое дерево. Большинство кандидатов пробуются, оцениваются и отвергаются; здесь только два доживают до того, чтобы стать вопросами в дереве, которое мы построим, а остальные четыре оцениваются и отбрасываются. И список конечен — никогда не длиннее числа различных значений в таблице, — поэтому следующий шаг может просто перебрать их все.

В коде вопросы определяются одним маленьким классом. Rule хранит индекс столбца и значение, а его метод holds решает, какое из двух сравнений применить, глядя на тип этого значения:

FEATURES = ["stress_test", "vessels", "disease"]

class Rule:
    """One yes/no test: a column, and the value it is compared against."""

    def __init__(self, column, value):
        self.column = column
        self.value = value

    def holds(self, row):
        observed = row[self.column]
        if isinstance(self.value, (int, float)):
            return observed >= self.value     # numeric: threshold
        return observed == self.value         # categorical: equality

    def __repr__(self):
        operator = ">=" if isinstance(self.value, (int, float)) else "=="
        return f"Is {FEATURES[self.column]} {operator} {self.value}?"

Эту логику сопоставления можно реализовать по-разному. Наша — проверка isinstance внутри holds, и именно она позволяет этому дереву работать с текстовым и числовым столбцом бок о бок вообще без предобработки. Библиотеки идут не все этим путём.

Scikit-learn требует числовые входы. Для неупорядоченного признака stress_test кодирование one-hot сохраняет категории, не задавая им порядок:

stress_testis_normalis_fixedis_reversable
normal→100
fixed→010
reversable→001

Тогда дерево спрашивает is_fixed >= 0.5 там, где наше спрашивает stress_test == fixed, — то же разбиение, размазанное по трём столбцам. Само 0.5 ничего не значит: столбец содержит только 0 и 1, так что любой разрез между ними отделяет те же строки, а sklearn ставит пороги в середину между двумя соседними значениями. Столбец с четырьмя категориями просто стал бы четырьмя такими столбцами 0/1, о каждом из которых спрашивают всё так же на 0.5 — кодировка растёт вширь, а каждый вопрос остаётся проверкой «да/нет» одного значения.

LightGBM вместо этого делит по подмножествам: он проверяет группу категорий разом, что по-прежнему один вопрос об одном столбце, — разница в том, что проверяемое значение является множеством, а не одной категорией:

ours:    Is stress_test == fixed?
theirs:  Is stress_test in {normal, reversable}?

LightGBM и XGBoost поддерживают разбиения категорий на группы, перебирая границы в порядке, заданном статистиками категорий. CatBoost для многих признаков использует упорядоченные целевые статистики, а для некоторых признаков с малым числом категорий — one-hot. Выбор зависит от настроек; см. документацию CatBoost.

Измеряем разнородность набора данных

Для выбора разбиения сначала измеряем смешанность меток с помощью критерия Джини. Затем вычисляем его уменьшение с учётом размеров групп. В коде это gain; строго говоря, прирост информации обычно означает соответствующее уменьшение энтропии.

Сначала посмотрим на примесь Джини — меру того, насколько перемешана коллекция: одно число, говорящее, всё ли в ней одного рода или это мешанина из многих.

Допустим, вам нужно сказать, какая из двух коллекций перемешана сильнее. Одного взгляда на картинку ниже достаточно, чтобы понять, что второй набор разнообразнее: четыре рода вместо двух и распределены ровнее. Глаз решает это мгновенно.

Which set is more diverse?

Теперь допустим, что состав ни одного из наборов нам неизвестен — ни подсчётов, ни списка родов, только возможность запустить руку и что-нибудь вынуть. Сможем ли мы выразить разнообразие числом?

Независимо выбираем два элемента с возвращением, отмечаем, различаются ли их типы, и повторяем. Первый элемент возвращаем до второго выбора: именно поэтому ниже перемножаются вероятности.

На рисунке показана иллюстративная последовательность из десяти пар для каждого набора:

Estimating diversity by drawing pairs
SameDifferentDifferentSameSameDifferentSameSameDifferentSameDifferent: 4 out of 10estimate = 0.40DifferentSameDifferentDifferentDifferentSameSameDifferentDifferentDifferentDifferent: 7 out of 10estimate = 0.70

Четыре пары левого набора содержат два разных рода; у правого таких семь. Поделите на число вытягиваний — и получите оценку; шляпка над d^\hat{d} обозначает величину, оценённую по выборке, в отличие от вычисленной по всей генеральной совокупности:

d^=pairs of different kindspairs drawn⇒410=0.40,710=0.70\hat{d} = \frac{\text{pairs of different kinds}}{\text{pairs drawn}} \qquad\Rightarrow\qquad \frac{4}{10} = 0.40, \qquad \frac{7}{10} = 0.70

Большая доля несовпадений указывает на большее разнообразие. Дополнительные независимые пары обычно улучшают оценку, но ни десять, ни сто пар не гарантируют точного результата.

Вообще говоря, сэмплировать не нужно вовсе: когда вы знаете состав набора, немного теории вероятностей даёт это точное значение напрямую. Вычислите шанс, что два выбора совпадут, и вычтите его из 1.

Возьмём левый набор. Семь из десяти его элементов — синие квадраты, так что один выбор оказывается квадратом с вероятностью 0.7, а вероятность вытянуть два квадрата подряд — 0.7×0.7=0.490.7 \times 0.7 = 0.49. Круги дают 0.3×0.3=0.090.3 \times 0.3 = 0.09. Это единственные два способа совпасть, так что совпадение происходит в 0.49+0.09=0.580.49 + 0.09 = 0.58 случаев. Но нам нужно обратное — как часто два выбора оказываются разными, — а поскольку каждое вытягивание либо совпадает, либо нет, это единица минус шанс совпадения: 1−0.58=1 - 0.58 = 0.42.

Правый набор — тот же расчёт, только с четырьмя родами вместо двух:

роддоляоба выбора попадают сюда
квадрат0.40.16
круг0.30.09
треугольник0.20.04
звезда0.10.01
совпадение 0.30

Два выбора совпадают в 30% случаев, значит различаются в 0.70 случаев — совпадая с семью из десяти, которые дало сэмплирование, но без единого вытягивания.

Ту же логику можно показать геометрически. Разложите каждую упорядоченную пару выборов как клетку сетки — первый выбор по горизонтали, второй по вертикали. Десять элементов дают сто клеток, и эта сетка — все возможные исходы:

Every pair of picks, one cell each
4102First elementSecond element4104103103102102101101101P(Both different)=P(Any pair)−P(Both equal)= 1−P(Both blue)−P(Both red)−P(Both green)−P(Both yellow)= 1−4102−3102−2102−1102= 1 − 0.16 − 0.09 − 0.04 − 0.01= 0.70
All 100 ordered pairs, one per cell. A cell is tinted when both picks are the same kind, so the matches clump into a square block per kind — side 4, 3, 2 and 1, giving 30 cells. The blue block is collapsed to show what a block is: a square of side 4/10, so area (4/10)². The 70 grey cells are the disagreements, and 0.70 is the Gini.

У этого точного значения есть имя. Вероятность того, что два случайно вытянутых из набора элемента окажутся разного рода, называется примесью Джини этого набора, и записывается она так:

Gini(S)=1−∑kpk2\text{Gini}(S) = 1 - \sum_{k} p_k^2

где pkp_k — доля набора, принадлежащая роду kk. Две половины — два способа сказать одно и то же: ∑kpk2\sum_k p_k^2 — вероятность совпадения вытягиваний (для каждого рода шанс, что оба попадут в него, сложенный по родам), а единица минус это — шанс, что они различаются.

Наши два набора, пропущенные через формулу, — это арифметика минутной давности в сжатом виде:

Gini(left)=1−(0.72+0.32)=1−0.58=0.42Gini(right)=1−(0.42+0.32+0.22+0.12)=1−0.30=0.70\begin{aligned} \text{Gini}(\text{left}) &= 1 - \left(0.7^2 + 0.3^2\right) &&= 1 - 0.58 &&= 0.42 \\ \text{Gini}(\text{right}) &= 1 - \left(0.4^2 + 0.3^2 + 0.2^2 + 0.1^2\right) &&= 1 - 0.30 &&= 0.70 \end{aligned}

Сумма ∑kpk2\sum_k p_k^2 используется также в индексе Симпсона и индексе Херфиндаля — Хиршмана. Критерий Джини равен единице минус эта сумма, а не самой сумме.

Джини на наших пяти строках

Теперь мы готовы вычислить примесь Джини для наших пяти строк. Сначала нужен подсчёт родов, которыми в нашем случае служат метки: где наборы выше содержали квадраты, круги, треугольники и звезду, куча строк содержит Yes и No. Так что сосчитайте их — сколько каких меток в данной куче, потому что каждая величина в этой статье вытекает из этого словаря.

def label_counts(rows):
    """Tally the labels in a pile — the label is always the last column."""
    counts = {}
    for row in rows:
        counts[row[-1]] = counts.get(row[-1], 0) + 1
    return counts

Первый прогон по всему набору, label_counts(training_data), даёт {'No': 2, 'Yes': 3} — наши пять пациентов, подсчитанные по родам.

Теперь, имея подсчёт под рукой, можно вычислить примесь Джини — пять строк Python:

def gini(rows):
    """Impurity of a pile: 0 when every row in it carries the same label."""
    if not rows:
        raise ValueError("Cannot measure impurity of an empty group")
    impurity = 1
    for count in label_counts(rows).values():
        share = count / len(rows)
        impurity -= share ** 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 — примерно настолько же перемешанно, насколько это возможно для пяти строк.

Уменьшение критерия Джини — оценка разбиения

Оценим кандидата: вычтем из неоднородности родителя взвешенную по размерам групп неоднородность дочерних узлов:

Записанный, он занимает одну строку:

Gain=Gini(S)−∣SL∣∣S∣Gini(SL)−∣SR∣∣S∣Gini(SR)\text{Gain} = \text{Gini}(S) - \frac{|S_L|}{|S|} \text{Gini}(S_L) - \frac{|S_R|}{|S|} \text{Gini}(S_R)

И состоит из четырёх шагов:

  1. разбить кучу SS вопросом, получив две кучи: SLS_L — строки, ответившие True, и SRS_R — строки, ответившие False;
  2. прогнать gini на каждой из них;
  3. свести эти два числа в одно, взвесив по тому, сколько строк ушло на каждую сторону: ∣SL∣∣S∣\frac{|S_L|}{|S|} даёт вес левой кучи, а ∣SR∣∣S∣\frac{|S_R|}{|S|} — правой, каждый из них есть доля строк родителя, ушедшая в эту сторону;
  4. вычесть это из примеси родителя, Gini(S)\text{Gini}(S).

Остаётся примесь, которую убрал вопрос: чем она выше, тем лучше вопрос.

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

Джини как ошибка прогноза

Если каждой строке листа присвоить один и тот же вектор вероятностей классов, эмпирические доли классов минимизируют среднюю сумму квадратов ошибок по индикаторам всех классов. Минимум равен 1−∑kpk21-\sum_k p_k^2, то есть критерию Джини. Это многоклассовая ошибка Брайера с суммированием по классам; бинарная версия только для вероятности положительного класса вдвое меньше. Жадные разбиения уменьшают эту ошибку по одному узлу, не гарантируя глобально оптимального дерева.

Стоит подчеркнуть, зачем вообще нужно взвешивание из шага 3, потому что без него оценку легко обмануть. Два наших кандидата, Is stress_test == normal? и Is vessels >= 1?, каждый делит пять строк на одного идеально чистого потомка с Джини ровно 0 и одного всё ещё перемешанного. Отличаются они тем, сколько данных уносит этот чистый потомок: один отслаивает единственного пациента и оставляет позади четыре перемешанные строки, другой забирает двоих и оставляет три. Только взвешивание видит эту разницу. Оно заставляет чистого потомка считаться ровно на столько, сколько он весит, так что потомок из одной строки едва заметен, а счёт задаёт оставленный беспорядок.

Вот оба кандидата, расписанные полностью, каждый со своими двумя потомками, сведёнными двумя способами: посчитанными поровну и взвешенными по доле строк, которую держит каждый потомок:

Is stress_test == normal?its two children combined two ways
stress_testvesselsdiseasenormal0Nofixed0Yesreversable2Yesreversable1Yesfixed0Noimpurity = 0.48Is stress_test == normal?FalseTruefixed0Yesreversable2Yesreversable1Yesfixed0Noimpurity = 0.3754 rows of 5normal0Noimpurity = 01 row of 5counted equallygain = 0.48 − (0 + 0.375) ÷ 2 = 0.48 − 0.188 = 0.293weighted by rowsgain = 0.48 − (⅕ × 0 + ⅘ × 0.375) = 0.48 − 0.30 = 0.180

Теперь то же самое для другого кандидата. Is vessels >= 1? тоже отрезает идеально чистого потомка, но в нём два пациента, а не один, а куча, которую он оставляет, — три строки, а не четыре, и она грязнее: 0.444 вместо 0.375:

Is vessels >= 1?its two children combined two ways
stress_testvesselsdiseasenormal0Nofixed0Yesreversable2Yesreversable1Yesfixed0Noimpurity = 0.48Is vessels >= 1?FalseTruenormal0Nofixed0Yesfixed0Noimpurity = 0.4443 rows of 5reversable2Yesreversable1Yesimpurity = 02 rows of 5counted equallygain = 0.48 − (0 + 0.444) ÷ 2 = 0.48 − 0.222 = 0.258weighted by rowsgain = 0.48 − (⅖ × 0 + ⅗ × 0.444) = 0.48 − 0.27 = 0.213

Так что числа на двух рисунках показывают нечто более сильное, чем изменение масштаба. Посчитанные поровну, Is stress_test == normal? даёт 0.293, а Is vessels >= 1? — 0.258, так что побеждает первый вопрос. Взвешенные, они выходят 0.180 и 0.213, и побеждает уже второй. Взвешивание не просто уменьшает оценки — оно переворачивает порядок, а поскольку это корень, два ответа дают деревья, различающиеся сверху донизу.

Вот как мы реализуем разбиение и его прирост информации, вместе со взвешиванием. split_rows выполняет шаг 1, раскладывая строки по двум кучам, которые делает вопрос, а split_gain выполняет шаги 2–4, оценивая то, что получилось, против того, что было:

def split_rows(rows, rule):
    """Sort every row into the pile where the rule holds, and the pile where it does not."""
    true_pile, false_pile = [], []
    for row in rows:
        (true_pile if rule.holds(row) else false_pile).append(row)
    return true_pile, false_pile

def split_gain(parent_impurity, true_pile, false_pile):
    """What went in, minus the two piles that came out, each weighed by its share."""
    share = len(true_pile) / (len(true_pile) + len(false_pile))
    return parent_impurity - share * gini(true_pile) - (1 - share) * gini(false_pile)

Прогоните их на двух кандидатах выше — и они вернут 0.180 и 0.213, те же числа, что рисунки посчитали вручную, только теперь вычисленные, а не нарисованные.

Механизм — разбить, потом рекурсировать

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

  1. Попробовать каждый вопрос, который допускают данные, — каждый признак, каждое значение, которое этот признак принимает.
  2. Оценить каждый вопрос по тому, насколько он расслаивает метки в куче, — это прирост информации, построенный на примеси Джини, ровно как мы только что вывели.
  3. Если ни один вопрос не помогает, остановиться: куча становится листом, а её подсчёты меток — предсказанием.
  4. Иначе разбить кучу лучшим вопросом на две кучи поменьше.
  5. Запустить ту же процедуру на каждой из двух куч.

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

Жадность означает выбор лучшего разбиения текущего узла без просмотра будущих шагов и пересмотра прежних решений. Разбиение без немедленного выигрыша может открыть полезные следующие разбиения, как в XOR. Наше правило остановки их пропустит.

Быть рекурсивным — это то, что вырезает прямоугольники с рисунка во вступлении: каждый вызов владеет одной областью пространства признаков (строками, пережившими вопросы выше) и либо подразделяет эту область, либо запечатывает её как лист. Прямоугольники — это кучи на дне рекурсии.

Выбор корневого разбиения — и ничья

Прежде чем строить всю рекурсивную конструкцию, быстро посмотрим на реализацию части, работающей в одном узле, — поиска лучшего разбивающего вопроса. В корне этот узел держит всех пятерых пациентов, а ищет функция choose_split, которая пробует каждое значение каждого признака и оставляет лучшее.

Это два вложенных цикла — каждый столбец снаружи, каждое различное значение этого столбца внутри, — и каждая порождённая ими пара проходит четыре шага:

  1. построить Rule из столбца и значения;
  2. передать его в split_rows, который раскладывает строки по двум кучам;
  3. оценить эти кучи через split_gain;
  4. сравнить оценку с лучшей на данный момент и оставить вопрос, если он побеждает.

Когда оба цикла заканчиваются, функция возвращает вопрос, всё ещё держащий лучшую оценку.

def choose_split(rows):
    parent_impurity = gini(rows)
    winning_gain, winning_rule = 0, None

    for column in range(len(rows[0]) - 1):
        for value in sorted(set(row[column] for row in rows), reverse=True):
            rule = Rule(column, value)
            true_pile, false_pile = split_rows(rows, rule)

            if not true_pile or not false_pile:
                continue  # this rule doesn't divide the data

            gain = split_gain(parent_impurity, true_pile, false_pile)

            if gain >= winning_gain:
                winning_gain, winning_rule = gain, rule

    return winning_gain, winning_rule

Один вызов в корне оценивает каждый вопрос, порождённый генератором, и возвращает вот это:

вопросприрост
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. Один потомок держит всю кучу, а другой — ничего, и это не разбиение, а копия: ничего не разделено, значит нечего оценивать. Защита not true_pile or not false_pile отбрасывает его до того, как split_gain его увидит. Его прирост всё равно был бы ровно 0 — один потомок ничего не весит, другой равен родителю, — но пропуск ещё и не даёт вернуть «неразбиение» как лучший вопрос, когда ничто другое не набирает больше нуля.

Теперь посмотрите на верх этой таблицы, потому что там самое интересное. Два разных вопроса вернули одинаковую оценку 0.2133. Когда так случается, поздний кандидат перезаписывает раннего. Такое поведение — деталь реализации, и в нашем алгоритме оно следует из двух вещей: столбцы просматриваются в порядке индексов, так что stress_test (столбец 0) добирается первым, а затем его тихо вытесняет столь же хороший vessels (столбец 1); и сравнение написано через >=, а не >, что и позволяет вытеснению случиться:

if gain >= winning_gain:

Здесь вопросы с одинаковой оценкой выделяют одни и те же две группы обучающих пациентов, поэтому дальнейшие разбиения этих групп не меняются. Но прогноз для новой комбинации, например normal и vessels=3, может измениться. Перебираем различные значения по убыванию для воспроизводимости; >= по-прежнему оставляет последнего кандидата с равной оценкой.

Рекурсия — строим дерево

Теперь, когда отдельный узел умеет находить свой вопрос, мы готовы построить всё дерево — рекурсию, которая гоняет этот поиск по куче за кучей и сохраняет находки. Для хранения нужен один класс на каждый тип узла из вступительной диаграммы: Leaf держит подсчёты меток тех строк, что до него дошли, а Node держит вопрос и две ветви. На языке учебника вопрос — это правило разбиения, один предикат по одному признаку, откуда и берёт имя класс Rule, а узел решения — это правило, встроенное в блок-схему, где двум ветвям есть куда вести свои ответы «да» и «нет». Готовое дерево можно мыслить как серию правил разбиения. Начиная с верхушки дерева и применяясь по пути вниз — choose_split учит эти правила, а узлы — место, где живут выбранные.

class Leaf:
    def __init__(self, rows):
        self.counts = label_counts(rows)

class Node:
    def __init__(self, rule, if_true, if_false):
        self.rule = rule
        self.if_true = if_true
        self.if_false = if_false

def grow_tree(rows):
    if not rows:
        raise ValueError("Training rows must not be empty")
    gain, rule = choose_split(rows)

    if gain <= 1e-12 or rule is None:
        return Leaf(rows)          # base case: no rule helps anymore

    true_pile, false_pile = split_rows(rows, rule)
    return Node(rule, grow_tree(true_pile), grow_tree(false_pile))

Когда мы запускаем это на пяти пациентах, получается вот такое дерево — нарисованное так, что кучи видны на каждой ветви:

The five patients, flowing down the finished tree
stress_testvesselsdiseasenormal0Nofixed0Yesreversable2Yesreversable1Yesfixed0NoIs vessels >= 1?FalseTruereversable2Yesreversable1Yespure — all Yes,done in one questionnormal0Nofixed0Yesfixed0Nostill mixedIs stress_test == fixed?FalseTruenormal0Nopure — no mixturefixed0Yesfixed0Nomixed — identical featuresno question separates them

У дерева глубина 2 и три листа. Корень направляет две обучающие строки с vessels >= 1 в чистый лист Yes. Это описание двух строк, а не медицинское правило.

На этом стоит задержаться: целая группа выпала из данных вообще без примеси — Джини 0, полученный одним вопросом. Уровнем ниже одинокий normal-пациент делает то же самое, — и заметьте, что у stress_test три значения, а дерево спрашивает лишь об одном из них. stress_test == fixed? отслаивает пациентов с фиксированным дефектом, а всё, что не fixed, едет по ветви False вместе, неразличённое. Здесь это оказывается единственный normal-пациент, потому что оба reversable-пациента ушли ещё в корне.

Из трёх листьев дерева два чистые; каждая обучающая строка, кроме сталкивающейся пары, попадает в группу с нулевой примесью, и рекурсия останавливается в каждой из них именно потому, что убирать больше нечего.

Остаётся третий лист, держащий одну метку Yes и одну No при одинаковом наборе значений признаков. Никакой вопрос не смог бы разделить этих двух пациентов — и никакая другая модель тоже, потому что то, что их различает, в данных попросту отсутствует. Это можно было бы разрешить, взяв больше предикторов, ведь одиннадцать выброшенных столбцов вполне могут содержать то, что этих двух пациентов разделяет.

Обучение выбирает структуру дерева, пороги и количества меток в листьях. После завершения рекурсивных вызовов модель готова к прогнозированию; градиентных обновлений и эпох здесь нет.

Классификация — считываем вероятность с листа

Предсказание — снова рекурсия, и короче обучающего кода:

def descend(row, node):
    if isinstance(node, Leaf):
        return node.counts

    branch = node.if_true if node.rule.holds(row) else node.if_false
    return descend(row, branch)

def as_percentages(counts):
    total = sum(counts.values())
    return {label: f"{count / total:.0%}" for label, count in counts.items()}

Каждый Node хранит один Rule — индекс столбца плюс значение, — а holds сравнивает запись строки в этом столбце с ним, возвращая обычный True или False. Этот булев результат — единственное, что нужно descend: True отправляет строку по if_true, False — по if_false, и рекурсия останавливается, как только попадает на Leaf.

Возьмём данные одного пациента, ['fixed', 0, 'Yes'], и посмотрим, как дерево предсказывает, есть ли у него сердечное заболевание:

  1. корень спрашивает Is vessels >= 1?; holds читает запись vessels пациента — это 0 — и, поскольку значение числовое, вычисляет 0 >= 1, получая False, так что строка идёт по ложной ветви;
  2. этот узел спрашивает Is stress_test == fixed?; holds читает запись stress_test пациента — 'fixed' — и, поскольку значение строковое, вычисляет 'fixed' == 'fixed', получая True, так что строка идёт по истинной ветви;
  3. эта ветвь — Leaf, так что descend возвращает хранящиеся там подсчёты: один Yes и один No.

Эти возвращённые descend подсчёты и есть предсказание в сыром виде. Их можно прочесть как одну метку, взяв ту, что чаще встречается в листе, — так делает predict в библиотеке вроде sklearn, и в чистых листьях это однозначно, где {'Yes': 2} означает Yes. А можно прочесть как вероятность, поделив каждый счётчик на сумму, — это predict_proba, и то, что здесь делает as_percentages.

Смешанный лист выдаёт 50/50, потому что содержит по одному примеру каждого класса. Это эмпирическая оценка по двум строкам, а не доказательство риска 50% в популяции и не полная оценка неопределённости. Противоречивые метки при одинаковых входах задают нижнюю границу обучающей ошибки детерминированного предиктора на этих признаках. Библиотеке, возвращающей один класс, также нужно правило разрешения равенства частот.

Виджет ниже — чуть более крупная игрушка (два числовых признака, пороги вместо наших смешанных типов), но механизм идентичен, и он показывает два взгляда на дерево сразу. Левая панель — разбиение; правая — обход. Два ползунка — значения признаков x1x_1 и x2x_2; их перетаскивание составляет новую строку и двигает её по пространству признаков. Момент, когда точка пересекает пунктирную линию, — ровно тот момент, когда меняется путь по дереву, потому что область и лист — один и тот же объект в разных нарядах.

One model, two views — move the point
Decision regions
AAAB047100610x₁x₂
Decision tree
YesNoYesNoYesNox₁ < 4?x₂ > 6?x₁ > 7?AAAB

Заметьте также, что это дерево использует признак, обозначенный x1x_1, дважды — один раз в корне и ещё раз двумя уровнями ниже при другом пороге. Это два разных вопроса по одному столбцу — тот же признак, другое значение, — потому что признак не расходуется от того, что по нему разбили: первый разрез разделяет, что может, а оставшиеся строки могут по-прежнему разделяться вдоль той же оси.

Один признак может использоваться в нескольких узлах. Каждый узел заново создаёт кандидатов по попавшим в него строкам, поэтому отсутствующее в них значение больше не рассматривается. Например, у ребёнка False нашего корня нет строки с 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" отдельной категорией.

Новая категория идёт в ту же ветвь False, что и любое другое значение, не прошедшее проверку равенства. Модель не обнаруживает новизну и не снижает показанный процент. В примере выше 100% No основано на листе всего с одной обучающей строкой.

Поддержка пропусков зависит от реализации. Некоторые библиотеки обучают направление по умолчанию, а исходный CART может использовать замещающие разбиения. Это не превращает произвольную строку вроде "NA" в распознаваемый пропуск: нужно представление, ожидаемое библиотекой.

Почему одно дерево — не конец истории

Это небольшая реализация по принципу CART. Её основные ограничения — неограниченный рост, чувствительность к обучающей выборке, упрощённая обработка категорий и отсутствие явных правил для пропусков.

Глубокое дерево может переобучиться. Создавая маленькие листья, оно улучшает обучение, но может подстраиваться под шум выборки. Остановка при малом приросте исключает бесполезные ближайшие разбиения, но не контролирует обобщение. Для этого нужны ограничения глубины, размера листьев и обрезка.

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

Теперь направьте его на набор данных о раке груди (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 (scikit-learn)
  depth      7
  leaves     19
  train acc  1.000
  test acc   0.912

Обучающая точность 1.000 означает воспроизведение всех 398 обучающих меток. Отдельный лист для каждой строки не нужен: здесь всего 19 листьев. Точность на 171 отложенной строке равна 0.912.

Теперь добавим единственную ручку, которой в нашей версии нет, — max_depth, жёсткий потолок на то, на сколько вопросов вглубь может уйти дерево, останавливающий разбиение независимо от того, остался ли прирост. Прогоняя то же дерево с рядом ограничений глубины и собирая результаты, получаем:

max_depthобучениетестразрыв
10.9300.895+0.035
20.9600.947+0.012
30.9670.947+0.020
50.9870.936+0.052
7 (без обрезки)1.0000.912+0.088

На этом разбиении глубина 2 даёт 0.947 на отложенных данных, а глубина 7 — 0.912, несмотря на лучшую обучающую точность. Это пример переобучения, а не доказательство универсальной оптимальности глубины 2. Если по этим результатам выбирать глубину, отложенный набор становится валидационным. Итоговую модель следует оценить на отдельном тестовом наборе.

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

Сортировка кандидатов устраняет различия между запусками из-за порядка множества Python. Статистическая нестабильность остаётся: другая обучающая выборка всё ещё может дать другое дерево.

Эти два поведения — две стороны компромисса смещения и дисперсии. В статистическом обучении смещение — это ошибка, вносимая приближением сложной реальности более простой моделью: модель, слишком жёсткая, чтобы представить закономерность, будет ошибаться, сколько данных ей ни дай. Дисперсия — это то, насколько сильно изменилась бы подогнанная модель, если оценить её по другому обучающему набору: переобучите на другой выборке пациентов, и метод с высокой дисперсией даст вам заметно другую модель, делающую заметно другие ошибки.

Как смещение и дисперсия складываются в ошибку предсказания

В регрессии с квадратичной ошибкой при фиксированном входе ожидаемая тестовая ошибка раскладывается на квадрат смещения, дисперсию прогноза по обучающим выборкам и условную дисперсию шума. Усреднение идёт по повторным обучающим выборкам и новым исходам:

prediction error=bias2⏟wrong assumptions+variance⏟sensitivity to the training set+noise⏟randomness in the data\text{prediction error} = \underbrace{\text{bias}^2}_{\text{wrong assumptions}} + \underbrace{\text{variance}}_{\text{sensitivity to the training set}} + \underbrace{\text{noise}}_{\text{randomness in the data}}

Смещение входит в квадрате, потому что это знаковая величина — насколько далеко среднее предсказание модели отстоит от истины, — которая иначе взаимно уничтожалась бы, а не накапливалась. Для ошибки классификации работают те же три источника, но складываются они не так аккуратно; интуиция переносится, арифметика — нет.

Смещение — разница между средним прогнозом и истинным условным средним. Дисперсия описывает изменчивость прогноза между обучающими выборками. У метода могут быть оба компонента; это не отдельные части каждой конкретной ошибки.

Увеличение глубины часто снижает смещение и повышает дисперсию, но это тенденция, а не гарантия для любого набора или показателя точности. Шум при заданных признаках ограничивает ожидаемое качество; противоречивая пара в малой выборке не определяет величину этого предела в популяции.

Есть устоявшиеся методы борьбы с этими проблемами: потолок глубины, минимальное число строк на лист, минимальный прирост, ради которого стоит разбивать, и обрезка ветвей постфактум. На практике их редко применяют к одинокому дереву; это ручки, которые крутят внутри ансамбля — модели, построенной из многих деревьев, чьи ответы объединяются в один, а именно этим и являются случайный лес и градиентный бустинг.

Предварительная обрезка ограничивает рост параметрами вроде max_depth и min_samples_leaf. Последующая обрезка сначала выращивает большое дерево, затем удаляет ветви. Обрезка CART по стоимости и сложности сочетает обучающую неоднородность со штрафом за каждый лист; в scikit-learn это ccp_alpha. Значение выбирают по валидации или кросс-валидации, сохраняя отдельный тестовый набор.

Случайный лес усредняет много деревьев, обученных со случайным выбором строк и признаков, снижая дисперсию, если их ошибки не полностью коррелированы. Градиентный бустинг последовательно добавляет деревья для уменьшения потерь ансамбля. Далее — статья о регрессионном дереве.

Где наша версия медленнее настоящей

Наш поиск разбиений многократно просматривает одни и те же строки. Производственные реализации могут переиспользовать счётчики или гистограммы. Они отличаются также порогами, поддержкой категорий и пропусков, правилами ничьих; скорость — не единственное различие.

Цикл проходит по каждому признаку и каждому его различному значению, так что число кандидатов равно признаки × значения. Затем каждый кандидат стоит полного прохода по данным: split_rows обходит каждую строку, раскладывая её по двум кучам, а split_gain вызывает gini на каждой куче, которая считает её метки с нуля. Это O(features×values×rows)O(\text{features} \times \text{values} \times \text{rows}). На пяти строках с тремя значениями на столбец — незаметно. На непрерывном признаке — скажем, холестерине — почти каждая строка несёт своё значение, так что число кандидатов растёт вместе с данными, а каждый кандидат по-прежнему стоит полного прохода: квадратично по числу строк и безнадёжно при сотне тысяч.

Возьмём пять строк столбца cholesterol — 210 (No), 233 (No), 250 (Yes), 286 (Yes), 300 (Yes). Генератор превращает их в пять вопросов-кандидатов, по одному на наблюдённое значение, из которых четыре действительно делят кучу:

кандидатниже порогана пороге или выше
>= 210пустовсе пять
>= 233210233, 250, 286, 300
>= 250210, 233250, 286, 300
>= 286210, 233, 250286, 300
>= 300210, 233, 250, 286300

Проследим двух из них, >= 233 и >= 250, через наш код.

Для >= 233 split_rows обходит все пять строк и кладёт 210 в список False, а остальные четыре — в список True. Затем split_gain вызывает gini на каждом, и gini обходит кучу из одной строки, считая метки, потом кучу из четырёх строк, считая метки. Пять посещений на разбиение, пять на подсчёт. Для >= 250 всё начинается заново с тех же пяти строк, и так далее по списку:

>= 233:  split_rows 5 rows → gini({210}) + gini({233,250,286,300})   = 10 visits
>= 250:  split_rows 5 rows → gini({210,233}) + gini({250,286,300})   = 10 visits
>= 286:  split_rows 5 rows → gini({210,233,250}) + gini({286,300})   = 10 visits
>= 300:  split_rows 5 rows → gini({210,233,250,286}) + gini({300})   = 10 visits

Сорок посещений строк, и между строками ничего не переносится — хотя каждая пара куч отличается от пары выше ровно одной строкой.

Вот те же пять строк, отсортированные по холестерину и несущие метку disease, которая у каждого пациента оказалась, — так их держала бы настоящая реализация:

cholesteroldisease
210No
233No
250Yes
286Yes
300Yes

Настоящие реализации получают оба ответа из первого же обхода. Отсортируйте строки по признаку, а затем, оценивая самого первого кандидата, пройдите по ним один раз, ведя нарастающий подсчёт увиденных меток: к концу этого единственного прохода ответ получен и для каждого последующего кандидата:

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}, — так что всё, что не ниже порога, находится выше. Ни один вопрос ничего не перечитывает: по две подстановки и вычитание на каждый, а Джини следует из четырёх целых чисел.

Для nn строк и dd признаков наивный поиск может требовать O(dn2)O(dn^2) операций в одном узле. Сортировка каждого числового признака с последующим проходом снижает стоимость примерно до O(dnlog⁡n)O(dn\log n) с учётом сортировки при фиксированном числе классов. Сам проход линеен. Это оценка для узла; стоимость всего дерева зависит и от его формы.