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

Марковские цепи и прямой процесс

Зашумление как цепь — и почему тысяча шагов складывается в одну формулу

Шаг 90 из 117 · ~26 мин

Что мы собираемся построить

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

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

q(xtxt1)=N(xt; 1βtxt1, βtI)\htmlData{k=markov}{q(x_t \mid x_{t-1})} = \mathcal{N}\big(x_t;\ \sqrt{1-\htmlData{k=beta}{\beta_t}}\, x_{t-1},\ \htmlData{k=beta}{\beta_t} I\big)

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

Обратите внимание на выбор коэффициентов: множитель у сигнала — 1βt\sqrt{1-\beta_t}, у шума — βt\sqrt{\beta_t}, и сумма их квадратов равна единице. Значит если Var(xt1)=1\operatorname{Var}(x_{t-1}) = 1, то и Var(xt)=1\operatorname{Var}(x_t) = 1. Такой процесс называют сохраняющим дисперсию, и это не косметика: без нормировки масштаб данных уезжал бы вместе с шумом, а сеть должна работать в одном диапазоне на всех шагах.

Тысяча шагов в одну формулу

Наивно, чтобы получить x100x_{100}, нужно сделать сто шагов. Но композиция гауссиан — снова гауссиана, и произведение сворачивается.

Обозначим αt=1βt\alpha_t = 1 - \beta_t и αˉt=stαs\bar\alpha_t = \prod_{s\le t}\alpha_s. Тогда

q(xtx0)=N(xt; αˉtx0, (1αˉt)I),xt=αˉtx0+1αˉtεq(x_t \mid x_0) = \mathcal{N}\big(x_t;\ \sqrt{\htmlData{k=abar}{\bar\alpha_t}}\, x_0,\ (1 - \htmlData{k=abar}{\bar\alpha_t}) I\big), \qquad x_t = \sqrt{\bar\alpha_t}\,x_0 + \sqrt{1-\bar\alpha_t}\,\varepsilon

Проверено численно: пять шагов с β=[0.1,0.15,0.2,0.05,0.3]\beta = [0.1, 0.15, 0.2, 0.05, 0.3] из x0=2x_0 = 2 по четырёхсот тысячам траекторий дают среднее 1.27451.2745 и дисперсию 0.59190.5919; замкнутая форма предсказывает 1.27591.2759 и 0.59300.5930.

Это самое важное практическое следствие во всей теме. Обучение не требует прогона цепи: чтобы получить обучающий пример на шаге tt, достаточно взять x0x_0, взять один ε\varepsilon и посчитать одно выражение. Без этого диффузия была бы неприменима — тысяча последовательных шагов на каждый пример в батче.

Заметьте, что αˉt\bar\alpha_t полностью описывает состояние процесса. Ни tt, ни отдельные β\beta дальше не понадобятся — только одно число.

Что видит модель

Потяните время и посмотрите на облако и на отношение сигнал/шум.

данные при t

сигнал/шум, log-шкала

время t 0.3
0.3956
√ᾱ
0.629
√(1−ᾱ)
0.7774
сигнал/шум
0.655
равенство при t
0.259
Исходное расписание DDPM. Обратите внимание, где сигнал и шум сравниваются: при t ≈ 0.26. То есть три четверти шагов процесс работает с почти чистым шумом — урок 070 разберёт, почему это плохо.

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

Два наблюдения, которые пригодятся дальше:

  1. Облако сжимается к нулю и расширяется до единичной дисперсии. Множитель αˉ\sqrt{\bar\alpha} стягивает данные к началу координат, 1αˉ\sqrt{1-\bar\alpha} добавляет разброс. При αˉ0\bar\alpha \to 0 остаётся N(0,I)\mathcal{N}(0, I) независимо от того, чем были данные.
  2. Время входит только через αˉ\bar\alpha. Поэтому «шаг 700 из 1000» — не информация; информация — это αˉ700\bar\alpha_{700}, и именно её (или эквивалентное ей число) подают в сеть.

Обратный процесс: почему он вообще существует

Прямой процесс задан. Нужен обратный: q(xt1xt)q(x_{t-1} \mid x_t). Он существует по правилу Байеса, но не выражается в замкнутой форме — для этого нужна маргинальная q(xt)q(x_t), то есть распределение всех данных, которого мы не знаем.

Здесь и появляется ключевой факт, из которого работает весь метод:

При малом βt\beta_t обратный шаг q(xt1xt)q(x_{t-1}\mid x_t) приближённо гауссов.

Это не тавтология. Прямой шаг гауссов по построению; обратный — вообще говоря нет. Но при малом шаге он становится гауссовым с точностью до O(βt2)O(\beta_t^2), а значит его можно параметризовать сетью, предсказывающей всего два числа на компоненту: среднее и дисперсию.

Отсюда сразу видна цена и причина существования тысячи шагов. Приближение тем точнее, чем меньше β\beta; меньше β\beta — больше шагов, чтобы дойти до шума. Тысяча шагов в DDPM — это не запас на всякий случай, а плата за гауссовость обратного ядра.

Обратное ядро при известном x0x_0

Одна вещь всё-таки считается точно, и на ней стоит обучение. Если знать x0x_0, то обратный шаг гауссов не приближённо, а ровно:

q(xt1xt,x0)=N(xt1; μ~t(xt,x0), β~tI)q(x_{t-1} \mid x_t, x_0) = \mathcal{N}\big(x_{t-1};\ \tilde\mu_t(x_t, x_0),\ \tilde\beta_t I\big)

μ~t=αt(1αˉt1)1αˉtxt+αˉt1βt1αˉtx0,β~t=1αˉt11αˉtβt\tilde\mu_t = \frac{\sqrt{\alpha_t}(1-\bar\alpha_{t-1})}{1-\bar\alpha_t}x_t + \frac{\sqrt{\bar\alpha_{t-1}}\beta_t}{1-\bar\alpha_t}x_0, \qquad \tilde\beta_t = \frac{1-\bar\alpha_{t-1}}{1-\bar\alpha_t}\beta_t

Формулы выглядят громоздко, а получаются механически: правило Байеса для гауссиан плюс приведение показателя к квадрату (блок 3, урок 090). Разбирать их наизусть не нужно; важно структурное наблюдение — среднее есть выпуклая комбинация xtx_t и x0x_0, то есть обратный шаг тянет текущее состояние в сторону чистых данных, и вес этой тяги задан расписанием.

Именно поэтому в обучении фигурирует x0x_0 (или, что то же самое, ε\varepsilon): цель обучения — приблизить это точное ядро, а не угадывать неизвестную маргинальную плотность. Урок 060 покажет, что предсказывать ε\varepsilon, предсказывать x0x_0 и предсказывать score — три записи одной задачи.

Итог

  • Прямой процесс задан руками, марковский, сохраняет дисперсию.
  • Он сворачивается в одну формулу: xt=αˉtx0+1αˉtεx_t = \sqrt{\bar\alpha_t}x_0 + \sqrt{1-\bar\alpha_t}\varepsilon, и без этого обучение было бы непрактично.
  • Всё состояние процесса — одно число αˉt\bar\alpha_t.
  • Обратный шаг гауссов приближённо, с точностью O(β2)O(\beta^2), — отсюда требование малых шагов и их количество.
  • Обратный шаг при известном x0x_0 гауссов точно, и на нём стоит обучение.

Источники

Проверки

0 из 2
  1. Прямой процесс и его свойства

    Отметьте все верные утверждения о прямом процессе диффузии.

  2. Замкнутая форма прямого процесса

    Реализуйте forward_state(betas, steps, x0, eps) — верните [alpha_bar, signal_scale, noise_scale, x_t, snr]:

    • alpha_bar = αˉ=s<steps(1βs)\bar\alpha = \prod_{s < \texttt{steps}} (1 - \beta_s) — произведение по первым steps элементам списка (при steps = 0 это пустое произведение, то есть единица);
    • signal_scale = αˉ\sqrt{\bar\alpha};
    • noise_scale = 1αˉ\sqrt{1 - \bar\alpha};
    • x_t = αˉx0+1αˉε\sqrt{\bar\alpha}\,x_0 + \sqrt{1-\bar\alpha}\,\varepsilon — состояние после steps шагов при данном единственном ε\varepsilon;
    • snr = αˉ1αˉ\dfrac{\bar\alpha}{1 - \bar\alpha} — отношение сигнала к шуму.

    Проверить себя можно так: сумма квадратов signal_scale и noise_scale обязана равняться единице при любых β\beta. Это и есть то, что называют сохранением дисперсии.

    функция forward_state

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

    Ctrl/⌘ + Enter