Несколько взглядов на кросс-энтропию

от автора

Всем привет! В этой статье я хотел бы попробовать осветить несколько взглядов на кросс-энтропию и попробовать сформировать некоторую интуицию с точки зрения теории информации.

Кросс-энтропия — одна из центральных метрик машинного обучения. Любая задача классификации так или иначе сводится к оптимизации модели через эту метрику. Однако новички нередко задаются вопросом, откуда берётся эта метрика и почему имеет такой вид, так как нередко вне профильных курсов она преподносится как данность.

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

Начнём с классического вывода: через метод максимального правдоподобия (далее будем называть его MLE). Вспомним курс статистики и что из себя вообще представляет функция правдоподобия.

У нас есть выборка D = \{(x_i, y_i)\}_{i=1}^{n} и параметрическая модель p_\theta. Функция правдоподобия — это вероятность увидеть ровно те данные, которые мы наблюдаем, как функция от параметров:

L(\theta) = p_\theta(D) = \prod_{i=1}^{n} p_\theta(y_i \mid x_i).

Тут p_\theta(y \mid x) мы рассматриваем не как функцию от y при фиксированных параметрах, а как функцию от \theta при фиксированных данных. Произведение берётся потому, что объекты выборки считаются независимыми (при фиксированных x_i):

\hat{\theta}_{\text{MLE}} = \arg\max_{\theta} L(\theta).

Рассмотрим простейший бинарный случай: y_i \in \{0, 1\}, модель выдаёт \hat{y}_i = p_\theta(y_i = 1 \mid x_i) \in (0,1). Это распределение Бернулли, и вероятность конкретного исхода записывается одной формулой:

p_\theta(y_i \mid x_i) = \hat{y}_i^{\,y_i} \, (1 - \hat{y}_i)^{\,1 - y_i}.

Трюк со степенями здесь чисто технический: при y_i = 1 второй множитель обращается в единицу и остаётся \hat{y}_i, при y_i = 0 — наоборот, остаётся 1 - \hat{y}_i. То есть формула просто выбирает вероятность того исхода, который реально произошёл.

Несложно обобщить на многоклассовый случай. Пусть классов K, модель выдаёт вектор вероятностей \hat{y}_i = (\hat{y}_{i1}, \dots, \hat{y}_{iK}), \sum_k \hat{y}_{ik} = 1, а целевую метку кодируем one-hot вектором y_i, где y_{ik} = 1 для истинного класса и 0 иначе. Тогда тот же трюк со степенями даёт

p_\theta(y_i \mid x_i) = \prod_{k=1}^{K} \hat{y}_{ik}^{\,y_{ik}},

и всё произведение снова схлопывается в один множитель — вероятность истинного класса.

После берём логарифм функции правдоподобия, так как с ним легче работать (произведение переходит в сумму + работает численно стабильнее). Логарифм монотонен, поэтому точка максимума не меняется. Домножим ещё на -1, чтобы вместо максимизации получить привычную минимизацию:

-\log L(\theta) = -\sum_{i=1}^{n} \log p_\theta(y_i \mid x_i).

Отсюда следует знакомая формула. Для бинарного случая

\mathcal{L} = -\sum_{i=1}^{n} \Big[\, y_i \log \hat{y}_i + (1 - y_i)\log(1 - \hat{y}_i) \,\Big],

и для многоклассового

\mathcal{L} = -\sum_{i=1}^{n} \sum_{k=1}^{K} y_{ik} \log \hat{y}_{ik} = -\sum_{i=1}^{n} \log \hat{y}_{i, c_i},

где c_i — индекс истинного класса i-го объекта. В правой части из-за one-hot кодирования вся внутренняя сумма сводится к одному слагаемому: в лосс входит только вероятность, приписанная правильному классу.

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

Вывод через теорию информации

Введём главный объект, с которым будем работать, — энтропию:

H(p) = -\sum_{x} p(x) \log p(x) = \mathbb{E}_{x \sim p}\big[-\log p(x)\big].

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

Хорошая иллюстрация — известная логическая задача про фальшивую монетку. Пусть есть 9 одинаковых на вид монет, одна из которых легче остальных, и чашечные весы. За сколько взвешиваний гарантированно найдём фальшивую?

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

H = \log_2 9 \approx 3.17 \text{ бита.}

Одно взвешивание — это канал с тремя возможными исходами: левая чаша легче, правая легче, равновесие. Больше \log_2 3 \approx 1.585 бита такой канал за раз не передаст, причём этот максимум достигается только тогда, когда все три исхода равновероятны. Значит, взвешиваний нужно не меньше, чем

\frac{\log_2 9}{\log_2 3} = 2.

И этот теоретический минимум действительно достижим: кладём по три монеты на каждую чашу, три откладываем в сторону. Каждый из трёх исходов имеет вероятность 1/3 и оставляет ровно три подозрительные монеты; вторым взвешиванием тем же приёмом находим фальшивую. Заметно, что энтропийная граница подсказывает и саму стратегию: делить нужно на равные части, потому что именно равновероятные исходы выжимают из взвешивания максимум бит. Классическая версия задачи с 12 монетами, где неизвестно, легче фальшивая или тяжелее, решается тем же способом: там 24 равновероятных исхода, \log_2 24 / \log_2 3 \approx 2.9, откуда честная нижняя граница в 3 взвешивания.

Для интересующихся: почитать про аксиоматический вывод энтропии Шеннона можно в оригинальной статье Шеннона 1948 года (раздел 6 и Приложение 2).

Внутри математического ожидания стоит величина -\log p(x), её называют собственной информацией, или «удивлением» (surprisal). Логика простая: если событие почти достоверно, p(x) \to 1, то узнать о том, что оно произошло, — это ноль новой информации, и -\log p(x) \to 0. Если событие крайне редкое, p(x) \to 0, то его наступление удивляет сильно, и -\log p(x) \to \infty. Энтропия — это просто среднее удивление. Если брать \log_2, всё меряется в привычных нам битах, если натуральный — в натах (от англ. natural). На оптимизацию выбор основания не влияет, так как это просто константный множитель.

Теперь введём ещё один объект — KL-дивергенцию:

D_{\mathrm{KL}}(p \,\|\, q) = \sum_{x} p(x) \log \frac{p(x)}{q(x)} = \mathbb{E}_{x \sim p}\left[\log \frac{p(x)}{q(x)}\right].

Она показывает расстояние между двумя взятыми распределениями p и q, точнее, показывает, насколько мы в среднем ошибаемся, когда думаем, что распределение это q, хотя на самом деле в реальности это p. Внутри ожидания стоит разность двух удивлений: \log \frac{p(x)}{q(x)} = (-\log q(x)) - (-\log p(x)), то есть «насколько сильнее меня удивил исход x, чем должен был бы».

У KL есть три свойства, которые стоит держать в голове:

  1. D_{\mathrm{KL}}(p \| q) \geq 0 всегда — это неравенство Гиббса, следствие выпуклости -\log и неравенства Йенсена. Доказательство можно глянуть вот тут.

  2. D_{\mathrm{KL}}(p \| q) = 0 тогда и только тогда, когда p = q (почти всюду). То есть ноль достигается ровно в одной точке — когда мы точно угадали оригинальное распределение.

  3. Это не метрика в строго математическом смысле. D_{\mathrm{KL}}(p \| q) \neq D_{\mathrm{KL}}(q \| p), и неравенство треугольника не выполняется. Поэтому расстояние тут — это скорее просто наименование; формально правильнее говорить «дивергенция».

Важное замечание по асимметрии: ожидание берётся по p, поэтому штрафуются только те точки, где у p есть масса. Если

p(x) > 0, а q(x) \to 0, под логарифмом возникает бесконечность, и значение дивергенции взрывается. Обратная ситуация нормальна: там, где p(x) = 0, значение q(x) вообще не проверяется. Отсюда известное поведение: прямая KL даёт mode-covering приближения (модель обязана накрыть всё, что реально встречается), обратная KL — mode-seeking (модель может залипнуть в одну моду).

Теперь достаточно легко можно обнаружить следующее тождество. Разобьём логарифм отношения на разность:

D_{\mathrm{KL}}(p \,\|\, q) = \mathbb{E}_p[\log p(x)] - \mathbb{E}_p[\log q(x)] = -H(p) + H(p, q),

откуда

H(p, q) = H(p) + D_{\mathrm{KL}}(p \,\|\, q),

где H(p, q) = -\sum_x p(x) \log q(x) — уже известная нам кросс-энтропия.

Так и зачем все эти сложности?

Перейдём к интерпретации: из тождества выше видно, что кросс-энтропия распадается на два разных слагаемых. H(p) — это энтропия самих данных, от параметров модели она не зависит. D_{\mathrm{KL}}(p \| q) — это то, с чем мы работаем: насколько наша модель q промахивается мимо реального распределения p. Фактически минимизация кросс-энтропии — это минимизация KL-дивергенции: поскольку H(p) не зависит от параметров модели, обе задачи имеют один и тот же оптимум и одни и те же градиенты, \nabla_\theta H(p, q_\theta) = \nabla_\theta D_{\mathrm{KL}}(p \| q_\theta). Вычитать H(p) при обучении попросту незачем. Другое дело, если нужно именно численное значение KL: тогда H(p) знать необходимо, а истинное p нам обычно недоступно, так что честную KL-дивергенцию мы посчитать не можем.

Отсюда следует достаточно явный вывод: абсолютное значение лосса мало о чём говорит. Лосс 0.3 — это плохо или хорошо? Ответ зависит от H(p). Если задача шумная и разметчики сами не сходятся, то H(p) может быть 0.25, и мы почти у идеала. Если задача детерминированная, H(p) = 0, и мы всё ещё далеко.

Но самое интересное, на мой взгляд, — это интерпретация через кодирование. Величина -\log_2 q(x) — это длина в битах, которую оптимальный код припишет символу x, если считать, что символы приходят из распределения q. Частым символам достаются короткие коды, редким — длинные. Тогда:

  • H(p) — средняя длина сообщения, если код построен под истинное распределение. Это теоретический минимум (теорема Шеннона об источнике).

  • H(p, q) — средняя длина, если код построен под q, а данные на самом деле идут из p.

  • D_{\mathrm{KL}}(p \| q) — переплата. Лишние биты, которые появляются из-за ошибок при приближении к реальному распределению.

Разберём на конкретном примере. Пусть источник выдаёт четыре символа с вероятностями

p = \left(\tfrac{1}{2},\ \tfrac{1}{4},\ \tfrac{1}{8},\ \tfrac{1}{8}\right) \quad \text{для } (A, B, C, D).

Оптимальный код (код Хаффмана) здесь такой:

Символ

p

Код

Длина

A

1/2

0

1

B

1/4

10

2

C

1/8

110

3

D

1/8

111

3

Длины ровно совпадают с -\log_2 p(x), и средняя длина сообщения равна

H(p) = \tfrac{1}{2}\cdot 1 + \tfrac{1}{4}\cdot 2 + \tfrac{1}{8}\cdot 3 + \tfrac{1}{8}\cdot 3 = 1.75 \text{ бита на символ.}

Теперь представим, что наша «модель» считает распределение равномерным: q = (\tfrac{1}{4}, \tfrac{1}{4}, \tfrac{1}{4}, \tfrac{1}{4}). Под такое q оптимальный код — фиксированные два бита на символ: 00, 01, 10, 11. Код корректный, сообщения декодируются. Но средняя длина теперь

H(p, q) = \sum_x p(x) \cdot 2 = 2 \text{ бита на символ,}

а переплата составляет

D_{\mathrm{KL}}(p \,\|\, q) = 2 - 1.75 = 0.25 \text{ бита на символ.}

Прямая проверка по формуле даёт то же самое: \tfrac{1}{2}\log_2\tfrac{1/2}{1/4} + \tfrac{1}{4}\log_2 1 + 2 \cdot \tfrac{1}{8}\log_2\tfrac{1/8}{1/4} = 0.5 + 0 - 0.25 = 0.25.

Получается, что мы недооценили частый символ A (дали ему 2 бита вместо 1) и переоценили относительно редкие C и D. На миллионе символов это 250 000 лишних бит. Модель классификации можно интерпретировать так же: обучая её кросс-энтропией, мы стараемся построить максимально экономный код для реальных меток. Уверенное и правильное предсказание — короткий код. Уверенное и неправильное — очень длинный: -\log_2 0.001 \approx 10 бит за один объект.

Калибровка

Из кодовой интерпретации почти сразу выпадает идея калибровки. Раз переплата D_{\mathrm{KL}}(p \| q) обнуляется тогда и только тогда, когда q = p, то оптимум кросс-энтропии достигается не на угадывании класса, а на сообщении истинных вероятностей. На языке статистики это называется строго правильным правилом оценивания (strictly proper scoring rule).

Сравним с accuracy: она не различает предсказания 0.51 и 0.99, потому что argmax в обоих случаях один и тот же. Кросс-энтропия же различает.

То есть, если модель выдаёт 0.9, то примерно в 90% таких случаев предсказание должно оказываться верным. Если верных 70%, модель переуверена, и её вероятности нельзя подставлять в бизнес-логику (пороги, ожидаемая стоимость ошибки, ранжирование по риску). Проверяется это диаграммой надёжности (reliability diagram) и метриками вроде ECE, а обрабатывается, например, температурным шкалированием: делим логиты на T и подбираем T на валидации, минимизируя ту же кросс-энтропию.

Вместо заключения

Итого, кросс-энтропия появляется в задачах классификации как один и тот же объект, возникающий из трёх идей: отрицательное логарифмическое правдоподобие в статистике, KL-дивергенция плюс константа в терминах расхождения распределений и средняя длина сообщения в терминах кодирования. Все три взгляда сходятся в одной точке: оптимум достигается тогда, когда модель сообщает истинные вероятности, а не тогда, когда она чаще угадывает класс. Мне кажется, именно это и стоит вынести из статьи, потому что отсюда естественно вырастает и калибровка, и более внимательное отношение к предсказаниям модели.

Также можете посмотреть статью в моём бложике, где можно потыкать интерактивные графики.

ссылка на оригинал статьи https://habr.com/ru/articles/1061392/