Назад
194

Flow Matching: обучение и дистилляция

194

📌 Время прочтения статьи — 30-60 минут. Предполагается, что читатель знаком с базовыми понятиями математического анализа и теории вероятностей.

Введение

Рисунок 1. Примеры возможностей генеративных моделей: из шума появляются изображения, видео, звук и другие сложные объекты

Flow Matching — современный подход генеративного моделирования. Основная идея — научить модель постепенно превращать простое распределение, например, случайный шум, в сложные данные — изображения, видео, звук, 3D-сцены или действия робота.

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

Один из способов его решения — дистилляция. Мы заменим многошаговый процесс генератором, который за одно вычисление нейронной сети произведёт нам нужные картинки.

Однако такую модель нельзя просто обучить регрессией по принципу: «Вот — вход, вот — правильный ответ». Мы разберём, почему это так и как из задачи дистилляции естественным образом возникает минимаксная постановка: две модели начинают играть друг с другом, а результатом игры становится качественная генеративная модель!

Сферы применения

Flow Matching сегодня — один из ключевых подходов в современных генеративных моделях, «равный» диффузионным моделям (в каком-то смысле считается, что это одно и то же).

Известные примеры применения — Stable Diffusion 3.5, FLUX.2 и Kandinsky от Сбера. Эти модели способны превращать текстовое описание в новое изображение — от реалистичной фотографии до сложной фантастической сцены.

Рисунок 2. Пример генерации изображений в Stable Diffusion 3.5

Но изображениями применение Flow Matching’а не ограничивается. С его помощью можно генерировать и более сложные объекты — например, видео вместе со звуком. MiniMax H3 получает текстовое описание сцены и постепенно превращает случайный шум не просто в одну картинку, а в целую последовательность кадров и соответствующий ей звук.

Рисунок 3. Пример генерации видео со звуком в MiniMax H3

Тот же принцип работает и с музыкой. Например, Music 3 использует Flow Matching для генерации музыкальных фрагментов по текстовому описанию. Можно за несколько минут сделать композицию и выложить её на музыкальную платформу.

Рисунок 4. Пример генерации аудио в MiniMax Music 3

Но генеративные модели идут ещё дальше: сегодня они умеют создавать уже не отдельные изображения или видео, а целые интерактивные миры, с которыми можно взаимодействовать в реальном времени. Например, семейство Genie от Google DeepMind позволяет сгенерировать виртуальное окружение из изображения или текстового описания, а затем исследовать его, управляя персонажем. Получается, можно придумать игру и почти сразу оказаться внутри неё!

Рисунок 5. Пример интерактивного мира, сгенерированного моделью Genie

Определение

Flow Matching — один из современных подходов генерации изображений, видео, звука. Его основная идея очень похожа на интуитивное «расшумление»: мы начинаем со случайного шума и постепенно превращаем его в осмысленный объект — например, в чёткое изображение.

Рисунок 6. Интуиция «расшумления»: шум постепенно превращается в осмысленное изображение

Формально у нас есть настоящие данные \(x_\text{data} \sim p_{\text{data}}\). Это могут быть, например, реальные изображения из интернета. Мы хотим научиться порождать новые объекты, которые похожи на данные из этого распределения, но не являются простым копированием обучающих примеров. То есть мы хотим выучить распределение картинок по набору примеров из него.

С другой стороны, нам нужно распределение, из которого легко начинать генерацию. Обычно для этого берут стандартное нормальное распределение: \(x_\text{noise} \sim \mathcal{N}(0, I).\) Из него легко семплировать случайные точки, поэтому оно обычно выступает начальным шумом.

Как происходит переход от шума к данным с помощью непрерывного времени \(t \in [0,1]\)? В момент \(t=0\) мы находимся в шуме: \(x_0 = x_\text{noise} \sim \mathcal{N}(0, I),\) а в момент \(t=1\) хотим оказаться в данных: \(x_1 = x_\text{data} \sim p_{\text{data}}.\)

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

Рисунок 7. Схема движения семплов из шумового распределения \(p_0\) к семплам из распределения данных \(p_1\); \(x_t\) — промежуточные состояния

На рисунке выше мы работаем в одномерном пространстве (по вертикали — координата, по горизонтали — время), либо в двумерном пространстве (шум — слева, картинки — справа, и мы движемся в этом двумерном пространстве во времени).

Но для двумерного пространства траектории могут быть и другие, например, петли:

Рисунок 8. Схема движения семплов из шумового распределения \(p_0\) к семплам из распределения данных \(p_1\) с образованием петель

Самый простой способ соединить шум \(x_0\) и объект \(x_1\) во времени — провести между ними прямую линию:

\[x_t = (1-t)x_0 + t x_1, \qquad t\in[0,1].\]

При \(t=0\) эта формула даёт \(x_t=x_0\), то есть шум. При \(t=1\) она определяет \(x_t=x_1\), то есть объект из данных. А при промежуточных значениях времени \(t\) точка находится где-то между шумом и изображением.

Скорость движения вдоль такой прямой траектории постоянна и равна:

\[\frac{d x_t}{dt} = x_1 - x_0.\]

То есть здесь направление движения (вектор скорости) — вектор от начального шума к конечному объекту.

Рисунок 9. Простой случай: каждая точка идёт к выбранному объекту по прямой траектории, а скорость вдоль неё равна \(x_1-x_0\)

Траектории точек (или частиц) могут быть совершенно разными. Вместо предварительного знания конечной точки \(x_1\) или точки на траектории (что кажется очень сложным, потому что нужно сопоставлять точки из шума с реальными картинками) попробуем описать само движение. Для этого удобно использовать векторное поле.

Векторное поле \(u(x,t)\) говорит нам, куда и с какой скоростью нужно двигаться из точки \(x\) в момент времени \(t\). Например, если сейчас мы находимся в точке \(x_t\), то поле \(u(x_t,t)\) показывает направление следующего шага.

Рисунок 10. Векторное поле — набор маленьких стрелок: в каждой точке оно говорит, куда сделать следующий шаг

На этой и следующей визуализациях векторы короче, чем они есть, потому что удобно брать далёкие \(x_0\) и \(x_1\). В таком случае \(x_1 - x_0\) будет большим и сложным для отображения. Тогда мы будем «домножать» векторы на некоторое \(d \tau\).

Если сделать небольшой шаг по времени \(dt\) — новое положение точки можно записать так:

\[x_{t+dt} = x_t + dt \cdot u(x_t, t).\]

Следовательно, мы берём текущую точку \(x_t\), смотрим на скорость, которую предсказывает поле, и немного передвигаемся в этом направлении.

Рисунок 11. Семплирование \(x_1\) от \(x_0\) в направлении по \(u(x_t, t)\)

Итак, задачу генерации можно сформулировать так: необходимо обучить нейросеть \(u_\theta(x,t)\), которая приближает правильное векторное поле. После этого мы можем взять случайный шум \(x_0 \sim \mathcal{N}(0,I)\) и постепенно двигать его по этому векторному полю от \(t=0\) до \(t=1\). Если поле обучено хорошо, то в конце этого движения мы получим реалистичное изображение.

В математике для формального обозначения движения по вектору скорости используют обыкновенные дифференциальные уравнения (ordinary differential equations – ODE). Мы говорим, что хотим найти \(x_1\), зная \(x_0 \sim \mathcal{N}(0,I)\) и закон \(\frac{d x_t}{d t} = u(x_t, t)\). Имеем начальную точку и закон скорости, значит, можем получить конечную точку, просто проинтегрировав \(x_s = x_0 + \int_{0}^s u(x_t, t) dt\) и \(x_1 = x_0 + \int_{0}^1 u(x_t, t) dt\), где \(x_0\) известен. Это мы и аппроксимируем через маленькие шаги \(dt\).

Если мы для генерации траектории каждый раз будем вычислять \(u(x_t, t)\) — мы не получим траектории, которые в определённый момент времени пересекутся и отойдут друг от друга. Дело в том, что при пересечении мы получим одинаковый \(x_t\), и далее обе траектории пойдут по \(u(x_t, t)\), то есть уже не разъединятся.

Таким образом, мы избавились от необходимости описывать, какие начальные и конечные точки принадлежат траекториям. Достаточно знать скорость в каждой точке.

Обучение

Вернёмся к простой конструкции с прямыми траекториями. Возьмём пару объектов \(x_0, x_1\), где \(x_0 \sim \mathcal{N}(0,I)\) — случайный шум, а \(x_1 \sim p_{\text{data}}\) — реальное изображение. Между ними можно провести прямую траекторию:

\[x_t = (1-t)x_0 + t x_1.\]

Мы можем соединить все \(x_0 \sim p(x_0)\) со всеми \(x_1 \sim p(x_1)\) прямыми траекториями (каждый с каждым). В таком случае для текущего \(x_t\) бесконечное число пар \(x_0, x_1\) задаёт прямую траекторию, проходящую через \(x_t\). То есть промежуточная точка появляется из разных пар: шумы \(x_0\) могли двигаться к различным изображениям \(x_1\) и в какой-то момент оказаться примерно в одном и том же месте, имея скорости \(x_1 - x_0\).

Итоговую скорость \(u(x_t, t)\) можно получить, если усреднить все скорости, которые проходят через точку \(x_t\) в момент времени \(t\). То есть мы можем выразить глобальную скорость \(u(x_t, t)\) через среднее локальных скоростей \(x_1 - x_0\) по всем парам \(x_0, x_1\).

Итак, Flow Matching показывает, что нужно предсказывать среднюю скорость всех частиц, которые проходят через эту точку в данный момент времени, чтобы получить скорость для искомых траекторий. Формально это записывается через условное матожидание:

\[u^*(x,t) = \mathbb{E}_{x_0,x_1} \left[ x_1 - x_0 \mid x_t = x \right].\]

Условное математическое ожидание

Математическое ожидание — обобщение понятия среднего элементов. Но элементов может быть бесконечно много. Как и в нашем случае 🙂

Условное математическое ожидание — среднее элементов, полученных при каких-то условиях. К примеру, если мы хотим найти среднюю продолжительность сна перед сессией, то сессия — это условие.

Формула выше означает следующее: мы смотрим на все пары \(x_0,x_1\), для которых промежуточная точка \(x_t\) оказалась равна \(x\), берём их скорости \(x_1-x_0\) и усредняем. Получившееся среднее направление и есть правильное векторное поле в точке \(x\) и времени \(t\). Его мы и будем использовать для семплирования.

Здесь время \(t\) фиксировано: мы усредняем по парам \((x_0,x_1)\), которые в этот момент проходят через точку \(x\), а не по разным моментам времени.

Рисунок 12. Прохождение траекторий от разных пар \((x_0,x_1)\) в одной точке; поле \(u^*(x,t)\) берёт среднее направление их скоростей, а мы в итоге получаем искомую траекторию

Прямые траектории при соединении \(x_0\) с \(x_1\) и конечные траектории при семплировании — не одно и то же. Просто нам достаточно учить не все скорости, проходящие через текущую точку, а их среднее.

Почему можно брать именно среднее? Интуитивно, нас интересует не судьба одной конкретной частицы, а то, как меняется всё распределение точек. Если через одну область пространства проходит много частиц с разными скоростями, то общее движение плотности определяется их средним потоком. Поэтому для генерации достаточно выучить не каждую отдельную траекторию, а усреднённое векторное поле \(u^*(x,t)\).

На практике мы обучаем нейросеть \(u_\theta(x,t)\) приближать поле. Для этого семплируем пару \(x_0,x_1\), выбираем случайное время \(t\), строим промежуточную точку \(x_t = (1-t)x_0 + t x_1\) и просим модель по \(x_t,t\) предсказать скорость \(x_1 - x_0.\)

Получается обычная задача регрессии:

\[\mathcal{L}(\theta) = \mathbb{E}_{x_0,x_1,t} \left[ \left\| u_\theta(x_t,t) - (x_1-x_0) \right\|^2 \right].\]

Минимум этой MSE-задачи при каждом фиксированном времени \(t\) как раз и равен условному среднему:

\[u_\theta(x,t) \approx u^*(x,t) = \mathbb{E}_{x_0,x_1} \left[ x_1-x_0 \mid x_t=x \right].\]

Условное среднее

Покажем, что минимум MSE-задачи равен условному среднему.

Распишем наш лосс как математическое ожидание от условного математического ожидания (вынесем \(x_t\)):

\[\begin{aligned}\mathcal{L}(\theta)&=\mathbb{E}_{x_0,x_1,t}\left[\left\|u_\theta(x_t,t)-(x_1-x_0)\right\|^2\right] \\&=\mathbb{E}_{t,x_t}\left[\mathbb{E}_{x_0,x_1 \mid x_t,t}\left[\left\|u_\theta(x_t,t)-(x_1-x_0)\right\|^2\right]\right].\end{aligned}\]

Отдельно обозначим \(\mu(x_t, t)\) как условное среднее \(\mu(x_t,t)=\mathbb{E}_{x_0,x_1 \mid x_t,t}\left[x_1-x_0\right]\). Мы хотим показать, что \(u_\theta(x_t, t)\) как раз выучит \(\mu(x_t, t)\).

Распишем отдельно то, что стоит под математическим ожиданием:

\[\begin{aligned} & \mathbb{E}_{x_0,x_1 \mid x_t,t} \left[ \left\| u_\theta(x_t,t)-(x_1-x_0) \right\|^2 \right] \\ &= \mathbb{E}_{x_0,x_1 \mid x_t,t} \left[ \left\| u_\theta(x_t,t)-\mu(x_t,t) + \mu(x_t,t)-(x_1-x_0) \right\|^2 \right] \\ &= \mathbb{E}_{x_0,x_1 \mid x_t,t} \left[ \left\| u_\theta(x_t,t)-\mu(x_t,t) \right\|^2 \right] \\ &\quad + \mathbb{E}_{x_0,x_1 \mid x_t,t} \left[ \left\| \mu(x_t,t)-(x_1-x_0) \right\|^2 \right] \\ &\quad + 2 \mathbb{E}_{x_0,x_1 \mid x_t,t} \left[ \left\langle u_\theta(x_t,t)-\mu(x_t,t), \mu(x_t,t)-(x_1-x_0) \right\rangle \right] \\ &= \left\| u_\theta(x_t,t)-\mu(x_t,t) \right\|^2 + \mathbb{E}_{x_0,x_1 \mid x_t,t} \left[ \left\| \mu(x_t,t)-(x_1-x_0) \right\|^2 \right] \\ &\quad + 2 \left\langle u_\theta(x_t,t)-\mu(x_t,t), \mathbb{E}_{x_0,x_1 \mid x_t,t} \left[ \mu(x_t,t)-(x_1-x_0) \right] \right\rangle \\ &= \left\| u_\theta(x_t,t)-\mu(x_t,t) \right\|^2 + \mathbb{E}_{x_0,x_1 \mid x_t,t} \left[ \left\| \mu(x_t,t)-(x_1-x_0) \right\|^2 \right] \\ &\quad + 2 \left\langle u_\theta(x_t,t)-\mu(x_t,t), \mu(x_t,t) - \mathbb{E}_{x_0,x_1 \mid x_t,t} \left[x_1-x_0 \right] \right\rangle \\ &= \left\| u_\theta(x_t,t)-\mu(x_t,t) \right\|^2 + \mathbb{E}_{x_0,x_1 \mid x_t,t} \left[ \left\| \mu(x_t,t)-(x_1-x_0) \right\|^2 \right] \\ &\quad + 2 \left\langle u_\theta(x_t,t)-\mu(x_t,t), \mu(x_t,t) - \mu(x_t, t) \right\rangle \\ &= \left\| u_\theta(x_t,t)-\mu(x_t,t) \right\|^2 + \mathbb{E}_{x_0,x_1 \mid x_t,t} \left[ \left\| \mu(x_t,t)-(x_1-x_0) \right\|^2 \right] \\ &= \left\| u_\theta(x_t,t) - \mu(x_t,t) \right\|^2 + \operatorname{tr} \operatorname{Cov}_{x_0,x_1 \mid x_t,t}(x_1-x_0). \end{aligned}\]

В итоге получим, что оптимум у \(u_\theta(x_t, t)\) равен \(\mu(x_t, t)\), так как второе слагаемое не зависит от \(u_\theta\). Что мы и хотели доказать.

Именно поэтому Flow Matching можно обучать просто: мы сами строим промежуточные точки \(x_t\) и знаем целевую скорость \(x_1-x_0\) для каждой обучающей пары.

Рисунок 13. Переход от траекторий к полю: в промежуточной точке много семпловых скоростей, усредняем их и получаем локальную стрелку векторного поля

В следующих трёх скрытых блоках можно посмотреть примеры траекторий Flow Matching’а на задаче перевода одного нормального распределения в два, а также подробнее почитать про интуицию метода.

Перевод одной гауссианы в две

Рассмотрим пример с отображением одного нормального распределения в смесь двух нормальных распределений. Для начала визуализируем семплы из \(x_0 \sim p_0\) и \(x_1 \sim p_1\), а также промежуточные семплы \(x_t \sim p_t\), которые были получены как \(x_t = (1 - t)x_0 + t x_1\)

Рисунок 14. Игрушечный пример с настоящими гауссианами: семплы из одного начального гауссиана переходят к смеси из двух гауссианов

Для каждой промежуточной точки мы можем найти \(x_0\) и \(x_1\), которые её породили, и соответствующую скорость \(x_1 - x_0\) (на рисунке для удобства восприятия мы ужали векторы по длине).

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

Итоговая скорость в момент времени \(t\) для точки \(x\) — среднее среди скоростей \(x_1 - x_0\), для которых \(x_t = x\) (для удобства восприятия векторы мы ужали по длине).

Рисунок 16. После усреднения семпловых скоростей получается итоговое поле: верхняя часть потока идёт к верхней моде, нижняя — к нижней
Появление средней скорости

Представим, что в момент времени \(t\) мы смотрим на некоторую точку \(x_t\). Она могла появиться из разных пар \((x_0, x_1)\). Для каждой такой пары есть своя скорость \(v = x_1 - x_0.\) Можно сказать, что она получена применением некоторой силы, толчка. Поэтому в одной и той же точке \(x_t\) может быть не одна «правильная сила», а целое распределение возможных сил.

Теперь вместо того, чтобы сразу заменять все эти силы на среднюю, будем прикладывать их по одной. На очень маленьком промежутке времени \(dt\) случайно выберем одну скорость из условного распределения:

\[p(x_1 - x_0 \mid x_t,t)\]

и сдвинем точку:

\[x_{t+dt} = x_t + dt \cdot (x_1 - x_0).\]

Почему такая процедура сохраняет правильное распределение? Потому что в исходном процессе происходит ровно то же самое, если смотреть только на точку \(x_t\). Среди всех частиц, которые оказались в \(x_t\) в момент времени \(t\), скорости распределены по тому же условному закону:

\[p(x_1 - x_0 \mid x_t,t).\]

Значит, если мы выбираем скорость из этого распределения, то за маленький шаг \(dt\) получаем такой же локальный перенос частиц, как и в исходном процессе с прямыми траекториями (где мы соединяли \(x_0\) и \(x_1\) и шли линейно от \(x_0\) к \(x_1\)).

Теперь разобьём маленький промежуток времени \(dt\) на множество микропромежутков времени \(d \tau\). На каждом микрошаге мы можем выбрать одну из возможных скоростей из \(p(x_1 - x_0 \mid x_t,t)\). Но отдельный сдвиг очень мал, потому что он умножается на \(d \tau\). Поэтому за каждый конкретный микрошаг мы почти не сдвигаемся, \(p(x_1 - x_0 \mid x_t,t)\) не меняется, а за множество таких шагов мы в итоге сдвинемся на среднее направление всех этих скоростей, умноженное на \(dt\).

Иными словами, если в точке \(x_t\) на частицу действует много возможных сил, то на бесконечно малом масштабе их можно заменить одной средней силой:

\[u^*(x,t) = \mathbb{E}_{x_0,x_1} \left[ x_1 - x_0 \mid x_t = x \right].\]

Именно это среднее поле и нужно выучить модели. Оно не пытается восстановить конкретную пару \((x_0,x_1)\), а описывает средний поток частиц, который переносит всё распределение шума к распределению данных.

Рисунок 17. Замена множества возможных микротолчков в точке \(x\) в момент времени \(t\) на малом масштабе одним средним сдвигом \(u^*(x, t)\)
Механическая аналогия

Эту идею можно представить с помощью механической аналогии. Пусть каждая пара \((x_0, x_1)\) задаёт один «удар по мячу»: футболист стоит в точке \(x_0\), ворота находятся в точке \(x_1\), и мяч летит от футболиста к воротам по прямой траектории:

\[x_t = (1-t)x_0 + t x_1.\]

Скорость такого мяча постоянна и равна:

\[x_1 - x_0.\]

Причём футболисты не знают, в какие конкретно ворота они должны попадать, поэтому каждый бьёт в каждые ворота.

Теперь посмотрим на некоторую точку \(x\) в момент времени \(t\). Через неё могут пролетать мячи, запущенные из разных \(x_0\) в разные \(x_1\). Поэтому в одной и той же точке \(x\) мы можем увидеть много различных скоростей: один мяч летит чуть левее, другой — чуть правее, третий — почти прямо.

Представим, что в точке \(x\) находится тяжёлое подвижное заграждение. Когда через него пролетают мячи, они слегка его толкают. Каждый отдельный толчок очень маленький: за короткое время \(dt\) заграждение успевает сдвинуться только на величину порядка \(dt\). Поэтому один случайный удар почти не меняет его положения.

Рисунок 18. Механическая аналогия: разные удары проходят через одну область, а заграждение сдвигается в направлении среднего толчка

Но если за это маленькое время через точку проходит много мячей с разными скоростями, то суммарный эффект определяется не одной конкретной скоростью, а средним направлением всех этих микротолчков. То есть заграждение будет двигаться так, как если бы на него действовала средняя скорость всех мячей, пролетающих через точку \(x\) в момент \(t\) (если бы «запульнули» все мячи сразу).

Именно эта средняя скорость и задаёт поле Flow Matching’а:

\[u^*(x,t) = \mathbb{E}_{x_0,x_1} \left[ x_1 - x_0 \mid x_t = x \right].\]

То есть модель не пытается понять, какой именно мяч сейчас пролетел через точку \(x\). Вместо этого она учится предсказывать средний толчок от всех возможных мячей, которые могли оказаться в точке в данный момент времени. Такое усреднённое движение и переносит всё распределение шума к распределению данных.

Алгоритм

На каждой итерации мы берём случайный шум \(x_0\), случайный объект из данных \(x_1\) и случайное время \(t\). Затем строим промежуточную точку \(x_t\) на прямой между шумом и объектом и просим модель предсказать её скорость. На практике всё то же самое делается сразу для целого батча объектов, поэтому алгоритм ниже записан в батчевой форме.

⚙️ Алгоритм: Обучение Flow Matching-модели

Вход: Семплы из распределения \(p_{\text{data}}\) (возможность семплировать), модель векторного поля \(u_\theta\), оптимизатор для параметров \(\theta\), количество итераций \(N\), размер батча \(B\)

Результат: Обученное поле \(u_\theta \approx u^*\)

Для \(i = 1\) до \(N\):

Семплируем батч шумов, объектов и времён: \(x_0^{(b)} \sim \mathcal{N}(0,I), x_1^{(b)} \sim p_{\text{data}}, t^{(b)} \sim \mathcal{U}[0,1], b=1,\ldots,B\)

Строим промежуточные точки: \(x_t^{(b)} \gets (1-t^{(b)})x_0^{(b)} + t^{(b)}x_1^{(b)}\)

Считаем целевые скорости прямых траекторий: \(v^{(b)} \gets x_1^{(b)} - x_0^{(b)}\)

Считаем среднюю ошибку предсказания скорости по батчу: \(\mathcal{L} \gets \frac{1}{B} \sum_{b=1}^{B} \left\| u_\theta(x_t^{(b)},t^{(b)}) - v^{(b)} \right\|_2^2\)

Обновляем параметры \(\theta\) шагом оптимизатора по \(\mathcal{L}\)

После обучения у нас есть модель \(u_\theta(x,t)\), которая приближает среднее поле скоростей. Теперь мы можем использовать её для генерации новых объектов.

Идея семплирования следующая — мы стартуем из случайного шума:

\[x_0 \sim \mathcal{N}(0,I)\]

И постепенно двигаем точку по выученному векторному полю от времени \(t=0\) до времени \(t=1\). Если разбить отрезок \([0,1]\) на \(K_{\text{ode}}\) маленьких шагов, то один шаг движения можно записать так:

\[x_{t+\Delta t} = x_t + \Delta t \cdot u_\theta(x_t,t),\]

Где:

\[\Delta t = \frac{1}{K_{\text{ode}}}.\]

То есть на каждом шаге модель говорит нам, в каком направлении нужно немного сдвинуть текущую точку. После \(K_{\text{ode}}\) таких шагов мы получаем финальный объект \(x_1\), который должен быть похож на данные. Если нужно сгенерировать множество объектов, тот же цикл обычно выполняется параллельно для батча начальных шумов.

⚙️ Алгоритм: Семплирование с помощью Flow Matching’а

Вход: Обученное поле скоростей \(u_\theta\), количество шагов семплирования \(K_{\text{ode}}\)

Результат: Сгенерированный объект \(\hat{x}\)

Семплируем начальный шум: \(x \sim \mathcal{N}(0,I)\)

Задаём размер шага по времени: \(\Delta t \gets \frac{1}{K_{\text{ode}}}\)

Для \(k = 0\) до \(K_{\text{ode}}-1\)

Текущее время: \(t_k \gets k\Delta t\)

Предсказываем скорость в текущей точке: \(v_k \gets u_\theta(x,t_k)\)

Делаем маленький шаг по полю: \(x \gets x + \Delta t \cdot v_k\)

\(\hat{x} \gets x\)

Здесь используется самый простой численный метод — метод Эйлера. Он буквально повторяет нашу интуицию: посмотреть на текущую скорость, сделать маленький шаг в этом направлении, затем снова посмотреть на скорость и снова сделать шаг. Чем больше \(K\), тем точнее мы следуем выученному полю, но тем медленнее становится генерация.

В следующем скрытом блоке представлен формальный вывод Flow Matching’а через пробные функции, а также разбор уравнения непрерывности. Это позволит перейти от интуиции к математическому пониманию метода.

Формальный вывод

Итак, давайте более формально с помощью пробных функций покажем, почему Flow Matching учит именно поле:

\[u^*(x,t) = \mathbb{E}_{x_0\sim\mathcal{N}(0,I),\,x_1\sim p_{\text{data}}} \left[ x_1 - x_0 \mid x_t = x \right],\]

И почему, начиная с \(x_0\) и двигаясь по \(u^*(x_t, t)\), мы получим то же самое распределение у \(x_t\), как если бы мы просто семплировали его как \(x_t = (1-t)x_0 + t x_1\).

Рассмотрим случайный процесс:

\[x_t = (1-t)x_0 + t x_1,\]

Где:

\[x_0 \sim \mathcal{N}(0,I), \qquad x_1 \sim p_{\text{data}}.\]

Для каждой конкретной пары \((x_0,x_1)\) скорость вдоль прямой траектории равна:

\[\frac{d x_t}{dt} = x_1 - x_0.\]

Пусть \(p_t(x)\) — плотность распределения промежуточных точек \(x\) в момент времени \(t\). Мы хотим понять, как она меняется со временем.

Возьмём произвольную гладкую пробную функцию с компактным носителем:

\[\varphi : \mathbb{R}^d \to \mathbb{R}\]

Тогда для фиксированного \(t\):

\[\mathbb{E}_{x_0\sim\mathcal{N}(0,I),\,x_1\sim p_{\text{data}}} \left[ \varphi(x_t) \right] = \int \varphi(x)\,p_t(x)\,dx.\]

Продифференцируем это ожидание по времени:

\[\frac{d}{dt} \mathbb{E}_{x_0\sim\mathcal{N}(0,I),\,x_1\sim p_{\text{data}}} \left[ \varphi(x_t) \right] = \mathbb{E}_{x_0\sim\mathcal{N}(0,I),\,x_1\sim p_{\text{data}}} \left[ \frac{d}{dt}\varphi(x_t) \right].\]

По правилу цепочки:

\[\frac{d}{dt}\varphi(x_t) = \nabla \varphi(x_t)\cdot \frac{d x_t}{dt} = \nabla \varphi(x_t)\cdot (x_1-x_0).\]

Значит:

\[\frac{d}{dt} \mathbb{E}_{x_0\sim\mathcal{N}(0,I),\,x_1\sim p_{\text{data}}} \left[ \varphi(x_t) \right] = \mathbb{E}_{x_0\sim\mathcal{N}(0,I),\,x_1\sim p_{\text{data}}} \left[ \nabla \varphi(x_t)\cdot (x_1-x_0) \right].\]

Теперь используем условное математическое ожидание по текущей точке \(x_t\):

\[\mathbb{E}_{x_0\sim\mathcal{N}(0,I),\,x_1\sim p_{\text{data}}} \left[ \nabla \varphi(x_t)\cdot (x_1-x_0) \right] = \mathbb{E}_{x_t} \left[ \mathbb{E}_{x_0\sim\mathcal{N}(0,I),\,x_1\sim p_{\text{data}}} \left[ \nabla \varphi(x_t)\cdot (x_1-x_0) \mid x_t \right] \right].\]

Так как \(\nabla\varphi(x_t)\) зависит только от \(x_t\), её можно вынести из внутреннего условного ожидания:

\[\mathbb{E}_{x_0\sim\mathcal{N}(0,I),\,x_1\sim p_{\text{data}}} \left[ \nabla \varphi(x_t)\cdot (x_1-x_0) \mid x_t \right] = \nabla \varphi(x_t) \cdot \mathbb{E}_{x_0\sim\mathcal{N}(0,I),\,x_1\sim p_{\text{data}}} \left[ x_1-x_0 \mid x_t \right].\]

Введём среднее поле скоростей:

\[u^*(x,t) = \mathbb{E}_{x_0\sim\mathcal{N}(0,I),\,x_1\sim p_{\text{data}}} \left[ x_1-x_0 \mid x_t=x \right].\]

Тогда:

\[\frac{d}{dt} \mathbb{E}_{x_0\sim\mathcal{N}(0,I),\,x_1\sim p_{\text{data}}} \left[ \varphi(x_t) \right] = \mathbb{E}_{x_t} \left[ \nabla \varphi(x_t)\cdot u^*(x_t,t) \right].\]

Переходим к интегралу по плотности \(p_t\) и получаем:

\[\frac{d}{dt} \mathbb{E}_{x_0\sim\mathcal{N}(0,I),\,x_1\sim p_{\text{data}}} \left[ \varphi(x_t) \right] = \int \nabla \varphi(x)\cdot u^*(x,t)\,p_t(x)\,dx.\]

Теперь проинтегрируем по частям. Так как \(\varphi\) имеет компактный носитель, граничные члены исчезают (на бесконечности они равны 0):

\[\int \nabla \varphi(x)\cdot u^*(x,t)\,p_t(x)\,dx = - \int \varphi(x)\, \nabla\cdot \left( p_t(x)u^*(x,t) \right) dx.\]

С другой стороны:

\[\frac{d}{dt} \mathbb{E}_{x_0\sim\mathcal{N}(0,I),\,x_1\sim p_{\text{data}}} \left[ \varphi(x_t) \right] = \frac{d}{dt} \int \varphi(x)p_t(x)\,dx = \int \varphi(x) \frac{\partial p_t(x)}{\partial t} dx.\]

Значит:

\[\int \varphi(x) \frac{\partial p_t(x)}{\partial t} dx = - \int \varphi(x) \nabla\cdot \left( p_t(x)u^*(x,t) \right) dx.\]

Или:

\[\int \varphi(x) \left[ \frac{\partial p_t(x)}{\partial t} + \nabla\cdot \left( p_t(x)u^*(x,t) \right) \right] dx = 0.\]

Это верно для любой пробной функции \(\varphi\), следовательно, мы получаем уравнение непрерывности для плотности \(p_t\):

\[\frac{\partial p_t}{\partial t} + \nabla\cdot \left( p_t u^* \right) = 0.\]

Итак, мы показали — плотности \(p_t\) из прямых траекторий:

\[x_t = (1-t)x_0 + t x_1,\]

Удовлетворяют уравнению непрерывности с полем \(u^*(x,t)\).

Но почему можно двигаться по среднему полю?

Рассмотрим новый процесс генерации. Возьмём начальную точку:

\[\tilde{x}_0 \sim \mathcal{N}(0,I)\]

И будем двигать её по ODE:

\[\frac{d \tilde{x}_t}{dt} = u^*(\tilde{x}_t,t).\]

Пусть \(q_t(x)\) — плотность распределения случайной величины \(\tilde{x}_t\).

Теперь отдельно покажем, что эта плотность \(q_t\) тоже удовлетворяет уравнению непрерывности.

Снова возьмем гладкую пробную функцию \(\varphi\):

\[\mathbb{E}_{\tilde{x}_0\sim\mathcal{N}(0,I)} \left[ \varphi(\tilde{x}_t) \right] = \int \varphi(x)q_t(x)\,dx.\]

Продифференцируем это ожидание по времени:

\[\frac{d}{dt} \mathbb{E}_{\tilde{x}_0\sim\mathcal{N}(0,I)} \left[ \varphi(\tilde{x}_t) \right] = \mathbb{E}_{\tilde{x}_0\sim\mathcal{N}(0,I)} \left[ \frac{d}{dt}\varphi(\tilde{x}_t) \right].\]

По правилу цепочки:

\[\frac{d}{dt}\varphi(\tilde{x}_t) = \nabla\varphi(\tilde{x}_t)\cdot \frac{d \tilde{x}_t}{dt}.\]

Так как \(\tilde{x}_t\) движется по ODE:

\[\frac{d \tilde{x}_t}{dt} = u^*(\tilde{x}_t,t),\]

Получаем:

\[\frac{d}{dt} \mathbb{E}_{\tilde{x}_0\sim\mathcal{N}(0,I)} \left[ \varphi(\tilde{x}_t) \right] = \mathbb{E}_{\tilde{x}_0\sim\mathcal{N}(0,I)} \left[ \nabla\varphi(\tilde{x}_t)\cdot u^*(\tilde{x}_t,t) \right].\]

Теперь перепишем это ожидание как интеграл по плотности \(q_t\):

\[\mathbb{E}_{\tilde{x}_0\sim\mathcal{N}(0,I)} \left[ \nabla\varphi(\tilde{x}_t)\cdot u^*(\tilde{x}_t,t) \right] = \int \nabla\varphi(x)\cdot u^*(x,t)q_t(x)\,dx.\]

Интегрируем по частям:

\[\int \nabla\varphi(x)\cdot u^*(x,t)q_t(x)\,dx = - \int \varphi(x) \nabla\cdot \left( q_t(x)u^*(x,t) \right) dx.\]

С другой стороны:

\[\frac{d}{dt} \mathbb{E}_{\tilde{x}_0\sim\mathcal{N}(0,I)} \left[ \varphi(\tilde{x}_t) \right] = \frac{d}{dt} \int \varphi(x)q_t(x)\,dx = \int \varphi(x) \frac{\partial q_t(x)}{\partial t} dx.\]

Значит:

\[\int \varphi(x) \frac{\partial q_t(x)}{\partial t} dx = - \int \varphi(x) \nabla\cdot \left( q_t(x)u^*(x,t) \right) dx.\]

Или:

\[\int \varphi(x) \left[ \frac{\partial q_t(x)}{\partial t} + \nabla\cdot \left( q_t(x)u^*(x,t) \right) \right] dx = 0.\]

Так как это верно для любой пробной функции \(\varphi\), получаем уравнение непрерывности для плотности \(q_t\):

\[\frac{\partial q_t}{\partial t} + \nabla\cdot \left( q_t u^* \right) = 0.\]

Теперь сравним два процесса. Для исходных прямых траекторий мы получили:

\[\frac{\partial p_t}{\partial t} + \nabla\cdot \left( p_t u^* \right) = 0, \qquad p_0=\mathcal{N}(0,I).\]

А для движения по среднему полю:

\[\frac{\partial q_t}{\partial t} + \nabla\cdot \left( q_t u^* \right) = 0, \qquad q_0=\mathcal{N}(0,I).\]

Это одно и то же уравнение с одинаковым начальным условием. При этом, в уравнении явно показывается изменение плотности \(\frac{\partial p_t}{\partial t}\) по текущим \(p_t\) и \(u^*\). Значит, если его решение единственно, то:

\[q_t = p_t\]

для всех \(t\in[0,1]\).

Зачем нужна единственность решения уравнения

Рассмотрим одномерное поле \(\frac{dx}{dt}=2\sqrt{|x|}\). Из точки \(x(0)=0\) существуют разные решения. Можно навсегда остаться в нуле \(x(t)=0\), а можно начать двигаться \(x(t)=t^2\).

Более того, можно некоторое время стоять, а затем начать движение \(x(t)= \begin{cases} 0, & 0 \le t\le \tau,\\ (t-\tau)^2, & t>\tau. \end{cases}, \tau \ge 0\)

Все эти траектории имеют одну и ту же начальную точку \(x(0)=0\) и удовлетворяют одинаковому полю. Причина в том, что функция \(u(x)=2\sqrt{|x|}\) не является локально липшицевой в окрестности точки \(x=0\).

На практике наш сетап удовлетворяет условию единственности (можно вывести из теоремы Коши — Липшица), примем это за факт не будем к нему возвращаться далее в статье 🙂

Именно это и объясняет, почему можно двигаться по среднему полю. Мы не просто «на глаз» заменили разные скорости на среднюю. Мы показали, что исходный процесс с прямыми траекториями и новый процесс, который движется по ODE, порождают одну и ту же эволюцию плотности:

\[\frac{d \tilde{x}_t}{dt} = u^*(\tilde{x}_t,t),\]

Интуитивно это означает, что в одной точке \(x\) в момент времени \(t\) разные частицы могут иметь различные скорости. Одни летят левее, другие — правее. Одни быстрее, другие — медленнее. Но изменение плотности зависит не от индивидуальной скорости каждой частицы, а от суммарного потока массы через эту точку. Он равен:

\[p_t(x)u^*(x,t).\]

Поэтому Flow Matching учит макроскопическое поле скоростей \(u^*(x,t)\), а не отдельные траектории между конкретными парами \((x_0,x_1)\). Двигаясь по этому среднему полю из шума, мы получаем ту же эволюцию распределений, что и при движении всех обучающих пар по прямым траекториям.

Итак, первая часть про обучение Flow Matching завершена. Мы показали, что суть Flow Matching’а — выучивание средних скоростей. Самое время обновить чай и перейти ко второй части — дистилляции в одношаговый генератор! 😇

Мотивация

Итак, мы разобрали, что такое Flow Matching. Давайте посмотрим, почему его не очень удобно использовать на практике 😢

После обучения Flow Matching-модели генерация требует нескольких шагов. Мы стартуем из шума:

\[x_0 \sim \mathcal{N}(0,I)\]

И численно движемся по выученному полю \(u_\theta(x,t)\) от \(t=0\) до \(t=1\):

\[x_{t+\Delta t} = x_t + \Delta t \cdot u_\theta(x_t,t).\]

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

\[\Delta t = 0.01,\]

То на отрезке \([0,1]\) получится \(100\) шагов. То есть \(100\) обращений к нейросети для генерации одного изображения. А это очень много.

Здесь возникают две основные проблемы:

  1. Вычислительная стоимость. Каждый шаг требует отдельного запуска нейросети. Даже если один проход модели работает быстро — десятки или сотни проходов могут сделать генерацию слишком медленной для практического использования.
  2. Ошибка дискретизации. В идеале мы хотели бы двигаться по непрерывной траектории. Но на практике мы делаем конечные шаги размера \(\Delta t\). Если шаг слишком большой, траектория может заметно отклоняться от той, которую задаёт непрерывное поле. Это может ухудшать качество сгенерированного изображения.
Рисунок 19. Отклонение дискретной траектории от непрерывной при крупных шагах и приход в другую конечную точку \(\hat{x}_1\). Здесь мы считаем пространство одномерным: по вертикальной оси — время, по горизонтальной — координата

Дистилляция

Отсюда возникает идея — научить отдельную модель, которая будет сразу выучивать генерацию картинок. То есть мы хотим обучить генератор \(G_\phi\) сразу строить финальный объект:

\[\hat{x}_1 = G_\phi(z), \qquad z \sim \mathcal{N}(0,I_m).\]

Такой подход называется дистилляцией: мы используем обученную многошаговую Flow Matching-модель как учителя и пытаемся передать её поведение более быстрой одношаговой модели-генератору.

Обычно под дистилляцией понимают обучение маленькой модели по выходам большой, но мы имеем в виду обучение одношагового генератора (сама модель-генератор может быть больше модели учителя).

Рисунок 20. Дистилляция учит одношаговый генератор \(G_\phi\) делать распределение \(p_G\) ближе к распределению данных \(p_1\), которое получается многошаговым учителем

Важное уточнение — мы не хотим выучить отображение из \(x_0\) в \(x_1\), которое бы повторяло предсказание Flow Matching’а. У нас \(G_\phi\) — отдельный генератор со своим латентным пространством \(z\), не связанным с \(x_0\) (может иметь даже другую размерность). На самом деле мы сделаем что-то похожее на GAN, дискриминатор которого имеет определённый вид и некоторую стохастичность из-за \(x_0\).

Есть и другие методы, например, consistency models, которые ставят своей задачей выучить генератор \(G: x_0 → x_1\) повторять \(\hat{x}_1\), полученный по \(x_0\), двигаясь по \(u_\theta(x_t, t)\) или \(u^*(x_t, t)\), но это уже тема для отдельного поста 🙂.

Построение поля по сгенерированным картинкам

Теперь мы хотим заменить многошаговое движение по Flow Matching-полю одним прямым переходом от шума к изображению.

Для этого построим генератор:

\[G_\phi : \mathbb{R}^m \to \mathbb{R}^d,\]

Он получает на вход латентный шум (его размерность может отличаться от размерности исходного пространства \(x_0\) и \(x_1\)):

\[z \sim \mathcal{N}(0,I_m)\]

И сразу выдаёт объект в пространстве данных:

\[\hat{x}_1 = G_\phi(z) \in \mathbb{R}^d.\]

Здесь \(m\) — размерность латентного шума генератора, а \(d\) — размерность данных. Например, если мы генерируем изображение размером \(32\times32\times3\), то \(d=3072\), а размерность \(m\) можно выбрать отдельно.

Во время генерации нам нужен только один запуск генератора: мы семплируем \(z\), считаем \(G_\phi(z)\) и сразу получаем изображение. Но для его обучения нужно построить вспомогательное поле, как и для Flow Matching’а.

Идея следующая — траектории поля, полученного по генератору, должны быть похожи на настоящие Flow Matching-траектории (которые мы получили по полю, построенному из реальных данных). Тогда генерации и исходные картинки совпадут.

Рисунок 21. Обучение генератора с использованием двух независимых шумов: \(z\) задаёт финальную точку \(\hat{x}_1=G_\phi(z)\), а \(x_0\) — начало траектории в пространстве данных

Начнём с построения поля для сгенерированных картинок. Повторим (практически) вывод, который мы делали для Flow Matching’а, но для сгенерированных изображений, не реальных. Для этого семплируем два независимых шума:

\[x_0 \sim \mathcal{N}(0,I_d), \qquad z \sim \mathcal{N}(0,I_m).\]

Шум \(z\) подаётся в генератор и задаёт сгенерированный объект:

\[\hat{x}_1 = G_\phi(z) \in \mathbb{R}^d.\]

А независимый шум \(x_0\) используется как начальная точка траектории в пространстве данных. Между \(x_0\) и \(\hat{x}_1\) строим прямую:

\[\hat{x}_t = (1-t)x_0 + t\hat{x}_1, \qquad t\in[0,1].\]

Так как и \(x_0\), и \(\hat{x}_1\) лежат в \(\mathbb{R}^d\), вся траектория \(\hat{x}_t\) тоже находится в пространстве данных:

\[\hat{x}_t \in \mathbb{R}^d.\]

Скорость вдоль этой прямой равна:

\[\frac{\partial \hat{x}_t}{\partial t} = \hat{x}_1 - x_0 = G_\phi(z)-x_0.\]

💡Важно: \(x_0\) и \(z\) — разные шумы.

Шум \(z\) живёт в латентном пространстве \(\mathbb{R}^m\) и нужен генератору для финальной точки \(\hat{x}_1=G_\phi(z)\). Шум \(x_0\) живёт в пространстве данных \(\mathbb{R}^d\) и нужен для Flow Matching-траектории от стандартного шума к распределению генератора.

Поэтому при \(t=0\) мы имеем обычный шум в пространстве данных:

\[\hat{x}_0 = x_0 \sim \mathcal{N}(0,I_d),\]

А при \(t=1\) получаем объект из генератора:

\[\hat{x}_1 = G_\phi(z).\]

Теперь у сгенерированных траекторий тоже есть своё среднее поле скоростей. В одной и той же точке \(x\) в момент времени \(t\) могут проходить разные траектории — они построены из различных пар \((x_0,z)\). Поэтому определим поле генератора как условное среднее:

\[u^\phi(x,t) = \mathbb{E}_{x_0\sim\mathcal{N}(0,I_d),\,z\sim\mathcal{N}(0,I_m)} \left[ G_\phi(z)-x_0 \mid \hat{x}_t = x \right].\]

Оно описывает, как в среднем движутся точки, если мы берём независимый шум \(x_0\), сгенерированный объект \(G_\phi(z)\), а затем соединяем их прямой траекторией.

То есть мы строим поле Flow Matching’а, но не на реальных данных, а на сгенерированных.

Рисунок 22. Поле генератора \(u^\phi(x,t)\) — средняя скорость траекторий, которые соединяют шум \(x_0\) со сгенерированными объектами \(G_\phi(z)\)

Лосс дистилляции

С другой стороны, у нас уже есть поле учителя \(u^*(x,t)\), заранее обученное с помощью Flow Matching’а. Оно описывает правильное движение от шума к настоящим данным.

Поэтому цель у нас — подобрать генератор \(G_\phi\) так, чтобы поле его траекторий \(u^\phi(x,t)\) совпадало с полем учителя \(u^*(x,t)\). Формально это можно записать так:

\[\min_{G_\phi} \mathbb{E}_{t,\,x_0\sim\mathcal{N}(0,I_d),\,z\sim\mathcal{N}(0,I_m)} \left[ \left\| u^*(\hat{x}_t,t) - u^\phi(\hat{x}_t,t) \right\|_2^2 \right],\]

Где:

\[\hat{x}_t = (1-t)x_0 + tG_\phi(z).\]

То есть мы семплируем \(\hat{x}_t\) по сгенерированным изображениям и учим приближать поля (скорости) в точке. При этом реальные картинки нам не нужны, достаточно предобученной модели Flow Matching’а — предобученного поля.

Рисунок 23. Приближение \(u^*\) и \(u^\phi\) семплированием из \(p_z\) и \(p_0\); \(p_1\) показаны для наглядности

Мы хотим обучить генератор так, чтобы его поле совпадало с полем на реальных данных. Еще мы знаем, что распределения в начальный момент времени \(x_0\) и \(\hat{x}_0\) совпадают (это просто гауссианы). Тогда совпадут и траектории, по которым мы движемся, и сгенерированные по ним изображения (обсудили в формальном выводе Flow Matching).

Но это пока идеальная математическая цель. На практике поле \(u^\phi(x,t)\) неизвестно явно, потому что оно само является условным средним по траекториям генератора. Поэтому теперь нам нужно придумать, как приблизить цель и получить удобный алгоритм обучения! 😉

Как вычислять лосс дистилляции (делать tractable)

В предыдущем разделе мы получили идеальную цель дистилляции:

\[\min_{G_\phi} \mathbb{E}_{t,\,x_0\sim\mathcal{N}(0,I_d),\,z\sim\mathcal{N}(0,I_m)} \left[ \left\| u^*(\hat{x}_t,t) - u^\phi(\hat{x}_t,t) \right\|_2^2 \right],\]

Где:

\[\hat{x}_1 = G_\phi(z), \qquad \hat{x}_t = (1-t)x_0 + t\hat{x}_1,\]

А поле генератора определяется как условное среднее:

\[u^\phi(x,t) = \mathbb{E}_{x_0\sim\mathcal{N}(0,I_d),\,z\sim\mathcal{N}(0,I_m)} \left[ \hat{x}_1 - x_0 \mid \hat{x}_t = x \right].\]

Проблема в том, что \(u^\phi(x,t)\) нельзя посчитать напрямую. Это условное математическое ожидание скорости по всем парам \((x_0,z)\) или, эквивалентно, \((x_0, \hat{x}_1)\), траектории через точку \(x\) в момент времени \(t\). Поэтому нам нужен способ приблизить это поле.

Для этого введём дополнительную Flow Matching-модель \(u_\psi(x,t)\). Её задача — учиться предсказывать скорость траекторий генератора \(\hat{x}_1 - x_0\), чтобы использовать эту модель для вычисления лосса.

То есть \(u_\psi\) обучается, как и Flow Matching-модель, только вместо настоящих данных \(x_1\sim p_{\text{data}}\) мы используем сгенерированные объекты:

\[\hat{x}_1 = G_\phi(z).\]

При фиксированном генераторе \(G_\phi\) модель \(u_\psi\) можно обучать обычной MSE-регрессией:

\[\min_{u_\psi} \mathbb{E}_{t,\,x_0\sim\mathcal{N}(0,I_d),\,z\sim\mathcal{N}(0,I_m)} \left[ \left\| u_\psi(\hat{x}_t,t) - (\hat{x}_1-x_0) \right\|_2^2 \right].\]

Как и раньше, минимум MSE соответствует условному среднему, поэтому в оптимуме получаем:

\[u_\psi(x,t) \approx u^\phi(x,t).\]

Такую модель \(u_\psi(x,t)\) мы будем называть фейк-моделью — она обучена на сгенерированных (фейковых) данных.

💡Здесь важно не перепутать три разных поля:

\[u^*(x,t)\]

— это фиксированное поле учителя, полученное из настоящих данных (мы выучили его заранее с помощью Flow Matching-лосса);

\[u^\phi(x,t)\]

— истинное условное среднее поле траекторий генератора (его невозможно вычислить, потому что нам нужно посчитать матож по бесконечному числу семплов);

\[u_\psi(x,t)\]

— обучаемая фейк-модель, которая приближает \(u^\phi(x,t)\).

Выпишем отдельно текущие обозначения. К ним можно будет возвращаться при необходимости 😇

ОбозначениеСмысл
\(x_0\)шум в пространстве данных, \(x_0\sim\mathcal{N}(0,I_d)\)
\(x_1\)настоящий объект из данных, \(x_1\sim p_{\text{data}}\)
\(x_t\)промежуточная точка между \(x_0\) и \(x_1\), \(x_t = (1-t)x_0 + t x_1\)
\(u^*(x,t)\)поле реальных данных или поле учителя, предобученное с помощью Flow Matching’а (\(u_\theta(x_t, t)\))
\(z\)латентный шум генератора, \(z\sim\mathcal{N}(0,I_m)\)
\(\hat{x}_1=G_\phi(z)\)объект, сгенерированный одношаговым генератором
\(\hat{x}_t\)промежуточная точка между \(x_0\) и \(\hat{x}_1\), \(\hat{x}_t = (1 - t) x_0 + t \hat{x}_1\)
\(u^\phi(x,t)\)истинное среднее поле траекторий генератора (матож, который нельзя просто получить)
\(u_\psi(x,t)\)вспомогательная фейк-модель, приближающая \(u^\phi(x,t)\)

Теперь возникает идея: если \(u_\psi\) приближает поле генератора — можно сравнивать его с полем учителя \(u^*\). То есть мы хотим, чтобы поле генератора стало ближе к полю учителя. Поэтому генератор должен менять свои выходы так, чтобы траектории от \(x_0\) к \(G_\phi(z)\) имели такое же среднее поле скоростей, как у учителя.

Тут и появляется игра \(u_\psi\) и \(G_\phi\) — первый учит поле, созданное вторым, а второй меняется в зависимости от того, что выучил первый. Можно записать игру между генератором \(G_\phi\) и вспомогательным полем \(u_\psi\):

\[\min_{G_\phi}\max_{u_\psi} \mathbb{E}_{t,\,x_0\sim\mathcal{N}(0,I_d),\,z\sim\mathcal{N}(0,I_m)} \left[ \left\| u^*(\hat{x}_t,t) - (\hat{x}_1-x_0) \right\|_2^2 - \left\| u_\psi(\hat{x}_t,t) - (\hat{x}_1-x_0) \right\|_2^2 \right].\]

Рисунок 24. Схема метода. Справа в лоссе — не разность \(u\) и \(v\), а разность \(x_t + (1 - t) u\) и \(x_1\) для наглядности (так как \(x_1 = x_t + (1 - t) (x_1 - x_0)\), то с точностью до веса это одно и то же)

То есть сначала фейк-модель \(u_\psi\) максимизирует значение внутри, и как только получается оптимум, мы делаем небольшой шаг оптимизации по \(G_\phi\). То есть \(G_\phi\) должен оптимизировать некоторый максимум по \(u_\psi\).

Мы видим, что в новой игре \(u_\psi\) также учится по лоссу Flow Matching’а на сгенерированных данных:

\[\min_{u_\psi} \mathbb{E}_{t,\,x_0\sim\mathcal{N}(0,I_d),\,z\sim\mathcal{N}(0,I_m)} \left[ \left\| u_\psi(\hat{x}_t,t) - (\hat{x}_1-x_0) \right\|_2^2 \right].\]

Потому что первый аргумент не зависит от \(u_\psi\) и \(\max -x = \min x\).

В полученной игре первое слагаемое сравнивает текущую скорость траектории генератора с полем учителя \(u^*\). Второе слагаемое вычитает ошибку лучшего поля, которое может быть выучено на траекториях самого генератора.

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

Рисунок 25. Сравнение минимаксным лоссом поле учителя \(u^*\), поле генератора \(u^\phi\) и вспомогательного поля \(u_\psi\), чтобы приблизить средний поток генератора к учителю
Сравнение минимаксной целью средних полей

Если вспомогательная фейк-модель \(u_\psi\) идеально выучила поле генератора — минимаксная цель действительно превращается в «выучивание»:

\[u^\phi(x,t) \approx u^*(x,t).\]

Зафиксируем точку \(\hat{x}_t\) и время \(t\). Обозначим скорость траектории генератора через:

\[v = \hat{x}_1 - x_0.\]

Тогда поле генератора — условное среднее этой скорости:

\[u^\phi(\hat{x}_t,t) = \mathbb{E}_{x_0,z} \left[ v \mid \hat{x}_t,t \right].\]

Для любого вектора \(a(\hat{x}_t,t)\) верно стандартное разложение MSE (мы уже делали это, когда показывали, почему при использовании MSE мы выучиваем условное среднее):

\[\mathbb{E}_{x_0,z} \left[ \left\| a(\hat{x}_t,t)-v \right\|_2^2 \mid \hat{x}_t,t \right] = \left\| a(\hat{x}_t,t) - u^\phi(\hat{x}_t,t) \right\|_2^2 + \operatorname{tr} \operatorname{Cov}_{x_0,z} \left( v \mid \hat{x}_t,t \right).\]

Теперь подставим вместо \(a\) поле учителя \(u^*\):

\[\begin{aligned} & \mathbb{E}_{x_0,z} \left[ \left\| u^*(\hat{x}_t,t)-v \right\|_2^2 \mid \hat{x}_t,t \right] \\ &= \left\| u^*(\hat{x}_t,t) - u^\phi(\hat{x}_t,t) \right\|_2^2 + \operatorname{tr} \operatorname{Cov}_{x_0,z} \left( v \mid \hat{x}_t,t \right). \end{aligned}\]

А теперь — вместо \(a\) идеальное вспомогательное поле. Если \(u_\psi\) обучено идеально:

\[u_\psi(\hat{x}_t,t) = u^\phi(\hat{x}_t,t).\]

Поэтому:

\[\begin{aligned} & \mathbb{E}_{x_0,z} \left[ \left\| u_\psi(\hat{x}_t,t)-v \right\|_2^2 \mid \hat{x}_t,t \right] \\ &= \operatorname{tr} \operatorname{Cov}_{x_0,z} \left( v \mid \hat{x}_t,t \right). \end{aligned}\]

Вычтем второе равенство из первого. Дисперсионные члены сократятся:

\[\begin{aligned} & \mathbb{E}_{x_0,z} \left[ \left\| u^*(\hat{x}_t,t)-v \right\|_2^2 - \left\| u_\psi(\hat{x}_t,t)-v \right\|_2^2 \mid \hat{x}_t,t \right] \\ &= \left\| u^*(\hat{x}_t,t) - u^\phi(\hat{x}_t,t) \right\|_2^2. \end{aligned}\]

Теперь возьмём ожидание по всем \(t\), \(x_0\) и \(z\). Получим:

\[\begin{aligned} & \mathbb{E}_{t,\hat{x}_t} \mathbb{E}_{x_0,z | t, \hat{x}_t} \left[ \left\| u^*(\hat{x}_t,t)-(\hat{x}_1-x_0) \right\|_2^2 - \left\| u_\psi(\hat{x}_t,t)-(\hat{x}_1-x_0) \right\|_2^2 \right] \\ &= \mathbb{E}_{t,\hat{x}_t} \mathbb{E}_{x_0,z | t, \hat{x}_t} \left[ \left\| u^*(\hat{x}_t,t) - u^\phi(\hat{x}_t,t) \right\|_2^2 \right] \\ &= \mathbb{E}_{t,\,x_0,z} \left[ \left\| u^*(\hat{x}_t,t) - u^\phi(\hat{x}_t,t) \right\|_2^2 \right]. \end{aligned}\]

Поэтому такая разность двух MSE-ошибок является удобным способом обучать генератор. Она заставляет каждую отдельную скорость \(\hat{x}_1-x_0\) совпадать не \(u^*\), а условные средние поля:

\[u^\phi(x,t) \approx u^*(x,t).\]

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

Рисунок 26. Минимаксная цель убирает разброс отдельных траекторий и оставляет сравнение средних полей
Линеаризация

Теперь покажем другой способ получить тот же минимаксный лосс. Этот вывод основан на приёме линеаризации (linearization trick), который используется в нашей статье RealUID.

Напомним, что идеальная цель дистилляции состоит в согласовании двух средних полей:

\[\min_{G_\phi} \mathbb{E}_{t,\,x_0\sim\mathcal{N}(0,I_d),\,z\sim\mathcal{N}(0,I_m)} \left[ \left\| u^*(\hat{x}_t,t) - u^\phi(\hat{x}_t,t) \right\|_2^2 \right],\]

Где:

\[\hat{x}_1 = G_\phi(z), \hat{x}_t = (1-t)x_0 + t\hat{x}_1,\]

А поле генератора равно:

\[u^\phi(x,t) = \mathbb{E}_{x_0\sim\mathcal{N}(0,I_d),\,z\sim\mathcal{N}(0,I_m)} \left[ \hat{x}_1-x_0 \mid \hat{x}_t=x \right].\]

Проблема в том, что здесь стоит разность с условным математическим ожиданием \(u^\phi\). Напрямую считать и дифференцировать такое выражение нельзя — нужно бесконечное число раз семплировать.

Запишем разность полей в другой форме. Так как \(u^*(\hat{x}_t,t)\) при фиксированных \(\hat{x}_t\) и \(t\) является константой, имеем:

\[\begin{aligned} u^*(\hat{x}_t,t)-u^\phi(\hat{x}_t,t) &= u^*(\hat{x}_t,t) - \mathbb{E}_{x_0,z} \left[ \hat{x}_1-x_0 \mid \hat{x}_t,t \right] \\ &= \mathbb{E}_{x_0,z} \left[ u^*(\hat{x}_t,t)-(\hat{x}_1-x_0) \mid \hat{x}_t,t \right]. \end{aligned}\]

Обозначим:

\[\zeta = u^*(\hat{x}_t,t)-(\hat{x}_1-x_0).\]

Тогда разность средних полей можно записать как:

\[u^*(\hat{x}_t,t)-u^\phi(\hat{x}_t,t) = \mathbb{E}_{x_0,z} \left[ \zeta \mid \hat{x}_t,t \right].\]

Значит, внутри идеальной цели стоит выражение вида:

\[\left\| \mathbb{E}_{x_0,z} \left[ \zeta \mid \hat{x}_t,t \right] \right\|_2^2.\]

Теперь используем простое тождество (приём линеаризации):

\[\|a\|_2^2 = \max_s \left[ -\|s\|_2^2 + 2\langle s,a\rangle \right].\]

Максимум достигается при \(s=a\), поэтому и слева, и справа производная по \(a\) равна \(2a\) (\(2s = 2a\) для выражения справа). Равенство значений и производных по \(a\) показывают, что мы можем эквивалентно подставлять выражение справа вместо выражения слева при подсчете градиентов.

Применим это тождество к:

\[a = \mathbb{E}_{x_0,z} \left[ \zeta \mid \hat{x}_t,t \right].\]

Тогда:

\[\left\| \mathbb{E}_{x_0,z} \left[ \zeta \mid \hat{x}_t,t \right] \right\|_2^2 = \max_s \left[ -\|s(\hat{x}_t,t)\|_2^2 + 2 \left\langle s(\hat{x}_t,t), \mathbb{E}_{x_0,z} \left[ \zeta \mid \hat{x}_t,t \right] \right\rangle \right].\]

Так как \(s(\hat{x}_t,t)\) зависит только от \(\hat{x}_t\) и \(t\), его можно внести внутрь условного ожидания:

\[\begin{aligned} & -\|s(\hat{x}_t,t)\|_2^2 + 2 \left\langle s(\hat{x}_t,t), \mathbb{E}_{x_0,z} \left[ \zeta \mid \hat{x}_t,t \right] \right\rangle \\ &= \mathbb{E}_{x_0,z} \left[ -\|s(\hat{x}_t,t)\|_2^2 + 2 \left\langle s(\hat{x}_t,t), \zeta \right\rangle \mid \hat{x}_t,t \right]. \end{aligned}\]

Теперь в выражении больше нет нормы от условного математического ожидания. Мы заменили её на максимум по вспомогательной функции \(s\), а внутри ожидания осталось линейное выражение по \(\zeta\). Именно поэтому этот шаг называется приёмом линеаризации.

Параметризуем вспомогательную функцию через дополнительную нейросеть \(u_\psi\):

\[s_\psi(x,t) = u^*(x,t)-u_\psi(x,t).\]

Тогда в оптимуме \(s_\psi\) должен приближать:

\[u^*(x,t)-u^\phi(x,t),\]

Значит, \(u_\psi\) должен приближать поле генератора \(u^\phi\).

Подставим:

\[s_\psi(\hat{x}_t,t) = u^*(\hat{x}_t,t)-u_\psi(\hat{x}_t,t)\]

И:

\[\zeta = u^*(\hat{x}_t,t)-(\hat{x}_1-x_0)\]

В максимизируемое выражение. Получаем:

\[\begin{aligned} & -\left\| u^*(\hat{x}_t,t)-u_\psi(\hat{x}_t,t) \right\|_2^2 \\ &\quad + 2 \left\langle u^*(\hat{x}_t,t)-u_\psi(\hat{x}_t,t), u^*(\hat{x}_t,t)-(\hat{x}_1-x_0) \right\rangle. \end{aligned}\]

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

\[a = u^*(\hat{x}_t,t), \qquad b = u_\psi(\hat{x}_t,t), \qquad v = \hat{x}_1-x_0.\]

Тогда:

\[-\|a-b\|_2^2 + 2\langle a-b,a-v\rangle = \|a-v\|_2^2 - \|b-v\|_2^2.\]

Возвращаемся к исходным обозначениям, получаем:

\[\begin{aligned} & -\left\| u^*(\hat{x}_t,t)-u_\psi(\hat{x}_t,t) \right\|_2^2 \\ &\quad + 2 \left\langle u^*(\hat{x}_t,t)-u_\psi(\hat{x}_t,t), u^*(\hat{x}_t,t)-(\hat{x}_1-x_0) \right\rangle \\ &= \left\| u^*(\hat{x}_t,t)-(\hat{x}_1-x_0) \right\|_2^2 - \left\| u_\psi(\hat{x}_t,t)-(\hat{x}_1-x_0) \right\|_2^2. \end{aligned}\]

В итоге вместо:

\[\min_{G_\phi} \mathbb{E}_{t,\,x_0\sim\mathcal{N}(0,I_d),\,z\sim\mathcal{N}(0,I_m)} \left[ \left\| u^*(\hat{x}_t,t) - u^\phi(\hat{x}_t,t) \right\|_2^2 \right],\]

Получаем минимаксную задачу:

\[\min_{G_\phi}\max_{u_\psi} \mathbb{E}_{t,\,x_0\sim\mathcal{N}(0,I_d),\,z\sim\mathcal{N}(0,I_m)} \left[ \left\| u^*(\hat{x}_t,t)-(\hat{x}_1-x_0) \right\|_2^2 - \left\| u_\psi(\hat{x}_t,t)-(\hat{x}_1-x_0) \right\|_2^2 \right].\]

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

\[\hat{x}_1-x_0.\]

Роль \(u_\psi\) следующая: он выучивает среднее поле текущего генератора. После этого генератор меняется так, чтобы оно приблизилось к полю учителя \(u^*\).

Рисунок 27. Линеаризация — превращение трудной цели со средним полем в удобную игру генератора и вспомогательного поля

Алгоритм обучения

Теперь соберём всё в практический алгоритм обучения.

У нас есть три модели:

  1. Предобученное поле учителя \(u^*(x,t)\), полученное обычным Flow Matching’ом;
  2. Одношаговый генератор \(G_\phi(z)\), который мы хотим обучить;
  3. Вспомогательное поле \(u_\psi(x,t)\), которое приближает среднее поле текущего генератора.
Шаг 1: обновляем вспомогательное поле \(u_\psi\)

Здесь мы хотим обновлять только параметры \(u_\psi\). Генератор \(G_\phi\) используется только для получения текущих сгенерированных объектов, но сам генератор на этом шаге не обновляется.

Семплируем:

\[x_0 \sim \mathcal{N}(0,I_d), z \sim \mathcal{N}(0,I_m), t \sim \mathcal{U}[0,1].\]

Затем строим:

\[\hat{x}_1 = \operatorname{stopgrad}(G_\phi(z)), \hat{x}_t = (1-t)x_0 + t\hat{x}_1.\]

Скорость прямой траектории генератора равна:

\[v = \hat{x}_1 - x_0.\]

Теперь обучаем \(u_\psi\) предсказывать эту скорость:

\[\mathcal{L}_\psi = \mathbb{E}_{t,\,x_0\sim\mathcal{N}(0,I_d),\,z\sim\mathcal{N}(0,I_m)} \left[ \left\| u_\psi(\hat{x}_t,t) - v \right\|_2^2 \right].\]

Это MSE-регрессия. Поэтому при достаточно хорошем обучении фейк-модель \(u_\psi\) приближает условное среднее поле генератора:

\[u_\psi(x,t) \approx u^\phi(x,t).\]

Делаем такие шаги обновления несколько раз.

Шаг 2: обновляем генератор \(G_\phi\)

Здесь мы хотим обновлять только параметры генератора \(G_\phi\). Поля \(u^*\) и \(u_\psi\) используются как фиксированные функции.

Семплируем:

\[x_0 \sim \mathcal{N}(0,I_d), z \sim \mathcal{N}(0,I_m), t \sim \mathcal{U}[0,1].\]

Теперь считаем выход генератора без stop-gradient:

\[\hat{x}_1 = G_\phi(z), \hat{x}_t = (1-t)x_0 + t\hat{x}_1, v = \hat{x}_1 - x_0.\]

После этого минимизируем лосс генератора:

\[\mathcal{L}_G = \mathbb{E}_{t,\,x_0\sim\mathcal{N}(0,I_d),\,z\sim\mathcal{N}(0,I_m)} \left[ \left\| u^*(\hat{x}_t,t) - v \right\|_2^2 - \left\| u_\psi(\hat{x}_t,t) - v \right\|_2^2 \right].\]

Интуитивно этот шаг меняет генератор так, чтобы среднее поле его траекторий стало ближе к полю учителя \(u^*\).

Теперь посмотрим на псевдокод (в батчевой форме), чтобы ещё лучше закрепить материал 😎

⚙️ Алгоритм: Дистилляция Flow Matching’а в одношаговый генератор

Вход: Предобученное поле учителя \(u^*\), генератор \(G_\phi\), вспомогательное поле \(u_\psi\), оптимизаторы для \(\phi\) и \(\psi\), количество итераций \(N\), количество шагов обновления \(u_\psi\) на один шаг генератора \(K_\psi\), размер батча \(B\)

Результат: Одношаговый генератор \(G_\phi\)

Для \(i = 1\) до \(N\)

Шаг 1: обновляем вспомогательное поле \(u_\psi\)

На этом шаге обновляем только \(u_\psi\)

Для \(k = 1\) до \(K_\psi\)

Семплируем батч: \(x_0^{(b)} \sim \mathcal{N}(0,I_d), z^{(b)} \sim \mathcal{N}(0,I_m), t^{(b)} \sim \mathcal{U}[0,1], b=1,\ldots,B\)

Считаем выход генератора со stop-gradient: \(\hat{x}_1^{(b)} \gets \operatorname{stopgrad}(G_\phi(z^{(b)})).\)

Строим промежуточные точки: \(\hat{x}_t^{(b)} \gets (1-t^{(b)})x_0^{(b)} + t^{(b)}\hat{x}_1^{(b)}.\)

Целевые скорости: \(v^{(b)} \gets \hat{x}_1^{(b)} - x_0^{(b)}.\)

Считаем средний лосс для \(u_\psi\) по батчу: \(\mathcal{L}_\psi \gets \frac{1}{B} \sum_{b=1}^{B} \left\| u_\psi(\hat{x}_t^{(b)},t^{(b)}) - v^{(b)} \right\|_2^2.\)

Обновляем \(\psi\) шагом оптимизатора по \(\mathcal{L}_\psi\)

Шаг 2: обновляем генератор \(G_\phi\)

На этом шаге обновляем только \(G_\phi\) Поля \(u^*\) и \(u_\psi\) используются как фиксированные функции Важно: не используем torch.no_grad() для \(u^*(\hat{x}_t^{(b)},t^{(b)})\) и \(u_\psi(\hat{x}_t^{(b)},t^{(b)})\) или .detach()для \(\hat{x}_t^{(b)}\).

Семплируем батч: \(x_0^{(b)} \sim \mathcal{N}(0,I_d), z^{(b)} \sim \mathcal{N}(0,I_m), t^{(b)} \sim \mathcal{U}[0,1], b=1,\ldots,B\)

Считаем выход генератора без stop-gradient: \(\hat{x}_1^{(b)} \gets G_\phi(z^{(b)}).\)

Строим промежуточные точки: \(\hat{x}_t^{(b)} \gets (1-t^{(b)})x_0^{(b)} + t^{(b)}\hat{x}_1^{(b)}.\)

Целевые скорости текущих траекторий генератора: \(v^{(b)} \gets \hat{x}_1^{(b)} - x_0^{(b)}.\)

Считаем средний лосс генератора по батчу: \(\mathcal{L}_G \gets \frac{1}{B} \sum_{b=1}^{B} \Big[ \left\| u^*(\hat{x}_t^{(b)},t^{(b)}) - v^{(b)} \right\|_2^2 - \left\| u_\psi(\hat{x}_t^{(b)},t^{(b)}) - v^{(b)} \right\|_2^2 \Big].\)

Обновляем \(\phi\) шагом оптимизатора по \(\mathcal{L}_G\)

После обучения для генерации больше не нужны ни поле учителя \(u^*\), ни вспомогательное поле \(u_\psi\). Мы просто семплируем латентный шум:

\[z \sim \mathcal{N}(0,I_m)\]

И один раз применяем генератор:

\[\hat{x}_1 = G_\phi(z).\]

Так многошаговая Flow Matching-генерация заменяется одним шагом.

Следующий скрытый блок — подробный практический алгоритм дистилляции.

Обучение на практике

Рассмотрим нюансы при обучении.

Поле учителя \(u^*\) во время дистилляции не обучается. Оно уже знает, как правильно переносить шум в распределение данных. Поэтому мы один раз переводим его в режим eval и замораживаем параметры для скорости:

\[u^*\texttt{.eval()}, u^*\texttt{.requires_grad_(False)}.\]

Однако здесь есть важная тонкость. Заморозить параметры модели — не то же самое, что запретить градиент через её вход. При обновлении генератора параметры \(u^*\) не должны меняться, но градиент должен проходить через значение \(u^*(\hat{x}_t,t)\) к точке \(\hat{x}_t\), потому что \(\hat{x}_t\) зависит от выхода генератора \(G_\phi(z)\).

Поэтому при обновлении генератора нельзя считать \(u^*(\hat{x}_t,t)\) внутри torch.no_grad() и применять .detach()для \(\hat{x}_t\).

Рисунок 28. Заморозка параметров \(u^*\) и \(u_\psi\) при обновлении генератора, но градиент должен пройти через их вход \(\hat{x}_t\) обратно к \(G_\phi\)

💡

Частая ошибка: при обновлении генератора нельзя использовать torch.no_grad() для вычисления \(u^*(\hat{x}_t,t)\) и \(u_\psi(\hat{x}_t,t)\) или применять .detach()для \(\hat{x}_t\). Параметры этих моделей заморожены, но градиент должен проходить через их вход \(\hat{x}_t\) к генератору \(G_\phi\).

В практическом коде величины ниже обычно являются батчами. Например, для изображений \(x_0\), \(\hat{x}_1\), \(\hat{x}_t\) и \(v\) имеют форму \([B, C, H, W]\), а время \(t\) удобно хранить в форме \([B, 1, 1, 1]\), чтобы оно автоматически broadcast’илось по каналам и пикселям.

Ещё раз приведем алгоритм обучения, но с более подробным техническим описанием.

Шаг 1: обновляем вспомогательное поле \(u_\psi\)

Здесь мы хотим обновлять параметры \(u_\psi\). Генератор \(G_\phi\) используется только для получения текущих сгенерированных объектов, но сам генератор на этом шаге не обновляется.

Поэтому:

  • \(u_\psi\) переводим в режим train и разрешаем градиенты по его параметрам;
  • \(G_\phi\) переводим в режим eval и замораживаем параметры;
  • выход генератора \(\hat{x}_1=G_\phi(z)\) считаем со stop-gradient, например, через torch.no_grad() или .detach();
  • \(u^*\) на этом шаге не используем.

Семплируем:

\[x_0 \sim \mathcal{N}(0,I_d), z \sim \mathcal{N}(0,I_m), t \sim \mathcal{U}[0,1].\]

Затем строим:

\[\hat{x}_1 = \operatorname{stopgrad}(G_\phi(z)), \hat{x}_t = (1-t)x_0 + t\hat{x}_1.\]

Скорость прямой траектории генератора равна:

\[v = \hat{x}_1 - x_0.\]

Теперь обучаем \(u_\psi\) предсказывать скорость:

\[\mathcal{L}_\psi = \mathbb{E}_{t,\,x_0\sim\mathcal{N}(0,I_d),\,z\sim\mathcal{N}(0,I_m)} \left[ \left\| u_\psi(\hat{x}_t,t) - v \right\|_2^2 \right].\]

Это обычная MSE-регрессия. Поэтому при достаточно хорошем обучении \(u_\psi\) приближает условное среднее поле генератора:

\[u_\psi(x,t) \approx u^\phi(x,t).\]

Шаг 2: обновляем генератор \(G_\phi\)

Здесь мы хотим обновлять только параметры генератора \(G_\phi\). Поля \(u^*\) и \(u_\psi\) используются как фиксированные функции.

Поэтому:

  • \(G_\phi\) переводим в режим train и разрешаем градиенты по его параметрам;
  • \(u^*\) оставляем в режиме eval с замороженными параметрами;
  • \(u_\psi\) переводим в режим eval и замораживаем его параметры;
  • не используем torch.no_grad() вокруг \(u^*(\hat{x}_t,t)\) и \(u_\psi(\hat{x}_t,t)\);
  • не делаем \(\hat{x}_t.\text{detach}()\);
  • проводим градиент через \(u^*(\hat{x}_t,t)\), \(u_\psi(\hat{x}_t,t)\), \(\hat{x}_t\), \(\hat{x}_1\) и далее в параметры \(G_\phi\).

Семплируем:

\[x_0 \sim \mathcal{N}(0,I_d), z \sim \mathcal{N}(0,I_m), t \sim \mathcal{U}[0,1].\]

Теперь считаем выход генератора без stop-gradient:

\[\hat{x}_1 = G_\phi(z), \hat{x}_t = (1-t)x_0 + t\hat{x}_1, v = \hat{x}_1 - x_0.\]

После этого минимизируем лосс генератора:

\[\mathcal{L}_G = \mathbb{E}_{t,\,x_0\sim\mathcal{N}(0,I_d),\,z\sim\mathcal{N}(0,I_m)} \left[ \left\| u^*(\hat{x}_t,t) - v \right\|_2^2 - \left\| u_\psi(\hat{x}_t,t) - v \right\|_2^2 \right].\]

Интуитивно этот шаг меняет генератор так, чтобы среднее поле его траекторий стало ближе к полю учителя \(u^*\).

⚙️ Алгоритм: Дистилляция Flow Matching’а в одношаговый генератор

Вход: Предобученное поле учителя \(u^*\), генератор \(G_\phi\), вспомогательное поле \(u_\psi\), оптимизаторы для \(\phi\) и \(\psi\), количество итераций \(N\), количество шагов обновления \(u_\psi\) на один шаг генератора \(K_\psi\), размер батча \(B\)

Результат: Одношаговый генератор \(G_\phi\)

Переводим \(u^*\) в режим eval
Замораживаем параметры \(u^*\): requires_grad_(False)

Для \(i = 1\) до \(N\)

Шаг 1: обновляем вспомогательное поле \(u_\psi\)

На этом шаге обновляем только \(u_\psi\) Переводим \(u_\psi\) в режим train
Разрешаем градиенты для \(u_\psi\): requires_grad_(True)
Переводим \(G_\phi\) в режим eval
Замораживаем параметры \(G_\phi\): requires_grad_(False)

Для \(k = 1\) до \(K_\psi\)

Семплируем батч: \(x_0^{(b)} \sim \mathcal{N}(0,I_d), z^{(b)} \sim \mathcal{N}(0,I_m), t^{(b)} \sim \mathcal{U}[0,1], b=1,\ldots,B\)

Считаем выход генератора со stop-gradient: \(\hat{x}_1^{(b)} \gets \operatorname{stopgrad}(G_\phi(z^{(b)})).\)

Строим промежуточные точки: \(\hat{x}_t^{(b)} \gets (1-t^{(b)})x_0^{(b)} + t^{(b)}\hat{x}_1^{(b)}.\)

Целевые скорости: \(v^{(b)} \gets \hat{x}_1^{(b)} - x_0^{(b)}.\)

Считаем средний лосс для \(u_\psi\) по батчу: \(\mathcal{L}_\psi \gets \frac{1}{B} \sum_{b=1}^{B} \left\| u_\psi(\hat{x}_t^{(b)},t^{(b)}) - v^{(b)} \right\|_2^2.\)

Обновляем \(\psi\) шагом оптимизатора по \(\mathcal{L}_\psi\)

Шаг 2: обновляем генератор \(G_\phi\)

На этом шаге обновляем только \(G_\phi\) Поля \(u^*\) и \(u_\psi\) используются как фиксированные функции Важно: не используем torch.no_grad() для \(u^*(\hat{x}_t^{(b)},t^{(b)})\) и \(u_\psi(\hat{x}_t^{(b)},t^{(b)})\)

Переводим \(G_\phi\) в режим train
Разрешаем градиенты для \(G_\phi\): requires_grad_(True) Переводим \(u_\psi\) в режим eval Замораживаем параметры \(u_\psi\): requires_grad_(False) \(u^*\) остаётся в режиме eval и с замороженными параметрами

Семплируем батч: \(x_0^{(b)} \sim \mathcal{N}(0,I_d), z^{(b)} \sim \mathcal{N}(0,I_m), t^{(b)} \sim \mathcal{U}[0,1], b=1,\ldots,B\)

Считаем выход генератора без stop-gradient: \(\hat{x}_1^{(b)} \gets G_\phi(z^{(b)}).\)

Строим промежуточные точки: \(\hat{x}_t^{(b)} \gets (1-t^{(b)})x_0^{(b)} + t^{(b)}\hat{x}_1^{(b)}.\)

Целевые скорости текущих траекторий генератора: \(v^{(b)} \gets \hat{x}_1^{(b)} - x_0^{(b)}.\)

Считаем средний лосс генератора по батчу: \(\mathcal{L}_G \gets \frac{1}{B} \sum_{b=1}^{B} \Big[ \left\| u^*(\hat{x}_t^{(b)},t^{(b)}) - v^{(b)} \right\|_2^2 - \left\| u_\psi(\hat{x}_t^{(b)},t^{(b)}) - v^{(b)} \right\|_2^2 \Big].\)

Обновляем \(\phi\) шагом оптимизатора по \(\mathcal{L}_G\)

После обучения для генерации больше не нужны ни поле учителя \(u^*\), ни вспомогательное поле \(u_\psi\). Мы просто семплируем латентный шум:

\[z \sim \mathcal{N}(0,I_m)\]

И один раз применяем генератор:

\[\hat{x}_1 = G_\phi(z).\]

Так, многошаговая Flow Matching-генерация заменяется одним шагом.

Ура, мы на финишной прямой! Мы рассмотрели, что такое Flow Matching и как его дистиллировать. Осталось только обсудить детали обучения и посмотреть на результаты! 😎

Пара деталей

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

Что применимо и для обучения, и для дистилляции Flow Matching’а:

  • EMA

На практике мы не просто учим модель, мы сохраняем EMA-версию и уже используем её.

EMA-версия модели — скользящее среднее по параметрам модели. Из-за того, что обучение происходит нестабильно, мы берём среднее последних весов модели.

Это стандартная практика в связке с Adam/AdamW, которая дат стабильность и сильный прирост к качеству. Для начала можно взять EMA = 0.999-0.9999.

  • Gradient clipping

Ещё одна техника — обрезка значений градиентов. Это заметно улучшает стабильность обучения и является стандартом. Норму градиента для начала можно ограничить единицей.

  • Warmup

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

Warmup можно использовать для обеих моделей и выравнивать число warmup шагов из-за \(K\), чтобы они закончили его одновременно.

Warmup можно поставить, например, на 5000 шагов обучения. У основной модели базовый lr = 2e-4.

  • Dropout

Можно представить, что на каждом шаге обучения мы случайно «выключаем» часть нейронов. Поэтому модель должна решать задачу другим набором нейронов и не привыкать к одному пути вычислений. Это работает как регуляризация и помогает бороться с переобучением.

Позже мы обсудим, почему это может не сработать при оптимизации \(u_\psi\).

Dropout для начала можно поставить 0.0-0.2.

  • Реализация U-Net

Модель U-Net (которую мы использовали в статье RealUID) можно взять из torchcfm. В этом репозитории понятный код, есть веса предобученных моделей для разных задач.

Подробнее про U-Net можно почитать тут.

Специфичные детали для дистилляции:

  • Параметризация генератора

Если латентный шум имеет ту же размерность, что и данные, то есть \(m=d\), генератор удобно параметризовать как \(G_\phi(z) = z + g_\phi(z,0)\) или \(G_\phi(z) = z + g_\phi(z,\tau)\), где \(g_\phi\) — нейросеть той же архитектуры, что и поле Flow Matching’а \(u_\theta\), и \(\tau\) может быть любым. В этом случае \(g_\phi\) инициализируем весами предобученной Flow Matching-модели.

Если же \(m \neq d\) — такая остаточная форма \(z+g_\phi(z,0)\) уже не подходит напрямую, потому что \(z\) и \(G_\phi(z)\) живут в пространствах разной размерности. В этом случае генератор — отображение \(G_\phi : \mathbb{R}^m \to \mathbb{R}^d\).

  • Инициализация вспомогательного поля

Вспомогательное поле \(u_\psi\) удобно инициализировать весами предобученного поля учителя \(u^*\). В начале обучения генератор ещё «плохой», поэтому такая инициализация может сделать оптимизацию более стабильной.

Генератор также можно инициализировать весами \(u^*\).

  • Оптимизаторы

Для обеих сетей можно использовать Adam/AdamW. Для вспомогательного поля \(u_\psi\) часто полезно отключить первый момент, то есть взять Adam с \(\beta_1 = 0\). Для генератора \(G_\phi\) — применять Adam с \(\beta_1 \ge 0\). Такая настройка используется в минимаксных задачах, где одна модель играет роль критика или вспомогательного поля.

Интуиция следующая: \(\beta_1\) — momentum, то есть некоторая инерция при оптимизации. И так как внешняя игра очень сильно меняет внутреннюю, то эта инерция будет мешать, то есть модель Flow Matching не успеет за генератором.

Также рекомендуется для начала отключать dropout для вспомогательного поля \(u_\psi\) — это замедляет обучение модели.

  • Несколько шагов \(u_\psi\) на один шаг генератора

На практике полезно делать несколько обновлений \(u_\psi\) на один шаг генератора. Рабочее соотношение: \(K = 5\). То есть сначала \(5\) раз обновляем \(u_\psi\), а затем один раз обновляем \(G_\phi\). Это помогает вспомогательному полю лучше отслеживать текущее распределение траекторий генератора.

Также вместе с K можно адаптировать lr у \(u_\psi\). Вместе с \(K\) это позволит \(u_\psi\) успевать учиться за генератором.

  • Низкая скорость обучения генератора

Стоит ставить lr у генератора поменьше (например, 3e-5), иначе он будет учиться нестабильно из-за сложности лосса.

При этом lr у фейк-модели можно брать как 3e-5, так и 2e-4. Остальные гиперпараметры можно взять стандартные.

  • Заморозка моделей и stop-gradient

Очень важно правильно управлять градиентами. Когда мы обновляем \(u_\psi\), генератор \(G_\phi\) используется только для получения текущих сгенерированных объектов. Поэтому на этом шаге параметры генератора не обновляются: G_phi.eval(), G_phi.requires_grad_(False). Выход генератора можно считать со stop-gradient: \(\hat{x}_1 = \operatorname{stopgrad}(G_\phi(z))\).

Например, в PyTorch это можно сделать через torch.no_grad() или .detach(). Когда мы обновляем генератор \(G_\phi\), наоборот, параметры \(u_\psi\) и \(u^*\) замораживаются: u_psi.eval(), u_psi.requires_grad_(False), u_star.eval(), u_star.requires_grad_(False).

Но здесь нельзя использовать torch.no_grad() вокруг \(u^*(\hat{x}_t,t)\) и \(u_\psi(\hat{x}_t,t)\) и делать x_t_hat.detach(). Причина следующая: параметры \(u^*\) и \(u_\psi\) не должны обновляться, но градиенту нужно проходить через их вход \(\hat{x}_t\), потому что \(\hat{x}_t = (1-t)x_0 + tG_\phi(z)\) зависит от генератора. Если остановить этот градиент, генератор не получит правильный обучающий сигнал.

  • Постоянная фиксация учителя

Предобученное поле \(u^*\) во время всей дистилляции остаётся в режиме eval и не обновляется. Оно играет роль фиксированного учителя, который задаёт правильное направление движения от шума к данным.

Результаты обучения

Рассмотрим результаты обучения и дистилляции Flow Matching-модели на датасете CIFAR-10 — картинки размером 32 х 32 пикселя. Всего есть 10 разных классов — самолёты, машины, птицы и др.

Чтобы проверить, насколько дистилляция работает хорошо, мы будем использовать метрику FID (Fréchet Inception Distance) — она сравнивает набор фото с набором сгенерированных картинок. Чем наборы картинок ближе, тем лучше.

FID-метрика

Сначала по каждому изображению \(x_i\) из предобученной нейросети (например, Inception-v3) извлекаются числовые признаки \(h_i\). Затем мы смотрим на эти признаки как на семплированные из многомерного нормального распределения. При этом, признаки реальных и сгенерированных картинок семплируются из двух разных нормальных распределений: \(\mathcal{N(\mu, \Sigma)}\) и \(\mathcal{N(\hat{\mu}, \hat{\Sigma})}\).

Для двух нормальных распределений мы знаем метрику: \[d_F\left(\mathcal{N}(\mu,\Sigma),\mathcal{N}(\hat{\mu},\hat{\Sigma})\right)^2=\left\|\mu-\hat{\mu}\right\|_2^2+\operatorname{tr}\left(\Sigma+\hat{\Sigma}-2\left(\Sigma\hat{\Sigma}\right)^{\frac{1}{2}}\right).\]

Среднее и матрицу ковариаций распределения можно оценить по семплам: \(\mu\) — среднее семплов, матрица ковариации — через \(\Sigma=\frac{1}{N-1}(X-\mathbf{1}\mu)^T(X-\mathbf{1}\mu)\).

В статье RealUID FID для Flow Matching’а мы получили равным 3.57 на 100 шагах генерации, при этом после дистилляции FID уменьшился до 2.58. То есть мы не только в несколько раз ускоряем генерацию (1 шаг вместо 100), но и улучшаем её качество.

Мы сгенерировали несколько примеров из дистиллированной модели на CIFAR-10, чтобы увидеть результат:

Рисунок 29. Примеры генерации дистиллированной модели, выученной на CIFAR-10

Также приведём примеры генерации лиц на датасете CelebA размером 64 х 64:

Рисунок 30. Примеры генерации дистиллированной модели, выученной на CelebA

Заключение

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

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

И, наконец, мы рассмотрели идею дистилляции: заменить многошаговую генерацию одним запуском отдельного генератора \(G_\phi\), который обучается с помощью \(\min\max\) игры. В итоге мы не только ускорили генерацию, но и улучшили качество картинок.

Полезные ссылки

  1. Статья про Flow Matching.
  2. Статья про ошибку, которую минимизирует Flow Matching.
  3. TorchCFM — реализация Flow Matching’а.
  4. SiD, DMD, FGM — методы дистилляции.
  5. GAN как min-max игра.
  6. RealUID, code — наша статья про универсальный метод дистилляции с реальными данными + код к ней.
  7. Consistency models, FACM, DuMo, MeanFlow, π-Flow, LADD — другие методы дистилляции и генерации в один или несколько шагов.
  8. Rectified Flows — иной способ ускорить генерацию Flow Matching’а.
  9. Diffusion Meets Flow Matching — пост про связь диффузионных моделей и flow matching’а.
  10. DDPM, Score-Based Generative Modeling through Stochastic Differential Equations — про диффузионные модели (сильно связаны с Flow Matching’ом).
  11. VAE, LVAE, VDVAE, NCSN, Deep Unsupervised Learning using Nonequilibrium Thermodynamics, Variational Diffusion Models, Latent Diffusion, NFDM — для лучшего понимания диффузионных моделей (три взгляда на них — variational, score и flow).
  12. Статья про continuous normalizing flows (может помочь лучше понять Flow Matching).
  13. Лекции BayesGroup по continuous normalizing flows и диффузионным моделям.
  14. The Annotated Transformer — для лучшего понимания устройства используемых моделей.
  15. nn.labml.ai — для понимания реализации различных методов машинного обучения, в том числе — U-Net, Stable Diffusion.

Пожалуйста, оцените статью, это поможет нам создавать качественный материал. Спасибо!

Полезные ссылки на посты DeepSchool

  1. Пост 1 и пост 2 про дистиляцию диффузии
  2. Consistency models
  3. Про Rectified Flows и InstaFlow
  4. Введение в диффузионные модели
  5. Про модели Genie
Computer Vision Rocket

Как готовить качественные датасеты и обучать модели для задач CV рассказываем на нашем курсе CV Rocket. Изучайте подробности на сайте и записывайтесь в лист ожидания!

0/0

Телеграм-канал

DeepSchool

Короткие посты по теории ML/DL, полезные
библиотеки и фреймворки, вопросы с собеседований
и советы, которые помогут в работе

Открыть Телеграм

Увидели ошибку?

Напишите нам в Telegram!