Диффузия и потоки

Classifier-free guidance

Одна линейная комбинация, которая управляет соответствием условию — и чем за неё платят

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

Задача

Условная модель p(xc)p(x\mid c) обучается так же, как безусловная, — просто с cc на входе. Проблема в том, что получается: сэмплы соответствуют условию слабее, чем хотелось бы. Модель честно воспроизводит p(xc)p(x\mid c), а хочется чего-то более выраженного.

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

Формула

ε~=εθ(xt)+w(εθ(xt,c)εθ(xt))\tilde\varepsilon = \htmlData{k=uncond}{\varepsilon_\theta(x_t)} + \htmlData{k=w}{w}\,\htmlData{k=delta}{\big(\varepsilon_\theta(x_t, c) - \varepsilon_\theta(x_t)\big)}

Обучение — один трюк: во время обучения условие с вероятностью около 10%10\% заменяется пустым. Одна и та же сеть выучивает обе функции, и на генерации её вызывают дважды. Отсюда и цена: два прогона на шаг вместо одного.

Почему разность в скобках — это направление условия? В score-виде

logp(xc)logp(x)=logp(xc)p(x)=logp(cx)\nabla\log p(x\mid c) - \nabla\log p(x) = \nabla \log \frac{p(x\mid c)}{p(x)} = \nabla\log p(c\mid x)

по правилу Байеса, поскольку logp(c)\log p(c) от xx не зависит. То есть — это градиент логарифма правдоподобия классификатора, взятый без всякого классификатора. В этом весь смысл названия.

Из какого распределения мы сэмплируем

Подставим и соберём:

s~=logp(x)+wlogp(cx)=log[p(x)p(cx)w]\tilde{s} = \nabla\log p(x) + w\,\nabla\log p(c\mid x) = \nabla\log\Big[p(x)\,p(c\mid x)^{w}\Big]

То есть guidance сэмплирует из распределения, где правдоподобие условия возведено в степень ww. При w=1w = 1 это ровно p(xc)p(x\mid c); при w>1w > 1 — что-то более сосредоточенное.

Оценивать это стоит трезво: это не более точная выборка из p(xc)p(x\mid c), а точная выборка из другого распределения. Метрики соответствия условию растут, метрики разнообразия падают, и компромисс здесь не артефакт реализации, а определение метода.

Что именно происходит с распределением

Здесь стоит поправить популярное «guidance делает распределение острее». Проверим два случая.

Случай гауссиан. Пусть p(xc)=N(μ,σ2)p(x\mid c) = \mathcal{N}(\mu, \sigma^2) и p(x)=N(0,σ2)p(x) = \mathcal{N}(0, \sigma^2). Тогда

s~=xσ2+wμσ2=xwμσ2\tilde{s} = -\frac{x}{\sigma^2} + w\frac{\mu}{\sigma^2} = -\frac{x - w\mu}{\sigma^2}

то есть направляемое распределение — это N(wμ,σ2)\mathcal{N}(w\mu, \sigma^2). Среднее уехало в ww раз дальше, дисперсия не изменилась вовсе. Никакого «острее» здесь нет.

Случай смеси. Возьмём безусловную смесь двух мод с весами 0.5/0.50.5/0.5 и условную с 0.8/0.20.8/0.2 (обе с σ=0.6\sigma = 0.6). Численно:

wwсреднеест. откл.масса на левой моде
000.0000.0001.6161.6160.5000.500
110.900-0.9001.3421.3420.7960.796
221.330-1.3300.9240.9240.9380.938
331.468-1.4680.6930.6930.9830.983
551.520-1.5200.5830.5830.9990.999
10101.535-1.5350.5660.5661.0001.000

Вот теперь разнообразие действительно падает — но конкретным образом: guidance выбирает моду, а не сжимает её. Стандартное отклонение падает с 1.621.62 до 0.570.57 и дальше не идёт: 0.570.57 — это почти ширина одной компоненты (σ=0.6\sigma = 0.6). Всё, что guidance убрал, — это разнообразие между модами.

Точная формулировка, которую стоит унести: guidance перераспределяет массу к тем областям, где условие выполняется лучше. Выглядит это как «резкость», если моды соответствуют разным вариантам ответа, и не выглядит никак, если распределение унимодально.

Покрутите γ\gamma в виджете — там та же операция, применённая к смеси трёх гауссиан.

сила guidance γ 1
шаг ε 0.05
температура T 1
макс. |score|
17.82
длина пути
98.2
доля времени у мод
25.3%
Здесь γ умножает весь score, то есть сэмплирует из p^γ — упрощённая версия guidance, где условие совпадает с самим распределением. Следите за долей времени у мод: при γ = 1 блуждание навещает все три, при больших γ оседает в одной.

Практика

wwчто получается
00безусловная генерация
11честная условная модель, соответствие слабое
383{-}8рабочий диапазон для изображений
>15> 15пересыщенные цвета, артефакты, потеря разнообразия

Верхняя строка объясняет типичный артефакт. Сдвиг среднего (гауссов случай выше) в ww раз означает выход за пределы области, где модель обучалась: при больших ww значения пикселей вылезают за допустимый диапазон, и результат приходится обрезать. Отсюда и приёмы вроде динамического ограничения, и вариант с изменением ww по ходу сэмплирования — большой на ранних шагах, где решается композиция, и малый на поздних, где решаются детали.

Итог

  • Guidance — одна линейная комбинация двух предсказаний, обучаемая случайным обнулением условия.
  • Разность предсказаний равна logp(cx)\nabla\log p(c\mid x) — классификатор без классификатора.
  • Сэмплирование идёт из p(x)p(cx)wp(x)p(c\mid x)^w, то есть из другого распределения, а не более точно из условного.
  • Гауссианы равной ширины при guidance сдвигаются, но не сужаются; сужение возникает из выбора моды в мультимодальном распределении.
  • Цена: два прогона сети на шаг и потеря разнообразия, растущая с ww.

Источники

Проверки

0 из 2
  1. Что делает guidance

    Отметьте все верные утверждения о classifier-free guidance.

  2. Направляемое предсказание

    Реализуйте guided(eps_uncond, eps_cond, w, alpha_bar) — верните [eps_guided, score_guided, overshoot, amplification]:

    • eps_guided = εu+w(εcεu)\varepsilon_u + w(\varepsilon_c - \varepsilon_u);
    • score_guided = eps_guided1αˉ-\dfrac{\texttt{eps\_guided}}{\sqrt{1-\bar\alpha}} — то же предсказание в виде score (урок 060);
    • overshoot = eps_guidedεc\texttt{eps\_guided} - \varepsilon_c — насколько направляемое предсказание ушло дальше честного условного;
    • amplification = eps_guidedεuεcεu\dfrac{|\texttt{eps\_guided} - \varepsilon_u|}{|\varepsilon_c - \varepsilon_u|} — во сколько раз усилено направление условия. Если знаменатель равен нулю, верните ноль.

    Проверить себя можно так: при w=1w = 1 второй аргумент возвращается без изменений, а overshoot обязан быть нулевым.

    функция guided

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

    Ctrl/⌘ + Enter