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

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

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

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

Начнём с классического вывода: через метод максимального правдоподобия (далее будем называть его 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}}=argmax_{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 для истинного класса и Несколько взглядов на кросс-энтропию - 21 иначе. Тогда тот же трюк со степенями даёт

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_2tfrac{1/2}{1/4} + tfrac{1}{4}log_2 1 + 2 cdot tfrac{1}{8}log_2tfrac{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-дивергенция плюс константа в терминах расхождения распределений и средняя длина сообщения в терминах кодирования. Все три взгляда сходятся в одной точке: оптимум достигается тогда, когда модель сообщает истинные вероятности, а не тогда, когда она чаще угадывает класс. Мне кажется, именно это и стоит вынести из статьи, потому что отсюда естественно вырастает и калибровка, и более внимательное отношение к предсказаниям модели.

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

Автор: MMKuz

Источник