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

Практика: DDPM с нуля

Обучить диффузию на MNIST, добавить DDIM и flow matching — и проверить каждое утверждение ветки

Шаг 101 из 117 · ~50 мин

Ветка в двух формулах

L=Eεθ(xt,t)ε2,dx=[f12g2logpt]dt\htmlData{k=train}{\mathcal{L} = \mathbb{E}\big\|\varepsilon_\theta(x_t, t) - \varepsilon\big\|^2}, \qquad \htmlData{k=sample}{dx = \Big[f - \tfrac12 g^2 \nabla\log p_t\Big]dt}

Слева — урок 060, справа — урок 030. Между ними ровно одно связующее утверждение: предсказанный шум и есть score с точностью до множителя. Всё остальное в ветке — расписания, guidance, потоки — это варианты того, как выбрать путь и как по нему пройти.

Что реализовать

1. Прямой процесс и расписания

Начните с того, что не требует обучения вовсе.

  • реализуйте alpha_bar(t) для линейного и косинусного расписаний;
  • нарисуйте xtx_t для одного изображения при t=0,0.25,0.5,0.75,1t = 0, 0.25, 0.5, 0.75, 1;
  • посчитайте, при каком tt отношение сигнал/шум равно единице.

Ожидание из урока 070: 0.260.26 для линейного и 0.500.50 для косинусного. Если получилось иначе — ошибка в накоплении произведения, и её лучше найти сейчас, чем после суток обучения.

данные при t

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

время t 0.26
0.4967
√ᾱ
0.7047
√(1−ᾱ)
0.7095
сигнал/шум
0.987
равенство при t
0.259
Точка равенства 0.26. Сверьте с этим числом свою реализацию: она чувствительна к тому, накапливаете ли вы произведение (1−β) или сумму логарифмов, и к тому, с какого индекса начинается отсчёт.

Шум преобладает. Если ваш выбор t для отладки лежит здесь, ожидать осмысленной реконструкции не стоит.

2. Обучение

Модель — маленький U-Net на MNIST. Обязательные детали:

  • tt подаётся как эмбеддинг, синусоидальный (блок 6, урок 110) или через log\log SNR;
  • на каждый пример один случайный tt и один ε\varepsilon;
  • цель — упрощённая, без весов ELBO (урок 060).

Проверка, которую стоит сделать до обучения: при случайной инициализации лосс обязан быть около единицы. Причина в том, что цель — MSE до εN(0,I)\varepsilon \sim \mathcal{N}(0, I), а необученная сеть предсказывает примерно нуль; значит ошибка равна дисперсии шума. Значение сильно другое означает ошибку в масштабировании входа.

3. Три сэмплера

Один и тот же обученный набор весов, три способа пройти обратный путь:

сэмплерформулировкашагов
DDPMстохастическая10001000
DDIMдетерминированная, каждый kk-й шаг2010020{-}100
Хойн по probability-flow ОДУдетерминированная, второй порядок205020{-}50

Из урока 030: третий имеет смысл только потому, что второй детерминированный. Проверьте это измерением — сравните качество Хойна и Эйлера при равном числе вызовов сети, а не при равном числе шагов (урок 020).

4. Пять измерений

Как и в практике блока 6, работающая генерация ничего не доказывает. Проверьте утверждения ветки на своей реализации.

Измерение 1: замкнутая форма совпадает с цепью. Прогоните цепь q(xtxt1)q(x_t\mid x_{t-1}) сто раз и сравните выборочные среднее и дисперсию с αˉx0\sqrt{\bar\alpha}x_0 и 1αˉ1-\bar\alpha. Ожидание — совпадение в пределах статистической ошибки. Это проверяет самый фундаментальный кирпич.

Измерение 2: score и ε\varepsilon — одно и то же. Возьмите предсказание сети, переведите его в score по s=ε/1αˉs = -\varepsilon/\sqrt{1-\bar\alpha}, затем восстановите x0x_0 по формуле Tweedie. Сравните с прямым переводом εx0\varepsilon \to x_0. Ожидание: побитовое совпадение с точностью арифметики.

Измерение 3: DDIM детерминирован. Дважды сгенерируйте из одного и того же стартового шума. Ожидание: одинаковые изображения. У DDPM — разные. Если DDIM даёт разные, где-то остался случайный член.

Измерение 4: инверсия работает. DDIM обратим: пройдите от шума к изображению, затем обратно к шуму, и сравните с исходным шумом. Ожидание: расхождение мало и растёт при уменьшении числа шагов — это ошибка интегрирования, а не ошибка модели.

Измерение 5: guidance уменьшает разнообразие. Сгенерируйте по пятьдесят образцов при w=1,3,8w = 1, 3, 8 и посчитайте попарное расстояние внутри каждого набора. Ожидание из урока 080: монотонное падение. Заодно посмотрите на долю пикселей, вышедших за допустимый диапазон, — она растёт.

5. Flow matching для сравнения

Та же архитектура, другая цель:

L=Evθ(xt,t)(x1x0)2,xt=(1t)x0+tx1\mathcal{L} = \mathbb{E}\big\|v_\theta(x_t, t) - (x_1 - x_0)\big\|^2, \qquad x_t = (1-t)x_0 + t x_1

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

Ожидание из урока 090, сформулированное осторожно: путей прямее, шагов меньше, но не «один шаг». Если хотите увидеть настоящее спрямление, добавьте minibatch OT (урок 110): решите задачу назначения внутри батча и составляйте пары по ней, а не случайно.

шагов сэмплера 4
длина пути
2.702
прямое расстояние
2.702
кривизна
×1
ошибка ломаной
0.0e+0
Ошибка ломаной нулевая при любом числе шагов — но помните оговорку из урока 090: это про условный путь, а сэмплер идёт по маргинальному полю.

Прямой путь: число шагов перестаёт влиять на геометрию.

Куда смотреть, если не получается

симптомчто проверять первым
лосс на первом шаге сильно не единицанормировку входа, масштаб ε\varepsilon
сэмплы — шумзнак в обратном шаге, порядок обхода tt
сэмплы — серое пятноподачу tt в сеть (модель усредняет по всем уровням шума)
DDIM даёт разные результаты из одного шумаоставшийся случайный член
сэмплы приемлемы при 1000 шагов и мусор при 50использование стохастической формулы вместо ОДУ
guidance ломает картинку при w>5w > 5выход за диапазон, нужна обрезка

Третья строка — самая коварная и стоит объяснения. Если tt до сети не доходит (забытый эмбеддинг, потерянный аргумент), обучение не падает: сеть выучивает средний по всем tt шумоподавитель. Лосс при этом заметно выше оптимального, но снижается, а генерация даёт размытое нечто. Проверяется за минуту: подайте одно и то же xtx_t при двух разных tt и убедитесь, что выход изменился.

Чем закончить

Прогоните все пять измерений на обученной модели. Особенно второе и третье: они проверяют понимание того, что предсказание шума, score и восстановление x0x_0 — один объект, и что детерминированное сэмплирование действительно детерминировано. Красивые картинки получаются и при ошибках; эти проверки — нет.

Источники

Проверки

0 из 2
  1. Отладка диффузионной модели

    Отметьте все верные утверждения о сборке и отладке диффузионной модели.

  2. Шаг DDIM

    Реализуйте ddim_step(alpha_bar_t, alpha_bar_prev, x_t, eps) — один шаг детерминированного сэмплера. Он состоит из двух действий:

    1. восстановить чистые данные по текущему состоянию и предсказанному шуму: x^0=xt1αˉtεαˉt\hat{x}_0 = \dfrac{x_t - \sqrt{1-\bar\alpha_t}\,\varepsilon}{\sqrt{\bar\alpha_t}};
    2. заново зашумить их до предыдущего уровня тем же самым ε\varepsilon: xprev=αˉprevx^0+1αˉprevεx_{\text{prev}} = \sqrt{\bar\alpha_{\text{prev}}}\,\hat{x}_0 + \sqrt{1-\bar\alpha_{\text{prev}}}\,\varepsilon.

    Верните [x0_hat, x_prev, delta, signal_gain]:

    • x0_hat и x_prev — из формул выше;
    • delta = xprevxtx_{\text{prev}} - x_t — насколько сдвинулось состояние;
    • signal_gain = αˉprevαˉt\sqrt{\bar\alpha_{\text{prev}}} - \sqrt{\bar\alpha_t} — насколько прибавилось сигнала за шаг.

    Обратите внимание: случайности здесь нет вовсе — тот же ε\varepsilon используется дважды. Именно поэтому шаг детерминирован, и именно поэтому его можно обратить.

    функция ddim_step

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

    Ctrl/⌘ + Enter