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

Flow matching

Прямые пути вместо кривых — и почему это меняет число шагов, а не качество

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

Сменить вопрос

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

Простейший разумный ответ: прямая.

xt=(1t)x0+tx1,vt=dxtdt=x1x0\htmlData{k=path}{x_t = (1-t)\,x_0 + t\,x_1}, \qquad \htmlData{k=velocity}{v_t = \frac{dx_t}{dt} = x_1 - x_0}

Здесь x0x_0 — шум, x1x_1 — данные. Обратите внимание на скорость: она не зависит от tt. Вдоль всего пути она одна и та же, и это ровно то, что означает «прямая» — постоянная скорость и нулевая кривизна.

Цель обучения

L=Et,x0,x1vθ(xt,t)(x1x0)2\mathcal{L} = \mathbb{E}_{t,\, x_0,\, x_1}\big\|v_\theta(x_t, t) - (x_1 - x_0)\big\|^2

Опять обычный MSE, и снова та же логика, что в уроке 060: сеть учится предсказывать величину, известную только на обучении (пара x0,x1x_0, x_1 известна, потому что мы её сами составили), и в оптимуме выдаёт маргинальное поле скоростей — среднее по всем парам, проходящим через данную точку.

Стоит проговорить, что здесь усредняется. Через точку xtx_t проходит много пар, и их скорости различны; сеть выдаёт их среднее. Значит маргинальное поле не прямое — прямыми были условные пути, а усреднение их искривляет. Это важная оговорка, и урок 110 покажет, что с ней делают.

Насколько прямее

Сравните два семейства путей. Смотрите на «кривизну» — отношение длины пути к прямому расстоянию — и на ошибку при малом числе шагов.

шагов сэмплера 6
длина пути
2.83
прямое расстояние
2.83
кривизна
×1
ошибка ломаной
4.4e-16
Кривизна ровно единица: путь и есть отрезок. Ошибка ломаной равна нулю при любом числе шагов — сколько бы узлов ни было, они лежат на той же прямой. Отсюда и обещание генерации за несколько шагов.

Кривизна равна единице: путь — отрезок, и ломаная из любого числа звеньев совпадает с ним точно.

Разница читается прямо: у прямого пути ошибка дискретизации тождественно нулевая, потому что любая точка ломаной лежит на самом пути. У кривого — растёт при уменьшении числа шагов.

Отдельно стоит заметить, что кривизна диффузионного пути зависит от пары концов. Для показанных трёх пар она около 1.071.111.07{-}1.11, а для пары, где шум и данные почти противоположны (x0λx1x_0 \approx -\lambda x_1), путь оказывается почти прямым — кривизна 1.00021.0002. Причина геометрическая: при таких концах обе координаты пути пропорциональны одному и тому же вектору, и кривая вырождается в отрезок. То есть «диффузионные пути кривые» — утверждение о типичной паре, а не о каждой.

Здесь нужна честная оговорка, и она существенна. Нулевая ошибка относится к условному пути — тому, который мы задали, зная обе точки. Сэмплер идёт по маргинальному полю, а оно, как сказано выше, кривое. Так что «flow matching генерирует за один шаг» — неверно; верно то, что его пути существенно прямее, и это переводится в меньшее число шагов при том же качестве.

Что общего с диффузией

Больше, чем кажется. Оба метода:

  • задают путь от шума к данным явной формулой;
  • обучают сеть предсказывать нечто, вычислимое по паре (x0,x1)(x_0, x_1), обычным MSE;
  • полагаются на то, что оптимум MSE — условное среднее;
  • на генерации решают ОДУ (или СДУ) численно.

Более того, диффузию можно записать как flow matching с другим путём: вместо (1t)x0+tx1(1-t)x_0 + t x_1 взять αˉtx1+1αˉtx0\sqrt{\bar\alpha_t}\,x_1 + \sqrt{1-\bar\alpha_t}\,x_0. Это и есть та кривая, которую рисует виджет во втором режиме.

диффузияflow matching
путьзадан физическим процессомвыбран нами
коэффициентыαˉ\sqrt{\bar\alpha}, 1αˉ\sqrt{1-\bar\alpha}1t1-t, tt
сумма квадратов коэффициентов11не единица
что предсказывает сетьε\varepsilon (или x0x_0, или score)скорость vv
кривизна условного путибольше единицыровно единица
шагов на генерациюдесятки–сотниединицы–десятки

Третья строка стоит внимания: у flow matching масштаб вдоль пути не сохраняется — при t=0.5t = 0.5 дисперсия равна 14+14=12\frac14 + \frac14 = \frac12 вместо единицы. Это не ошибка, а другое соглашение; сеть просто видит входы разного масштаба, и нормировка внутри неё это разбирает.

Rectified flow

Оговорка про кривое маргинальное поле имеет красивое решение. Обучим модель, сгенерируем ею пары (шум,результат)(\text{шум}, \text{результат}) и обучим новую модель на этих парах. Теперь пары не случайны: каждая точка шума сопоставлена с тем результатом, к которому её приводит первая модель. Пересечений путей меньше, усреднение искажает слабее, и маргинальное поле выпрямляется.

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

Итог

  • Flow matching выбирает путь, а не выводит его: линейная интерполяция даёт постоянную скорость.
  • Цель — MSE по скорости, устроенный так же, как MSE по шуму в диффузии.
  • Условные пути прямые точно; маргинальное поле — нет, потому что усреднение искривляет.
  • Диффузия — тот же метод с другим путём, и это видно по сравнению кривизны в виджете.
  • Rectified flow спрямляет маргинальное поле переобучением на собственных парах.

Источники

Проверки

0 из 2
  1. Прямые пути и их пределы

    Отметьте все верные утверждения о flow matching.

  2. Геометрия двух путей

    Реализуйте path_geometry(x0, y0, x1, y1, t). Точка (x0,y0)(x_0, y_0) — шум, (x1,y1)(x_1, y_1) — данные. Верните [flow_x, flow_y, diffusion_x, diffusion_y, curvature]:

    • flow_x, flow_y — линейный путь (1t)p0+tp1(1-t)\,p_0 + t\,p_1 в момент t;
    • diffusion_x, diffusion_y — путь диффузии αˉp1+1αˉp0\sqrt{\bar\alpha}\,p_1 + \sqrt{1-\bar\alpha}\,p_0, где αˉ\bar\alpha берётся из косинусного расписания, прочитанного назад по времени: αˉ(t)=f(1t)f(0)\bar\alpha(t) = \dfrac{f(1-t)}{f(0)}, f(u)=cos2 ⁣(u+s1+sπ2)f(u) = \cos^2\!\left(\dfrac{u+s}{1+s}\cdot\dfrac{\pi}{2}\right), s=0.008s = 0.008;
    • curvature — длина диффузионного пути, делённая на прямое расстояние между концами. Длину считайте ломаной по двум тысячам равных отрезков по tt от 00 до 11 — так посчитан эталон.

    Проверить себя можно так: линейный путь при t=0.5t = 0.5 обязан дать середину отрезка, а curvature — быть не меньше единицы при любых концах.

    функция path_geometry

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

    Ctrl/⌘ + Enter