Теория информации

Прямая и обратная KL

Одна формула, два направления — и два совершенно разных поведения при обучении

Шаг 57 из 117 · ~28 мин

Два направления

Пусть pp — то, что есть (данные, истинный posterior), а qq — то, чем мы приближаем, и qq проще: например, одна гауссиана. Тогда «подогнать qq под pp» можно двумя способами, и они дают разные ответы.

KL(pq)=Ep ⁣[logpq]KL(qp)=Eq ⁣[logqp]\htmlData{k=fwd}{\text{KL}(p \,\|\, q) = \mathbb{E}_{p}\!\left[\log \frac{p}{q}\right]} \qquad\qquad \htmlData{k=rev}{\text{KL}(q \,\|\, p) = \mathbb{E}_{q}\!\left[\log \frac{q}{p}\right]}

Разница целиком в том, по какому распределению берётся ожидание, и из этого следует всё остальное.

Логика штрафов

Посмотрите, где каждая дивергенция становится большой.

, KL(pq)\text{KL}(p\|q): ожидание берётся по pp, поэтому вклад есть только там, где p>0p > 0. Если в такой точке q0q \approx 0, логарифм logpq\log \frac{p}{q} взрывается. Штраф за то, что qq пропустила массу pp. Обратное не наказывается: где p=0p = 0, слагаемое равно нулю, и что там делает qq, дивергенции безразлично.

, KL(qp)\text{KL}(q\|p): ожидание по qq, вклад только там, где q>0q > 0. Если там p0p \approx 0, взрыв. Штраф за массу qq в местах, где pp её не предусмотрел. А пропуск моды pp бесплатен — туда qq просто не заглянет.

Отсюда названия:

прямая KL(pq)\text{KL}(p\|q)обратная KL(qp)\text{KL}(q\|p)
поведениеmass-covering, «покрывающая»mode-seeking, «выбирающая моду»
боитсяпропустить массу ppпопасть туда, где pp пусто
результат на бимодальном ppразмазывается между модамисадится на одну моду
требуетсэмплы из pp (данные)вычислять pp в точке

Последняя строка объясняет, кто где используется, и это не вопрос вкуса.

Числа

Возьмём p=12N(2,0.52)+12N(2,0.52)p = \frac12\mathcal{N}(-2, 0.5^2) + \frac12\mathcal{N}(2, 0.5^2) — две узкие моды — и будем приближать её одной гауссианой q=N(μ,σ2)q = \mathcal{N}(\mu, \sigma^2). Оптимумы найдены сеточным поиском:

что минимизируемμ\muσ\sigmaзначение
прямую KL(pq)\text{KL}(p\|q)0.000.002.052.050.7240.724 ната
обратную KL(qp)\text{KL}(q\|p)2.002.000.490.490.6930.693 ната

Прямая KL встала посередине между модами, накрыв обе — в точке, где у pp плотность минимальна. Обратная выбрала правую моду и воспроизвела её ширину почти точно (0.490.49 против истинных 0.50.5).

Два перекрёстных числа показывают, насколько это не мелочь:

  • прямая KL на решении обратной (μ=2\mu = 2, σ=0.5\sigma = 0.5) равна 15.3115.31 — против 0.7240.724 в своём оптимуме. Пропуск моды для неё катастрофа;
  • обратная KL на решении прямой (μ=0\mu = 0, σ=2.06\sigma = 2.06) равна 2.092.09 — против 0.6930.693. Хуже втрое, но не катастрофа.

И красивая деталь: оптимум обратной KL равен 0.693ln20.693 \approx \ln 2. Не совпадение — qq покрыла одну моду из двух, то есть ровно половину массы pp, а цена «не знать про половину» и есть ln2\ln 2.

Два независимых подтверждения

Оптимум прямой KL в гауссовском семействе можно получить без всякого поиска. Раскроем:

KL(pq)=H(p)constEp[logq]\text{KL}(p\|q) = \underbrace{-H(p)}_{\text{const}} - \mathbb{E}_p[\log q]

Остаётся максимизировать Ep[logq]\mathbb{E}_p[\log q] — то есть сделать MLE гауссианы по «выборке», распределённой как pp. Ответ известен из блока 3: совпадение первых двух моментов.

μ=Ep[X]=0,σ2=Varp(X)=4+0.25=4.25\mu^* = \mathbb{E}_p[X] = 0, \qquad \sigma^{*2} = \operatorname{Var}_p(X) = 4 + 0.25 = 4.25

то есть σ=2.0616\sigma^* = 2.0616 — а сетка дала 2.052.05 при шаге 0.0260.026. Совпадает.

Минимизация прямой KL в экспоненциальном семействе — это в точности сопоставление моментов. Тот же факт, что делал softmax параметризацией категориального в блоке 3, здесь объясняет, почему прямая KL «размазывает»: она согласует средние, а не форму.

Кто что использует

Различие полностью практическое.

Обратная KL — вариационный вывод и VAE. В KL(qp)\text{KL}(q\|p) ожидание берётся по qq, из которого мы умеем сэмплировать по построению. Прямую KL здесь считать нельзя: сэмплов из истинного posterior у нас нет — за ними мы и пришли. Отсюда известное свойство: VAE недооценивает дисперсию posterior и склонен к смазанным реконструкциям. Это не баг реализации, а поведение выбранного направления.

Прямая KL — обучение с учителем. Кросс-энтропийный лосс — это KL(p^данныеpмодель)\text{KL}(\hat p_{\text{данные}} \| p_{\text{модель}}). Ожидание по данным, которых у нас как раз много. Соответственно модель наказывается за нулевую вероятность на наблюдённом примере — и это именно то, чего мы хотим от классификатора.

Обе — RLHF и PPO. Там KL-штраф удерживает политику около исходной модели, и направление выбирают из тех же соображений: считать надо то, из чего умеешь сэмплировать.

методнаправлениепочему так
обучение с учителемпрямаяесть сэмплы из данных
VAE, вариационный выводобратнаяесть сэмплы только из qq
дистилляцияобычно прямаяучитель даёт всё распределение
GANни та, ни другая (JS, Вассерштейн)носители не пересекаются

Между двумя крайностями есть непрерывное семейство — α\alpha-дивергенции, где прямая и обратная KL суть предельные случаи α1\alpha \to 1 и α0\alpha \to 0. На практике почти всегда берут край, но полезно знать, что выбор не бинарный.

Источники

  • Bishop — Pattern Recognition and Machine Learning, гл. 10.1 — Прямая и обратная KL, вариационное приближение
  • Minka — Divergence Measures and Message Passing — Семейство альфа-дивергенций и поведение крайних случаев

Проверки

0 из 2
  1. Какое направление где

    Отметьте все верные утверждения о прямой KL(pq)\text{KL}(p\|q) и обратной KL(qp)\text{KL}(q\|p), где pp — цель, а qq — приближение.

  2. Какого кандидата выберет каждое направление

    Дано целевое распределение p_weights и список кандидатов candidates — все ненормированные. Нормируйте всё и верните [forward_index, reverse_index, forward_value, reverse_value]:

    • forward_index — индекс кандидата, минимизирующего KL(pq)\text{KL}(p \| q);
    • reverse_index — индекс кандидата, минимизирующего KL(qp)\text{KL}(q \| p);
    • forward_value, reverse_value — сами минимумы, в битах.

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

    Обратите внимание, что один и тот же кандидат почти никогда не выигрывает в обоих направлениях — в этом и смысл задачи.

    функция pick_direction

    Загрузка редактора…

    Ctrl/⌘ + Enter